From f0a6eed2c052d0f22796a549a2de3026e0c82e49 Mon Sep 17 00:00:00 2001 From: Spicy-cream <2667314914@qq.com> Date: Fri, 11 Sep 2026 05:10:33 -0700 Subject: [PATCH] feat: add common transport-level handshake for RDMA and UBSHM --- BUILD.bazel | 1 - CMakeLists.txt | 2 +- Makefile | 2 +- docs/cn/handshake_common_design.md | 499 ++ src/brpc/adapter_transport.cpp | 571 ++ src/brpc/adapter_transport.h | 101 + src/brpc/global.cpp | 14 +- src/brpc/handshake/handshake_adapter.cpp | 55 + src/brpc/handshake/handshake_adapter.h | 73 + src/brpc/handshake/handshake_frame.cpp | 235 + src/brpc/handshake/handshake_frame.h | 92 + src/brpc/handshake/handshake_io.cpp | 150 + src/brpc/handshake/handshake_io.h | 89 + src/brpc/handshake/rdma_handshake.cpp | 576 ++ src/brpc/handshake/rdma_handshake.h | 154 + .../rdma_handshake_constants.h | 31 +- src/brpc/handshake/ubshm_handshake.cpp | 407 ++ src/brpc/handshake/ubshm_handshake.h | 84 + src/brpc/input_messenger.h | 2 + src/brpc/policy/rdma_handshake_protocol.cpp | 14 +- src/brpc/policy/rdma_handshake_protocol.h | 27 +- .../policy/transport_handshake_protocol.cpp | 37 + .../policy/transport_handshake_protocol.h | 44 + src/brpc/rdma/rdma_endpoint.cpp | 953 +--- src/brpc/rdma/rdma_endpoint.h | 178 +- src/brpc/rdma/rdma_handshake.cpp | 514 -- src/brpc/rdma/rdma_handshake.h | 185 - src/brpc/rdma/rdma_handshake_server.cpp | 219 - src/brpc/rdma/rdma_handshake_server.h | 48 - src/brpc/{rdma => }/rdma_handshake.proto | 0 src/brpc/rdma_transport.cpp | 118 +- src/brpc/rdma_transport.h | 34 +- src/brpc/socket.h | 20 +- src/brpc/transport_factory.cpp | 10 +- src/brpc/transport_factory.h | 5 +- src/brpc/transport_handshake.cpp | 338 ++ src/brpc/transport_handshake.h | 197 + src/brpc/ubshm/ub_endpoint.cpp | 645 +-- src/brpc/ubshm/ub_endpoint.h | 125 +- src/brpc/ubshm_transport.cpp | 108 +- src/brpc/ubshm_transport.h | 21 +- test/brpc_rdma_unittest.cpp | 5042 +++++++++-------- test/brpc_transport_handshake_unittest.cpp | 463 ++ test/brpc_ubring_unittest.cpp | 40 + 44 files changed, 7399 insertions(+), 5124 deletions(-) create mode 100644 docs/cn/handshake_common_design.md create mode 100644 src/brpc/adapter_transport.cpp create mode 100644 src/brpc/adapter_transport.h create mode 100644 src/brpc/handshake/handshake_adapter.cpp create mode 100644 src/brpc/handshake/handshake_adapter.h create mode 100644 src/brpc/handshake/handshake_frame.cpp create mode 100644 src/brpc/handshake/handshake_frame.h create mode 100644 src/brpc/handshake/handshake_io.cpp create mode 100644 src/brpc/handshake/handshake_io.h create mode 100644 src/brpc/handshake/rdma_handshake.cpp create mode 100644 src/brpc/handshake/rdma_handshake.h rename src/brpc/{rdma => handshake}/rdma_handshake_constants.h (69%) create mode 100644 src/brpc/handshake/ubshm_handshake.cpp create mode 100644 src/brpc/handshake/ubshm_handshake.h create mode 100644 src/brpc/policy/transport_handshake_protocol.cpp create mode 100644 src/brpc/policy/transport_handshake_protocol.h delete mode 100644 src/brpc/rdma/rdma_handshake.cpp delete mode 100644 src/brpc/rdma/rdma_handshake.h delete mode 100644 src/brpc/rdma/rdma_handshake_server.cpp delete mode 100644 src/brpc/rdma/rdma_handshake_server.h rename src/brpc/{rdma => }/rdma_handshake.proto (100%) create mode 100644 src/brpc/transport_handshake.cpp create mode 100644 src/brpc/transport_handshake.h create mode 100644 test/brpc_transport_handshake_unittest.cpp diff --git a/BUILD.bazel b/BUILD.bazel index 727af8574a..fd315ffa28 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -522,7 +522,6 @@ filegroup( srcs = glob([ "src/brpc/*.proto", "src/brpc/policy/*.proto", - "src/brpc/rdma/*.proto", ]), visibility = ["//visibility:public"], ) diff --git a/CMakeLists.txt b/CMakeLists.txt index 47c1458625..622f587963 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -619,7 +619,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) file(MAKE_DIRECTORY ${PROJECT_BINARY_DIR}/output/include/brpc) set(PROTOC_FLAGS ${PROTOC_FLAGS} -I${PROTOBUF_INCLUDE_DIR}) compile_proto(PROTO_HDRS PROTO_SRCS ${PROJECT_BINARY_DIR} diff --git a/Makefile b/Makefile index 86de388448..90be5516c3 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 THRIFT_SOURCES = $(foreach d,$(BRPC_DIRS),$(wildcard $(addprefix $(d)/thrift*,$(SRCEXTS)))) EXCLUDE_SOURCES = $(foreach d,$(BRPC_DIRS),$(wildcard $(addprefix $(d)/event_dispatcher_*,$(SRCEXTS)))) BRPC_SOURCES_ALL = $(foreach d,$(BRPC_DIRS),$(wildcard $(addprefix $(d)/*,$(SRCEXTS)))) diff --git a/docs/cn/handshake_common_design.md b/docs/cn/handshake_common_design.md new file mode 100644 index 0000000000..169866dc6d --- /dev/null +++ b/docs/cn/handshake_common_design.md @@ -0,0 +1,499 @@ +# RDMA、URMA、UBSHM 公共握手设计(修订版) + +## 1. 文档目的 + +本文在现有公共 framing、`HandshakeSession` 和协议字段 adapter 的基础上,进一步明确四个职责边界: + +1. `Socket` 只感知一个顶层 `AdapterTransport`,不感知 TCP、RDMA、URMA、UBSHM,也不感知握手 phase。 +2. `AdapterTransport` 自动完成连接升级;升级成功或回退 TCP 后,统一通知 `Socket` 建链完成。 +3. 具体 Transport 只提供握手所需的资源操作接口;握手顺序、状态转换和 fallback 由上层统一编排。 +4. Transport 只负责数据传输和资源生命周期,不负责 TCP 建链、wire framing、握手状态机或握手流程编排。 + +本文只定义架构和迁移方向,不改变 RDMA/URMA/UBSHM 当前 wire format。 + +## 2. 当前问题 + +现有实现已经抽取了 `HandshakeSession`、`HandshakeCodec` 和公共 framing,但职责仍未完全收敛: + +- `ProcessHandshakeAtClient` 仍位于具体 Transport 或 endpoint 路径中,client 的握手入口不统一。 +- `RdmaTransport`、`UBShmTransport` 等仍通过各自的 `S_*` 常量暴露握手阶段;`HandshakePhases` 需要填入不同 transport 的 phase,公共 session 无法真正统一。 +- server handshake adapter 需要通过 `Socket` 找到 `AdapterTransport`,再找到具体 Transport,并且直接驱动资源创建、激活和 fallback。 +- `AdapterTransport` 虽然已经是 `Socket` 的顶层对象,但握手的 client/server 入口、连接完成通知和数据面切换仍分散在多个层次。 +- “握手成功”与“Socket 建链完成”不是同一个明确事件,导致 TCP fallback、升级成功和异常退出的发布顺序难以验证。 + +根因是把“传输能力”和“建链流程”混在了一起。Transport 是被流程调用的参与者,不应成为流程的拥有者。 + +## 3. 目标架构 + +### 3.1 分层 + +```text +Socket + | + v +AdapterTransport Socket 唯一感知的 Transport + |-- TcpTransport TCP 控制面和 fallback 数据面 + |-- HighSpeedTransport RDMA / URMA / UBSHM 数据面 + |-- ConnectionUpgradeCoordinator 统一的建链与升级编排 + | |-- HandshakeSession 公共状态机、I/O、framing + | |-- HandshakeCodec 协议 wire 字段编解码 + | |-- TransportUpgradeOps 被调用的资源阶段接口 + | `-- SocketConnectionNotifier 建链完成/失败通知 + `-- ActiveTransport 根据统一状态选择数据面 +``` + +这里的 `ConnectionUpgradeCoordinator` 可以先作为 `AdapterTransport` 的内部实现,不要求立即新增独立公开类;重要的是职责必须集中在该层,而不是分散到具体 Transport。 + +### 3.2 依赖方向 + +```text +Socket -> AdapterTransport -> TcpTransport / HighSpeedTransport +Socket -> AdapterTransport -> ConnectionUpgradeCoordinator +ConnectionUpgradeCoordinator -> TransportUpgradeOps +ConnectionUpgradeCoordinator -> HandshakeSession +TransportUpgradeOps -> transport-specific endpoint/resource +``` + +禁止以下反向依赖: + +- `Socket` 直接调用具体 Transport 或 handshake adapter。 +- `TcpTransport` 调用具体高速度 Transport 的握手函数。 +- `RdmaTransport`、`UrmaTransport`、`UBShmTransport` 调用 `ProcessHandshakeAtClient`、`RunServerHandshake` 等流程函数。 +- 公共 `HandshakeSession` 依赖 `ibv_*`、`urma_*`、UBRING 类型。 +- 具体 Transport 通过自定义 phase 影响 `Socket` 的建链判断。 + +## 4. 统一状态模型 + +### 4.1 握手状态由上层拥有 + +公共状态只描述连接升级流程,不描述某一种资源如何创建: + +```cpp +namespace brpc { +namespace handshake { + +enum class Phase { + kUninitialized, + kPreparing, + kHelloSending, + kHelloWaiting, + kNegotiating, + kAckSending, + kAckWaiting, + kEstablished, + kFallbackTcp, + kFailed, +}; + +enum class StepResult { + kOk, + kFallback, + kNeedMore, + kNotMine, + kError, +}; + +} // namespace handshake +} // namespace brpc +``` + +`Socket` 和 `AdapterTransport` 只读取 `Phase` 的终态:`kEstablished`、`kFallbackTcp`、`kFailed`。中间状态只由 coordinator/session 使用。 + +### 4.2 删除 `HandshakePhases` + +不再使用: + +```cpp +struct HandshakePhases { + int prepare_local; + int hello_send; + int hello_wait; + int negotiate; + int ack_send; + int ack_wait; +}; +``` + +原因是这些数字并不是协议字段,也不是数据面状态;它们只是历史实现中的日志/状态映射。公共 session 应使用固定的 `Phase`,具体 Transport 的资源子状态留在自身实现中,不向上暴露。 + +例如 RDMA 的 `S_ALLOC_QPCQ`、`S_BRINGUP_QP`、UBSHM 的 `S_ALLOC_SHM` 都只能是具体 Transport 的内部 debug 状态;它们不能再填入公共 `HandshakePhases`,也不能作为 `Socket` 判断建链完成的依据。 + +## 5. 模块职责 + +### 5.1 Socket + +`Socket` 只保存和调用一个 `Transport*`,实际对象始终为 `AdapterTransport`。它只关心以下事件: + +- TCP socket 是否连接成功; +- `AdapterTransport` 是否报告建链完成; +- 当前数据面读写是否可用; +- 连接是否失败或 EOF。 + +Socket 不直接读取 handshake phase,也不负责决定 TCP 还是高速度 Transport。 + +### 5.2 AdapterTransport + +`AdapterTransport` 是顶层 Transport 和连接升级协调器的宿主,负责: + +- 创建并持有 TCP Transport 和候选高速度 Transport; +- 建立 TCP 控制连接; +- 自动启动 client/server 的升级编排; +- 将 TCP 可读事件交给 coordinator,直到升级进入终态; +- 在 `kEstablished` 与 `kFallbackTcp` 之间选择 active data transport; +- 对升级结果执行一次性的连接完成通知; +- 保证 fallback 数据回放和状态发布的内存顺序。 + +建议的内部接口如下: + +```cpp +class AdapterTransport : public Transport { +public: + std::shared_ptr Connect() override; + void ProcessEvent(bthread_attr_t attr) override; + + // 由上层 Socket/连接流程使用,不暴露具体 Transport 类型。 + handshake::Phase connection_phase() const; + +private: + void StartClientUpgrade(); + void ProcessUpgradeReadable(); + void CompleteConnection(handshake::Phase terminal_phase); + void ActivateTcp(); + void ActivateHighSpeed(); + + handshake::HandshakeSession _handshake; + std::unique_ptr _tcp_transport; + std::unique_ptr _high_speed_transport; + std::unique_ptr _upgrade; +}; +``` + +`StartClientUpgrade` 和 `ProcessUpgradeReadable` 是上层编排入口。它们不能下沉到 `RdmaTransport`、`UBShmTransport` 或 endpoint。 + +### 5.3 ConnectionUpgradeCoordinator / HandshakeSession + +该模块拥有连接升级的完整流程: + +- 选择协议 codec; +- 通过 TCP 发送和接收 hello/ack; +- 增量 framing、半包和非本协议数据回放; +- 按固定顺序调用 Transport 能力接口; +- 将协议不支持、资源失败转换为 fallback; +- 发布 `kEstablished`、`kFallbackTcp` 或 `kFailed`; +- 在终态时通知 `AdapterTransport` 完成建链。 + +`HandshakeSession` 不知道 RDMA、URMA、UBSHM 的资源类型,只调用抽象的阶段接口。 + +### 5.4 具体 Transport + +具体 Transport 只负责: + +- 数据面读写、事件和 completion; +- 资源创建、导入、激活、停用和释放; +- 提供本端 hello 所需的只读能力/字段; +- 接收上层解析后的远端参数; +- 报告资源阶段成功、不可用或失败。 + +具体 Transport 不负责: + +- TCP fd 读写; +- magic/length framing; +- hello/ack 的收发顺序; +- server 增量解析; +- fallback 决策; +- `Socket` 建链完成通知; +- 公共 handshake phase 的推进。 + +## 6. Transport 能力接口 + +Transport 提供的是“被上层调用的资源能力”,不是 handshake driver。建议抽象为以下接口;具体命名可根据现有类调整: + +```cpp +class TransportUpgradeOps { +public: + virtual ~TransportUpgradeOps() = default; + + // 返回本端能力和协议 adapter 所需的字段来源。 + virtual handshake::StepResult PrepareLocal() = 0; + + // 上层完成 wire parse/validate 后交付远端参数。 + virtual handshake::StepResult ApplyRemote(const RemoteParameters&) = 0; + + // 创建、导入或建立数据面资源。 + virtual handshake::StepResult PrepareResources() = 0; + virtual handshake::StepResult NegotiateResources() = 0; + + // 握手 ACK 已发送/确认后切换数据面。 + virtual handshake::StepResult Activate() = 0; + + // fallback 或失败时关闭本次升级准备的资源。 + virtual void Deactivate() = 0; +}; +``` + +说明: + +- `PrepareLocal`、`ApplyRemote`、`PrepareResources`、`NegotiateResources` 的具体数量可以按现有资源模型合并,但调用顺序由 coordinator 固定。 +- 只有 `TransportUpgradeOps` 的实现可以访问 endpoint 和 transport-specific 类型。 +- 如果某种 Transport 不支持升级,应返回 `kFallback`,而不是创建一套假的握手状态机。 +- `Activate` 成功后,coordinator 才能发布 `kEstablished`。 +- 资源失败默认进入 `kFallbackTcp`;不可恢复的协议错误才进入 `kFailed`,具体策略由 coordinator 统一决定。 + +## 7. 协议 adapter 与公共编排的边界 + +每一种 wire protocol 保留自己的 `HandshakeCodec` 和字段 adapter: + +```cpp +struct HandshakeCodec { + int protocol_version; + FrameSpec hello_frame; + FrameSpec ack_frame; + std::function build_hello; + std::function parse_hello; + std::function build_ack; + std::function parse_ack; +}; +``` + +adapter 只处理以下内容: + +- magic、版本、字段序列化和反序列化; +- payload 合法性检查; +- 将字段转换为 `RemoteParameters`; +- 生成本端 hello 和 ack。 + +adapter 不再实现完整的 `RunServerHandshake` 或 `ProcessHandshakeAtClient`。这些函数中的流程代码应迁移到 coordinator;adapter 只提供 codec 和 `TransportUpgradeOps` 所需的字段转换。 + +server 端的 `HandshakeAdapter::ExecuteServerHandshake` 如果暂时需要保留以兼容 `InputMessenger`,其实现只能是薄适配层: + +```text +InputMessenger -> AdapterTransport/Coordinator -> HandshakeSession + | + `-> protocol codec + TransportUpgradeOps +``` + +它不应再根据 `SocketMode` 选择不同的 phase,也不应直接调用具体 Transport 的握手流程。 + +## 8. Client 建链时序 + +```mermaid +sequenceDiagram + participant S as Socket + participant A as AdapterTransport + participant C as Coordinator + participant T as TcpTransport + participant H as HandshakeSession + participant U as TransportUpgradeOps + + S->>A: Connect() + A->>T: 建立 TCP 控制连接 + T-->>A: TCP connected + A->>C: StartClientUpgrade() + C->>U: PrepareLocal() + C->>H: 发送 local hello + H->>T: WriteFrame() + T-->>H: 接收 remote hello + C->>U: ApplyRemote() / PrepareResources() + C->>U: NegotiateResources() + alt 升级成功 + C->>H: 发送 enabled ACK + C->>U: Activate() + C->>A: kEstablished + A->>S: ConnectionReady(high-speed) + else 对端不支持或资源失败 + C->>H: 发送 disabled ACK(如协议要求) + C->>U: Deactivate() + C->>A: kFallbackTcp + A->>S: ConnectionReady(tcp) + end +``` + +关键约束:`ConnectionReady` 只发送一次,并且必须发生在 active transport 已设置、fallback 缓冲区已回放、终态已 release-store 之后。 + +## 9. Server 建链时序 + +server 收到 TCP 数据后,由 `AdapterTransport` 的上层 coordinator 处理;具体 Transport 不参与入口选择: + +```text +TCP readable + -> AdapterTransport::ProcessUpgradeReadable + -> HandshakeSession::RunServer(input) + -> 根据 magic 选择 codec + -> 增量读取完整 hello + -> codec parse/validate + -> TransportUpgradeOps::ApplyRemote/NegotiateResources + -> 发送 ACK + -> Activate 或 Deactivate + -> AdapterTransport::CompleteConnection +``` + +对于 `IOBuf` 增量输入: + +- `kNeedMore` 时不得消费不完整 frame; +- magic 不匹配时必须把已检查的数据回放给 TCP parser; +- 已确认属于升级协议但资源失败时不能把完整握手 frame 当作普通 TCP 数据; +- ACK 没有 magic 时,选中的 codec 必须保存在 coordinator/session context 中,而不能依赖重新探测。 + +## 10. 连接终态与 Socket 通知 + +### 10.1 终态定义 + +| 终态 | active data transport | Socket 结果 | 说明 | +|---|---|---|---| +| `kEstablished` | 高速度 Transport | 建链完成 | `Activate()` 成功后发布 | +| `kFallbackTcp` | TCP Transport | 建链完成 | 协议不支持或资源不可用 | +| `kFailed` | 无 | 建链失败 | I/O、协议或不可恢复错误 | + +`kFallbackTcp` 不是失败。它表示控制连接成功且连接可继续使用 TCP。 + +### 10.2 发布顺序 + +统一采用以下顺序: + +```text +1. 设置 active transport +2. 回放 fallback 时已读但不属于握手的数据(仅 TCP fallback) +3. 发布终态 phase(release) +4. 通知 Socket::ConnectionReady / ConnectionFailed +``` + +事件线程观察到终态后,才能读取 active transport 和回放数据。禁止在 phase 发布后再修改 active transport。 + +### 10.3 TCP 数据保护 + +进入 `kEstablished` 后,TCP 只作为控制连接存在;如果收到额外 TCP 应用数据,应按协议错误处理,不能静默交给高速度数据面。进入 `kFallbackTcp` 后,后续数据全部交给 TCP Transport。 + +## 11. 现有实现迁移方案 + +### Phase 1:统一公共状态 + +- 用 `handshake::Phase` 替换 `HandshakePhases` 的六个 transport-specific 数值。 +- `HandshakeSession` 只写入统一 phase。 +- 保留具体 Transport 内部状态用于日志和资源调试,但不再通过 `handshake_phase()` 暴露给 Socket。 + +### Phase 2:收拢 client 入口 + +- 将所有 `ProcessHandshakeAtClient` 的调用点迁移到 `AdapterTransport::StartClientUpgrade`。 +- 删除具体 Transport 中的 client handshake driver。 +- 具体 Transport 改为实现 `TransportUpgradeOps`,只提供资源阶段操作。 +- `AdapterTransport::Connect` 在 TCP connected 后自动启动 coordinator。 + +### Phase 3:收拢 server 入口 + +- `InputMessenger` 只把 handshake 输入交给 `AdapterTransport`/coordinator。 +- `RdmaServerHandshakeAdapter`、`UBShmServerHandshakeAdapter` 等降级为 codec/字段 adapter 或薄兼容层。 +- 删除 adapter 内按 `SocketMode` 选择 `S_ACK_WAIT` 等 phase 的逻辑。 +- URMA 接入时直接实现统一的 codec 和 `TransportUpgradeOps`,不复制一套 server driver。 + +### Phase 4:统一连接完成通知 + +- 为 `AdapterTransport` 增加单一的 `CompleteConnection` 路径。 +- 升级成功、TCP fallback、失败分别通过统一终态通知 Socket。 +- 增加断言,确保 `ConnectionReady` 只发生一次,且发生在终态发布之后。 + +### Phase 5:删除旧路径 + +- 删除具体 Transport 的 `ProcessHandshakeAtClient`、`RunServerHandshake` 和握手 phase 映射。 +- 删除 Socket 对具体 Transport 类型和 transport-specific phase 的依赖。 +- 清理只为旧握手路径存在的 endpoint 回调。 + +## 12. 兼容性与测试要求 + +### 12.1 Wire compatibility + +重构不得改变: + +- RDMA v2/v3、URMA v2/v3、UBSHM v2 的 magic; +- frame length 的字节序和语义; +- hello/ack 的字段布局和 ACK bit; +- fallback hello 的兼容行为。 + +### 12.2 公共模块测试 + +至少覆盖: + +- 2 字节和 4 字节 magic; +- fixed、U16 total length、U32 body length; +- 半包、粘包和多余数据; +- `kNeedMore` 不消费不完整输入; +- magic 不匹配的 push-back; +- I/O error、EOF、超时和协议错误; +- 终态发布顺序和一次性连接通知。 + +### 12.3 集成矩阵 + +| 场景 | 预期结果 | +|---|---| +| RDMA v2 ↔ RDMA v2 | 高速度 Transport 建链 | +| RDMA v3 ↔ RDMA v3 | 高速度 Transport 建链 | +| URMA v2/v3 ↔ 对应版本 | 高速度 Transport 建链 | +| UBSHM ↔ UBSHM | UBSHM 建链 | +| 高速度 client ↔ TCP server | TCP fallback,Socket 建链成功 | +| TCP client ↔ 高速度 server | TCP 建链成功 | +| hello 合法但资源创建失败 | TCP fallback | +| 已识别握手协议后发生协议错误 | 建链失败,不回放为 TCP 数据 | + +## 13. 验收标准 + +设计落地后应满足: + +1. `Socket::_transport` 永远只指向 `AdapterTransport`,Socket 源码中没有具体 Transport 类型判断。 +2. client 和 server 都只有一个握手编排入口,代码库中不再存在具体 Transport 的 `ProcessHandshakeAtClient`。 +3. 公共握手状态只有一套 `handshake::Phase`,不存在 RDMA/UBSHM/URMA 到公共 phase 的数字映射。 +4. 具体 Transport 不读写 TCP fd,不解析 magic/length,不调用 `HandshakeSession::RunClient/RunServer`。 +5. 升级成功和 TCP fallback 都由 `AdapterTransport` 自动完成 active transport 切换,并向 Socket 发出一次建链完成通知。 +6. 在 wire compatibility 测试通过的前提下,握手流程、fallback 和连接完成通知可以通过公共 coordinator 单元测试验证。 + +## 14. 结论 + +最终边界如下: + +```text +Socket + 只感知 AdapterTransport 和 ConnectionReady + +AdapterTransport / Coordinator + 拥有 TCP-first、握手状态机、升级/fallback 编排、active transport 切换 + +HandshakeSession / Codec + 拥有公共 I/O、framing、协议字段编解码和统一状态 + +TransportUpgradeOps + 提供具体传输所需的资源阶段接口 + +RDMA / URMA / UBSHM Transport + 只拥有各自资源、数据面和 completion +``` + +一句话概括:**握手是上层连接编排,Transport 是被编排的数据传输能力;Socket 只通过 AdapterTransport 观察连接结果。** +## 15. 整体架构图 + +```mermaid +flowchart TB + S["Socket\n只感知 AdapterTransport"] + A["AdapterTransport\nTCP-first / 建链编排 / active transport"] + C["Connection Upgrade Coordinator\nStartClientUpgrade / ProcessUpgradeReadable\nCompleteConnection"] + H["HandshakeSession\n统一 framing / I/O / handshake::Phase"] + RA["Protocol Codec Adapter\nRDMA v2/v3 / URMA / UBSHM"] + O["TransportUpgradeOps\nPrepare / Negotiate / Activate / Deactivate"] + T["Concrete Transport\nTCP / RDMA / URMA / UBSHM"] + E["Endpoint / Resource\n只提供传输资源能力"] + IM["InputMessenger\nserver 增量输入"] + READY["ConnectionReady\nESTABLISHED / FALLBACK_TCP"] + FAIL["ConnectionFailed\nFAILED"] + + S -->|Connect| A + A -->|TCP connected| C + IM -->|ProcessUpgradeReadable| A + A --> C + C <--> H + C --> RA + C --> O + O --> T + T --> E + C -->|terminal phase| A + A --> READY + A --> FAIL +``` + +核心边界:`HandshakeSession` 负责公共握手流程,协议 Adapter 只负责 wire codec,`TransportUpgradeOps` 只负责资源能力;具体 Transport 不读取握手输入、不编排握手阶段,也不向 Socket 暴露 transport-specific phase。 diff --git a/src/brpc/adapter_transport.cpp b/src/brpc/adapter_transport.cpp new file mode 100644 index 0000000000..f0e2534376 --- /dev/null +++ b/src/brpc/adapter_transport.cpp @@ -0,0 +1,571 @@ +// 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 "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_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 { + +class AdapterConnect : public AppConnect { +public: + explicit AdapterConnect(const std::shared_ptr& app_connect) + : _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; +}; + +} // 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()); +} + +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 { + const char* first = static_cast(source->fetch1()); + handshake::HandshakeAdapter* adapter = + first != NULL && *first == 'U' + ? 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); + rdma::ParsedHello remote{}; + handshake::ClientHandshakeCallbacks callbacks{}; + callbacks.codec = protocol->MakeCodec(&remote); + callbacks.transport.prepare_resources = [&]() { + if (transport->PrepareUpgradeResources() == 0) { + return handshake::STEP_OK; + } + errno = 0; + return handshake::STEP_FALLBACK; + }; + callbacks.transport.negotiate_resources = [&]() { + return transport->NegotiateUpgradeResources(remote, false) == 0 + ? handshake::STEP_OK : handshake::STEP_FALLBACK; + }; + callbacks.transport.set_high_speed_active = [transport]() { + transport->ActivateUpgrade(); + }; + callbacks.transport.set_tcp_active = [transport]() { + transport->DeactivateUpgrade(); + }; + callbacks.transport.on_failed = [&]() { + 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)); + }; + const handshake::StepResult result = adapter->_handshake.RunClient(callbacks); + 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::HelloMessage remote{}; + ubring::UBShmHandshakeAdapter wire; + handshake::ClientHandshakeCallbacks callbacks{}; + callbacks.codec = wire.MakeCodec(); + callbacks.codec.build_hello = [&](bool enabled, std::string* payload) { + CHECK(enabled); + return wire.BuildHello(true, local_shm_len, shm_name_str.c_str(), + payload); + }; + callbacks.codec.parse_hello = [&](const std::string& payload) { + return wire.ParseHello(payload, &remote); + }; + callbacks.transport.prepare_resources = [&]() { + return transport->PrepareUpgradeResources( + &local_trx_shm, shm_name_str.c_str()) == 0 + ? handshake::STEP_OK : handshake::STEP_FALLBACK; + }; + callbacks.transport.negotiate_resources = [&]() { + return transport->NegotiateUpgradeResources( + &local_trx_shm, shm_name_str.c_str()) == 0 + ? handshake::STEP_OK : handshake::STEP_FALLBACK; + }; + callbacks.transport.set_high_speed_active = [transport]() { + transport->ActivateUpgrade(); + }; + callbacks.transport.set_tcp_active = [transport]() { + transport->DeactivateUpgrade(); + }; + callbacks.transport.on_failed = [&]() { + 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)); + }; + const handshake::StepResult result = adapter->_handshake.RunClient(callbacks); + if (result == handshake::STEP_OK) { + 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 = InputMessenger::OnNewMessages; +#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 (_high_speed_transport) { + return std::make_shared(_default_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::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) { +#if BRPC_WITH_RDMA + AdapterTransport* adapter = Get(socket); + if (adapter->_mode == SOCKET_MODE_RDMA && + adapter->_handshake.phase() == handshake::ESTABLISHED) { + adapter->CheckUnexpectedTcpData(); + return; + } +#endif + + 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 != 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..1cbc9803bf --- /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() const { return _high_speed_transport != NULL; } + + 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..53beb06083 --- /dev/null +++ b/src/brpc/handshake/handshake_io.cpp @@ -0,0 +1,150 @@ +// 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, + 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 != EAGAIN) { + 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, 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 != 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..3cf19d0901 --- /dev/null +++ b/src/brpc/handshake/rdma_handshake.cpp @@ -0,0 +1,576 @@ +// 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"; + } +} + +handshake::HandshakeCodec RdmaHandshakeAdapter::MakeCodec( + ParsedHello* remote) { + handshake::HandshakeCodec codec{}; + codec.protocol_version = ProtocolVersion(); + codec.hello_frame = HelloFrameSpec(); + codec.ack_frame = RdmaAckFrameSpec(); + codec.build_hello = [this](bool enabled, std::string* payload) { + return BuildLocalHello(enabled, payload); + }; + codec.parse_hello = [this, remote](const std::string& payload) { + return ParseRemoteHello(payload, remote); + }; + codec.build_ack = [](bool enabled, std::string* payload) { + const uint32_t flags_be = butil::HostToNet32( + enabled ? HELLO_ACK_RDMA_OK : 0); + payload->assign(reinterpret_cast(&flags_be), + sizeof(flags_be)); + return handshake::STEP_OK; + }; + codec.parse_ack = [](const std::string& payload, bool* enabled) { + if (payload.size() != HELLO_ACK_LEN) { + errno = EPROTO; + return handshake::STEP_ERROR; + } + uint32_t flags_be = 0; + memcpy(&flags_be, payload.data(), sizeof(flags_be)); + *enabled = (butil::NetToHost32(flags_be) & HELLO_ACK_RDMA_OK) != 0; + return handshake::STEP_OK; + }; + return codec; +} + +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; + +static HandshakeCodec MakeRdmaFallbackCodec(int version) { + HandshakeCodec codec{}; + codec.protocol_version = version; + codec.hello_frame = rdma::RdmaHelloFrameSpec(version); + codec.ack_frame = rdma::RdmaAckFrameSpec(); + codec.parse_hello = [](const std::string&) { + return STEP_FALLBACK; + }; + codec.build_hello = [version](bool enabled, std::string* payload) { + 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; + }; + codec.build_ack = [](bool enabled, std::string* payload) { + const uint32_t flags_be = butil::HostToNet32( + enabled ? rdma::HELLO_ACK_RDMA_OK : 0); + payload->assign(reinterpret_cast(&flags_be), + sizeof(flags_be)); + return STEP_OK; + }; + codec.parse_ack = [](const std::string& payload, bool* enabled) { + if (payload.size() != rdma::HELLO_ACK_LEN) { + errno = EPROTO; + return STEP_ERROR; + } + *enabled = false; + return STEP_OK; + }; + return codec; +} + +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); + ServerHandshakeCallbacks callbacks{}; + callbacks.fallback_on_not_mine = false; + callbacks.codecs.push_back(MakeRdmaFallbackCodec(2)); + callbacks.codecs.push_back(MakeRdmaFallbackCodec(3)); + callbacks.input = &input; + callbacks.transport.prepare_resources = []() { return STEP_OK; }; + callbacks.transport.negotiate_resources = []() { return STEP_OK; }; + callbacks.transport.set_high_speed_active = []() {}; + callbacks.transport.set_tcp_active = []() {}; + callbacks.transport.on_failed = []() {}; + return GetSession(socket)->RunServer(callbacks); +} + +#if BRPC_WITH_RDMA +StepResult RdmaServerHandshakeAdapter::RunRdmaServerHandshake( + butil::IOBuf* source, Socket* socket) { + RdmaTransport* transport = RdmaTransport::Get(socket); + CHECK(transport->GetRdmaEp() != NULL); + + rdma::ParsedHello remote{}; + std::vector > protocols = + transport->CreateServerHandshakeAdapters(); + IOBufHandshakeInput input(source); + ServerHandshakeCallbacks callbacks{}; + callbacks.fallback_on_not_mine = false; + callbacks.input = &input; + for (size_t i = 0; i < protocols.size(); ++i) { + HandshakeCodec codec = protocols[i]->MakeCodec(&remote); + const std::function parse_hello = + codec.parse_hello; + codec.parse_hello = [transport, parse_hello]( + const std::string& payload) { + const StepResult result = parse_hello(payload); + if (result == STEP_FALLBACK) { + transport->DeactivateUpgrade(); + } + return result; + }; + callbacks.codecs.push_back(codec); + } + callbacks.transport.prepare_resources = [&]() { + if (transport->PrepareUpgradeResources() < 0) { + PLOG(WARNING) + << "Fail to allocate rdma resources, fallback to tcp:" + << socket->description(); + transport->DeactivateUpgrade(); + return STEP_FALLBACK; + } + return STEP_OK; + }; + callbacks.transport.negotiate_resources = [&]() { + if (transport->NegotiateUpgradeResources(remote, true) < 0) { + PLOG(WARNING) + << "Fail to negotiate rdma resources, fallback to tcp:" + << socket->description(); + transport->DeactivateUpgrade(); + return STEP_FALLBACK; + } + return STEP_OK; + }; + callbacks.validate_established = [&]() { + if (!source->empty()) { + return STEP_ERROR; + } + return STEP_OK; + }; + callbacks.transport.set_high_speed_active = [transport]() { + transport->ActivateUpgrade(); + }; + callbacks.transport.set_tcp_active = [transport]() { + transport->DeactivateUpgrade(); + }; + callbacks.transport.on_failed = []() {}; + return GetSession(socket)->RunServer(callbacks); +} +#endif + +StepResult RdmaServerHandshakeAdapter::RunServerStep( + butil::IOBuf* source, Socket* socket) { +#if BRPC_WITH_RDMA + if (AdapterTransport::Get(socket)->upgrade_capable()) { + 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..7552454c34 --- /dev/null +++ b/src/brpc/handshake/rdma_handshake.h @@ -0,0 +1,154 @@ +// 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: + RdmaHandshakeAdapter(RdmaEndpoint* ep, int version) + : _ep(ep), _version(version) {} + virtual ~RdmaHandshakeAdapter() = default; + + int ProtocolVersion() const { return _version; } + handshake::HandshakeCodec MakeCodec(ParsedHello* remote); + + virtual const handshake::FrameSpec& HelloFrameSpec() const = 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; + +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..d7d4bb7531 --- /dev/null +++ b/src/brpc/handshake/ubshm_handshake.cpp @@ -0,0 +1,407 @@ +// 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 = 2; +static const uint16_t IMPL_VERSION = 1; +#endif // BRPC_WITH_UBRING +static const uint32_t ACK_OK = 0x1; + +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 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)); +} + +handshake::HandshakeCodec UBShmHandshakeAdapter::MakeCodec() const { + handshake::HandshakeCodec codec{}; + codec.protocol_version = 2; + codec.hello_frame = handshake::ubshm_wire::HelloFrameSpec(); + codec.ack_frame = handshake::ubshm_wire::AckFrameSpec(); + codec.build_ack = [](bool enabled, std::string* payload) { + const uint32_t flags_be = butil::HostToNet32( + enabled ? handshake::ubshm_wire::ACK_OK : 0); + payload->assign(reinterpret_cast(&flags_be), + sizeof(flags_be)); + return handshake::STEP_OK; + }; + codec.parse_ack = [](const std::string& payload, bool* enabled) { + if (payload.size() != handshake::ubshm_wire::ACK_LEN) { + errno = EPROTO; + return handshake::STEP_ERROR; + } + uint32_t flags_be = 0; + memcpy(&flags_be, payload.data(), sizeof(flags_be)); + *enabled = (butil::NetToHost32(flags_be) & + handshake::ubshm_wire::ACK_OK) != 0; + return handshake::STEP_OK; + }; + return codec; +} + +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; + memcpy(message.shm_name, shm_name, SHM_MAX_NAME_BUFF_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; + } + return NegotiationValid(*message) ? + handshake::STEP_OK : handshake::STEP_FALLBACK; +} + +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); +}; + + +static HandshakeCodec MakeUBShmFallbackCodec() { + HandshakeCodec codec{}; + codec.protocol_version = 2; + codec.hello_frame = ubshm_wire::HelloFrameSpec(); + codec.ack_frame = ubshm_wire::AckFrameSpec(); + codec.parse_hello = [](const std::string&) { + return STEP_FALLBACK; + }; + codec.build_hello = [](bool enabled, std::string* payload) { + 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; + }; + codec.build_ack = [](bool enabled, std::string* payload) { + const uint32_t flags_be = butil::HostToNet32( + enabled ? ubshm_wire::ACK_OK : 0); + payload->assign(reinterpret_cast(&flags_be), + sizeof(flags_be)); + return STEP_OK; + }; + codec.parse_ack = [](const std::string& payload, bool* enabled) { + if (payload.size() != ubshm_wire::ACK_LEN) { + errno = EPROTO; + return STEP_ERROR; + } + *enabled = false; + return STEP_OK; + }; + return codec; +} + +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); + ServerHandshakeCallbacks callbacks{}; + callbacks.fallback_on_not_mine = false; + callbacks.codecs.push_back(MakeUBShmFallbackCodec()); + callbacks.input = &input; + callbacks.transport.prepare_resources = []() { return STEP_OK; }; + callbacks.transport.negotiate_resources = []() { return STEP_OK; }; + callbacks.transport.set_high_speed_active = []() {}; + callbacks.transport.set_tcp_active = []() {}; + callbacks.transport.on_failed = []() {}; + return GetSession(socket)->RunServer(callbacks); +} + +#if BRPC_WITH_UBRING +StepResult UBShmServerHandshakeAdapter::RunUBShmServerHandshake( + butil::IOBuf* source, Socket* socket) { + UBShmTransport* transport = UBShmTransport::Get(socket); + CHECK(transport->GetUBShmEp() != NULL); + + ubring::HelloMessage remote{}; + ubring::UBShmHandshakeAdapter wire; + IOBufHandshakeInput input(source); + ServerHandshakeCallbacks callbacks{}; + callbacks.fallback_on_not_mine = false; + callbacks.input = &input; + HandshakeCodec codec = wire.MakeCodec(); + codec.parse_hello = [&](const std::string& payload) { + const StepResult result = wire.ParseHello(payload, &remote); + if (result == STEP_OK || result == STEP_FALLBACK) { + LOG_IF(INFO, ubring::FLAGS_ub_trace_verbose) + << "server receive handshake message : " + << remote.toString(); + } + if (result == STEP_FALLBACK) { + transport->DeactivateUpgrade(); + } + return result; + }; + codec.build_hello = [&](bool enabled, std::string* payload) { + const uint64_t len = enabled + ? static_cast(ubring::FLAGS_data_queue_size) * + MB_TO_BYTE + : 0; + return wire.BuildHello( + enabled, len, enabled ? remote.shm_name : NULL, payload); + }; + callbacks.codecs.push_back(codec); + callbacks.transport.prepare_resources = [&]() { + if (!ubring::IsUBAvailable()) { + transport->DeactivateUpgrade(); + return STEP_FALLBACK; + } + ubring::SHM remote_trx_shm = { + NULL, remote.len, 0, {0}, + static_cast(socket->fd())}; + strncpy(remote_trx_shm.name, remote.shm_name, + SHM_MAX_NAME_BUFF_LEN); + + 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; + }; + callbacks.transport.negotiate_resources = []() { return STEP_OK; }; + callbacks.validate_established = [&]() { + if (!source->empty() || + !transport->UpgradeActive()) { + return STEP_ERROR; + } + return STEP_OK; + }; + callbacks.transport.set_high_speed_active = [transport]() { + transport->ActivateUpgrade(); + }; + callbacks.transport.set_tcp_active = [transport]() { + transport->DeactivateUpgrade(); + }; + callbacks.transport.on_failed = []() {}; + const StepResult result = GetSession(socket)->RunServer(callbacks); + if (result == STEP_OK) { + 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()) { + 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..1bf1bb9d9b --- /dev/null +++ b/src/brpc/handshake/ubshm_handshake.h @@ -0,0 +1,84 @@ +// 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 v2 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]; +}; + +class UBShmHandshakeAdapter { +public: + UBShmHandshakeAdapter() = default; + + handshake::HandshakeCodec MakeCodec() const; + 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; + 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 8a4b53d8d5..97b163365b 100644 --- a/src/brpc/input_messenger.h +++ b/src/brpc/input_messenger.h @@ -35,6 +35,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 @@ -97,6 +98,7 @@ class InputMessageClosure { class InputMessenger : public SocketUser { friend class TcpTransport; friend class RdmaTransport; +friend class AdapterTransport; friend class rdma::RdmaEndpoint; friend class ubring::UBShmEndpoint; friend class InputMessengerProcessor; 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..0f3a4ee6c4 100644 --- a/src/brpc/rdma/rdma_endpoint.cpp +++ b/src/brpc/rdma/rdma_endpoint.cpp @@ -17,29 +17,29 @@ #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); namespace brpc { namespace rdma { -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_comp_channel* (*IbvCreateCompChannel)(ibv_context*); extern int (*IbvDestroyCompChannel)(ibv_comp_channel*); @@ -110,30 +110,14 @@ RdmaResource::~RdmaResource() { } RdmaEndpoint::RdmaEndpoint(Socket* s) - : _socket(s) - , _state(UNINIT) - , _handshake_version(0) - , _resource(nullptr) - , _send_cq_events(0) - , _recv_cq_events(0) - , _cq_sid(INVALID_SOCKET_ID) - , _sq_size(FLAGS_rdma_sq_size) - , _rq_size(FLAGS_rdma_rq_size) - , _remote_recv_block_size(0) - , _accumulated_ack(0) - , _unsolicited(0) - , _unsolicited_bytes(0) - , _sq_current(0) - , _sq_unsignaled(0) - , _sq_sent(0) - , _rq_received(0) - , _local_window_capacity(0) - , _remote_window_capacity(0) - , _sq_imm_window_size(0) - , _remote_rq_window_size(0) - , _sq_window_size(0) - , _new_rq_wrs(0) -{ + : _socket(s), _resource(nullptr), + _send_cq_events(0), _recv_cq_events(0), _cq_sid(INVALID_SOCKET_ID), + _sq_size(FLAGS_rdma_sq_size), _rq_size(FLAGS_rdma_rq_size), + _remote_recv_block_size(0), _accumulated_ack(0), _unsolicited(0), + _unsolicited_bytes(0), _sq_current(0), _sq_unsignaled(0), _sq_sent(0), + _rq_received(0), _local_window_capacity(0), _remote_window_capacity(0), + _sq_imm_window_size(0), _remote_rq_window_size(0), _sq_window_size(0), + _new_rq_wrs(0) { if (_sq_size < MIN_QP_SIZE) { _sq_size = MIN_QP_SIZE; } @@ -146,623 +130,73 @@ 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); -} +RdmaEndpoint::~RdmaEndpoint() { Reset(); } void RdmaEndpoint::Reset() { - DeallocateResources(); - - _state.store(UNINIT, butil::memory_order_relaxed); - _handshake_version = 0; - _outgoing_ece.reset(); - _resource = nullptr; - _send_cq_events = 0; - _recv_cq_events = 0; - _cq_sid = INVALID_SOCKET_ID; - _sbuf.clear(); - _rbuf.clear(); - _rbuf_data.clear(); - _input_processor.Reset(); - _remote_recv_block_size = 0; - _accumulated_ack = 0; - _unsolicited = 0; - _unsolicited_bytes = 0; - _sq_current = 0; - _sq_unsignaled = 0; - _sq_sent = 0; - _rq_received = 0; - _local_window_capacity = 0; - _remote_window_capacity = 0; - _sq_imm_window_size = 0; - _remote_rq_window_size.store(0, butil::memory_order_relaxed); - _sq_window_size.store(0, butil::memory_order_relaxed); - _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; + DeallocateResources(); + + _outgoing_ece.reset(); + _resource = nullptr; + _send_cq_events = 0; + _recv_cq_events = 0; + _cq_sid = INVALID_SOCKET_ID; + _sbuf.clear(); + _rbuf.clear(); + _input_processor.Reset(); + _rbuf_data.clear(); + _remote_recv_block_size = 0; + _accumulated_ack = 0; + _unsolicited = 0; + _unsolicited_bytes = 0; + _sq_current = 0; + _sq_unsignaled = 0; + _sq_sent = 0; + _rq_received = 0; + _local_window_capacity = 0; + _remote_window_capacity = 0; + _sq_imm_window_size = 0; + _remote_rq_window_size.store(0, butil::memory_order_relaxed); + _sq_window_size.store(0, butil::memory_order_relaxed); + _new_rq_wrs.store(0, butil::memory_order_relaxed); } -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; - }); +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; + _sq_imm_window_size = RESERVED_WR_NUM; + _remote_rq_window_size.store(_local_window_capacity, + butil::memory_order_relaxed); + _sq_window_size.store(_local_window_capacity, butil::memory_order_relaxed); } -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::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; + } } -void RdmaEndpoint::ApplyRemoteHello(const ParsedHello& 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; - _sq_imm_window_size = RESERVED_WR_NUM; - _remote_rq_window_size.store(_local_window_capacity, butil::memory_order_relaxed); - _sq_window_size.store(_local_window_capacity, butil::memory_order_relaxed); +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; } -// 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; - } - - 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); - } - - 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 { if (BAIDU_UNLIKELY(g_skip_rdma_init)) { @@ -778,6 +212,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 +235,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 +280,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 +334,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 +379,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 +397,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 +432,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 +479,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 +543,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 +567,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 +599,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 +673,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 +708,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 +740,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 +758,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 +772,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 +808,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 | + 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)); + 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 +824,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 +878,7 @@ static int DrainCq(ibv_cq* cq) { } void RdmaEndpoint::DeallocateResources() { - if (_resource == nullptr) { + if (!_resource) { return; } if (FLAGS_rdma_use_polling) { @@ -1470,7 +905,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; @@ -1480,16 +915,16 @@ void RdmaEndpoint::DeallocateResources() { DeallocateCq(_resource->send_cq); DeallocateCq(_resource->recv_cq); - if (_resource->comp_channel != nullptr) { - if (_cq_sid != INVALID_SOCKET_ID) { + if (nullptr != _resource->comp_channel) { // 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 +1005,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 +1025,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 +1058,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 +1177,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 +1214,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 +1310,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 +1340,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]; + 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]; + 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]; + 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]; + 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..894637380e 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_ep = new (std::nothrow) rdma::RdmaEndpoint(socket); + if (!_rdma_ep) { + const int saved_errno = errno != 0 ? errno : ENOMEM; + PLOG(ERROR) << "Fail to create RdmaEndpoint"; + socket->SetFailed(saved_errno, "Fail to create RdmaEndpoint: %s", + berror(saved_errno)); } - _tcp_transport = std::make_shared(); - _tcp_transport->Init(socket, options); + _rdma_state = RDMA_UNKNOWN; } void RdmaTransport::Release() { @@ -70,31 +68,49 @@ int RdmaTransport::Reset(int32_t expected_nref) { } std::shared_ptr RdmaTransport::Connect() { - if (_default_connect == nullptr) { - return std::make_shared(); - } - return _default_connect; + return _default_connect; +} + +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); } + 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); - } + 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 +124,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 +137,6 @@ int RdmaTransport::WaitEpollOut(butil::atomic *_epollout_butex, } } } - } else { - return _tcp_transport->WaitEpollOut(_epollout_butex, pollin, duetime); } return 0; } @@ -163,15 +176,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 +200,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 +230,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..1590374189 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,27 @@ 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; } private: + void SetHighSpeedAvailable(bool available); + static bool OptionsAvailableForRdma(const ChannelOptions* opt); static bool OptionsAvailableOverRdma(const ServerOptions* opt); @@ -60,7 +81,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 6f6f52fb2b..579bc9cb79 100644 --- a/src/brpc/socket.h +++ b/src/brpc/socket.h @@ -56,21 +56,20 @@ class ChannelBalancer; } namespace rdma { class RdmaEndpoint; -class RdmaConnect; -class RdmaHandshakeClientV2; -class RdmaHandshakeServerV2; -class RdmaHandshakeClientV3; -class RdmaHandshakeServerV3; } 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); @@ -328,14 +327,9 @@ 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 HealthCheckTask; friend class OnAppHealthCheckDone; friend class HealthCheckManager; @@ -344,6 +338,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 TransportFactory; @@ -970,8 +966,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 36fdaaed05..3355e0764d 100644 --- a/src/brpc/transport_factory.cpp +++ b/src/brpc/transport_factory.cpp @@ -16,7 +16,7 @@ // under the License. #include "brpc/transport_factory.h" -#include "brpc/tcp_transport.h" +#include "brpc/adapter_transport.h" #include "brpc/rdma_transport.h" #include "brpc/ubshm_transport.h" @@ -43,16 +43,16 @@ int TransportFactory::ContextInitOrDie(SocketMode mode, bool serverOrNot, const 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 else if (mode == SOCKET_MODE_RDMA) { - return std::unique_ptr(new RdmaTransport()); + return std::unique_ptr(new AdapterTransport(mode)); } #endif #if BRPC_WITH_UBRING else if (mode == SOCKET_MODE_UBRING) { - return std::unique_ptr(new UBShmTransport()); + return std::unique_ptr(new AdapterTransport(mode)); } #endif else { @@ -60,4 +60,4 @@ std::unique_ptr TransportFactory::CreateTransport(SocketMode mode) { return nullptr; } } -} // namespace brpc \ No newline at end of file +} // namespace brpc diff --git a/src/brpc/transport_factory.h b/src/brpc/transport_factory.h index d933a130e1..add249c438 100644 --- a/src/brpc/transport_factory.h +++ b/src/brpc/transport_factory.h @@ -22,7 +22,8 @@ #include "brpc/transport.h" namespace brpc { -// TransportFactory to create transport instance with socket_mode {TCP, RDMA} +// Creates the top-level AdapterTransport for all socket modes. The adapter +// selects TcpTransport or a concrete accelerated Transport internally. class TransportFactory { public: static int ContextInitOrDie(SocketMode mode, bool serverOrNot, const void* _options); @@ -31,4 +32,4 @@ class TransportFactory { }; } // namespace brpc -#endif //BRPC_TRANSPORT_FACTORY_H \ No newline at end of file +#endif //BRPC_TRANSPORT_FACTORY_H diff --git a/src/brpc/transport_handshake.cpp b/src/brpc/transport_handshake.cpp new file mode 100644 index 0000000000..f97d7b7741 --- /dev/null +++ b/src/brpc/transport_handshake.cpp @@ -0,0 +1,338 @@ +// 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 "butil/logging.h" +#include "butil/object_pool.h" + +namespace brpc { +namespace handshake { + +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, const std::function& on_failed) { + if (on_failed) { + on_failed(); + } + session->MarkFailed(); + return STEP_ERROR; +} + +static StepResult FinishWithFallback( + HandshakeSession* session, const std::function& set_tcp_active) { + session->PublishFallback([&set_tcp_active]() { + if (set_tcp_active) { + set_tcp_active(); + } + }); + 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(const HandshakeCodec& codec, + bool enabled) { + CHECK(codec.build_hello); + std::string payload; + const StepResult result = codec.build_hello(enabled, &payload); + if (result != STEP_OK) { + return result; + } + return ConvertFrameResult( + FrameCodec::WriteFrame(_io, codec.hello_frame, payload)); +} + +StepResult HandshakeSession::ReceiveHello(const HandshakeCodec& codec, + HandshakeInput* input, + bool push_back_on_not_mine, + bool* magic_matched) { + CHECK(codec.parse_hello); + std::string payload; + const FrameResult frame_result = input != NULL + ? FrameCodec::ParseBufferedFrame( + input, codec.hello_frame, &payload, magic_matched) + : FrameCodec::ReadFrame( + _io, codec.hello_frame, push_back_on_not_mine, &payload); + const StepResult result = ConvertFrameResult(frame_result); + if (result != STEP_OK) { + return result; + } + set_protocol_version(codec.protocol_version); + return codec.parse_hello(payload); +} + +StepResult HandshakeSession::SendAck(const HandshakeCodec& codec, + bool enabled) { + CHECK(codec.build_ack); + std::string payload; + const StepResult result = codec.build_ack(enabled, &payload); + if (result != STEP_OK) { + return result; + } + return ConvertFrameResult( + FrameCodec::WriteFrame(_io, codec.ack_frame, payload)); +} + +StepResult HandshakeSession::ReceiveAck(const HandshakeCodec& codec, + HandshakeInput* input, + bool* enabled) { + CHECK(codec.parse_ack); + std::string payload; + const FrameResult frame_result = input != NULL + ? FrameCodec::ParseBufferedFrame(input, codec.ack_frame, &payload) + : FrameCodec::ReadFrame(_io, codec.ack_frame, false, &payload); + const StepResult result = ConvertFrameResult(frame_result); + if (result != STEP_OK) { + return result; + } + return codec.parse_ack(payload, enabled); +} + +StepResult HandshakeSession::SelectAndReceiveHello( + const std::vector& codecs, HandshakeInput* input, + bool push_back_on_not_mine, const HandshakeCodec** selected) { + CHECK(!codecs.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, codecs.size()); + *selected = &codecs.front(); + return ReceiveHello(**selected, NULL, push_back_on_not_mine); + } + + bool need_more = false; + for (size_t i = 0; i < codecs.size(); ++i) { + bool magic_matched = false; + const StepResult result = ReceiveHello( + codecs[i], input, false, &magic_matched); + if (result == STEP_NOT_MINE) { + continue; + } + if (result == STEP_NEED_MORE) { + if (magic_matched) { + *selected = &codecs[i]; + set_protocol_version(codecs[i].protocol_version); + return STEP_NEED_MORE; + } + need_more = true; + continue; + } + *selected = &codecs[i]; + return result; + } + return need_more ? STEP_NEED_MORE : STEP_NOT_MINE; +} + +StepResult HandshakeSession::RunClient( + const ClientHandshakeCallbacks& callbacks) { + CHECK(callbacks.transport.prepare_resources); + CHECK(callbacks.transport.negotiate_resources); + CHECK(callbacks.transport.set_high_speed_active); + CHECK(callbacks.transport.set_tcp_active); + // 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 = callbacks.transport.prepare_resources(); + if (result == STEP_FALLBACK) { + return FinishWithFallback(this, callbacks.transport.set_tcp_active); + } + if (result != STEP_OK) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + + SetPhase(HELLO_SEND); + if (SendHello(callbacks.codec, true) != STEP_OK) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + + SetPhase(HELLO_WAIT); + result = ReceiveHello(callbacks.codec, 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, callbacks.transport.on_failed); + } + bool enabled = result == STEP_OK; + + if (enabled) { + SetPhase(NEGOTIATING); + result = callbacks.transport.negotiate_resources(); + if (result != STEP_OK && result != STEP_FALLBACK) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + enabled = result == STEP_OK; + } + + SetPhase(ACK_SEND); + if (SendAck(callbacks.codec, enabled) != STEP_OK) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + + if (enabled) { + callbacks.transport.set_high_speed_active(); + MarkEstablished(); + return STEP_OK; + } + return FinishWithFallback(this, callbacks.transport.set_tcp_active); +} + +StepResult HandshakeSession::RunServer( + const ServerHandshakeCallbacks& callbacks) { + CHECK(!callbacks.codecs.empty()); + CHECK(callbacks.transport.prepare_resources); + CHECK(callbacks.transport.negotiate_resources); + CHECK(callbacks.transport.set_high_speed_active); + CHECK(callbacks.transport.set_tcp_active); + + // 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; + } + + const HandshakeCodec* selected = NULL; + if (phase() != ACK_WAIT) { + const int previous_phase = phase(); + _local_enabled = false; + SetPhase(HELLO_WAIT); + StepResult result = SelectAndReceiveHello( + callbacks.codecs, callbacks.input, + callbacks.fallback_on_not_mine, &selected); + if (result == STEP_NOT_MINE) { + if (callbacks.fallback_on_not_mine) { + return FinishWithFallback(this, callbacks.transport.set_tcp_active); + } + 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, callbacks.transport.on_failed); + } + bool enabled = result == STEP_OK; + + if (enabled) { + SetPhase(PREPARING); + result = callbacks.transport.prepare_resources(); + if (result != STEP_OK && result != STEP_FALLBACK) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + enabled = result == STEP_OK; + } + + if (enabled) { + SetPhase(NEGOTIATING); + result = callbacks.transport.negotiate_resources(); + if (result != STEP_OK && result != STEP_FALLBACK) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + enabled = result == STEP_OK; + } + + SetPhase(HELLO_SEND); + CHECK(selected != NULL); + _local_enabled = enabled; + if (SendHello(*selected, enabled) != STEP_OK) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + SetPhase(ACK_WAIT); + } else { + for (size_t i = 0; i < callbacks.codecs.size(); ++i) { + if (callbacks.codecs[i].protocol_version == protocol_version()) { + selected = &callbacks.codecs[i]; + break; + } + } + CHECK(selected != NULL); + } + + // 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, callbacks.input, &peer_enabled); + if (result == STEP_NEED_MORE) { + return STEP_NEED_MORE; + } + if (result == STEP_ERROR || result == STEP_NOT_MINE) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + if (result == STEP_FALLBACK) { + return FinishWithFallback(this, callbacks.transport.set_tcp_active); + } + if (!peer_enabled) { + return FinishWithFallback(this, callbacks.transport.set_tcp_active); + } + if (!_local_enabled) { + errno = EPROTO; + return FinishWithFailure(this, callbacks.transport.on_failed); + } + if (callbacks.validate_established && + callbacks.validate_established() != STEP_OK) { + return FinishWithFailure(this, callbacks.transport.on_failed); + } + + callbacks.transport.set_high_speed_active(); + 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..3899a3aa44 --- /dev/null +++ b/src/brpc/transport_handshake.h @@ -0,0 +1,197 @@ +// 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 + +#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, + ESTABLISHED = 0x100, + FALLBACK_TCP = 0x200, + FAILED = 0x300, +}; + +enum StepResult { + STEP_OK = 0, + STEP_FALLBACK, + STEP_NEED_MORE, + STEP_NOT_MINE, + STEP_ERROR, +}; + +// A protocol describes only its fields and resource-independent wire values. +// HandshakeSession owns framing and I/O through FrameCodec. The callbacks may +// retain strongly typed parsed state in their protocol adapter. +struct HandshakeCodec { + int protocol_version; + FrameSpec hello_frame; + FrameSpec ack_frame; + std::function build_hello; + std::function parse_hello; + std::function build_ack; + std::function parse_ack; +}; + +// Resource-specific operations supplied by a Transport and invoked by the +// common coordinator. Wire I/O and field codec invocation remain owned by +// HandshakeSession. +struct TransportUpgradeOps { + std::function prepare_resources; + std::function negotiate_resources; + std::function set_high_speed_active; + std::function set_tcp_active; + std::function on_failed; +}; + +struct ClientHandshakeCallbacks { + HandshakeCodec codec; + TransportUpgradeOps transport; +}; + +// The server driver is independent of the input mode. A parser callback can +// return STEP_NEED_MORE, while a blocking callback waits before returning. +struct ServerHandshakeCallbacks { + bool fallback_on_not_mine; + // Buffered parsers may offer multiple codecs (RDMA v2/v3). Blocking + // server handshakes currently provide exactly one codec. + std::vector codecs; + HandshakeInput* input; + TransportUpgradeOps transport; + std::function validate_established; +}; + +// 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_relaxed); + } + + 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(const ClientHandshakeCallbacks& callbacks); + StepResult RunServer(const ServerHandshakeCallbacks& callbacks); + +private: + StepResult SendHello(const HandshakeCodec& codec, bool enabled); + StepResult ReceiveHello(const HandshakeCodec& codec, + HandshakeInput* input, + bool push_back_on_not_mine, + bool* magic_matched = NULL); + StepResult SendAck(const HandshakeCodec& codec, bool enabled); + StepResult ReceiveAck(const HandshakeCodec& codec, + HandshakeInput* input, bool* enabled); + StepResult SelectAndReceiveHello( + const std::vector& codecs, HandshakeInput* input, + bool push_back_on_not_mine, const HandshakeCodec** 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 19bca4f964..5fa9ea5f2c 100644 --- a/src/brpc/ubshm/ub_endpoint.cpp +++ b/src/brpc/ubshm/ub_endpoint.cpp @@ -19,23 +19,20 @@ #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,95 +41,22 @@ 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_poller_yield, false, "Yield thread in RDMA polling mode."); DEFINE_bool(ub_edisp_unsched, false, "Disable event dispatcher schedule"); -DEFINE_bool(ub_disable_bthread, false, "Disable bthread in UBRing polling mode."); +DEFINE_bool(ub_disable_bthread, false, "Disable bthread in RDMA"); +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 = 2; -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 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 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)); -} +static butil::Mutex *g_ubring_resource_mutex = NULL; 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>(); -} + : _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); -} +UBShmEndpoint::~UBShmEndpoint() { Reset(); } void UBShmEndpoint::Reset() { DeallocateResources(); @@ -140,467 +64,6 @@ void UBShmEndpoint::Reset() { delete _ub_ring; _ub_ring = nullptr; _poller_sid = INVALID_SOCKET_ID; - _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); - 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; - } - - if (remote_msg.msg_len > HELLO_MSG_LEN_MIN) { - // TODO: Read Hello Message customized data - // Just for future use, should not happen now - } - - if (!HelloNegotiationValid(remote_msg)) { - LOG(WARNING) << "Fail to negotiate with server, fallback to tcp:" - << s->description(); - ub_transport->_ub_state = UBShmTransport::UB_OFF; - } else { - 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; - } - uint32_t* tmp = (uint32_t*)data; - *tmp = butil::HostToNet32(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->_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); - 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 (remote_msg.msg_len > HELLO_MSG_LEN_MIN) { - // TODO: Read Hello Message customized header - // Just for future use, should not happen now - } - - 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; - } - - 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* tmp = (uint32_t*)data; - uint32_t flags = butil::NetToHost32(*tmp); - if (flags & ACK_MSG_UB_OK) { - if (ub_transport->_ub_state == UBShmTransport::UB_OFF) { - LOG(WARNING) << "Fail to parse Hello Message length 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->_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 { @@ -641,9 +104,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)) { @@ -663,7 +128,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; @@ -677,18 +143,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; @@ -702,11 +180,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); @@ -734,7 +223,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; @@ -755,9 +244,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) { @@ -790,12 +280,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() { @@ -831,7 +320,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) { @@ -840,17 +329,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()); @@ -858,12 +348,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) { @@ -882,8 +372,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; @@ -908,8 +398,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 a29a0927f7..90ed94bad1 100644 --- a/src/brpc/ubshm/ub_endpoint.h +++ b/src/brpc/ubshm/ub_endpoint.h @@ -20,61 +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); -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; @@ -113,9 +87,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, @@ -124,33 +95,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_MAP_REMOTE_SHM = 0x4, - C_ACK_SEND = 0x5, - S_HELLO_WAIT = 0x11, - S_ALLOC_SHM = 0x12, - S_HELLO_SEND = 0x13, - S_ACK_WAIT = 0x14, - 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); @@ -160,57 +105,31 @@ 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; - // 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 { @@ -222,8 +141,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 @@ -238,8 +156,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..db1e6c6352 100644 --- a/src/brpc/ubshm_transport.cpp +++ b/src/brpc/ubshm_transport.cpp @@ -17,10 +17,16 @@ #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 +34,27 @@ 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_ep = new (std::nothrow) ubring::UBShmEndpoint(socket); + if (!_ub_ep) { + const int saved_errno = errno != 0 ? errno : ENOMEM; + errno = saved_errno; + PLOG(ERROR) << "Fail to create UBShmEndpoint"; + socket->SetFailed(saved_errno, "Fail to create UBShmEndpoint: %s", + berror(saved_errno)); } - _tcp_transport = std::make_shared(); - _tcp_transport->Init(socket, options); + _ub_state = UB_UNKNOWN; } void UBShmTransport::Release() { @@ -64,27 +74,47 @@ 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); } + +void UBShmTransport::FinishUpgrade() { + if (_ub_ep != NULL && _ub_ep->_ub_ring != NULL) { + _ub_ep->_ub_ring->UbrUnlinkLocalShm(); + } +} + +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) { - if (_ub_ep && _ub_state != UB_OFF) { + CHECK(_ub_ep != NULL); return _ub_ep->CutFromIOBufList(buf, ndata); } - return _tcp_transport->CutFromIOBufList(buf, ndata); -} int UBShmTransport::WaitEpollOut(butil::atomic *_epollout_butex, bool pollin, const timespec duetime) { @@ -94,14 +124,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 +144,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 +183,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 +207,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 +231,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..6de2e8917c 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,21 @@ 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; } private: + void SetHighSpeedAvailable(bool available); + static bool OptionsAvailableForUB(const ChannelOptions* opt); static bool OptionsAvailableOverUB(const ServerOptions* opt); private: @@ -57,7 +73,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..69b1cd8222 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,19 +59,21 @@ 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 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 (*IbvDestroyQp)(ibv_qp*); +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 (*IbvDestroyQp)(ibv_qp *); extern butil::atomic g_rdma_available; extern bool g_skip_rdma_init; extern bool g_fail_resource_alloc_for_test; @@ -169,139 +168,122 @@ 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, - google::protobuf::Closure* done) { - Controller* cntl = static_cast(cntl_base); - ClosureGuard done_guard(done); - g_echo_served.fetch_add(1, butil::memory_order_relaxed); - if (req->server_fail()) { - cntl->SetFailed(req->server_fail(), "Server fail1"); - cntl->SetFailed(req->server_fail(), "Server fail2"); - return; - } - if (req->close_fd()) { - usleep(1); - LOG(INFO) << "close fd..."; - cntl->CloseConnection("Close connection according to request"); - return; - } - if (req->sleep_us() > 0) { - LOG(INFO) << "sleep " << req->sleep_us() << "us..."; - bthread_usleep(req->sleep_us()); - } - res->set_message("MyEchoService"); - if (req->code() != 0) { - res->add_code_list(req->code()); - } - cntl->response_attachment().append(cntl->request_attachment()); - } + void Echo(google::protobuf::RpcController *cntl_base, + const ::test::EchoRequest *req, ::test::EchoResponse *res, + google::protobuf::Closure *done) { + Controller *cntl = static_cast(cntl_base); + ClosureGuard done_guard(done); + g_echo_served.fetch_add(1, butil::memory_order_relaxed); + if (req->server_fail()) { + cntl->SetFailed(req->server_fail(), "Server fail1"); + cntl->SetFailed(req->server_fail(), "Server fail2"); + return; + } + if (req->close_fd()) { + usleep(1); + LOG(INFO) << "close fd..."; + cntl->CloseConnection("Close connection according to request"); + return; + } + if (req->sleep_us() > 0) { + LOG(INFO) << "sleep " << req->sleep_us() << "us..."; + bthread_usleep(req->sleep_us()); + } + res->set_message("MyEchoService"); + if (req->code() != 0) { + res->add_code_list(req->code()); + } + cntl->response_attachment().append(cntl->request_attachment()); + } }; class RdmaTest : public ::testing::Test { protected: - RdmaTest() { - butil::ip_t ip; - EXPECT_EQ(0, butil::str2ip(g_ip.c_str(), &ip)); - butil::EndPoint ep(ip, PORT); - g_ep = ep; - EXPECT_EQ(0, _server_list.save(butil::endpoint2str(g_ep).c_str())); - _naming_url = std::string("File://") + _server_list.fname(); - _server.AddService(&_svc, SERVER_DOESNT_OWN_SERVICE); - } - ~RdmaTest() { } + RdmaTest() { + butil::ip_t ip; + EXPECT_EQ(0, butil::str2ip(g_ip.c_str(), &ip)); + butil::EndPoint ep(ip, PORT); + g_ep = ep; + EXPECT_EQ(0, _server_list.save(butil::endpoint2str(g_ep).c_str())); + _naming_url = std::string("File://") + _server_list.fname(); + _server.AddService(&_svc, SERVER_DOESNT_OWN_SERVICE); + } + ~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) { - ServerOptions options; - options.enabled_protocols = "baidu_std"; - options.socket_mode = use_rdma ? SOCKET_MODE_RDMA : SOCKET_MODE_TCP; - options.idle_timeout_sec = 5; - options.max_concurrency = 0; - options.internal_port = -1; - EXPECT_EQ(0, _server.Start(PORT, &options)); - } - - void StopServer() { - _server.Stop(0); - _server.Join(); + void StartServer(bool use_rdma = true) { + ServerOptions options; + options.enabled_protocols = "baidu_std"; + options.socket_mode = use_rdma ? SOCKET_MODE_RDMA : SOCKET_MODE_TCP; + options.idle_timeout_sec = 5; + options.max_concurrency = 0; + options.internal_port = -1; + EXPECT_EQ(0, _server.Start(PORT, &options)); + } + + void StopServer() { + _server.Stop(0); + _server.Join(); + } + + Socket *GetSocketFromServer(size_t index) { + std::vector sids; + _server._am->ListConnections(&sids); + if (index >= sids.size()) { + return nullptr; } - - Socket* GetSocketFromServer(size_t index) { - std::vector sids; - _server._am->ListConnections(&sids); - if (index >= sids.size()) { - return nullptr; - } - SocketUniquePtr s; - if (Socket::Address(sids[index], &s) == 0) { - return s.get(); - } - return nullptr; + SocketUniquePtr s; + if (Socket::Address(sids[index], &s) == 0) { + return s.get(); } + return nullptr; + } - // Accepting the connection happens in the server threads, so poll for it - // rather than sleeping. Returns nullptr if it never showed up. - Socket* WaitForServerSocket() { - Socket* s = nullptr; - WaitUntil([this, &s] { return (s = GetSocketFromServer(0)) != nullptr; }); - return s; - } + // 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; }); - } + bool WaitForServerSocketGone() { + return WaitUntil([this] { return GetSocketFromServer(0) == nullptr; }); + } - butil::TempFile _server_list; - std::string _naming_url; + butil::TempFile _server_list; + std::string _naming_url; - Server _server; - MyEchoService _svc; + Server _server; + MyEchoService _svc; }; -// Shorthand for the RDMA transport behind a Socket, which every endpoint state -// check below has to go through. -static RdmaTransport* RdmaTransportOf(Socket* s) { - return static_cast(s->_transport.get()); -} -static RdmaTransport* RdmaTransportOf(const SocketUniquePtr& s) { - return RdmaTransportOf(s.get()); +// Shorthand for the RDMA transport behind a Socket. +static RdmaTransport *RdmaTransportOf(Socket *s) { + 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; - }); - return state; +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 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; - }); +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,21 +304,20 @@ 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(); - _saved_handshake_version = rdma::FLAGS_rdma_client_handshake_version; - rdma::FLAGS_rdma_client_handshake_version = GetParam(); - } - void TearDown() override { - rdma::FLAGS_rdma_client_handshake_version = _saved_handshake_version; - RdmaTest::TearDown(); - } + void SetUp() override { + RdmaTest::SetUp(); + _saved_handshake_version = rdma::FLAGS_rdma_client_handshake_version; + rdma::FLAGS_rdma_client_handshake_version = GetParam(); + } + void TearDown() override { + rdma::FLAGS_rdma_client_handshake_version = _saved_handshake_version; + RdmaTest::TearDown(); + } private: - int _saved_handshake_version = 2; + int _saved_handshake_version = 2; }; TEST_F(RdmaTest, stale_cq_callback_does_not_poll_new_generation) { @@ -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; @@ -374,394 +356,568 @@ TEST_F(RdmaTest, stale_cq_callback_does_not_poll_new_generation) { } TEST_F(RdmaTest, client_close_before_hello_send) { - StartServer(); + 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); - StopServer(); + 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(); } 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); - - 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); - - StopServer(); + StartServer(); + + 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_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(); } TEST_F(RdmaTest, client_close_during_hello_send) { - StartServer(); - - 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); - 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); - 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); - // 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()); - - StopServer(); + 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(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_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_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_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(); } TEST_F(RdmaTest, client_hello_msg_invalid_len) { - StartServer(); - - 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); - memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd1, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - 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); - memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd2, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - 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()); - - StopServer(); + 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(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_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_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_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_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(); } TEST_F(RdmaTest, client_hello_msg_invalid_version) { - StartServer(); - - 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); - memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd1, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - memcpy(data + 4, &len, 2); - memset(data + 6, 0, 34); - memcpy(data + 6, &ver, 2); // hello_ver == 1, impl_ver == 0 - // Write the 36B base starting at data + 4 (NOT data). Pre-Step-1 this - // UT mistakenly wrote `data, 36` which included the leftover "RDMA" - // magic at data[0..4); the server parsed it as msg_len = 0x5244 and - // happened to fall through to NegotiationValid (which then failed on - // 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); - uint32_t flags = 0; - ASSERT_TRUE(WriteAll(sockfd1, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); - 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); - memcpy(data, "RDMA", 4); - ASSERT_TRUE(WriteAll(sockfd2, data, 4)); // Write magic string. - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::S_HELLO_WAIT, RdmaTransportOf(s)); - 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 - // 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)); - sockfd2.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - StopServer(); + 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(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_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 + // Write the 36B base starting at data + 4 (NOT data). Pre-Step-1 this + // UT mistakenly wrote `data, 36` which included the leftover "RDMA" + // magic at data[0..4); the server parsed it as msg_len = 0x5244 and + // happened to fall through to NegotiationValid (which then failed on + // 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_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_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); + 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_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 `write(sockfd1, data + 4, 36)` for why we + // write from data + 4 instead of data. + 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); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + StopServer(); } TEST_F(RdmaTest, client_hello_msg_invalid_sq_rq_block_size) { - StartServer(); - - Socket* s = nullptr; - uint32_t flags = butil::HostToNet32(0); - rdma::v2_wire::HelloMessage msg{}; - uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - 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 = 10; - msg.rq_size = 16; - 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)); - sockfd1.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - 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)); - sockfd2.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - 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)); - sockfd3.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - StopServer(); + 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{}; + uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; + 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 = 10; + msg.rq_size = 16; + msg.block_size = 8192; + memcpy(data, "RDMA", 4); + msg.Serialize(data + 4); + 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); + 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(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); + 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(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); + usleep(100000); // wait for server to handle the msg + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + StopServer(); } TEST_F(RdmaTest, client_close_after_qp_build) { - StartServer(); - - Socket* s = nullptr; - uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - MakeV2ClientHello(data); - - 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()); - - StopServer(); + 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]; + 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(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(); } TEST_F(RdmaTest, client_close_during_ack_send) { - StartServer(); - - Socket* s = nullptr; - uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - MakeV2ClientHello(data); - - 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)); - 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()); - - StopServer(); + 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]; + 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(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_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(); } TEST_F(RdmaTest, client_close_after_ack_send) { - StartServer(); - - Socket* s = nullptr; - uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - MakeV2ClientHello(data); - - 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)); - 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)); - 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()); - - StopServer(); + 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]; + 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(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_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_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(); } TEST_F(RdmaTest, client_send_data_on_tcp_after_ack_send) { - StartServer(); - - Socket* s = nullptr; - uint8_t data[rdma::HELLO_V2_MSG_LEN_MIN]; - MakeV2ClientHello(data); - - 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)); - 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)); - 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()); - - StopServer(); + 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]; + 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(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_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_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()); @@ -959,1366 +1102,1464 @@ TEST_F(RdmaTest, fd_and_qp_input_streams_are_separate) { } TEST_F(RdmaTest, server_miss_before_hello_send) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - 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); + 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); - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); + usleep(100000); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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); - bthread_id_join(cntl.call_id()); + butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); + ASSERT_TRUE(acc_fd >= 0); + bthread_id_join(cntl.call_id()); - ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); + ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); } TEST_F(RdmaTest, server_close_before_hello_send) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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)); + 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()); +} - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); +TEST_F(RdmaTest, server_miss_during_magic_str) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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(2, write(acc_fd, "RD", 2)); + usleep(100000); + bthread_id_join(cntl.call_id()); + + ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); +} - 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); +TEST_F(RdmaTest, server_close_during_magic_str) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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(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()); +} - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); +TEST_F(RdmaTest, server_hello_invalid_magic_str) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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(4, write(acc_fd, "ABCD", 4)); + usleep(100000); + ASSERT_EQ(handshake::FAILED, AdapterTransport::Get(s.get())->handshake_phase()); + bthread_id_join(cntl.call_id()); + + ASSERT_EQ(EPROTO, cntl.ErrorCode()); +} - 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)); - close(acc_fd); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FAILED, RdmaTransportOf(s)); - bthread_id_join(cntl.call_id()); +TEST_F(RdmaTest, server_miss_during_hello_msg) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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(4, write(acc_fd, "RDMA", 4)); + 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()); +} - ASSERT_EQ(EEOF, cntl.ErrorCode()); +TEST_F(RdmaTest, server_close_during_hello_msg) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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(4, write(acc_fd, "RDMA", 4)); + 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); + usleep(100000); + ASSERT_EQ(handshake::FAILED, AdapterTransport::Get(s.get())->handshake_phase()); + bthread_id_join(cntl.call_id()); + + ASSERT_EQ(EEOF, cntl.ErrorCode()); } -TEST_F(RdmaTest, server_miss_during_magic_str) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); +TEST_F(RdmaTest, server_hello_invalid_msg_len) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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)); + 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)); + usleep(100000); + ASSERT_EQ(handshake::FAILED, AdapterTransport::Get(s.get())->handshake_phase()); + bthread_id_join(cntl.call_id()); + + ASSERT_EQ(EPROTO, cntl.ErrorCode()); +} - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); +TEST_F(RdmaTest, server_hello_invalid_version) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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)); + 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)); + 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)); + bthread_id_join(cntl.call_id()); + + ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); +} - 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); +TEST_F(RdmaTest, server_hello_invalid_sq_rq_size) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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)); + + rdma::v2_wire::HelloMessage msg{}; + msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; + msg.hello_ver = 1; + msg.impl_ver = 1; + msg.sq_size = 0; + msg.rq_size = 0; + msg.block_size = 8192; + msg.qp_num = 0; + 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)); + + 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)); + bthread_id_join(cntl.call_id()); + + ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); +} - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); +TEST_F(RdmaTest, server_miss_after_ack) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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)); + + rdma::v2_wire::HelloMessage msg{}; + 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); + 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(4, read(acc_fd, data, 4)); + uint32_t *tmp = (uint32_t *)data; + ASSERT_EQ(1, butil::NetToHost32(*tmp)); + bthread_id_join(cntl.call_id()); + + ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); +} - 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. - bthread_id_join(cntl.call_id()); +TEST_F(RdmaTest, server_close_after_ack) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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)); + + rdma::v2_wire::HelloMessage msg{}; + 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); + 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(4, read(acc_fd, data, 4)); + uint32_t *tmp = (uint32_t *)data; + ASSERT_EQ(1, butil::NetToHost32(*tmp)); + close(acc_fd); + bthread_id_join(cntl.call_id()); + + ASSERT_EQ(EEOF, cntl.ErrorCode()); +} - ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); +TEST_F(RdmaTest, server_send_data_on_tcp_after_ack) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &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)); + + rdma::v2_wire::HelloMessage msg{}; + 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); + 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, server_close_during_magic_str) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); +TEST_F(RdmaTest, v2_client_hello_bytes_baseline) { + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); + + 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)); + + // [0..4) magic + ASSERT_EQ(0, memcmp(data, "RDMA", 4)); + // [4..6) msg_len, big-endian uint16 == 40 + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + (size_t)(((uint16_t)data[4] << 8) | (uint16_t)data[5])); + // [6..8) hello_ver, big-endian uint16 == rdma::HELLO_V2_VERSION + ASSERT_EQ(rdma::HELLO_V2_VERSION, + (uint16_t)(((uint16_t)data[6] << 8) | (uint16_t)data[7])); + // [8..10) impl_ver, big-endian uint16 == rdma::IMPL_V2_VERSION + ASSERT_EQ(rdma::IMPL_V2_VERSION, + (uint16_t)(((uint16_t)data[8] << 8) | (uint16_t)data[9])); + + rdma::v2_wire::HelloMessage msg{}; + 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); + + bthread_id_join(cntl.call_id()); +} - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); +TEST_F(RdmaTest, v2_server_hello_bytes_baseline) { + StartServer(); + + 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); + Socket *s = GetSocketFromServer(0); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + + // 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; + 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(); + + 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)); + 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(0, memcmp(reply, "RDMA", 4)); + ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, + (size_t)(((uint16_t)reply[4] << 8) | (uint16_t)reply[5])); + ASSERT_EQ(rdma::HELLO_V2_VERSION, + (uint16_t)(((uint16_t)reply[6] << 8) | (uint16_t)reply[7])); + ASSERT_EQ(rdma::IMPL_V2_VERSION, + (uint16_t)(((uint16_t)reply[8] << 8) | (uint16_t)reply[9])); + + rdma::v2_wire::HelloMessage reply_msg{}; + 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); + + // 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))); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); + + sockfd.reset(-1); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + StopServer(); +} + +TEST_F(RdmaTest, v2_server_preserves_coalesced_ack_after_extension) { + StartServer(); + + 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); + 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{}; + msg.msg_len = 48; + 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(); + + uint8_t buf[52]; + memcpy(buf, "RDMA", 4); + msg.Serialize(buf + 4); + memset(buf + 40, 0x00, 8); // 8B zero tail + // 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); + 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); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + StopServer(); +} - 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); +TEST_F(RdmaTest, v2_server_rejects_oversized_msg_len) { + StartServer(); + + 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); + Socket *s = GetSocketFromServer(0); + ASSERT_TRUE(s != nullptr); + 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 + // (and definitely before attempting to drain) any "tail". + rdma::v2_wire::HelloMessage msg{}; + msg.msg_len = 4097; + 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(); + + 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)); + usleep(100000); + + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + sockfd.reset(-1); + usleep(100000); + + StopServer(); +} - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); +// RAII for FLAGS_rdma_client_handshake_version: lets us flip the +// client-side handshake version for a single test and restore it on +// scope exit so subsequent tests stay on the v2 default. +class HandshakeVersionFlag { +public: + explicit HandshakeVersionFlag(int v) + : _saved(rdma::FLAGS_rdma_client_handshake_version) { + rdma::FLAGS_rdma_client_handshake_version = v; + } + ~HandshakeVersionFlag() { + rdma::FLAGS_rdma_client_handshake_version = _saved; + } - 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)); - bthread_id_join(cntl.call_id()); +private: + int _saved; +}; - ASSERT_EQ(EEOF, cntl.ErrorCode()); +// Build a v3 wire packet from an RdmaHello: "RDM3" + pb_size_be + body. +std::string MakeV3Packet(const rdma::RdmaHello &msg) { + std::string body; + EXPECT_TRUE(msg.SerializeToString(&body)); + std::string packet; + packet.reserve(4 + 4 + body.size()); + packet.append("RDM3", 4); + uint32_t pb_size_be = butil::HostToNet32(static_cast(body.size())); + packet.append(reinterpret_cast(&pb_size_be), 4); + packet.append(body); + return packet; } -TEST_F(RdmaTest, server_hello_invalid_magic_str) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); +// Build a fully-valid RdmaHello: all 6 required fields are set, with +// values that pass RdmaHelloV3Wire::RdmaHelloValid(). +// - block_size = 8192 (>= MIN_BLOCK_SIZE) +// - sq_size / rq_size = 16 (>= MIN_QP_SIZE) +// - gid = exactly 16B (sizeof(ibv_gid)) +// - qp_num = 0 (allowed because g_skip_rdma_init in UT) +rdma::RdmaHello MakeValidV3Hello() { + rdma::RdmaHello msg; + msg.set_block_size(8192); + msg.set_sq_size(16); + 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_qp_num(0); + return msg; +} - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); +TEST_F(RdmaTest, v3_client_hello_bytes_baseline) { + HandshakeVersionFlag _hsv(3); + + butil::fd_guard sockfd(butil::tcp_listen(g_ep)); + EXPECT_TRUE(sockfd >= 0); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + 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); + + butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); + ASSERT_TRUE(acc_fd >= 0); + + // [0..4) magic "RDM3" + uint8_t magic[4]; + ASSERT_EQ(4, read(acc_fd, magic, 4)); + ASSERT_EQ(0, memcmp(magic, "RDM3", 4)); + + // [4..8) pb_size, big-endian uint32, must be in (0, 4096] + uint8_t size_buf[4]; + ASSERT_EQ(4, read(acc_fd, size_buf, 4)); + uint32_t pb_size = + butil::NetToHost32(*reinterpret_cast(size_buf)); + ASSERT_GT(pb_size, 0u); + ASSERT_LE(pb_size, 4096u); + + // [8..8+pb_size) RdmaHello protobuf body. + std::string body(pb_size, '\0'); + ASSERT_EQ((ssize_t)pb_size, read(acc_fd, &body[0], pb_size)); + rdma::RdmaHello msg; + ASSERT_TRUE(msg.ParseFromString(body)); + + // All 6 required fields must be present (ParseFromString would + // have already returned false otherwise). + ASSERT_TRUE(msg.has_block_size()); + ASSERT_TRUE(msg.has_sq_size()); + ASSERT_TRUE(msg.has_rq_size()); + ASSERT_TRUE(msg.has_lid()); + ASSERT_TRUE(msg.has_gid()); + ASSERT_TRUE(msg.has_qp_num()); + // gid wire encoding must be exactly 16 bytes (sizeof(ibv_gid)). + ASSERT_EQ(sizeof(ibv_gid), msg.gid().size()); + + // Let the RPC time out and release resources. + bthread_id_join(cntl.call_id()); +} - 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); +TEST_F(RdmaTest, v3_server_hello_bytes_baseline) { + StartServer(); + + 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); + 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())); + 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]; + ASSERT_EQ(4, read(sockfd, reply_magic, 4)); + ASSERT_EQ(0, memcmp(reply_magic, "RDM3", 4)); + + uint8_t size_buf[4]; + ASSERT_EQ(4, read(sockfd, size_buf, 4)); + uint32_t pb_size = + butil::NetToHost32(*reinterpret_cast(size_buf)); + ASSERT_GT(pb_size, 0u); + ASSERT_LE(pb_size, 4096u); + + std::string body(pb_size, '\0'); + ASSERT_EQ((ssize_t)pb_size, read(sockfd, &body[0], pb_size)); + rdma::RdmaHello reply; + ASSERT_TRUE(reply.ParseFromString(body)); + ASSERT_TRUE(reply.has_block_size()); + ASSERT_TRUE(reply.has_sq_size()); + ASSERT_TRUE(reply.has_rq_size()); + ASSERT_TRUE(reply.has_gid()); + ASSERT_EQ(sizeof(ibv_gid), reply.gid().size()); + + // 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))); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); + + sockfd.reset(-1); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + StopServer(); +} - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); +TEST_F(RdmaTest, v3_server_rejects_zero_pb_size) { + StartServer(); - 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(4, write(acc_fd, "ABCD", 4)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FAILED, RdmaTransportOf(s)); - bthread_id_join(cntl.call_id()); + 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); + Socket *s = GetSocketFromServer(0); + ASSERT_TRUE(s != nullptr); - ASSERT_EQ(EPROTO, cntl.ErrorCode()); -} + // "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)); + usleep(100000); -TEST_F(RdmaTest, server_miss_during_hello_msg) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); - - 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(4, write(acc_fd, "RDMA", 4)); - ASSERT_EQ(2, write(acc_fd, "00", 2)); - bthread_id_join(cntl.call_id()); - - ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); -} - -TEST_F(RdmaTest, server_close_during_hello_msg) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); - - 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(4, write(acc_fd, "RDMA", 4)); - ASSERT_EQ(2, write(acc_fd, "00", 2)); - close(acc_fd); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FAILED, RdmaTransportOf(s)); - bthread_id_join(cntl.call_id()); - - ASSERT_EQ(EEOF, cntl.ErrorCode()); -} - -TEST_F(RdmaTest, server_hello_invalid_msg_len) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); - - 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)); - 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)); - bthread_id_join(cntl.call_id()); - - ASSERT_EQ(EPROTO, cntl.ErrorCode()); -} - -TEST_F(RdmaTest, server_hello_invalid_version) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); - - 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)); - 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(4, read(acc_fd, data, 4)); - uint32_t* tmp = (uint32_t*)data; - ASSERT_EQ(0, butil::NetToHost32(*tmp)); - bthread_id_join(cntl.call_id()); - - ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); -} - -TEST_F(RdmaTest, server_hello_invalid_sq_rq_size) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); - - 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)); - - rdma::v2_wire::HelloMessage msg{}; - msg.msg_len = rdma::HELLO_V2_MSG_LEN_MIN; - msg.hello_ver = 1; - msg.impl_ver = 1; - msg.sq_size = 0; - msg.rq_size = 0; - msg.block_size = 8192; - msg.qp_num = 0; - 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_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); - ASSERT_EQ(4, read(acc_fd, data, 4)); - uint32_t* tmp = (uint32_t*)data; - ASSERT_EQ(0, butil::NetToHost32(*tmp)); - bthread_id_join(cntl.call_id()); - - ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); -} - -TEST_F(RdmaTest, server_miss_after_ack) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); - - 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)); - - rdma::v2_wire::HelloMessage msg{}; - 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); - 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(4, read(acc_fd, data, 4)); - uint32_t* tmp = (uint32_t*)data; - ASSERT_EQ(1, butil::NetToHost32(*tmp)); - bthread_id_join(cntl.call_id()); - - ASSERT_EQ(ERPCTIMEDOUT, cntl.ErrorCode()); -} - -TEST_F(RdmaTest, server_close_after_ack) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); - - 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)); - - rdma::v2_wire::HelloMessage msg{}; - 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); - 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(4, read(acc_fd, data, 4)); - uint32_t* tmp = (uint32_t*)data; - ASSERT_EQ(1, butil::NetToHost32(*tmp)); - close(acc_fd); - bthread_id_join(cntl.call_id()); - - ASSERT_EQ(EEOF, cntl.ErrorCode()); -} - -TEST_F(RdmaTest, server_send_data_on_tcp_after_ack) { - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::C_HELLO_WAIT, RdmaTransportOf(s)); - - 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)); - - rdma::v2_wire::HelloMessage msg{}; - 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); - 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)); - 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); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - - 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)); - - // [0..4) magic - ASSERT_EQ(0, memcmp(data, "RDMA", 4)); - // [4..6) msg_len, big-endian uint16 == 40 - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, - (size_t)(((uint16_t)data[4] << 8) | (uint16_t)data[5])); - // [6..8) hello_ver, big-endian uint16 == rdma::HELLO_V2_VERSION - ASSERT_EQ(rdma::HELLO_V2_VERSION, - (uint16_t)(((uint16_t)data[6] << 8) | (uint16_t)data[7])); - // [8..10) impl_ver, big-endian uint16 == rdma::IMPL_V2_VERSION - ASSERT_EQ(rdma::IMPL_V2_VERSION, - (uint16_t)(((uint16_t)data[8] << 8) | (uint16_t)data[9])); - - rdma::v2_wire::HelloMessage msg{}; - 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); - - bthread_id_join(cntl.call_id()); -} - -TEST_F(RdmaTest, v2_server_hello_bytes_baseline) { - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - - // Send a well-formed v2 hello so the server enters S_ACK_WAIT. - rdma::v2_wire::HelloMessage msg{}; - 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(); - - 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)); - - // 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(0, memcmp(reply, "RDMA", 4)); - ASSERT_EQ(rdma::HELLO_V2_MSG_LEN_MIN, - (size_t)(((uint16_t)reply[4] << 8) | (uint16_t)reply[5])); - ASSERT_EQ(rdma::HELLO_V2_VERSION, - (uint16_t)(((uint16_t)reply[6] << 8) | (uint16_t)reply[7])); - ASSERT_EQ(rdma::IMPL_V2_VERSION, - (uint16_t)(((uint16_t)reply[8] << 8) | (uint16_t)reply[9])); - - rdma::v2_wire::HelloMessage reply_msg{}; - 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); - - // 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)); - - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - StopServer(); -} - -TEST_F(RdmaTest, v2_server_drains_tail_then_reads_ack) { - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - - // Build a v2 hello with msg_len = 48 (40 base + 8B zero tail). - rdma::v2_wire::HelloMessage msg{}; - msg.msg_len = 48; - 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(); - - uint8_t buf[48]; - 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). - uint32_t flags = butil::HostToNet32(1); - ASSERT_EQ(sizeof(flags), write(sockfd, &flags, sizeof(flags))); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::ESTABLISHED, RdmaTransportOf(s)); - - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - StopServer(); -} - -TEST_F(RdmaTest, v2_server_rejects_oversized_msg_len) { - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - - // 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 - // (and definitely before attempting to drain) any "tail". - rdma::v2_wire::HelloMessage msg{}; - msg.msg_len = 4097; - 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(); - - 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()); - - sockfd.reset(-1); - - StopServer(); -} - -// RAII for FLAGS_rdma_client_handshake_version: lets us flip the -// client-side handshake version for a single test and restore it on -// scope exit so subsequent tests stay on the v2 default. -class HandshakeVersionFlag { -public: - explicit HandshakeVersionFlag(int v) - : _saved(rdma::FLAGS_rdma_client_handshake_version) { - rdma::FLAGS_rdma_client_handshake_version = v; - } - ~HandshakeVersionFlag() { - rdma::FLAGS_rdma_client_handshake_version = _saved; - } -private: - int _saved; -}; - -// Build a v3 wire packet from an RdmaHello: "RDM3" + pb_size_be + body. -std::string MakeV3Packet(const rdma::RdmaHello& msg) { - std::string body; - EXPECT_TRUE(msg.SerializeToString(&body)); - std::string packet; - packet.reserve(4 + 4 + body.size()); - packet.append("RDM3", 4); - uint32_t pb_size_be = - butil::HostToNet32(static_cast(body.size())); - packet.append(reinterpret_cast(&pb_size_be), 4); - packet.append(body); - return packet; -} - -// Build a fully-valid RdmaHello: all 6 required fields are set, with -// values that pass RdmaHelloV3Wire::RdmaHelloValid(). -// - block_size = 8192 (>= MIN_BLOCK_SIZE) -// - sq_size / rq_size = 16 (>= MIN_QP_SIZE) -// - gid = exactly 16B (sizeof(ibv_gid)) -// - qp_num = 0 (allowed because g_skip_rdma_init in UT) -rdma::RdmaHello MakeValidV3Hello() { - rdma::RdmaHello msg; - msg.set_block_size(8192); - msg.set_sq_size(16); - 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_qp_num(0); - return msg; -} - - -TEST_F(RdmaTest, v3_client_hello_bytes_baseline) { - HandshakeVersionFlag _hsv(3); - - butil::fd_guard sockfd(butil::tcp_listen(g_ep)); - EXPECT_TRUE(sockfd >= 0); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - 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); - - butil::fd_guard acc_fd(accept(sockfd, nullptr, nullptr)); - ASSERT_TRUE(acc_fd >= 0); - - // [0..4) magic "RDM3" - uint8_t magic[4]; - ASSERT_EQ(4, read(acc_fd, magic, 4)); - ASSERT_EQ(0, memcmp(magic, "RDM3", 4)); - - // [4..8) pb_size, big-endian uint32, must be in (0, 4096] - uint8_t size_buf[4]; - ASSERT_EQ(4, read(acc_fd, size_buf, 4)); - uint32_t pb_size = - butil::NetToHost32(*reinterpret_cast(size_buf)); - ASSERT_GT(pb_size, 0u); - ASSERT_LE(pb_size, 4096u); - - // [8..8+pb_size) RdmaHello protobuf body. - std::string body(pb_size, '\0'); - ASSERT_EQ((ssize_t)pb_size, read(acc_fd, &body[0], pb_size)); - rdma::RdmaHello msg; - ASSERT_TRUE(msg.ParseFromString(body)); - - // All 6 required fields must be present (ParseFromString would - // have already returned false otherwise). - ASSERT_TRUE(msg.has_block_size()); - ASSERT_TRUE(msg.has_sq_size()); - ASSERT_TRUE(msg.has_rq_size()); - ASSERT_TRUE(msg.has_lid()); - ASSERT_TRUE(msg.has_gid()); - ASSERT_TRUE(msg.has_qp_num()); - // gid wire encoding must be exactly 16 bytes (sizeof(ibv_gid)). - ASSERT_EQ(sizeof(ibv_gid), msg.gid().size()); - - // Let the RPC time out and release resources. - bthread_id_join(cntl.call_id()); -} - -TEST_F(RdmaTest, v3_server_hello_bytes_baseline) { - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - - // 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)); - - // Read server's reply hello: 4B magic + 4B pb_size + body. - uint8_t reply_magic[4]; - ASSERT_EQ(4, read(sockfd, reply_magic, 4)); - ASSERT_EQ(0, memcmp(reply_magic, "RDM3", 4)); - - uint8_t size_buf[4]; - ASSERT_EQ(4, read(sockfd, size_buf, 4)); - uint32_t pb_size = - butil::NetToHost32(*reinterpret_cast(size_buf)); - ASSERT_GT(pb_size, 0u); - ASSERT_LE(pb_size, 4096u); - - std::string body(pb_size, '\0'); - ASSERT_EQ((ssize_t)pb_size, read(sockfd, &body[0], pb_size)); - rdma::RdmaHello reply; - ASSERT_TRUE(reply.ParseFromString(body)); - ASSERT_TRUE(reply.has_block_size()); - ASSERT_TRUE(reply.has_sq_size()); - ASSERT_TRUE(reply.has_rq_size()); - ASSERT_TRUE(reply.has_gid()); - ASSERT_EQ(sizeof(ibv_gid), reply.gid().size()); - - // 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)); - - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - StopServer(); -} - -TEST_F(RdmaTest, v3_server_rejects_zero_pb_size) { - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - 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()); - - sockfd.reset(-1); - StopServer(); + sockfd.reset(-1); + StopServer(); } TEST_F(RdmaTest, v3_server_rejects_oversized_pb_size) { - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - - uint8_t buf[8]; - memcpy(buf, "RDM3", 4); - // pb_size just above the allowed maximum -> rejected. - uint32_t pb_size_be = - 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()); - - sockfd.reset(-1); - StopServer(); + StartServer(); + + 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); + Socket *s = GetSocketFromServer(0); + ASSERT_TRUE(s != nullptr); + + uint8_t buf[8]; + memcpy(buf, "RDM3", 4); + // pb_size just above the allowed maximum -> rejected. + uint32_t pb_size_be = + 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)); + usleep(100000); + + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + sockfd.reset(-1); + StopServer(); } TEST_F(RdmaTest, v3_server_rejects_invalid_pb_bytes) { - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - - // "RDM3" + pb_size = 8 + 8 bytes of 0xff (invalid protobuf body). - uint8_t buf[16]; - memcpy(buf, "RDM3", 4); - uint32_t pb_size_be = butil::HostToNet32(8); - memcpy(buf + 4, &pb_size_be, 4); - memset(buf + 8, 0xff, 8); - ASSERT_EQ(16, write(sockfd, buf, 16)); - ASSERT_TRUE(WaitForServerSocketGone()); - - sockfd.reset(-1); - StopServer(); + StartServer(); + + 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); + Socket *s = GetSocketFromServer(0); + ASSERT_TRUE(s != nullptr); + + // "RDM3" + pb_size = 8 + 8 bytes of 0xff (invalid protobuf body). + uint8_t buf[16]; + memcpy(buf, "RDM3", 4); + uint32_t pb_size_be = butil::HostToNet32(8); + memcpy(buf + 4, &pb_size_be, 4); + memset(buf + 8, 0xff, 8); + ASSERT_EQ(16, write(sockfd, buf, 16)); + usleep(100000); + + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + sockfd.reset(-1); + StopServer(); } TEST_F(RdmaTest, v3_server_invalid_sq_size_falls_back) { - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - 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())); - - // 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); - - // 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)); - 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)); - - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - StopServer(); + StartServer(); + + 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); + 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_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 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)); + 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))); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); + + sockfd.reset(-1); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + StopServer(); } // RAII guard to toggle FLAGS_rdma_ece for a single test and restore it. class EceFlagGuard { public: - explicit EceFlagGuard(bool v) : _saved(rdma::FLAGS_rdma_ece) { - rdma::FLAGS_rdma_ece = v; - } - ~EceFlagGuard() { - rdma::FLAGS_rdma_ece = _saved; - } + explicit EceFlagGuard(bool v) : _saved(rdma::FLAGS_rdma_ece) { + rdma::FLAGS_rdma_ece = v; + } + ~EceFlagGuard() { rdma::FLAGS_rdma_ece = _saved; } + private: - bool _saved; + 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(); - ece->set_vendor_id(vendor_id); - ece->set_options(options); - ece->set_comp_mask(comp_mask); - return msg; + rdma::RdmaHello msg = MakeValidV3Hello(); + rdma::RdmaEce *ece = msg.mutable_ece(); + ece->set_vendor_id(vendor_id); + ece->set_options(options); + ece->set_comp_mask(comp_mask); + return msg; } // Read the server's v3 reply hello (4B magic + 4B pb_size + body) and parse // it into `reply`. Asserts the framing along the way. -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)); - 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_TRUE(reply->ParseFromString(reply_body)); +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)); + 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_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(); - - 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))); - Socket* s = WaitForServerSocket(); - 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)); - - rdma::RdmaHello reply; - ReadServerV3Reply(sockfd, &reply); - - // 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)); - - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - StopServer(); + EceFlagGuard ece_flag_guard(true); + StartServer(); + + 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); + 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())); + usleep(100000); + + ASSERT_EQ(handshake::ACK_WAIT, AdapterTransport::Get(s)->handshake_phase()); + + rdma::RdmaHello reply; + ReadServerV3Reply(sockfd, &reply); + + // 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))); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); + + sockfd.reset(-1); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + StopServer(); } // When ECE negotiation is disabled, the server reply must NOT advertise ECE, // even if the client advertised it (FillLocalRdmaHello degrade branch #1). TEST_F(RdmaTest, v3_server_reply_has_no_ece_when_disabled) { - EceFlagGuard ece_flag_guard(false); - StartServer(); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - - rdma::RdmaHello msg = MakeValidV3HelloWithEce(0x02c9, 0x1, 0x0); - std::string packet = MakeV3Packet(msg); - ASSERT_TRUE(WriteAll(sockfd, packet.data(), packet.size())); - - // 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)); - - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - StopServer(); + EceFlagGuard ece_flag_guard(false); + StartServer(); + + 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); + 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())); + usleep(100000); + + rdma::RdmaHello reply; + ReadServerV3Reply(sockfd, &reply); + EXPECT_FALSE(reply.has_ece()); + + uint32_t flags = butil::HostToNet32(0); + ASSERT_EQ((ssize_t)sizeof(flags), write(sockfd, &flags, sizeof(flags))); + usleep(100000); + + sockfd.reset(-1); + 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(); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - - rdma::RdmaHello msg = MakeValidV3HelloWithEce(0x02c9, 0x1, 0x0); - std::string packet = MakeV3Packet(msg); - ASSERT_TRUE(WriteAll(sockfd, packet.data(), packet.size())); - - // 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)); - - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - StopServer(); + EceFlagGuard ece_flag_guard(true); + StartServer(); + + 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); + 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())); + usleep(100000); + + rdma::RdmaHello reply; + ReadServerV3Reply(sockfd, &reply); + EXPECT_FALSE(reply.has_ece()); + + uint32_t flags = butil::HostToNet32(0); + ASSERT_EQ((ssize_t)sizeof(flags), write(sockfd, &flags, sizeof(flags))); + usleep(100000); + + sockfd.reset(-1); + usleep(100000); + StopServer(); } class ResourceAllocFailGuard { public: - explicit ResourceAllocFailGuard(bool v) - : _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; - } + explicit ResourceAllocFailGuard(bool v) + : _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; } + private: - bool _saved; + bool _saved; }; -TEST_F(RdmaTest, client_alloc_resource_fail_fallback_tcp) { - StartServer(); - ResourceAllocFailGuard alloc_fail_guard(true); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - Controller cntl; - test::EchoRequest req; - test::EchoResponse res; - req.set_message(__FUNCTION__); - req.set_sleep_us(200000); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); - - 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); - // The socket must not be failed, otherwise it can no longer carry TCP. - ASSERT_FALSE(s->Failed()); - - // The RPC still completes over TCP. - bthread_id_join(cntl.call_id()); - ASSERT_EQ(0, cntl.ErrorCode()) << cntl.ErrorText(); - - StopServer(); -} - -TEST_F(RdmaTest, server_alloc_resource_fail_fallback_tcp) { - StartServer(); - ResourceAllocFailGuard alloc_fail_guard(true); - - 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))); - Socket* s = WaitForServerSocket(); - ASSERT_TRUE(s != nullptr); - ASSERT_EQ(rdma::RdmaEndpoint::UNINIT, RdmaTransportOf(s)->_rdma_ep->_state); - - // Send a well-formed v2 hello: the negotiation succeeds - // but the resource allocation does not. - rdma::v2_wire::HelloMessage msg{}; - 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(); - - 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(RdmaTransport::RDMA_OFF, RdmaTransportOf(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)); - ASSERT_FALSE(s->Failed()); - - sockfd.reset(-1); - ASSERT_TRUE(WaitForServerSocketGone()); - - StopServer(); -} - -TEST_F(RdmaTest, try_global_disable_rdma) { - StartServer(); - rdma::g_rdma_available.store(false, butil::memory_order_relaxed); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl; - test::EchoRequest req; - test::EchoResponse res; - - req.set_message(__FUNCTION__); - req.set_sleep_us(200000); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl._single_server_id, &s)); - ASSERT_RDMA_STATE(rdma::RdmaEndpoint::FALLBACK_TCP, RdmaTransportOf(s)); - bthread_id_join(cntl.call_id()); - ASSERT_EQ(0, cntl.ErrorCode()); - - StopServer(); - rdma::g_rdma_available.store(true, butil::memory_order_relaxed); +TEST_F(RdmaTest, client_alloc_resource_fail_fallback_tcp) { + StartServer(); + ResourceAllocFailGuard alloc_fail_guard(true); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + req.set_message(__FUNCTION__); + 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_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()); + + // The RPC still completes over TCP. + bthread_id_join(cntl.call_id()); + ASSERT_EQ(0, cntl.ErrorCode()) << cntl.ErrorText(); + + StopServer(); +} + +TEST_F(RdmaTest, server_alloc_resource_fail_fallback_tcp) { + StartServer(); + ResourceAllocFailGuard alloc_fail_guard(true); + + 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_TRUE(s != nullptr); + ASSERT_EQ(handshake::UNINITIALIZED, + AdapterTransport::Get(s)->handshake_phase()); + + // Send a well-formed v2 hello: the negotiation succeeds + // but the resource allocation does not. + rdma::v2_wire::HelloMessage msg{}; + 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(); + + 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)); + 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))); + usleep(100000); + ASSERT_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s)->handshake_phase()); + ASSERT_FALSE(s->Failed()); + + sockfd.reset(-1); + usleep(100000); + ASSERT_EQ(nullptr, GetSocketFromServer(0)); + + StopServer(); +} + +TEST_F(RdmaTest, try_global_disable_rdma) { + StartServer(); + rdma::g_rdma_available.store(false, butil::memory_order_relaxed); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + + req.set_message(__FUNCTION__); + 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_EQ(handshake::FALLBACK_TCP, + AdapterTransport::Get(s.get())->handshake_phase()); + bthread_id_join(cntl.call_id()); + ASSERT_EQ(0, cntl.ErrorCode()); + + StopServer(); + rdma::g_rdma_available.store(true, butil::memory_order_relaxed); } TEST_F(RdmaTest, server_option_invalid) { - Server server; - ServerOptions options; - options.socket_mode = SOCKET_MODE_RDMA; + Server server; + ServerOptions options; + options.socket_mode = SOCKET_MODE_RDMA; - // rtmp and rdma are incompatible - options.rtmp_service = (RtmpService*)1; - ASSERT_EQ(-1, server.Start(PORT, &options)); + // rtmp and rdma are incompatible + options.rtmp_service = (RtmpService *)1; + ASSERT_EQ(-1, server.Start(PORT, &options)); - // nshead and rdma are incompatible - options.rtmp_service = nullptr; - options.nshead_service = (NsheadService*)1; - ASSERT_EQ(-1, server.Start(PORT, &options)); + // nshead and rdma are incompatible + options.rtmp_service = nullptr; + options.nshead_service = (NsheadService *)1; + ASSERT_EQ(-1, server.Start(PORT, &options)); - // mongo and rdma are incompatible - options.nshead_service = nullptr; - options.mongo_service_adaptor = (MongoServiceAdaptor*)1; - ASSERT_EQ(-1, server.Start(PORT, &options)); + // mongo and rdma are incompatible + options.nshead_service = nullptr; + options.mongo_service_adaptor = (MongoServiceAdaptor *)1; + ASSERT_EQ(-1, server.Start(PORT, &options)); - // ssl and rdma are incompatible - options.mongo_service_adaptor = nullptr; - options.mutable_ssl_options()->default_cert.certificate = "test"; - ASSERT_EQ(-1, server.Start(PORT, &options)); + // ssl and rdma are incompatible + options.mongo_service_adaptor = nullptr; + options.mutable_ssl_options()->default_cert.certificate = "test"; + ASSERT_EQ(-1, server.Start(PORT, &options)); } TEST_F(RdmaTest, channel_option_invalid) { - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; - // rtmp and rdma are incompatible - chan_options.protocol = "rtmp"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // rtmp and rdma are incompatible + chan_options.protocol = "rtmp"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - chan_options.protocol = "streaming_rpc"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + chan_options.protocol = "streaming_rpc"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // nshead and rdma are incompatible - chan_options.protocol = "nshead"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - chan_options.protocol = "nshead_mcpack"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // nshead and rdma are incompatible + chan_options.protocol = "nshead"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + chan_options.protocol = "nshead_mcpack"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // nova_pbrpc and rdma are incompatible - chan_options.protocol = "nova_pbrpc"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // nova_pbrpc and rdma are incompatible + chan_options.protocol = "nova_pbrpc"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // public_pbrpc and rdma are incompatible - chan_options.protocol = "public_pbrpc"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // public_pbrpc and rdma are incompatible + chan_options.protocol = "public_pbrpc"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // redis and rdma are incompatible - chan_options.protocol = "redis"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // redis and rdma are incompatible + chan_options.protocol = "redis"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // memcache and rdma are incompatible - chan_options.protocol = "memcache"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // memcache and rdma are incompatible + chan_options.protocol = "memcache"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // ubrpc and rdma are incompatible - chan_options.protocol = "ubrpc_compack"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // ubrpc and rdma are incompatible + chan_options.protocol = "ubrpc_compack"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // itp and rdma are incompatible - chan_options.protocol = "itp"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // itp and rdma are incompatible + chan_options.protocol = "itp"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // esp and rdma are incompatible - chan_options.protocol = "esp"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // esp and rdma are incompatible + chan_options.protocol = "esp"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // hulu_pbrpc and rdma are incompatible - chan_options.protocol = "hulu_pbrpc"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // hulu_pbrpc and rdma are incompatible + chan_options.protocol = "hulu_pbrpc"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // sofa_pbrpc and rdma are incompatible - chan_options.protocol = "sofa_pbrpc"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // sofa_pbrpc and rdma are incompatible + chan_options.protocol = "sofa_pbrpc"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // http and rdma are incompatible - chan_options.protocol = "http"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // http and rdma are incompatible + chan_options.protocol = "http"; + ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); - // ssl and rdma are incompatible - chan_options.protocol = "baidu_std"; - chan_options.mutable_ssl_options()->sni_name = "test"; - ASSERT_EQ(-1, channel.Init(g_ep, &chan_options)); + // ssl and rdma are incompatible + chan_options.protocol = "baidu_std"; + chan_options.mutable_ssl_options()->sni_name = "test"; + 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; @@ -2368,77 +2609,99 @@ 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; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 5000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - ASSERT_NO_FATAL_FAILURE(SendEchoRpcsInRounds(channel)); - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + 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(); } TEST_P(RdmaRpcTest, tcp_client_to_tcp_server) { - StartServer(false); - - Channel channel; - ChannelOptions chan_options; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 5000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - ASSERT_NO_FATAL_FAILURE(SendEchoRpcsInRounds(channel)); - - StopServer(); + StartServer(false); + + Channel channel; + ChannelOptions chan_options; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + 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(); } TEST_P(RdmaRpcTest, tcp_client_to_rdma_server) { - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 5000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - ASSERT_NO_FATAL_FAILURE(SendEchoRpcsInRounds(channel)); - - StopServer(); + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + 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(); } TEST_P(RdmaRpcTest, rdma_client_to_tcp_server) { - StartServer(false); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 5000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - ASSERT_NO_FATAL_FAILURE(SendEchoRpcsInRounds(channel)); - - StopServer(); + StartServer(false); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + 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(); } TEST_P(RdmaRpcTest, tcp_client_to_rdma_server_short_connection) { @@ -2459,7 +2722,6 @@ TEST_P(RdmaRpcTest, tcp_client_to_rdma_server_short_connection) { 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; @@ -2506,365 +2768,365 @@ 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); +void DumpRdmaEndpointInfo(Socket *client, Socket *server) { + 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) { - if (!FLAGS_rdma_test_enable) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 50000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - - LOG(INFO) << "send 0 attachment"; - for (int i = 0; i < RPC_NUM; ++i) { - req[i].set_message(__FUNCTION__); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - for (int i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); - Socket* m = GetSocketFromServer(0); - DumpRdmaEndpointInfo(s.get(), m); - } - ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; - } - - LOG(INFO) << "send 4KB attachment"; - butil::IOBuf attach; - attach.resize(4096); - for (int i = 0; i < RPC_NUM; ++i) { - cntl[i].Reset(); - cntl[i].request_attachment().append(attach); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - for (int i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); - Socket* m = GetSocketFromServer(0); - DumpRdmaEndpointInfo(s.get(), m); - } - ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; - } - - LOG(INFO) << "send 1MB attachment"; - attach.resize(1048576); - for (int i = 0; i < RPC_NUM; ++i) { - cntl[i].Reset(); - cntl[i].request_attachment().append(attach); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - for (int i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); - Socket* m = GetSocketFromServer(0); - DumpRdmaEndpointInfo(s.get(), m); - } - ASSERT_TRUE(0 == cntl[i].ErrorCode() || - EOVERCROWDED == cntl[i].ErrorCode()) << "req[" << i << "] " << berror(cntl[i].ErrorCode()); - } - - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl[0]._single_server_id, &s)); - Socket* m = GetSocketFromServer(0); - DumpRdmaEndpointInfo(s.get(), m); - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 50000; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + + LOG(INFO) << "send 0 attachment"; + for (int i = 0; i < RPC_NUM; ++i) { + req[i].set_message(__FUNCTION__); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + for (int i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); + Socket *m = GetSocketFromServer(0); + DumpRdmaEndpointInfo(s.get(), m); + } + ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; + } + + LOG(INFO) << "send 4KB attachment"; + butil::IOBuf attach; + attach.resize(4096); + for (int i = 0; i < RPC_NUM; ++i) { + cntl[i].Reset(); + cntl[i].request_attachment().append(attach); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + for (int i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); + Socket *m = GetSocketFromServer(0); + DumpRdmaEndpointInfo(s.get(), m); + } + ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; + } + + LOG(INFO) << "send 1MB attachment"; + attach.resize(1048576); + for (int i = 0; i < RPC_NUM; ++i) { + cntl[i].Reset(); + cntl[i].request_attachment().append(attach); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + for (int i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); + Socket *m = GetSocketFromServer(0); + DumpRdmaEndpointInfo(s.get(), m); + } + ASSERT_TRUE(0 == cntl[i].ErrorCode() || EOVERCROWDED == cntl[i].ErrorCode()) + << "req[" << i << "] " << berror(cntl[i].ErrorCode()); + } + + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl[0]._single_server_id, &s)); + Socket *m = GetSocketFromServer(0); + DumpRdmaEndpointInfo(s.get(), m); + + StopServer(); } TEST_P(RdmaRpcTest, send_rpc_in_many_qp) { - if (!FLAGS_rdma_test_enable) { - return; - } - - butil::ip_t ip; - ASSERT_EQ(0, butil::str2ip(g_ip.c_str(), &ip)); - - Server server[100]; - MyEchoService svc[100]; - int num = 100; - butil::EndPoint server_eps[100]; - for (int i = 0; i < num; ++i) { - ServerOptions options; - options.socket_mode = SOCKET_MODE_RDMA; - options.idle_timeout_sec = 1; - options.max_concurrency = 0; - options.internal_port = -1; - server[i].AddService(&svc[i], SERVER_DOESNT_OWN_SERVICE); - ASSERT_EQ(0, server[i].Start(0, &options)); - server_eps[i] = butil::EndPoint(ip, server[i].listen_address().port); - } - - int port = 0; - butil::IOBuf attach; - attach.resize(4096); - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 100000; - chan_options.max_retry = 0; - Channel channel[RPC_NUM]; - Server* svr[RPC_NUM]; - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - for (int i = 0; i < RPC_NUM; ++i) { - svr[i] = &server[i % num]; - ASSERT_EQ(0, channel[i].Init(server_eps[(port++) % num], &chan_options)); - 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); - } - for (int i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { - SocketUniquePtr s; - EXPECT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); - if (s && svr[i] && svr[i]->_am) { - std::vector sids; - svr[i]->_am->ListConnections(&sids); - for (size_t j = 0; j < sids.size(); ++j) { - SocketUniquePtr m; - if (Socket::AddressFailedAsWell(sids[j], &m) == 0) { - DumpRdmaEndpointInfo(s.get(), m.get()); - } - } - } + if (!FLAGS_rdma_test_enable) { + return; + } + + butil::ip_t ip; + ASSERT_EQ(0, butil::str2ip(g_ip.c_str(), &ip)); + + Server server[100]; + MyEchoService svc[100]; + int num = 100; + butil::EndPoint server_eps[100]; + for (int i = 0; i < num; ++i) { + ServerOptions options; + options.socket_mode = SOCKET_MODE_RDMA; + options.idle_timeout_sec = 1; + options.max_concurrency = 0; + options.internal_port = -1; + server[i].AddService(&svc[i], SERVER_DOESNT_OWN_SERVICE); + ASSERT_EQ(0, server[i].Start(0, &options)); + server_eps[i] = butil::EndPoint(ip, server[i].listen_address().port); + } + + int port = 0; + butil::IOBuf attach; + attach.resize(4096); + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 100000; + chan_options.max_retry = 0; + Channel channel[RPC_NUM]; + Server *svr[RPC_NUM]; + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + for (int i = 0; i < RPC_NUM; ++i) { + svr[i] = &server[i % num]; + ASSERT_EQ(0, channel[i].Init(server_eps[(port++) % num], &chan_options)); + 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); + } + for (int i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { + SocketUniquePtr s; + EXPECT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); + if (s && svr[i] && svr[i]->_am) { + std::vector sids; + svr[i]->_am->ListConnections(&sids); + for (size_t j = 0; j < sids.size(); ++j) { + SocketUniquePtr m; + if (Socket::AddressFailedAsWell(sids[j], &m) == 0) { + DumpRdmaEndpointInfo(s.get(), m.get()); + } } - EXPECT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; + } } + EXPECT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; + } - for (int i = 0; i < num; ++i) { - server[i].Stop(0); - server[i].Join(); - } + for (int i = 0; i < num; ++i) { + server[i].Stop(0); + server[i].Join(); + } } TEST_P(RdmaRpcTest, send_rpcs_as_pooled_connection) { - if (!FLAGS_rdma_test_enable) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 30000; // it may very slow - chan_options.timeout_ms = 30000; - chan_options.max_retry = 0; - chan_options.connection_type = "pooled"; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - - butil::IOBuf attach; - attach.resize(4096); - for (int i = 0; i < RPC_NUM; ++i) { - req[i].set_message(__FUNCTION__); - cntl[i].request_attachment().append(attach); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - for (int i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); - Socket* m = GetSocketFromServer(0); - DumpRdmaEndpointInfo(s.get(), m); - } - ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; - } - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 30000; // it may very slow + chan_options.timeout_ms = 30000; + chan_options.max_retry = 0; + chan_options.connection_type = "pooled"; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + + butil::IOBuf attach; + attach.resize(4096); + for (int i = 0; i < RPC_NUM; ++i) { + req[i].set_message(__FUNCTION__); + cntl[i].request_attachment().append(attach); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + for (int i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); + Socket *m = GetSocketFromServer(0); + DumpRdmaEndpointInfo(s.get(), m); + } + ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; + } + + StopServer(); } TEST_P(RdmaRpcTest, send_rpcs_as_short_connection) { - if (!FLAGS_rdma_test_enable) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 30000; // it may very slow - chan_options.timeout_ms = 30000; - chan_options.max_retry = 0; - chan_options.connection_type = "short"; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - - butil::IOBuf attach; - attach.resize(4096); - for (int i = 0; i < RPC_NUM; ++i) { - req[i].set_message(__FUNCTION__); - cntl[i].request_attachment().append(attach); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - for (int i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { - SocketUniquePtr s; - ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); - Socket* m = GetSocketFromServer(0); - DumpRdmaEndpointInfo(s.get(), m); - } - ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; - } - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 30000; // it may very slow + chan_options.timeout_ms = 30000; + chan_options.max_retry = 0; + chan_options.connection_type = "short"; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + + butil::IOBuf attach; + attach.resize(4096); + for (int i = 0; i < RPC_NUM; ++i) { + req[i].set_message(__FUNCTION__); + cntl[i].request_attachment().append(attach); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + for (int i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + if (cntl[i].ErrorCode() == ERPCTIMEDOUT) { + SocketUniquePtr s; + ASSERT_EQ(0, Socket::Address(cntl[i]._single_server_id, &s)); + Socket *m = GetSocketFromServer(0); + DumpRdmaEndpointInfo(s.get(), m); + } + ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; + } + + StopServer(); } TEST_P(RdmaRpcTest, server_stop_during_rpc) { - if (!FLAGS_rdma_test_enable) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 3000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - - butil::IOBuf attach; - attach.resize(4096); - for (int i = 0; i < RPC_NUM; ++i) { - req[i].set_message(__FUNCTION__); - cntl[i].request_attachment().append(attach); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - - for (int i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - 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; - } + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 3000; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + + butil::IOBuf attach; + attach.resize(4096); + for (int i = 0; i < RPC_NUM; ++i) { + req[i].set_message(__FUNCTION__); + cntl[i].request_attachment().append(attach); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + + for (int i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + 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; + } } TEST_P(RdmaRpcTest, server_close_during_rpc) { - if (!FLAGS_rdma_test_enable) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 3000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - - butil::IOBuf attach; - attach.resize(4096); - for (int i = 0; i < RPC_NUM; ++i) { - req[i].set_message(__FUNCTION__); - cntl[i].request_attachment().append(attach); - if (i == RPC_NUM / 2) { - req[i].set_close_fd(true); - } - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - - 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; - } - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 3000; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + + butil::IOBuf attach; + attach.resize(4096); + for (int i = 0; i < RPC_NUM; ++i) { + req[i].set_message(__FUNCTION__); + cntl[i].request_attachment().append(attach); + if (i == RPC_NUM / 2) { + req[i].set_close_fd(true); + } + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + + 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; + } + + StopServer(); } TEST_P(RdmaRpcTest, client_close_during_rpc) { - if (!FLAGS_rdma_test_enable) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 3000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - - butil::IOBuf attach; - attach.resize(4096); - for (int i = 0; i < RPC_NUM; ++i) { - req[i].set_message(__FUNCTION__); - cntl[i].request_attachment().append(attach); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - - cntl[0].CloseConnection("Close connection"); - - 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; - } - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 3000; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + + butil::IOBuf attach; + attach.resize(4096); + for (int i = 0; i < RPC_NUM; ++i) { + req[i].set_message(__FUNCTION__); + cntl[i].request_attachment().append(attach); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + + cntl[0].CloseConnection("Close connection"); + + 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; + } + + StopServer(); } TEST_P(RdmaRpcTest, rdma_client_close_during_rpc_repeatedly) { @@ -2908,216 +3170,218 @@ TEST_P(RdmaRpcTest, rdma_client_close_during_rpc_repeatedly) { } TEST_P(RdmaRpcTest, verbs_error_handling) { - if (!FLAGS_rdma_test_enable) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - - Controller cntl; - test::EchoRequest req; - test::EchoResponse res; - req.set_message(__FUNCTION__); - req.set_sleep_us(200000); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, done); - - 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; - void* buf = malloc(8192); - sge.addr = (uint64_t)buf; - sge.length = 8192; - sge.lkey = 1; // incorrect lkey - wr.sg_list = &sge; - wr.num_sge = 1; - ibv_send_wr* bad = nullptr; - auto rdma_transport = RdmaTransportOf(s); - ibv_post_send(rdma_transport->_rdma_ep->_resource->qp, &wr, &bad); - bthread_id_join(cntl.call_id()); - ASSERT_EQ(ERDMA, cntl.ErrorCode()); - free(buf); - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + req.set_message(__FUNCTION__); + req.set_sleep_us(200000); + 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)); + ibv_send_wr wr; + memset(&wr, 0, sizeof(wr)); + ibv_sge sge; + void *buf = malloc(8192); + sge.addr = (uint64_t)buf; + sge.length = 8192; + sge.lkey = 1; // incorrect lkey + wr.sg_list = &sge; + wr.num_sge = 1; + ibv_send_wr *bad = nullptr; + 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()); + free(buf); + + StopServer(); } TEST_P(RdmaRpcTest, rdma_use_parallel_channel) { - if (!FLAGS_rdma_test_enable) { - return; - } + if (!FLAGS_rdma_test_enable) { + return; + } - StartServer(); + StartServer(); - const size_t NCHANS = 8; - Channel subchans[NCHANS]; - ParallelChannel channel; - ChannelOptions opts; - 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.Init(nullptr)); + const size_t NCHANS = 8; + Channel subchans[NCHANS]; + ParallelChannel channel; + ChannelOptions opts; + 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.Init(nullptr)); - Controller cntl; - test::EchoRequest req; - test::EchoResponse res; - req.set_message(__FUNCTION__); - ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, nullptr); + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + req.set_message(__FUNCTION__); + ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, nullptr); - ASSERT_EQ(0, cntl.ErrorCode()); - ASSERT_EQ(NCHANS, (size_t)cntl.sub_count()); + ASSERT_EQ(0, cntl.ErrorCode()); + ASSERT_EQ(NCHANS, (size_t)cntl.sub_count()); - StopServer(); + StopServer(); } TEST_P(RdmaRpcTest, rdma_use_selective_channel) { - if (!FLAGS_rdma_test_enable) { - return; - } + if (!FLAGS_rdma_test_enable) { + return; + } - StartServer(); + StartServer(); - const size_t NCHANS = 8; - SelectiveChannel channel; - ChannelOptions opts; - opts.socket_mode = SOCKET_MODE_RDMA; - ASSERT_EQ(0, channel.Init("rr", &opts)); - for (size_t i = 0; i < NCHANS; ++i) { - Channel* subchan = new Channel; - ASSERT_EQ(0, subchan->Init(_naming_url.c_str(), "rR", &opts)); - ASSERT_EQ(0, channel.AddChannel(subchan, nullptr)); - } + const size_t NCHANS = 8; + SelectiveChannel channel; + ChannelOptions opts; + opts.socket_mode = SOCKET_MODE_RDMA; + ASSERT_EQ(0, channel.Init("rr", &opts)); + for (size_t i = 0; i < NCHANS; ++i) { + Channel *subchan = new Channel; + ASSERT_EQ(0, subchan->Init(_naming_url.c_str(), "rR", &opts)); + ASSERT_EQ(0, channel.AddChannel(subchan, nullptr)); + } - Controller cntl; - test::EchoRequest req; - test::EchoResponse res; - req.set_message(__FUNCTION__); - ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, nullptr); + Controller cntl; + test::EchoRequest req; + test::EchoResponse res; + req.set_message(__FUNCTION__); + ::test::EchoService::Stub(&channel).Echo(&cntl, &req, &res, nullptr); - ASSERT_EQ(0, cntl.ErrorCode()) << cntl.ErrorText(); - ASSERT_EQ(1, cntl.sub_count()); + ASSERT_EQ(0, cntl.ErrorCode()) << cntl.ErrorText(); + ASSERT_EQ(1, cntl.sub_count()); - StopServer(); + 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) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 500; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - - butil::IOBuf attach; - void* data = malloc(4096);; - attach.append_user_data(data, 4096, nullptr); - req[0].set_message(__FUNCTION__); - cntl[0].request_attachment().append(attach); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[0], &req[0], &res[0], done); - bthread_id_join(cntl[0].call_id()); - ASSERT_EQ(ERDMAMEM, cntl[0].ErrorCode()); - attach.clear(); - sleep(2); // wait for client recover from EHOSTDOWN - cntl[0].Reset(); - - char* mr[2 * RPC_NUM]; - uint32_t lkey[2 * RPC_NUM]; - for (size_t i = 0; i < RPC_NUM; ++i) { - mr[2 * i] = (char*)malloc(4096); - 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]); - 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]); - req[i].set_message(__FUNCTION__); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - for (size_t i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; - rdma::DeregisterMemoryForRdma(mr[i]); - ASSERT_EQ(2 * (4096 - i), cntl[i].response_attachment().size()); - char tmp[8192]; - cntl[i].response_attachment().copy_to(tmp, 2 * (4096 - i)); - ASSERT_EQ(0, memcmp(mr[2 * i] + i, tmp, 4096 - i)); - ASSERT_EQ(0, memcmp(mr[2 * i + 1] + i, tmp + 4096 - i, 4096 - i)); - free(mr[2 * i]); - free(mr[2 * i + 1]); - } - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 500; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + + butil::IOBuf attach; + void *data = malloc(4096); + ; + attach.append_user_data(data, 4096, nullptr); + req[0].set_message(__FUNCTION__); + cntl[0].request_attachment().append(attach); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[0], &req[0], &res[0], done); + bthread_id_join(cntl[0].call_id()); + ASSERT_EQ(ERDMAMEM, cntl[0].ErrorCode()); + attach.clear(); + sleep(2); // wait for client recover from EHOSTDOWN + cntl[0].Reset(); + + char *mr[2 * RPC_NUM]; + uint32_t lkey[2 * RPC_NUM]; + for (size_t i = 0; i < RPC_NUM; ++i) { + mr[2 * i] = (char *)malloc(4096); + 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]); + 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]); + req[i].set_message(__FUNCTION__); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + for (size_t i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + ASSERT_EQ(0, cntl[i].ErrorCode()) << "req[" << i << "]"; + rdma::DeregisterMemoryForRdma(mr[i]); + ASSERT_EQ(2 * (4096 - i), cntl[i].response_attachment().size()); + char tmp[8192]; + cntl[i].response_attachment().copy_to(tmp, 2 * (4096 - i)); + ASSERT_EQ(0, memcmp(mr[2 * i] + i, tmp, 4096 - i)); + ASSERT_EQ(0, memcmp(mr[2 * i + 1] + i, tmp + 4096 - i, 4096 - i)); + free(mr[2 * i]); + free(mr[2 * i + 1]); + } + + StopServer(); } TEST_P(RdmaRpcTest, try_memory_pool_empty) { - if (!FLAGS_rdma_test_enable) { - return; - } - - StartServer(); - - Channel channel; - ChannelOptions chan_options; - chan_options.socket_mode = SOCKET_MODE_RDMA; - chan_options.connect_timeout_ms = 500; - chan_options.timeout_ms = 60000; - chan_options.max_retry = 0; - ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); - Controller cntl[RPC_NUM]; - test::EchoRequest req[RPC_NUM]; - test::EchoResponse res[RPC_NUM]; - - butil::IOBuf iobuf[RPC_NUM]; - for (int i = 0; i < 1024; ++i) { - if (iobuf[i].resize(1048576 * 8)) { - // 8MB for each iobuf - break; - } - } - - for (int i = 0; i < RPC_NUM; ++i) { - req[i].set_message(__FUNCTION__); - cntl[i].request_attachment().append(iobuf[i]); - google::protobuf::Closure* done = DoNothing(); - ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); - } - for (int i = 0; i < RPC_NUM; ++i) { - bthread_id_join(cntl[i].call_id()); - } - - StopServer(); + if (!FLAGS_rdma_test_enable) { + return; + } + + StartServer(); + + Channel channel; + ChannelOptions chan_options; + chan_options.socket_mode = SOCKET_MODE_RDMA; + chan_options.connect_timeout_ms = 500; + chan_options.timeout_ms = 60000; + chan_options.max_retry = 0; + ASSERT_EQ(0, channel.Init(g_ep, &chan_options)); + Controller cntl[RPC_NUM]; + test::EchoRequest req[RPC_NUM]; + test::EchoResponse res[RPC_NUM]; + + butil::IOBuf iobuf[RPC_NUM]; + for (int i = 0; i < 1024; ++i) { + if (iobuf[i].resize(1048576 * 8)) { + // 8MB for each iobuf + break; + } + } + + for (int i = 0; i < RPC_NUM; ++i) { + req[i].set_message(__FUNCTION__); + cntl[i].request_attachment().append(iobuf[i]); + google::protobuf::Closure *done = DoNothing(); + ::test::EchoService::Stub(&channel).Echo(&cntl[i], &req[i], &res[i], done); + } + for (int i = 0; i < RPC_NUM; ++i) { + bthread_id_join(cntl[i].call_id()); + } + + StopServer(); } // Run every TEST_P(RdmaRpcTest, ...) above twice: once with the @@ -3126,27 +3390,25 @@ 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); - }); - -#endif // if BRPC_WITH_RDMA - -int main(int argc, char* argv[]) { - testing::InitGoogleTest(&argc, argv); - GFLAGS_NAMESPACE::ParseCommandLineFlags(&argc, &argv, true); +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 + +int main(int argc, char *argv[]) { + testing::InitGoogleTest(&argc, argv); + GFLAGS_NAMESPACE::ParseCommandLineFlags(&argc, &argv, true); #if BRPC_WITH_RDMA - rdma::FLAGS_rdma_trace_verbose = true; - rdma::FLAGS_rdma_memory_pool_max_regions = 2; - FLAGS_log_idle_connection_close = true; - if (!FLAGS_rdma_test_enable) { - // skip UT requiring rdma runtime environment - rdma::g_rdma_available.store(true, butil::memory_order_relaxed); - rdma::g_skip_rdma_init = true; - } -#endif // if BRPC_WITH_RDMA - return RUN_ALL_TESTS(); + rdma::FLAGS_rdma_trace_verbose = true; + rdma::FLAGS_rdma_memory_pool_max_regions = 2; + FLAGS_log_idle_connection_close = true; + if (!FLAGS_rdma_test_enable) { + // skip UT requiring rdma runtime environment + rdma::g_rdma_available.store(true, butil::memory_order_relaxed); + rdma::g_skip_rdma_init = true; + } +#endif // if BRPC_WITH_RDMA + return RUN_ALL_TESTS(); } diff --git a/test/brpc_transport_handshake_unittest.cpp b/test/brpc_transport_handshake_unittest.cpp new file mode 100644 index 0000000000..2c084c3990 --- /dev/null +++ b/test/brpc_transport_handshake_unittest.cpp @@ -0,0 +1,463 @@ +// 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 "butil/fd_guard.h" +#include "butil/sys_byteorder.h" +#include "brpc/adapter_transport.h" +#include "brpc/policy/transport_handshake_protocol.h" +#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; +}; + +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); +} + +static HandshakeCodec MakeTestCodec(std::string* calls = NULL) { + HandshakeCodec codec{}; + codec.protocol_version = 7; + codec.hello_frame = FixedSpec("HS", 2, 4); + codec.ack_frame = FixedSpec(NULL, 0, 1); + codec.build_hello = [calls](bool enabled, std::string* payload) { + if (calls) *calls += "build "; + *payload = enabled ? "LO" : "NO"; + return STEP_OK; + }; + codec.parse_hello = [calls](const std::string& payload) { + if (calls) *calls += "parse "; + return payload == "OK" ? STEP_OK : STEP_FALLBACK; + }; + codec.build_ack = [calls](bool enabled, std::string* payload) { + if (calls) *calls += enabled ? "ack1 " : "ack0 "; + *payload = enabled ? "1" : "0"; + return STEP_OK; + }; + codec.parse_ack = [calls](const std::string& payload, bool* enabled) { + if (calls) *calls += "parse_ack "; + *enabled = payload == "1"; + return payload == "1" || payload == "0" ? STEP_OK : STEP_ERROR; + }; + return codec; +} + +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(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()); +} + +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; + + ClientHandshakeCallbacks callbacks{}; + callbacks.codec = MakeTestCodec(&calls); + callbacks.transport.prepare_resources = [&]() { + calls += "prepare "; + return STEP_OK; + }; + callbacks.transport.negotiate_resources = [&]() { + calls += "negotiate "; + return STEP_OK; + }; + callbacks.transport.set_high_speed_active = [&]() { calls += "activate"; }; + callbacks.transport.set_tcp_active = []() {}; + callbacks.transport.on_failed = []() {}; + + ASSERT_EQ(STEP_OK, session.RunClient(callbacks)); + 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_resource_failure_falls_back_before_io) { + MemoryHandshakeIO io; + HandshakeSession session; + session.SetIOForTest(&io); + bool tcp_active = false; + + ClientHandshakeCallbacks callbacks{}; + callbacks.codec = MakeTestCodec(); + callbacks.transport.prepare_resources = []() { return STEP_FALLBACK; }; + callbacks.transport.negotiate_resources = []() { return STEP_ERROR; }; + callbacks.transport.set_high_speed_active = []() {}; + callbacks.transport.set_tcp_active = [&]() { tcp_active = true; }; + callbacks.transport.on_failed = []() {}; + + ASSERT_EQ(STEP_FALLBACK, session.RunClient(callbacks)); + ASSERT_TRUE(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; + + ServerHandshakeCallbacks callbacks{}; + callbacks.fallback_on_not_mine = false; + callbacks.codecs.push_back(MakeTestCodec(&calls)); + callbacks.input = &input; + callbacks.transport.prepare_resources = [&]() { + calls += "prepare "; + return STEP_OK; + }; + callbacks.transport.negotiate_resources = [&]() { + calls += "negotiate "; + return STEP_OK; + }; + callbacks.validate_established = [&]() { + calls += "validate "; + return STEP_OK; + }; + callbacks.transport.set_high_speed_active = [&]() { calls += "activate"; }; + callbacks.transport.set_tcp_active = []() {}; + callbacks.transport.on_failed = []() {}; + + ASSERT_EQ(STEP_NEED_MORE, session.RunServer(callbacks)); + 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(callbacks)); + 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); + bool tcp_active = false; + + ServerHandshakeCallbacks callbacks{}; + callbacks.fallback_on_not_mine = false; + callbacks.codecs.push_back(MakeTestCodec()); + callbacks.input = &input; + callbacks.transport.prepare_resources = []() { return STEP_FALLBACK; }; + callbacks.transport.negotiate_resources = []() { return STEP_ERROR; }; + callbacks.transport.set_high_speed_active = []() {}; + callbacks.transport.set_tcp_active = [&]() { tcp_active = true; }; + callbacks.transport.on_failed = []() {}; + + ASSERT_EQ(STEP_FALLBACK, session.RunServer(callbacks)); + ASSERT_EQ("HSNO", io.output()); + ASSERT_TRUE(source.empty()); + ASSERT_TRUE(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); + bool tcp_active = false; + + ServerHandshakeCallbacks callbacks{}; + callbacks.fallback_on_not_mine = true; + callbacks.codecs.push_back(MakeTestCodec()); + callbacks.input = &input; + callbacks.transport.prepare_resources = []() { return STEP_ERROR; }; + callbacks.transport.negotiate_resources = []() { return STEP_ERROR; }; + callbacks.transport.set_high_speed_active = []() {}; + callbacks.transport.set_tcp_active = [&]() { tcp_active = true; }; + callbacks.transport.on_failed = []() {}; + + ASSERT_EQ(STEP_FALLBACK, session.RunServer(callbacks)); + ASSERT_TRUE(tcp_active); + ASSERT_EQ(4UL, source.size()); + ASSERT_EQ(FALLBACK_TCP, session.phase()); + + ASSERT_EQ(STEP_NOT_MINE, session.RunServer(callbacks)); + 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); + + ServerHandshakeCallbacks callbacks{}; + callbacks.fallback_on_not_mine = false; + callbacks.codecs.push_back(MakeTestCodec()); + callbacks.input = &input; + callbacks.transport.prepare_resources = []() { return STEP_OK; }; + callbacks.transport.negotiate_resources = []() { return STEP_OK; }; + callbacks.transport.set_high_speed_active = []() {}; + callbacks.transport.set_tcp_active = []() {}; + callbacks.transport.on_failed = []() {}; + + source.append("H", 1); + ASSERT_EQ(STEP_NEED_MORE, session.RunServer(callbacks)); + ASSERT_EQ(UNINITIALIZED, session.phase()); + + source.append("S", 1); + ASSERT_EQ(STEP_NEED_MORE, session.RunServer(callbacks)); + ASSERT_EQ(HELLO_WAIT, session.phase()); + ASSERT_EQ(7, session.protocol_version()); + ASSERT_EQ(2UL, source.size()); +} + +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_consumes_coalesced_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 ParseResult 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(); +} + +} // namespace handshake +} // namespace brpc diff --git a/test/brpc_ubring_unittest.cpp b/test/brpc_ubring_unittest.cpp index d3d0a7ce98..8f1c545dfa 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" @@ -145,6 +146,45 @@ TEST_F(HelloMessageTest, toString_contains_fields) { EXPECT_NE(std::string::npos, s.find("UBRING_test")); } +TEST(UBShmHandshakeAdapterTest, codec_preserves_v2_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(2, 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; + const brpc::handshake::HandshakeCodec codec = adapter.MakeCodec(); + ASSERT_EQ(brpc::handshake::FRAME_OK, + brpc::handshake::FrameCodec::Encode( + codec.hello_frame, payload, &frame)); + ASSERT_EQ(64, frame.size()); + EXPECT_EQ("UB", frame.substr(0, 2)); +} + +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;