diff --git a/src/brpc/adapter_transport.cpp b/src/brpc/adapter_transport.cpp index 777e2282bf..ffd8e46521 100644 --- a/src/brpc/adapter_transport.cpp +++ b/src/brpc/adapter_transport.cpp @@ -17,7 +17,9 @@ #include "brpc/adapter_transport.h" +#include #include +#include #include #include @@ -40,6 +42,12 @@ namespace brpc { namespace { +bool MatchesMagicPrefix(const char *prefix, size_t prefix_len, + const char *magic, size_t magic_len) { + const size_t compare_len = std::min(prefix_len, magic_len); + return memcmp(prefix, magic, compare_len) == 0; +} + class AdapterConnect : public AppConnect { public: explicit AdapterConnect(const std::shared_ptr& app_connect) @@ -161,13 +169,26 @@ ParseResult AdapterTransport::ProcessUpgradeReadable(butil::IOBuf* source) { _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); + } else if (!source->empty()) { + static const size_t MAX_MAGIC_LEN = 4; + char prefix[MAX_MAGIC_LEN] = {}; + const size_t prefix_len = std::min(source->size(), MAX_MAGIC_LEN); + source->copy_to(prefix, prefix_len); + + const bool matches_ub = + MatchesMagicPrefix(prefix, prefix_len, "UB", 2); + const bool matches_rdma = + MatchesMagicPrefix(prefix, prefix_len, "RDMA", 4) || + MatchesMagicPrefix(prefix, prefix_len, "RDM3", 4); + if (!matches_ub && !matches_rdma) { + result = ParseResult(PARSE_ERROR_TRY_OTHERS); + } else { + handshake::HandshakeAdapter* adapter = + matches_ub + ? handshake::GetUBShmServerHandshakeAdapter() + : handshake::GetRdmaServerHandshakeAdapter(); + result = adapter->ExecuteServerHandshake(source, _socket); + } } const int phase = _handshake.phase(); if (!connection_completed() && @@ -560,6 +581,9 @@ void AdapterTransport::CheckUnexpectedTcpData() { _socket->description().c_str()); return; } + if (errno == EINTR) { + continue; + } if (errno != EAGAIN) { const int saved_errno = errno; _socket->SetFailed(saved_errno, "Fail to read from %s: %s", diff --git a/src/brpc/handshake/handshake_io.cpp b/src/brpc/handshake/handshake_io.cpp index ca5e7c17c5..67cc95e37c 100644 --- a/src/brpc/handshake/handshake_io.cpp +++ b/src/brpc/handshake/handshake_io.cpp @@ -63,6 +63,7 @@ void SocketHandshakeIO::NotifyReadable() { template static int ReadExactLoop(butil::atomic* read_butex, + SocketId socket_id, size_t len, ReadOnce read_once) { size_t received = 0; while (received < len) { @@ -76,6 +77,11 @@ static int ReadExactLoop(butil::atomic* read_butex, if (errno != EAGAIN) { return -1; } + SocketUniquePtr alive; + if (Socket::Address(socket_id, &alive) != 0) { + errno = EFAILEDSOCKET; + return -1; + } if (bthread::butex_wait(read_butex, expected, &duetime) < 0 && errno != EWOULDBLOCK && errno != ETIMEDOUT) { return -1; @@ -94,7 +100,7 @@ int SocketHandshakeIO::ReadExact(void* data, size_t len) { CHECK(data != NULL); CHECK(_socket != NULL); const int fd = _socket->fd(); - return ReadExactLoop(_read_butex, len, + return ReadExactLoop(_read_butex, _socket->id(), len, [data, fd](size_t offset, size_t remaining) { return read(fd, static_cast(data) + offset, remaining); }); @@ -116,7 +122,7 @@ static int WriteAllLoop(size_t len, WriteOnce write_once, return -1; } if (errno == EINTR) { - continue; + continue; } if (errno != EAGAIN) { return -1; diff --git a/src/brpc/rdma/rdma_endpoint.cpp b/src/brpc/rdma/rdma_endpoint.cpp index 0f3a4ee6c4..1b4b802648 100644 --- a/src/brpc/rdma/rdma_endpoint.cpp +++ b/src/brpc/rdma/rdma_endpoint.cpp @@ -32,7 +32,6 @@ #include "butil/sys_byteorder.h" // HostToNet,NetToHost #include - DECLARE_int32(task_group_ntags); namespace brpc { @@ -109,7 +108,7 @@ RdmaResource::~RdmaResource() { } } -RdmaEndpoint::RdmaEndpoint(Socket* s) +RdmaEndpoint::RdmaEndpoint(Socket *s) : _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), @@ -133,70 +132,74 @@ RdmaEndpoint::RdmaEndpoint(Socket* s) _input_processor.Init(s, InputMessengerProcessor::STREAM_RDMA_QP); } -RdmaEndpoint::~RdmaEndpoint() { Reset(); } +RdmaEndpoint::~RdmaEndpoint() { + Reset(); +} void RdmaEndpoint::Reset() { - 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); + 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); } 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); + _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); } 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; - } + 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; + } } int RdmaEndpoint::QueryLocalEce(ibv_ece *ece) const { - if (ece == NULL || IbvQueryEce == NULL || _resource == NULL || - _resource->qp == NULL) { - return 1; - } - return IbvQueryEce(_resource->qp, ece) == 0 ? 0 : -1; + if (ece == NULL || IbvQueryEce == NULL || _resource == NULL || + _resource->qp == NULL) { + return 1; + } + return IbvQueryEce(_resource->qp, ece) == 0 ? 0 : -1; } -void RdmaEndpoint::SetOutgoingEce(const ibv_ece &ece) { _outgoing_ece = ece; } +void RdmaEndpoint::SetOutgoingEce(const ibv_ece &ece) { + _outgoing_ece = ece; +} bool RdmaEndpoint::IsWritable() const { if (BAIDU_UNLIKELY(g_skip_rdma_init)) { @@ -235,8 +238,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) { @@ -280,7 +283,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]; @@ -334,8 +337,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; @@ -379,8 +382,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; } @@ -397,16 +400,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)); } @@ -432,8 +435,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; } @@ -480,10 +483,10 @@ ssize_t RdmaEndpoint::HandleCompletion(ibv_wc& wc) { zerocopy = false; } if (zerocopy) { - _rbuf[_rq_received].cutn(&_input_processor.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 - _input_processor.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) { @@ -543,8 +546,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 @@ -599,36 +602,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"; @@ -673,12 +676,12 @@ int RdmaEndpoint::DoAllocateResources() { g_rdma_resource_list = g_rdma_resource_list->next; } } - if (!_resource) { + if (!_resource) { _resource = AllocateQpCq(_sq_size, _rq_size); } else { _resource->next = nullptr; } - if (!_resource) { + if (!_resource) { return -1; } @@ -708,10 +711,10 @@ 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; } @@ -723,8 +726,8 @@ int RdmaEndpoint::StartCqEvents() { if (BAIDU_UNLIKELY(g_skip_rdma_init)) { 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; } @@ -758,9 +761,9 @@ int RdmaEndpoint::BringUpQp(const RdmaConnectionInfo &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; @@ -772,7 +775,7 @@ int RdmaEndpoint::BringUpQp(const RdmaConnectionInfo &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 - // return it in the server negotiation response. + // return it in the server negotiation response. // Client: `remote->ece' is the server's reduced ECE; // just set it here. bool use_ece = true; @@ -808,11 +811,11 @@ int RdmaEndpoint::BringUpQp(const RdmaConnectionInfo &remote, bool is_server) { attr.rq_psn = 0; attr.max_dest_rd_atomic = 0; attr.min_rnr_timer = 0; // We do not allow rnr error - err = IbvModifyQp(_resource->qp, &attr, - (ibv_qp_attr_mask)(IBV_QP_STATE | IBV_QP_PATH_MTU | - IBV_QP_MIN_RNR_TIMER | IBV_QP_AV | - IBV_QP_MAX_DEST_RD_ATOMIC | - IBV_QP_DEST_QPN | IBV_QP_RQ_PSN)); + err = IbvModifyQp(_resource->qp, &attr, + (ibv_qp_attr_mask)(IBV_QP_STATE | IBV_QP_PATH_MTU | + IBV_QP_MIN_RNR_TIMER | IBV_QP_AV | + IBV_QP_MAX_DEST_RD_ATOMIC | + IBV_QP_DEST_QPN | IBV_QP_RQ_PSN)); if (err != 0) { LOG(WARNING) << "Fail to modify QP from INIT to RTR: " << berror(err); return -1; @@ -824,29 +827,29 @@ int RdmaEndpoint::BringUpQp(const RdmaConnectionInfo &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); } } @@ -878,7 +881,7 @@ static int DrainCq(ibv_cq* cq) { } void RdmaEndpoint::DeallocateResources() { - if (!_resource) { + if (!_resource) { return; } if (FLAGS_rdma_use_polling) { @@ -905,7 +908,7 @@ void RdmaEndpoint::DeallocateResources() { bool remove_consumer = true; _reclaim: if (!move_to_rdma_resource_list) { - if (nullptr != _resource->qp) { + if (nullptr != _resource->qp) { int err = IbvDestroyQp(_resource->qp); LOG_IF(WARNING, 0 != err) << "Fail to destroy QP: " << berror(err); _resource->qp = nullptr; @@ -915,16 +918,19 @@ void RdmaEndpoint::DeallocateResources() { DeallocateCq(_resource->send_cq); DeallocateCq(_resource->recv_cq); - if (nullptr != _resource->comp_channel) { + if (_resource->comp_channel != nullptr) { + if (_cq_sid != INVALID_SOCKET_ID) { // 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; @@ -1005,8 +1011,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); @@ -1025,15 +1031,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; @@ -1058,7 +1064,7 @@ void RdmaEndpoint::PollCq(Socket* m) { if (m->id() != ep->_cq_sid) { return; } - RdmaTransport *rdma_transport = RdmaTransport::Get(s.get()); + RdmaTransport *rdma_transport = RdmaTransport::Get(s.get()); CHECK(ep == rdma_transport->_rdma_ep); bool send = false; @@ -1177,8 +1183,8 @@ 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; } } @@ -1186,21 +1192,21 @@ void RdmaEndpoint::PollCq(Socket* m) { 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; + 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() { @@ -1214,8 +1220,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; } @@ -1310,8 +1316,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; @@ -1340,23 +1346,23 @@ void RdmaEndpoint::PollingModeRelease(bthread_tag_t tag) { } void RdmaEndpoint::PollerAddCqSid() { - 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}); - } + 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() { - 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}); - } + 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_transport.cpp b/src/brpc/rdma_transport.cpp index 81ebd26b85..c7933eeb3c 100644 --- a/src/brpc/rdma_transport.cpp +++ b/src/brpc/rdma_transport.cpp @@ -30,25 +30,25 @@ 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); + 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); _socket = socket; _default_connect = options.app_connect; - _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)); + _on_edge_trigger = nullptr; + _rdma_state = RDMA_UNKNOWN; + _rdma_ep = new (std::nothrow) rdma::RdmaEndpoint(socket); + if (!_rdma_ep) { + const int saved_errno = errno != 0 ? errno : ENOMEM; + errno = saved_errno; + PLOG(WARNING) << "Fail to create RdmaEndpoint, disable RDMA upgrade"; + _rdma_state = RDMA_OFF; } - _rdma_state = RDMA_UNKNOWN; } void RdmaTransport::Release() { @@ -68,54 +68,56 @@ int RdmaTransport::Reset(int32_t expected_nref) { } std::shared_ptr RdmaTransport::Connect() { - return _default_connect; + return _default_connect; } void RdmaTransport::SetHighSpeedAvailable(bool available) { - _rdma_state = available ? RDMA_ON : RDMA_OFF; + _rdma_state = available ? RDMA_ON : RDMA_OFF; } int RdmaTransport::PrepareUpgradeResources() { - return _rdma_ep->AllocateResources(); + return _rdma_ep->AllocateResources(); } int RdmaTransport::NegotiateUpgradeResources( const rdma::RdmaConnectionInfo &remote, bool server) { - _rdma_ep->ApplyRemoteInfo(remote); - return _rdma_ep->BringUpQp(remote, server); + _rdma_ep->ApplyRemoteInfo(remote); + return _rdma_ep->BringUpQp(remote, server); } int RdmaTransport::StartUpgradeEvents() { - return _rdma_ep->StartCqEvents(); + return _rdma_ep->StartCqEvents(); } std::unique_ptr RdmaTransport::CreateClientHandshakeAdapter() { - return rdma::CreateClientHandshakeAdapter(_rdma_ep); + return rdma::CreateClientHandshakeAdapter(_rdma_ep); } std::vector> RdmaTransport::CreateServerHandshakeAdapters() { - return rdma::CreateServerHandshakeAdapters(_rdma_ep); + return rdma::CreateServerHandshakeAdapters(_rdma_ep); } -void RdmaTransport::ActivateUpgrade() { SetHighSpeedAvailable(true); } +void RdmaTransport::ActivateUpgrade() { + SetHighSpeedAvailable(true); +} void RdmaTransport::DeactivateUpgrade() { - SetHighSpeedAvailable(false); - if (_rdma_ep != nullptr) { - _rdma_ep->Reset(); - } + SetHighSpeedAvailable(false); + if (_rdma_ep != nullptr) { + _rdma_ep->Reset(); + } } int RdmaTransport::CutFromIOBuf(butil::IOBuf *buf) { - butil::IOBuf *data[1] = {buf}; - return static_cast(CutFromIOBufList(data, 1)); + butil::IOBuf *data[1] = {buf}; + return static_cast(CutFromIOBufList(data, 1)); } ssize_t RdmaTransport::CutFromIOBufList(butil::IOBuf **buf, size_t ndata) { - CHECK(_rdma_ep != nullptr); - return _rdma_ep->CutFromIOBufList(buf, ndata); + CHECK(_rdma_ep != nullptr); + return _rdma_ep->CutFromIOBufList(buf, ndata); } int RdmaTransport::WaitEpollOut(butil::atomic *_epollout_butex, @@ -129,7 +131,7 @@ int RdmaTransport::WaitEpollOut(butil::atomic *_epollout_butex, if (errno != EAGAIN && errno != ETIMEDOUT) { const int saved_errno = errno; PLOG(WARNING) << "Fail to wait rdma window of " << _socket; - _socket->SetFailed(saved_errno, "Fail to wait rdma window of %s: %s", + _socket->SetFailed(saved_errno, "Fail to wait rdma window of %s: %s", _socket->description().c_str(), berror(saved_errno)); } @@ -181,16 +183,16 @@ void RdmaTransport::QueueMessage(InputMessageClosure& input_msg, // TODO(gejun): Join threads. bthread_t th; - bthread_attr_t tmp = - (FLAGS_usercode_in_pthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL) | - BTHREAD_NOSIGNAL; + bthread_attr_t tmp = + (FLAGS_usercode_in_pthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL) | + BTHREAD_NOSIGNAL; tmp.keytable_pool = _socket->keytable_pool(); tmp.tag = bthread_self_tag(); bthread_attr_set_name(&tmp, "ProcessInputMessage"); - if (!FLAGS_usercode_in_coroutine && - bthread_start_background(&th, &tmp, ProcessInputMessage, to_run_msg) == - 0) { + if (!FLAGS_usercode_in_coroutine && + bthread_start_background(&th, &tmp, ProcessInputMessage, to_run_msg) == + 0) { ++*num_bthread_created; } else { ProcessInputMessage(to_run_msg); @@ -205,18 +207,18 @@ void RdmaTransport::Debug(std::ostream &os) { int RdmaTransport::ContextInitOrDie(bool serverOrNot, const void* _options) { if (serverOrNot) { - if (!OptionsAvailableOverRdma( - static_cast(_options))) { + if (!OptionsAvailableOverRdma( + static_cast(_options))) { return -1; } rdma::GlobalRdmaInitializeOrDie(); - if (!rdma::InitPollingModeWithTag( - static_cast(_options)->bthread_tag)) { + if (!rdma::InitPollingModeWithTag( + static_cast(_options)->bthread_tag)) { return -1; } } else { - if (!OptionsAvailableForRdma( - static_cast(_options))) { + if (!OptionsAvailableForRdma( + static_cast(_options))) { return -1; } rdma::GlobalRdmaInitializeOrDie(); @@ -235,7 +237,7 @@ bool RdmaTransport::OptionsAvailableForRdma(const ChannelOptions* opt) { return false; } if (!rdma::SupportedByRdma(opt->protocol.name())) { - LOG(WARNING) << "Cannot use " << opt->protocol.name() << " over RDMA"; + LOG(WARNING) << "Cannot use " << opt->protocol.name() << " over RDMA"; return false; } return true; diff --git a/src/brpc/ubshm/ub_endpoint.cpp b/src/brpc/ubshm/ub_endpoint.cpp index 5fa9ea5f2c..131fbfcadf 100644 --- a/src/brpc/ubshm/ub_endpoint.cpp +++ b/src/brpc/ubshm/ub_endpoint.cpp @@ -33,7 +33,6 @@ #include "butil/logging.h" // CHECK, LOG #include - DECLARE_int32(task_group_ntags); namespace brpc { @@ -42,9 +41,9 @@ namespace ubring { extern bool g_skip_ub_init; DEFINE_int32(ub_poller_num, 1, "Poller number in ub polling mode."); -DEFINE_bool(ub_poller_yield, false, "Yield thread in RDMA polling mode."); +DEFINE_bool(ub_poller_yield, false, "Yield thread in UBRing polling mode."); DEFINE_bool(ub_edisp_unsched, false, "Disable event dispatcher schedule"); -DEFINE_bool(ub_disable_bthread, false, "Disable bthread in RDMA"); +DEFINE_bool(ub_disable_bthread, false, "Disable bthread in UBRing polling mode."); static const size_t MIN_ONCE_READ = 4096; static const size_t MAX_ONCE_READ = 524288; @@ -52,11 +51,14 @@ static const size_t IOBUF_IOV_MAX = 256; static butil::Mutex *g_ubring_resource_mutex = NULL; -UBShmEndpoint::UBShmEndpoint(Socket* s) +UBShmEndpoint::UBShmEndpoint(Socket *s) : _socket(s), _socket_id(s ? s->id() : INVALID_SOCKET_ID), - _ub_ring(nullptr), _poller_sid(INVALID_SOCKET_ID) {} + _ub_ring(nullptr), _poller_sid(INVALID_SOCKET_ID) { +} -UBShmEndpoint::~UBShmEndpoint() { Reset(); } +UBShmEndpoint::~UBShmEndpoint() { + Reset(); +} void UBShmEndpoint::Reset() { DeallocateResources(); @@ -104,11 +106,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)) { @@ -143,22 +145,22 @@ int UBShmEndpoint::AllocateClientResources(ubring::SHM *local_trx_shm, options.user = this; options.keytable_pool = _socket->_keytable_pool; if (Socket::Create(options, &_poller_sid) < 0) { - const int saved_errno = errno; + 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; + 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; + 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); @@ -180,22 +182,22 @@ int UBShmEndpoint::AllocateServerResources(ubring::SHM *remote_trx_shm, options.user = this; options.keytable_pool = _socket->_keytable_pool; if (Socket::Create(options, &_poller_sid) < 0) { - const int saved_errno = errno; + 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; + 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; + 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); @@ -223,7 +225,7 @@ void UBShmEndpoint::PollIn(UBShmEndpoint* ep, uint32_t ep_event) { if (Socket::Address(ep->_socket_id, &s) < 0) { return; } - UBShmTransport *ub_transport = UBShmTransport::Get(s.get()); + UBShmTransport *ub_transport = UBShmTransport::Get(s.get()); CHECK(ep == ub_transport->_ub_ep); InputMessageClosure last_msg; @@ -244,10 +246,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) { @@ -280,7 +282,7 @@ void UBShmEndpoint::PollOut(UBShmEndpoint* ep, uint32_t ep_event) { if (Socket::Address(ep->_socket_id, &s) < 0) { return; } - UBShmTransport *ub_transport = UBShmTransport::Get(s.get()); + UBShmTransport *ub_transport = UBShmTransport::Get(s.get()); CHECK(ep == ub_transport->_ub_ep); if (ep->IsWritable()) { s->WakeAsEpollOut(); @@ -320,7 +322,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 cq_sids; + std::unordered_set cq_sids; PollerSidOp op; if (poller->init_fn) { @@ -329,18 +331,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) { - cq_sids.emplace(op); + cq_sids.emplace(op); } else if (op.type == PollerSidOp::REMOVE) { - cq_sids.erase(op); + cq_sids.erase(op); } else if (op.type == PollerSidOp::MOD) { - cq_sids.erase(op); - cq_sids.emplace(op); + cq_sids.erase(op); + cq_sids.emplace(op); } } - for (auto cq : cq_sids) { + for (auto cq : cq_sids) { SocketUniquePtr s; - if (Socket::Address(cq.sid, &s) < 0) { + if (Socket::Address(cq.sid, &s) < 0) { continue; } UBShmEndpoint* ep = static_cast(s->user()); @@ -348,12 +350,12 @@ int UBShmEndpoint::PollingModeInitialize(bthread_tag_t tag, continue; } - if (cq.event & EPOLLIN) { - PollIn(ep, cq.event); + if (cq.event & EPOLLIN) { + PollIn(ep, cq.event); } - if (cq.event & EPOLLOUT) { - PollOut(ep, cq.event); + if (cq.event & EPOLLOUT) { + PollOut(ep, cq.event); } } if (poller->callback) { @@ -372,8 +374,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; diff --git a/src/brpc/ubshm_transport.cpp b/src/brpc/ubshm_transport.cpp index bcc0fec7a4..9b0a0203e0 100644 --- a/src/brpc/ubshm_transport.cpp +++ b/src/brpc/ubshm_transport.cpp @@ -27,7 +27,6 @@ #include "brpc/ubshm/ubr_trx.h" #include "brpc/ubshm_transport.h" - namespace brpc { DECLARE_bool(usercode_in_coroutine); DECLARE_bool(usercode_in_pthread); @@ -35,26 +34,25 @@ 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); + 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); _socket = socket; _default_connect = options.app_connect; - _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)); + _on_edge_trigger = nullptr; + _ub_state = UB_UNKNOWN; + _ub_ep = new (std::nothrow) ubring::UBShmEndpoint(socket); + if (!_ub_ep) { + const int saved_errno = errno != 0 ? errno : ENOMEM; + errno = saved_errno; + PLOG(WARNING) << "Fail to create UBShmEndpoint, disable UBSHM upgrade"; + _ub_state = UB_OFF; } - _ub_state = UB_UNKNOWN; } void UBShmTransport::Release() { @@ -78,48 +76,50 @@ std::shared_ptr UBShmTransport::Connect() { } void UBShmTransport::SetHighSpeedAvailable(bool available) { - _ub_state = available ? UB_ON : UB_OFF; - } + _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); + 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); + 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); + return _ub_ep->AllocateServerResources(remote_trx_shm, local_trx_shm); } -void UBShmTransport::ActivateUpgrade() { SetHighSpeedAvailable(true); } +void UBShmTransport::ActivateUpgrade() { + SetHighSpeedAvailable(true); +} void UBShmTransport::DeactivateUpgrade() { - SetHighSpeedAvailable(false); - if (_ub_ep != nullptr) { - _ub_ep->Reset(); - } + SetHighSpeedAvailable(false); + if (_ub_ep != nullptr) { + _ub_ep->Reset(); + } } void UBShmTransport::FinishUpgrade() { - if (_ub_ep != NULL && _ub_ep->_ub_ring != NULL) { - _ub_ep->_ub_ring->UbrUnlinkLocalShm(); - } + 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)); + butil::IOBuf *data[1] = {buf}; + return static_cast(CutFromIOBufList(data, 1)); } ssize_t UBShmTransport::CutFromIOBufList(butil::IOBuf **buf, size_t ndata) { - CHECK(_ub_ep != NULL); - return _ub_ep->CutFromIOBufList(buf, ndata); - } + CHECK(_ub_ep != NULL); + return _ub_ep->CutFromIOBufList(buf, ndata); +} int UBShmTransport::WaitEpollOut(butil::atomic *_epollout_butex, bool pollin, const timespec duetime) { @@ -129,13 +129,13 @@ int UBShmTransport::WaitEpollOut(butil::atomic *_epollout_butex, if (!_ub_ep->IsWritable()) { g_vars->nwaitepollout << 1; _ub_ep->PollerRegisterEpollOut(pollin); - const int wait_rc = - bthread::butex_wait(_epollout_butex, expected_val, &duetime); + const int wait_rc = + bthread::butex_wait(_epollout_butex, expected_val, &duetime); if (wait_rc < 0) { if (errno != EAGAIN && errno != ETIMEDOUT) { const int saved_errno = errno; PLOG(WARNING) << "Fail to wait ub window of " << _socket; - _socket->SetFailed(saved_errno, "Fail to wait ub window of %s: %s", + _socket->SetFailed(saved_errno, "Fail to wait ub window of %s: %s", _socket->description().c_str(), berror(saved_errno)); } @@ -188,16 +188,16 @@ void UBShmTransport::QueueMessage(InputMessageClosure& input_msg, // TODO(gejun): Join threads. bthread_t th; - bthread_attr_t tmp = - (FLAGS_usercode_in_pthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL) | - BTHREAD_NOSIGNAL; + bthread_attr_t tmp = + (FLAGS_usercode_in_pthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL) | + BTHREAD_NOSIGNAL; tmp.keytable_pool = _socket->keytable_pool(); tmp.tag = bthread_self_tag(); bthread_attr_set_name(&tmp, "ProcessInputMessage"); - if (!FLAGS_usercode_in_coroutine && - bthread_start_background(&th, &tmp, ProcessInputMessage, to_run_msg) == - 0) { + if (!FLAGS_usercode_in_coroutine && + bthread_start_background(&th, &tmp, ProcessInputMessage, to_run_msg) == + 0) { ++*num_bthread_created; } else { ProcessInputMessage(to_run_msg); @@ -212,8 +212,8 @@ int UBShmTransport::ContextInitOrDie(bool serverOrNot, const void* _options) { return -1; } ubring::GlobalUBInitializeOrDie(); - if (!ubring::InitPollingModeWithTag( - static_cast(_options)->bthread_tag)) { + if (!ubring::InitPollingModeWithTag( + static_cast(_options)->bthread_tag)) { return -1; } } else { @@ -236,7 +236,7 @@ bool UBShmTransport::OptionsAvailableForUB(const ChannelOptions* opt) { return false; } if (!ubring::SupportedByUB(opt->protocol.name())) { - LOG(WARNING) << "Cannot use " << opt->protocol.name() << " over UB"; + LOG(WARNING) << "Cannot use " << opt->protocol.name() << " over UB"; return false; } return true; diff --git a/test/brpc_transport_handshake_unittest.cpp b/test/brpc_transport_handshake_unittest.cpp index 032befe218..97dbb45c57 100644 --- a/test/brpc_transport_handshake_unittest.cpp +++ b/test/brpc_transport_handshake_unittest.cpp @@ -23,9 +23,11 @@ #include #include +#include "bthread/bthread.h" #include "butil/fd_guard.h" #include "butil/sys_byteorder.h" #include "brpc/adapter_transport.h" +#include "brpc/handshake/handshake_io.h" #include "brpc/policy/transport_handshake_protocol.h" #include "brpc/socket.h" #include "brpc/transport_handshake.h" @@ -68,6 +70,21 @@ class MemoryHandshakeIO : public HandshakeIO { std::string _pushed_back; }; +struct BlockingHandshakeRead { + SocketHandshakeIO* io; + int result; + int error; +}; + +static void* RunBlockingHandshakeRead(void* arg) { + BlockingHandshakeRead* read = + static_cast(arg); + char byte = 0; + read->result = read->io->ReadExact(&byte, sizeof(byte)); + read->error = errno; + return NULL; +} + static FrameSpec FixedSpec(const char* magic, size_t magic_len, size_t total_len) { return FrameSpec(magic, magic_len, total_len, total_len, @@ -113,6 +130,35 @@ static std::string MakeUBShmHello() { return frame; } +TEST(TransportHandshakeTest, blocked_read_stops_after_socket_failure) { + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard peer_fd(fds[1]); + + SocketOptions options; + options.fd = fds[0]; + SocketId id; + ASSERT_EQ(0, Socket::Create(options, &id)); + + SocketUniquePtr socket; + ASSERT_EQ(0, Socket::Address(id, &socket)); + + SocketHandshakeIO io(socket.get()); + BlockingHandshakeRead read = {&io, 0, 0}; + bthread_t tid; + bthread_attr_t attr = BTHREAD_ATTR_NORMAL; + ASSERT_EQ(0, bthread_start_background( + &tid, &attr, RunBlockingHandshakeRead, &read)); + + bthread_usleep(10000); + ASSERT_EQ(0, socket->SetFailed( + EFAILEDSOCKET, "cancel blocked handshake read")); + ASSERT_EQ(0, bthread_join(tid, NULL)); + + EXPECT_EQ(-1, read.result); + EXPECT_EQ(EFAILEDSOCKET, read.error); +} + TEST(HandshakeFrameTest, supports_two_byte_fixed_magic) { const FrameSpec spec = FixedSpec("UB", 2, 6); std::string frame; @@ -403,6 +449,72 @@ TEST(TransportHandshakeTest, server_enters_hello_phase_after_magic_matches) { ASSERT_EQ(2UL, source.size()); } +TEST(TransportHandshakeTest, + impossible_magic_prefixes_try_other_protocols_without_consuming) { + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard peer_fd(fds[1]); + SocketOptions options; + options.fd = fds[0]; + SocketId id; + ASSERT_EQ(0, Socket::Create(options, &id)); + SocketUniquePtr socket; + ASSERT_EQ(0, Socket::Address(id, &socket)); + + const char* const invalid_prefixes[] = { + "X", "UX", "RX", "RDN", "RDMB", + }; + for (size_t i = 0; i < arraysize(invalid_prefixes); ++i) { + butil::IOBuf source; + source.append(invalid_prefixes[i]); + const size_t original_size = source.size(); + const ParseResult result = policy::ParseTransportHandshake( + &source, socket.get(), false, NULL); + ASSERT_FALSE(result.is_ok()); + EXPECT_EQ(PARSE_ERROR_TRY_OTHERS, result.error()) + << invalid_prefixes[i]; + EXPECT_EQ(original_size, source.size()) + << invalid_prefixes[i]; + EXPECT_EQ(UNINITIALIZED, + AdapterTransport::Get(socket.get())->handshake_phase()); + EXPECT_EQ(nullptr, socket->parsing_context()); + } + socket->SetFailed(); +} + +TEST(TransportHandshakeTest, + partial_rdma_magic_waits_for_more_data_without_consuming) { + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard peer_fd(fds[1]); + SocketOptions options; + options.fd = fds[0]; + SocketId id; + ASSERT_EQ(0, Socket::Create(options, &id)); + SocketUniquePtr socket; + ASSERT_EQ(0, Socket::Address(id, &socket)); + + const char* const partial_prefixes[] = { + "R", "RD", "RDM", + }; + for (size_t i = 0; i < arraysize(partial_prefixes); ++i) { + butil::IOBuf source; + source.append(partial_prefixes[i]); + const size_t original_size = source.size(); + const ParseResult result = policy::ParseTransportHandshake( + &source, socket.get(), false, NULL); + ASSERT_FALSE(result.is_ok()); + EXPECT_EQ(PARSE_ERROR_NOT_ENOUGH_DATA, result.error()) + << partial_prefixes[i]; + EXPECT_EQ(original_size, source.size()) + << partial_prefixes[i]; + EXPECT_EQ(UNINITIALIZED, + AdapterTransport::Get(socket.get())->handshake_phase()); + EXPECT_EQ(nullptr, socket->parsing_context()); + } + socket->SetFailed(); +} + TEST(TransportHandshakeTest, plain_tcp_server_incrementally_rejects_ubshm_upgrade) { int fds[2]; @@ -458,7 +570,8 @@ TEST(TransportHandshakeTest, socket->SetFailed(); } -TEST(TransportHandshakeTest, plain_tcp_server_consumes_coalesced_ubshm_ack) { +TEST(TransportHandshakeTest, + plain_tcp_server_preserves_data_coalesced_behind_ubshm_ack) { int fds[2]; ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); butil::fd_guard peer_fd(fds[1]); @@ -473,11 +586,19 @@ TEST(TransportHandshakeTest, plain_tcp_server_consumes_coalesced_ubshm_ack) { source.append(MakeUBShmHello()); const uint32_t ack = 0; source.append(&ack, sizeof(ack)); + const char application_data[] = "application-data"; + source.append(application_data, sizeof(application_data) - 1); + const ParseResult result = policy::ParseTransportHandshake( &source, socket.get(), false, NULL); ASSERT_FALSE(result.is_ok()); ASSERT_EQ(PARSE_ERROR_TRY_OTHERS, result.error()); - ASSERT_TRUE(source.empty()); + ASSERT_EQ(sizeof(application_data) - 1, source.size()); + + char remaining[sizeof(application_data) - 1] = {}; + ASSERT_EQ(sizeof(remaining), + source.copy_to(remaining, sizeof(remaining))); + EXPECT_EQ(0, memcmp(application_data, remaining, sizeof(remaining))); ASSERT_EQ(FALLBACK_TCP, AdapterTransport::Get(socket.get())->handshake_phase()); ASSERT_EQ(nullptr, socket->parsing_context());