diff --git a/src/brpc/controller.cpp b/src/brpc/controller.cpp index 0bcfb4122d..8215e7f549 100644 --- a/src/brpc/controller.cpp +++ b/src/brpc/controller.cpp @@ -1478,7 +1478,7 @@ void Controller::HandleStreamConnection(Socket *host_socket) { return; } size_t stream_num = _request_streams.size(); - std::vector ptrs(stream_num); + std::vector ptrs(stream_num); if (!FailedInline()) { if (_remote_stream_settings == NULL) { if (!FailedInline()) { @@ -1486,7 +1486,7 @@ void Controller::HandleStreamConnection(Socket *host_socket) { } } else { for (size_t i = 0; i < stream_num; ++i) { - if (Socket::Address(_request_streams[i], &ptrs[i]) != 0) { + if (Stream::Address(_request_streams[i], &ptrs[i]) != 0) { if (!FailedInline()) { SetFailed(EREQUEST, "Request stream=%" PRIu64 " was closed before responded", _request_streams[i]); @@ -1509,14 +1509,14 @@ void Controller::HandleStreamConnection(Socket *host_socket) { } return; } - Stream* s = (Stream*)ptrs[0]->conn(); + Stream* s = ptrs[0].get(); s->SetConnected(_remote_stream_settings); if (stream_num > 1) { auto extra_stream_ids = std::move(*_remote_stream_settings->mutable_extra_stream_ids()); _remote_stream_settings->clear_extra_stream_ids(); for (size_t i = 1; i < stream_num; ++i) { if(!ptrs[i]) continue; - Stream* extra_stream = (Stream *) ptrs[i]->conn(); + Stream* extra_stream = ptrs[i].get(); _remote_stream_settings->set_stream_id(extra_stream_ids[i - 1]); extra_stream->SetHostSocket(host_socket); extra_stream->SetConnected(_remote_stream_settings); diff --git a/src/brpc/policy/baidu_rpc_protocol.cpp b/src/brpc/policy/baidu_rpc_protocol.cpp index 2c5a7e7224..fa6c0fd67c 100644 --- a/src/brpc/policy/baidu_rpc_protocol.cpp +++ b/src/brpc/policy/baidu_rpc_protocol.cpp @@ -353,11 +353,11 @@ void SendRpcResponse(int64_t correlation_id, Controller* cntl, meta.set_attachment_size(attached_size); } StreamId response_stream_id = INVALID_STREAM_ID; - SocketUniquePtr stream_ptr; + StreamUniquePtr stream_ptr; if (!response_stream_ids.empty()) { response_stream_id = response_stream_ids[0]; - if (Socket::Address(response_stream_id, &stream_ptr) == 0) { - Stream* s = (Stream *) stream_ptr->conn(); + if (Stream::Address(response_stream_id, &stream_ptr) == 0) { + Stream* s = stream_ptr.get(); StreamSettings *stream_settings = meta.mutable_stream_settings(); s->FillSettings(stream_settings); s->SetHostSocket(sock); @@ -419,13 +419,13 @@ void SendRpcResponse(int64_t correlation_id, Controller* cntl, // written user data would follower the RPC response. // Reuse stream_ptr to avoid address first stream id again if (stream_ptr) { - ((Stream*)stream_ptr->conn())->SetConnected(); + stream_ptr->SetConnected(); } for (size_t i = 1; i < response_stream_ids.size(); ++i) { StreamId extra_stream_id = response_stream_ids[i]; - SocketUniquePtr extra_stream_ptr; - if (Socket::Address(extra_stream_id, &extra_stream_ptr) == 0) { - Stream* extra_stream = (Stream *) extra_stream_ptr->conn(); + StreamUniquePtr extra_stream_ptr; + if (Stream::Address(extra_stream_id, &extra_stream_ptr) == 0) { + Stream* extra_stream = extra_stream_ptr.get(); extra_stream->SetHostSocket(sock); extra_stream->SetConnected(); } else { @@ -1086,12 +1086,12 @@ void PackRpcRequest(butil::IOBuf* req_buf, if (!request_stream_ids.empty()) { StreamSettings* stream_settings = meta.mutable_stream_settings(); StreamId request_stream_id = request_stream_ids[0]; - SocketUniquePtr ptr; - if (Socket::Address(request_stream_id, &ptr) != 0) { + StreamUniquePtr ptr; + if (Stream::Address(request_stream_id, &ptr) != 0) { return cntl->SetFailed(EREQUEST, "Stream=%" PRIu64 " was closed", request_stream_id); } - Stream* s = (Stream*) ptr->conn(); + Stream* s = ptr.get(); s->FillSettings(stream_settings); for (size_t i = 1; i < request_stream_ids.size(); ++i) { stream_settings->mutable_extra_stream_ids()->Add(request_stream_ids[i]); diff --git a/src/brpc/policy/streaming_rpc_protocol.cpp b/src/brpc/policy/streaming_rpc_protocol.cpp index b741acff5c..429d2bc282 100644 --- a/src/brpc/policy/streaming_rpc_protocol.cpp +++ b/src/brpc/policy/streaming_rpc_protocol.cpp @@ -102,11 +102,11 @@ ParseResult ParseStreamingMessage(butil::IOBuf* source, LOG(WARNING) << "Fail to Parse StreamFrameMeta from " << *socket; break; } - SocketUniquePtr ptr; - if (Socket::Address((SocketId)fm.stream_id(), &ptr) != 0) { - RPC_VLOG_IF(fm.frame_type() != FRAME_TYPE_RST - && fm.frame_type() != FRAME_TYPE_CLOSE - && fm.frame_type() != FRAME_TYPE_FEEDBACK) + StreamUniquePtr sptr; + if (Stream::Address((StreamId)fm.stream_id(), &sptr) != 0) { + RPC_VLOG_IF(fm.frame_type() != FRAME_TYPE_RST && + fm.frame_type() != FRAME_TYPE_CLOSE && + fm.frame_type() != FRAME_TYPE_FEEDBACK) << "Fail to find stream=" << fm.stream_id(); // It's normal that the stream is closed before receiving feedback frames from peer. // In this case, RST frame should not be sent to peer, otherwise on-fly data can be lost. @@ -116,16 +116,7 @@ ParseResult ParseStreamingMessage(butil::IOBuf* source, break; } meta_buf.clear(); // to reduce memory resident - // ptr->conn() returns the connection-level context attached to the - // socket. It may be NULL when the socket was found by ID but has no - // Stream object associated (e.g. during protocol probing or fuzz - // testing). Calling OnReceived on a null pointer would crash. - Stream* stream_conn = (Stream*)ptr->conn(); - if (stream_conn == NULL) { - LOG(FATAL) << "No stream object found"; - break; - } - stream_conn->OnReceived(fm, &payload, socket); + sptr->OnReceived(fm, &payload, socket); } while (0); // Hack input messenger @@ -136,7 +127,7 @@ void ProcessStreamingMessage(InputMessageBase* /*msg*/) { CHECK(false) << "Should never be called"; } -void SendStreamRst(Socket *sock, int64_t remote_stream_id) { +void SendStreamRst(Socket* sock, int64_t remote_stream_id) { CHECK(sock != NULL); StreamFrameMeta fm; fm.set_stream_id(remote_stream_id); @@ -148,7 +139,7 @@ void SendStreamRst(Socket *sock, int64_t remote_stream_id) { sock->Write(&out, &wopt); } -void SendStreamClose(Socket *sock, int64_t remote_stream_id, +void SendStreamClose(Socket* sock, int64_t remote_stream_id, int64_t source_stream_id) { CHECK(sock != NULL); StreamFrameMeta fm; diff --git a/src/brpc/socket.h b/src/brpc/socket.h index 7c53058941..3bc90918d9 100644 --- a/src/brpc/socket.h +++ b/src/brpc/socket.h @@ -353,7 +353,7 @@ friend class TransportFactory; // NOTE: User cannot create Socket from constructor. Use Create() // instead. It's public just because of requirement of ResourcePool. explicit Socket(Forbidden); - ~Socket() override; + ~Socket(); // Write `msg' into this Socket and clear it. The `msg' should be an // intact request or response. To prevent messages from interleaving diff --git a/src/brpc/stream.cpp b/src/brpc/stream.cpp index 2667614d9b..c799f2ffe1 100644 --- a/src/brpc/stream.cpp +++ b/src/brpc/stream.cpp @@ -42,23 +42,21 @@ BRPC_VALIDATE_GFLAG(stream_write_max_segment_size, PositiveInteger); const static butil::IOBuf *TIMEOUT_TASK = (butil::IOBuf*)-1L; -Stream::Stream() - : _host_socket(NULL) - , _fake_socket_weak_ref(NULL) +Stream::Stream(Forbidden f) + : VersionedRefWithId(f) + , _host_socket(NULL) , _connected(false) - , _closed(false) , _error_code(0) , _produced(0) , _remote_consumed(0) + , _socket_unconsumed_size(0) , _cur_buf_size(0) , _local_consumed(0) , _atomic_local_consumed(0) , _parse_rpc_response(false) , _pending_buf(NULL) , _start_idle_timer_us(0) - , _idle_timer(0) -{ - _connect_meta.on_connect = NULL; + , _idle_timer(0) { CHECK_EQ(0, bthread_mutex_init(&_connect_mutex, NULL)); CHECK_EQ(0, bthread_mutex_init(&_congestion_control_mutex, NULL)); } @@ -72,289 +70,262 @@ Stream::~Stream() { CHECK(_host_socket == NULL); bthread_mutex_destroy(&_connect_mutex); bthread_mutex_destroy(&_congestion_control_mutex); - bthread_id_list_destroy(&_writable_wait_list); } int Stream::Create(const StreamOptions &options, - const StreamSettings *remote_settings, + const StreamSettings* remote_settings, StreamId *id, bool parse_rpc_response) { - Stream* s = new Stream(); - s->_host_socket = NULL; - s->_fake_socket_weak_ref = NULL; - s->_connected = false; - s->_options = options; - s->_closed = false; - s->_error_code = 0; - s->_cur_buf_size = options.max_buf_size > 0 ? options.max_buf_size : 0; + return VersionedRefWithId::Create( + id, options, remote_settings, parse_rpc_response); +} + +int Stream::OnCreated(const StreamOptions& options, + const StreamSettings* remote_settings, + bool parse_rpc_response) { + _host_socket = NULL; + _connected.store(false, butil::memory_order_relaxed); + _options = options; + _error_code = 0; + _error_text.clear(); + _pending_writes.clear(); + _produced = 0; + _remote_consumed = 0; + _socket_unconsumed_size = 0; + _local_consumed = 0; + _atomic_local_consumed.store(0, butil::memory_order_relaxed); + _parse_rpc_response = parse_rpc_response; + _pending_buf = NULL; + _start_idle_timer_us = 0; + _idle_timer = 0; + _remote_settings.Clear(); + + _cur_buf_size = options.max_buf_size > 0 ? options.max_buf_size : 0; if (options.max_buf_size > 0 && options.min_buf_size > options.max_buf_size) { // set 0 if min_buf_size is invalid. - s->_options.min_buf_size = 0; + _options.min_buf_size = 0; LOG(WARNING) << "options.min_buf_size is larger than options.max_buf_size, it will be set to 0."; } - if (FLAGS_socket_max_streams_unconsumed_bytes > 0 && s->_options.min_buf_size > 0) { - s->_cur_buf_size = s->_options.min_buf_size; + if (FLAGS_socket_max_streams_unconsumed_bytes > 0 && _options.min_buf_size > 0) { + _cur_buf_size = _options.min_buf_size; } if (remote_settings != NULL) { - s->_remote_settings.MergeFrom(*remote_settings); - } - s->_parse_rpc_response = parse_rpc_response; - if (bthread_id_list_init(&s->_writable_wait_list, 8, 8/*FIXME*/)) { - delete s; - return -1; + _remote_settings.MergeFrom(*remote_settings); } + + CHECK_EQ(0, bthread_id_list_init(&_writable_wait_list, 8, 8/*FIXME*/)); + bthread::ExecutionQueueOptions q_opt; q_opt.bthread_attr = FLAGS_usercode_in_pthread ? BTHREAD_ATTR_PTHREAD : BTHREAD_ATTR_NORMAL; - if (bthread::execution_queue_start(&s->_consumer_queue, &q_opt, Consume, s) != 0) { + if (bthread::execution_queue_start(&_consumer_queue, &q_opt, Consume, this) != 0) { LOG(FATAL) << "Fail to create ExecutionQueue"; - delete s; return -1; } - SocketOptions sock_opt; - sock_opt.conn = s; - SocketId fake_sock_id; - if (Socket::Create(sock_opt, &fake_sock_id) != 0) { - s->BeforeRecycle(NULL); - return -1; - } - SocketUniquePtr ptr; - CHECK_EQ(0, Socket::Address(fake_sock_id, &ptr)); - s->_fake_socket_weak_ref = ptr.get(); - s->_id = fake_sock_id; - *id = s->id(); + + // The consumer queue holds one reference to this Stream. + AddReference(); return 0; } -void Stream::BeforeRecycle(Socket *) { - // No one holds reference now, so we don't need lock here - bthread_id_list_reset(&_writable_wait_list, ECONNRESET); - if (_connected) { - // Send CLOSE frame - RPC_VLOG << "Send close frame"; - CHECK(_host_socket != NULL); - policy::SendStreamClose(_host_socket, - _remote_settings.stream_id(), id()); +void Stream::OnFailed(int error_code, const std::string& error_text) { + bool connected = false; + { + // Record the error for on_failed callback fired in Consume(), and discard + // any writes buffered before connecting. + BAIDU_SCOPED_LOCK(_connect_mutex); + _error_code = error_code; + _error_text = error_text; + connected = _connected.load(butil::memory_order_relaxed); + _pending_writes.clear(); } - if (_host_socket) { - _host_socket->RemoveStream(id()); + // Wake up all threads blocked on writable. + bthread_id_list_reset(&_writable_wait_list, ECONNRESET); + + // Serialize the host Socket membership removal with SetHostSocket(). + // SetFailed() marks this Stream failed before entering OnFailed(), so a + // later SetHostSocket() observes Failed() and cannot add it back. + { + BAIDU_SCOPED_LOCK(_connect_mutex); + if (connected) { + RPC_VLOG << "Send close frame"; + CHECK(_host_socket != NULL); + policy::SendStreamClose( + _host_socket, _remote_settings.stream_id(), id()); + } + if (_host_socket != NULL) { + if (FLAGS_socket_max_streams_unconsumed_bytes > 0) { + BAIDU_SCOPED_LOCK(_congestion_control_mutex); + if (_socket_unconsumed_size != 0) { + _host_socket->_total_streams_unconsumed_size.fetch_sub( + _socket_unconsumed_size, butil::memory_order_relaxed); + _socket_unconsumed_size = 0; + } + } + _host_socket->RemoveStream(id()); + } } - // The instance is to be deleted in the consumer thread + // Stop the consumer queue. Consume() will fire on_failed/on_closed and + // release the reference held by the queue, which may recycle this instance. bthread::execution_queue_stop(_consumer_queue); } -ssize_t Stream::CutMessageIntoFileDescriptor(int /*fd*/, - butil::IOBuf **data_list, - size_t size) { +void Stream::BeforeRecycled() { + if (_pending_buf != NULL) { + delete _pending_buf; + _pending_buf = NULL; + } + + _pending_writes.clear(); + bthread_id_list_destroy(&_writable_wait_list); + if (_host_socket != NULL) { + DereferenceSocket(_host_socket); + _host_socket = NULL; + } +} + +std::string Stream::OnDescription() const { + BAIDU_SCOPED_LOCK(_connect_mutex); + if (_host_socket != NULL) { + return _host_socket->description(); + } else { + return "host_socket=NULL"; + } +} + +int Stream::WritePacked(const butil::IOBuf& data, + const StreamWriteOptions* options) { if (_host_socket == NULL) { CHECK(false) << "Not connected"; errno = EBADF; return -1; } if (!_remote_settings.writable()) { - LOG(WARNING) << "The remote side of Stream=" << id() + LOG(WARNING) << "The remote side of Stream=" << id() << "->" << _remote_settings.stream_id() << "@" << _host_socket->remote_side() << " doesn't have a handler"; errno = EBADF; return -1; } - butil::IOBuf out; - ssize_t len = 0; - ssize_t unwritten_data_size = 0; - for (size_t i = 0; i < size; ++i) { - butil::IOBuf *data = data_list[i]; - size_t length = data->length(); - if (length > FLAGS_stream_write_max_segment_size) { - if (unwritten_data_size) { - WriteToHostSocket(&out); - unwritten_data_size = 0; - out.clear(); - } - // segmenting large data into multiple parts - butil::IOBuf segment_buf; - bool has_continuation = true; - while (has_continuation) { - data->cutn(&segment_buf, FLAGS_stream_write_max_segment_size); - StreamFrameMeta fm; - fm.set_stream_id(_remote_settings.stream_id()); - fm.set_source_stream_id(id()); - fm.set_frame_type(FRAME_TYPE_DATA); - has_continuation = !data->empty(); - fm.set_has_continuation(has_continuation); - policy::PackStreamMessage(&out, fm, &segment_buf); - len += segment_buf.length(); - segment_buf.clear(); - WriteToHostSocket(&out); - out.clear(); - } - } else { - if (unwritten_data_size + length > FLAGS_stream_write_max_segment_size) { - WriteToHostSocket(&out); - unwritten_data_size = 0; - out.clear(); - } - unwritten_data_size += length; - StreamFrameMeta fm; - fm.set_stream_id(_remote_settings.stream_id()); - fm.set_source_stream_id(id()); - fm.set_frame_type(FRAME_TYPE_DATA); - fm.set_has_continuation(false); - policy::PackStreamMessage(&out, fm, data_list[i]); - len += length; - data_list[i]->clear(); - } - } - - if (!out.empty()) { - WriteToHostSocket(&out); - } - return len; -} - -void Stream::WriteToHostSocket(butil::IOBuf* b) { - BRPC_HANDLE_EOVERCROWDED(_host_socket->Write(b)); -} - -ssize_t Stream::CutMessageIntoSSLChannel(SSL*, butil::IOBuf**, size_t) { - CHECK(false) << "Stream does support SSL"; - errno = EINVAL; - return -1; -} -void* Stream::RunOnConnect(void *arg) { - ConnectMeta* meta = (ConnectMeta*)arg; - if (meta->ec == 0) { - meta->on_connect(Socket::STREAM_FAKE_FD, 0, meta->arg); - } else { - meta->on_connect(-1, meta->ec, meta->arg); - } - delete meta; - return NULL; -} + Socket::WriteOptions wopt; + wopt.write_in_background = options != NULL && options->write_in_background; -int Stream::Connect(Socket* ptr, const timespec*, - int (*on_connect)(int, int, void *), void *data) { - CHECK_EQ(ptr->id(), _id); - bthread_mutex_lock(&_connect_mutex); - if (_connect_meta.on_connect != NULL) { - CHECK(false) << "Connect is supposed to be called once"; - bthread_mutex_unlock(&_connect_mutex); + // Pack the whole message (splitting large data into multiple STRM frames) + // into a SINGLE IOBuf, then hand it to Socket::Write in one shot. + butil::IOBuf remaining(data); + butil::IOBuf out; + bool has_continuation = true; + do { + butil::IOBuf segment; + remaining.cutn(&segment, FLAGS_stream_write_max_segment_size); + has_continuation = !remaining.empty(); + StreamFrameMeta fm; + fm.set_stream_id(_remote_settings.stream_id()); + fm.set_source_stream_id(id()); + fm.set_frame_type(FRAME_TYPE_DATA); + fm.set_has_continuation(has_continuation); + policy::PackStreamMessage(&out, fm, &segment); + } while (has_continuation); + + if (BRPC_HANDLE_EOVERCROWDED(_host_socket->Write(&out, &wopt)) != 0) { + // Stream may be closed by peer before. + LOG(WARNING) << "Fail to write to host socket of stream=" << id() + << ", " << berror(); return -1; } - _connect_meta.on_connect = on_connect; - _connect_meta.arg = data; - if (_connected) { - ConnectMeta* meta = new ConnectMeta; - meta->on_connect = _connect_meta.on_connect; - meta->arg = _connect_meta.arg; - meta->ec = _connect_meta.ec; - bthread_mutex_unlock(&_connect_mutex); - bthread_t tid; - if (bthread_start_urgent(&tid, &BTHREAD_ATTR_NORMAL, RunOnConnect, meta) != 0) { - LOG(FATAL) << "Fail to start bthread, " << berror(); - RunOnConnect(meta); - } - return 0; - } - bthread_mutex_unlock(&_connect_mutex); return 0; } -void Stream::SetConnected() { - return SetConnected(NULL); -} - -void Stream::SetConnected(const StreamSettings* remote_settings) { - bthread_mutex_lock(&_connect_mutex); - if (_closed) { - bthread_mutex_unlock(&_connect_mutex); - return; - } - if (_connected) { - CHECK(false); - bthread_mutex_unlock(&_connect_mutex); - return; - } - CHECK(_host_socket != NULL); - if (remote_settings != NULL) { - CHECK(!_remote_settings.IsInitialized()); - _remote_settings.MergeFrom(*remote_settings); - } else { - CHECK(_remote_settings.IsInitialized()); - } - CHECK(_host_socket != NULL); - RPC_VLOG << "stream=" << id() << " is connected to stream_id=" - << _remote_settings.stream_id() << " at host_socket=" << *_host_socket; - _connected.store(true, butil::memory_order_release); - _connect_meta.ec = 0; - TriggerOnConnectIfNeed(); - if (remote_settings == NULL) { - // Start the timer at server-side - // Client-side timer would triggered in Consume after received the first - // message which is the very RPC response - StartIdleTimer(); - } else { - // send first feedback for client-side stream if it already consumed data - if (_remote_settings.need_feedback()) { - auto consumed_bytes = _atomic_local_consumed.load(butil::memory_order_acquire); - if (consumed_bytes > 0) - SendFeedback(consumed_bytes); - } - } +void Stream::WriteToHostSocket(butil::IOBuf* b) { + BRPC_HANDLE_EOVERCROWDED(_host_socket->Write(b)); } -void Stream::TriggerOnConnectIfNeed() { - if (_connect_meta.on_connect != NULL) { - ConnectMeta* meta = new ConnectMeta; - meta->on_connect = _connect_meta.on_connect; - meta->arg = _connect_meta.arg; - meta->ec = _connect_meta.ec; - bthread_mutex_unlock(&_connect_mutex); - bthread_t tid; - if (bthread_start_urgent(&tid, &BTHREAD_ATTR_NORMAL, RunOnConnect, meta) != 0) { - LOG(FATAL) << "Fail to start bthread, " << berror(); - RunOnConnect(meta); - } - return; +inline void Stream::RollbackProduced(size_t data_length) { + if (_cur_buf_size > 0) { + BAIDU_SCOPED_LOCK(_congestion_control_mutex); + _produced -= data_length; } - bthread_mutex_unlock(&_connect_mutex); } int Stream::AppendIfNotFull(const butil::IOBuf &data, const StreamWriteOptions* options) { + if (Failed()) { + errno = ECONNRESET; + return -1; + } + + size_t data_length = data.length(); if (_cur_buf_size > 0) { std::unique_lock lck(_congestion_control_mutex); if (_produced >= _remote_consumed + _cur_buf_size) { const size_t saved_produced = _produced; const size_t saved_remote_consumed = _remote_consumed; lck.unlock(); - RPC_VLOG << "Stream=" << _id << " is full" + RPC_VLOG << "Stream=" << id() << " is full" << "_produced=" << saved_produced << " _remote_consumed=" << saved_remote_consumed << " gap=" << saved_produced - saved_remote_consumed << " max_buf_size=" << _cur_buf_size; return 1; } - _produced += data.length(); + _produced += data_length; } - size_t data_length = data.length(); - butil::IOBuf copied_data(data); - Socket::WriteOptions wopt; - wopt.write_in_background = options != NULL && options->write_in_background; - const int rc = _fake_socket_weak_ref->Write(&copied_data, &wopt); - if (rc != 0) { - // Stream may be closed by peer before - LOG(WARNING) << "Fail to write to _fake_socket, " << berror(); - BAIDU_SCOPED_LOCK(_congestion_control_mutex); - _produced -= data_length; + // Fast path (the common case): once connected, write directly WITHOUT + // taking _connect_mutex. `_connected` is a one-way transition published + // by SetConnected() after flushing pending writes, so ordering is preserved + // and this path stays lock-free (besides the optional congestion window). + if (_connected.load(butil::memory_order_acquire)) { + if (WritePacked(data, options) != 0) { + RollbackProduced(data_length); + return -1; + } + if (FLAGS_socket_max_streams_unconsumed_bytes > 0) { + BAIDU_SCOPED_LOCK(_congestion_control_mutex); + if (!Failed()) { + _host_socket->_total_streams_unconsumed_size.fetch_add( + data_length, butil::memory_order_relaxed); + _socket_unconsumed_size += data_length; + } + } + return 0; + } + + // Slow path (rare): not connected yet, so the remote stream id is unknown. + // Buffer the raw data and options under `_connect_mutex`. SetConnected() + // will flush it. + { + BAIDU_SCOPED_LOCK(_connect_mutex); + if (Failed()) { + RollbackProduced(data_length); + return -1; + } + if (!_connected.load(butil::memory_order_acquire)) { + _pending_writes.emplace_back(data, options); + return 0; + } + // Connected between the two checks; fall through to a direct write. + // Ordering holds: reaching here means we acquired `_connect_mutex`, + // which SetConnected() releases only after it has flushed all pending + // writes (enqueued into the host socket) and published `_connected=true`. + // Hence, this direct write is necessarily enqueued after those pending writes. + } + + if (WritePacked(data, options) != 0) { + RollbackProduced(data_length); return -1; } if (FLAGS_socket_max_streams_unconsumed_bytes > 0) { - _host_socket->_total_streams_unconsumed_size += data_length; + BAIDU_SCOPED_LOCK(_congestion_control_mutex); + if (!Failed()) { + _host_socket->_total_streams_unconsumed_size.fetch_add( + data_length, butil::memory_order_relaxed); + _socket_unconsumed_size += data_length; + } } return 0; } @@ -362,7 +333,7 @@ int Stream::AppendIfNotFull(const butil::IOBuf &data, void Stream::SetRemoteConsumed(size_t new_remote_consumed) { CHECK(_cur_buf_size > 0); bthread_id_list_t tmplist; - bthread_id_list_init(&tmplist, 0, 0); + CHECK_EQ(0, bthread_id_list_init(&tmplist, 0, 0)); bthread_mutex_lock(&_congestion_control_mutex); if (_remote_consumed >= new_remote_consumed) { bthread_mutex_unlock(&_congestion_control_mutex); @@ -370,16 +341,28 @@ void Stream::SetRemoteConsumed(size_t new_remote_consumed) { } const bool was_full = _produced >= _remote_consumed + _cur_buf_size; - if (FLAGS_socket_max_streams_unconsumed_bytes > 0) { - _host_socket->_total_streams_unconsumed_size -= new_remote_consumed - _remote_consumed; - if (_host_socket->_total_streams_unconsumed_size > FLAGS_socket_max_streams_unconsumed_bytes) { + if (FLAGS_socket_max_streams_unconsumed_bytes > 0 && _host_socket != NULL) { + const size_t consumed_delta = new_remote_consumed - _remote_consumed; + const size_t accounted_delta = + std::min(consumed_delta, _socket_unconsumed_size); + if (accounted_delta != 0) { + _host_socket->_total_streams_unconsumed_size.fetch_sub( + accounted_delta, butil::memory_order_relaxed); + _socket_unconsumed_size -= accounted_delta; + } + const int64_t total_unconsumed = _host_socket->_total_streams_unconsumed_size.load( + butil::memory_order_relaxed); + if (total_unconsumed > FLAGS_socket_max_streams_unconsumed_bytes) { if (_options.min_buf_size > 0) { _cur_buf_size = _options.min_buf_size; } else { _cur_buf_size /= 2; } - LOG(INFO) << "stream consumers on socket " << _host_socket->id() << " is crowded, " << "cut stream " << id() << " buffer to " << _cur_buf_size; - } else if (_produced >= new_remote_consumed + _cur_buf_size && (_options.max_buf_size <= 0 || _cur_buf_size < (size_t)_options.max_buf_size)) { + LOG(INFO) << "stream consumers on socket " << _host_socket->id() + << " is crowded, cut stream " << id() + << " buffer to " << _cur_buf_size; + } else if (_produced >= new_remote_consumed + _cur_buf_size && + (_options.max_buf_size <= 0 || _cur_buf_size < (size_t)_options.max_buf_size)) { if (_options.max_buf_size > 0 && _cur_buf_size * 2 > (size_t)_options.max_buf_size) { _cur_buf_size = _options.max_buf_size; } else { @@ -496,15 +479,123 @@ int Stream::Wait(const timespec* due_time) { return rc; } +void Stream::SetConnected() { + return SetConnected(NULL); +} + +void Stream::SetConnected(const StreamSettings* remote_settings) { + bthread_mutex_lock(&_connect_mutex); + if (Failed()) { + bthread_mutex_unlock(&_connect_mutex); + return; + } + if (_connected.load(butil::memory_order_relaxed)) { + // SetConnected() may be driven more than once (and concurrently) for + // the same stream, notably for extra streams in batch creation. It must + // be idempotent: guarded by _connect_mutex, only the first call takes + // effect and later calls simply return. + bthread_mutex_unlock(&_connect_mutex); + return; + } + CHECK(_host_socket != NULL); + if (remote_settings != NULL) { + CHECK(!_remote_settings.IsInitialized()); + _remote_settings.MergeFrom(*remote_settings); + } else { + CHECK(_remote_settings.IsInitialized()); + } + RPC_VLOG << "stream=" << id() << " is connected to stream_id=" + << _remote_settings.stream_id() << " at host_socket=" << *_host_socket; + + // Flush writes buffered before connecting FIRST, while _connected is still + // false so concurrent AppendIfNotFull() take the slow path and block on + // _connect_mutex. Only after flushing do we publish _connected=true, so + // subsequent lock-free fast-path writes are strictly ordered after these + // pending writes. + std::vector pending; + pending.swap(_pending_writes); + for (size_t i = 0; i < pending.size(); ++i) { + if (Failed()) { + size_t unsent_size = 0; + for (size_t j = i; j < pending.size(); ++j) { + unsent_size += pending[j].data.length(); + } + RollbackProduced(unsent_size); + bthread_mutex_unlock(&_connect_mutex); + return; + } + + size_t len = pending[i].data.length(); + if (WritePacked(pending[i].data, &pending[i].options) != 0) { + int error_code = errno != 0 ? errno : EIO; + // The congestion window accounted for every pending write when it + // was accepted. Keep the successfully enqueued prefix accounted, + // but roll back the failed write and the unsent suffix. + size_t unsent_size = 0; + for (size_t j = i; j < pending.size(); ++j) { + unsent_size += pending[j].data.length(); + } + RollbackProduced(unsent_size); + bthread_mutex_unlock(&_connect_mutex); + VersionedRefWithId::SetFailed( + error_code, "Failed to flush pending writes during connection"); + return; + } + if (FLAGS_socket_max_streams_unconsumed_bytes > 0) { + BAIDU_SCOPED_LOCK(_congestion_control_mutex); + if (!Failed()) { + _host_socket->_total_streams_unconsumed_size.fetch_add( + len, butil::memory_order_relaxed); + _socket_unconsumed_size += len; + } + } + } + + // Check both before and after publishing. The second check closes the + // window in which SetFailed() can bump the version between the first check + // and the store. If failure happens after the second check, connection was + // published first and OnFailed() will observe and close it normally. + if (Failed()) { + bthread_mutex_unlock(&_connect_mutex); + return; + } + _connected.store(true, butil::memory_order_release); + if (Failed()) { + _connected.store(false, butil::memory_order_relaxed); + bthread_mutex_unlock(&_connect_mutex); + return; + } + bthread_mutex_unlock(&_connect_mutex); + + if (remote_settings == NULL) { + // Start the timer at server-side + // Client-side timer would triggered in Consume after received the first + // message which is the very RPC response + StartIdleTimer(); + } else { + // send first feedback for client-side stream if it already consumed data + if (_remote_settings.need_feedback()) { + auto consumed_bytes = _atomic_local_consumed.load(butil::memory_order_acquire); + if (consumed_bytes > 0) + SendFeedback(consumed_bytes); + } + } +} + int Stream::OnReceived(const StreamFrameMeta& fm, butil::IOBuf *buf, Socket* sock) { - if (_host_socket == NULL) { + if (!_connected.load(butil::memory_order_acquire)) { + // Before connection is published, let the locked slow path initialize + // the host socket or confirm that another thread already did so. if (SetHostSocket(sock) != 0) { return -1; } } + switch (fm.frame_type()) { case FRAME_TYPE_FEEDBACK: - SetRemoteConsumed(fm.feedback().consumed_size()); + if (_connected.load(butil::memory_order_acquire)) { + SetRemoteConsumed(fm.feedback().consumed_size()); + } CHECK(buf->empty()); break; case FRAME_TYPE_DATA: @@ -516,7 +607,7 @@ int Stream::OnReceived(const StreamFrameMeta& fm, butil::IOBuf *buf, Socket* soc _pending_buf->swap(*buf); } if (!fm.has_continuation()) { - butil::IOBuf *tmp = _pending_buf; + butil::IOBuf* tmp = _pending_buf; _pending_buf = NULL; int rc = bthread::execution_queue_execute(_consumer_queue, tmp); if (rc != 0) { @@ -583,12 +674,10 @@ int Stream::Consume(void *meta, bthread::TaskIterator& iter) { Stream* s = (Stream*)meta; s->StopIdleTimer(); if (iter.is_queue_stopped()) { - scoped_ptr recycled_stream(s); - // Indicating the queue was closed. - if (s->_host_socket) { - DereferenceSocket(s->_host_socket); - s->_host_socket = NULL; - } + // The consumer queue is stopped (the stream was SetFailed). Fire the + // user callbacks, then release the reference held by the queue (which + // was added in OnCreated). This may recycle the instance via + // BeforeRecycled(), so do not touch `s' afterwards. if (s->_options.handler != NULL) { int error_code; std::string error_text; @@ -603,8 +692,10 @@ int Stream::Consume(void *meta, bthread::TaskIterator& iter) { } s->_options.handler->on_closed(s->id()); } + DereferenceVersionedRefWithId(s); return 0; } + DEFINE_SMALL_ARRAY(butil::IOBuf*, buf_list, s->_options.messages_in_batch, 256); MessageBatcher mb(buf_list, s->_options.messages_in_batch, s); bool has_timeout_task = false; @@ -661,18 +752,24 @@ void Stream::SendFeedback(int64_t _consumed_bytes) { WriteToHostSocket(&out); } -int Stream::SetHostSocket(Socket *host_socket) { - std::call_once(_set_host_socket_flag, [this, host_socket]() { - SocketUniquePtr ptr; - host_socket->ReAddress(&ptr); - // TODO add *this to host socke - if (ptr->AddStream(id()) != 0) { - CHECK(false) << id() << " fail to add stream to host socket"; - return; - } - _host_socket = ptr.release(); - }); - return _host_socket != NULL ? 0 : -1; +int Stream::SetHostSocket(Socket* host_socket) { + BAIDU_SCOPED_LOCK(_connect_mutex); + if (Failed()) { + return -1; + } + if (_host_socket != NULL) { + return 0; + } + + SocketUniquePtr ptr; + host_socket->ReAddress(&ptr); + if (ptr->AddStream(id()) != 0) { + CHECK(false) << id() << " fail to add stream to host socket"; + return -1; + } + + _host_socket = ptr.release(); + return 0; } void Stream::FillSettings(StreamSettings *settings) { @@ -707,49 +804,50 @@ void Stream::StopIdleTimer() { } } -void Stream::Close(int error_code, const char* reason_fmt, ...) { - _fake_socket_weak_ref->SetFailed(); - bthread_mutex_lock(&_connect_mutex); - if (_closed) { - bthread_mutex_unlock(&_connect_mutex); +void Stream::CloseV(int error_code, const char* reason_fmt, va_list ap) { + if (Failed()) { return; } - _closed = true; - _error_code = error_code; + std::string error_text; + butil::string_vappendf(&error_text, reason_fmt, ap); + VersionedRefWithId::SetFailed(error_code, error_text); +} + +void Stream::Close(int error_code, const char* reason_fmt, ...) { va_list ap; va_start(ap, reason_fmt); - butil::string_vappendf(&_error_text, reason_fmt, ap); + CloseV(error_code, reason_fmt, ap); va_end(ap); +} - if (_connected) { - bthread_mutex_unlock(&_connect_mutex); - return; +int Stream::SetFailedV(StreamId id, int error_code, + const char* reason_fmt, va_list ap) { + StreamUniquePtr stream_ptr; + if (AddressFailedAsWell(id, &stream_ptr) == -1) { + // Don't care recycled stream. + return 0; } - _connect_meta.ec = ECONNRESET; - // Trigger on connect to release the reference of socket - return TriggerOnConnectIfNeed(); + stream_ptr->CloseV(error_code, reason_fmt, ap); + return 0; } int Stream::SetFailed(StreamId id, int error_code, const char* reason_fmt, ...) { - SocketUniquePtr ptr; - if (Socket::AddressFailedAsWell(id, &ptr) == -1) { - // Don't care recycled stream - return 0; - } - Stream* s = (Stream*)ptr->conn(); va_list ap; va_start(ap, reason_fmt); - s->Close(error_code, reason_fmt, ap); + int rc = SetFailedV(id, error_code, reason_fmt, ap); va_end(ap); - return 0; + return rc; } int Stream::SetFailed(const StreamIds& ids, int error_code, const char* reason_fmt, ...) { va_list ap; va_start(ap, reason_fmt); - for(size_t i = 0; i< ids.size(); ++i) { - Stream::SetFailed(ids[i], error_code, reason_fmt, ap); + for (auto id : ids) { + va_list ap_copy; + va_copy(ap_copy, ap); + SetFailedV(id, error_code, reason_fmt, ap_copy); + va_end(ap_copy); } va_end(ap); return 0; @@ -779,13 +877,13 @@ void Stream::HandleRpcResponse(butil::IOBuf* response_buffer) { policy::ProcessRpcResponse(msg); } -int StreamWrite(StreamId stream_id, const butil::IOBuf &message, +int StreamWrite(StreamId stream_id, const butil::IOBuf& message, const StreamWriteOptions* options) { - SocketUniquePtr ptr; - if (Socket::Address(stream_id, &ptr) != 0) { + StreamUniquePtr stream_ptr; + if (Stream::Address(stream_id, &stream_ptr) != 0) { return EINVAL; } - Stream* s = (Stream*)ptr->conn(); + Stream* s = stream_ptr.get(); const int rc = s->AppendIfNotFull(message, options); if (rc == 0) { return 0; @@ -795,15 +893,15 @@ int StreamWrite(StreamId stream_id, const butil::IOBuf &message, void StreamWait(StreamId stream_id, const timespec *due_time, void (*on_writable)(StreamId, void*, int), void *arg) { - SocketUniquePtr ptr; - if (Socket::Address(stream_id, &ptr) != 0) { + StreamUniquePtr stream_ptr; + if (Stream::Address(stream_id, &stream_ptr) != 0) { Stream::WritableMeta* wm = new Stream::WritableMeta; wm->id = stream_id; wm->arg= arg; wm->has_timer = false; wm->on_writable = on_writable; wm->error_code = EINVAL; - const bthread_attr_t* attr = + const bthread_attr_t* attr = FLAGS_usercode_in_pthread ? &BTHREAD_ATTR_PTHREAD : &BTHREAD_ATTR_NORMAL; bthread_t tid; @@ -813,16 +911,16 @@ void StreamWait(StreamId stream_id, const timespec *due_time, } return; } - Stream* s = (Stream*)ptr->conn(); + Stream* s = stream_ptr.get(); return s->Wait(on_writable, arg, due_time); } int StreamWait(StreamId stream_id, const timespec* due_time) { - SocketUniquePtr ptr; - if (Socket::Address(stream_id, &ptr) != 0) { + StreamUniquePtr stream_ptr; + if (Stream::Address(stream_id, &stream_ptr) != 0) { return EINVAL; } - Stream* s = (Stream*)ptr->conn(); + Stream* s = stream_ptr.get(); return s->Wait(due_time); } diff --git a/src/brpc/stream.h b/src/brpc/stream.h index 36c0def70f..73e7aff52f 100644 --- a/src/brpc/stream.h +++ b/src/brpc/stream.h @@ -19,17 +19,18 @@ #ifndef BRPC_STREAM_H #define BRPC_STREAM_H +#include #include "butil/iobuf.h" #include "butil/scoped_generic.h" -#include "brpc/socket_id.h" +#include "brpc/versioned_ref_with_id.h" namespace brpc { class Controller; -typedef SocketId StreamId; +typedef VRefId StreamId; using StreamIds = std::vector; -const StreamId INVALID_STREAM_ID = (StreamId)-1L; +const StreamId INVALID_STREAM_ID = INVALID_VREF_ID; namespace detail { struct StreamIdTraits; @@ -134,7 +135,7 @@ int StreamAccept(StreamIds& response_stream, Controller& cntl, // - EAGAIN: |stream_id| is created with positive |max_buf_size| and buf size // which the remote side hasn't consumed yet excceeds the number. // - EINVAL: |stream_id| is invalied or has been closed -int StreamWrite(StreamId stream_id, const butil::IOBuf &message, +int StreamWrite(StreamId stream_id, const butil::IOBuf& message, const StreamWriteOptions* options = NULL); // Write util the pending buffer size is less than |max_buf_size| or orrur diff --git a/src/brpc/stream_impl.h b/src/brpc/stream_impl.h index 284b33ca33..f9ae065dab 100644 --- a/src/brpc/stream_impl.h +++ b/src/brpc/stream_impl.h @@ -19,42 +19,42 @@ #ifndef BRPC_STREAM_IMPL_H #define BRPC_STREAM_IMPL_H -#include +#include +#include #include "bthread/bthread.h" #include "bthread/execution_queue.h" #include "brpc/socket.h" #include "brpc/stream.h" +#include "brpc/versioned_ref_with_id.h" #include "brpc/streaming_rpc_meta.pb.h" namespace brpc { -class BAIDU_CACHELINE_ALIGNMENT Stream : public SocketConnection { +// Stream is implemented on top of VersionedRefWithId, so that StreamId +// is a self-contained versioned reference id and no longer depends on a fake +// Socket. The instance is managed by a ResourcePool: it is reused rather than +// re-constructed, thus all per-stream state must be (re)initialized in +// OnCreated() and cleaned up in OnFailed()/BeforeRecycled(). +class BAIDU_CACHELINE_ALIGNMENT Stream : public VersionedRefWithId { public: - // |--------------------------------------------------| - // |----------- Implement SocketConnection -----------| - // |--------------------------------------------------| - - int Connect(Socket* ptr, const timespec* due_time, - int (*on_connect)(int, int, void *), void *data); - ssize_t CutMessageIntoFileDescriptor(int, butil::IOBuf **data_list, - size_t size); - ssize_t CutMessageIntoSSLChannel(SSL*, butil::IOBuf**, size_t); - void BeforeRecycle(Socket *); - - // --------------------- SocketConnection -------------- + // NOTE: Users cannot create Stream from constructor. Use Create() instead. + // It's public only because of the requirement of ResourcePool. + explicit Stream(Forbidden); + ~Stream(); + // Write `msg' into this stream. Returns 0 on success, 1 when the stream is + // full, -1 on error. int AppendIfNotFull(const butil::IOBuf& msg, const StreamWriteOptions* options = NULL); static int Create(const StreamOptions& options, - const StreamSettings *remote_settings, + const StreamSettings* remote_settings, StreamId *id, bool parse_rpc_response = true); - StreamId id() { return _id; } int OnReceived(const StreamFrameMeta& fm, butil::IOBuf *buf, Socket* sock); void SetRemoteSettings(const StreamSettings& remote_settings) { _remote_settings.MergeFrom(remote_settings); } - int SetHostSocket(Socket *host_socket); + int SetHostSocket(Socket* host_socket); void SetConnected(); void SetConnected(const StreamSettings *remote_settings); @@ -62,6 +62,7 @@ class BAIDU_CACHELINE_ALIGNMENT Stream : public SocketConnection { const timespec *due_time); int Wait(const timespec* due_time); void FillSettings(StreamSettings *settings); + static int SetFailed(StreamId id, int error_code, const char* reason_fmt, ...) __attribute__ ((__format__ (__printf__, 3, 4))); static int SetFailed(const StreamIds& ids, int error_code, const char* reason_fmt, ...) @@ -73,12 +74,21 @@ class BAIDU_CACHELINE_ALIGNMENT Stream : public SocketConnection { friend void StreamWait(StreamId stream_id, const timespec *due_time, void (*on_writable)(StreamId, void*, int), void *arg); friend class MessageBatcher; -friend struct butil::DefaultDeleter; - Stream(); - ~Stream(); - int Init(const StreamOptions options); +friend class VersionedRefWithId; + + // Initialize (or reset for a reused instance) the stream. + // Returns 0 on success, non-zero on failure. + int OnCreated(const StreamOptions& options, + const StreamSettings* remote_settings, + bool parse_rpc_response); + // Called once when SetFailed() succeeds. Performs the close actions + // (wake up waiters, send CLOSE frame, stop the consumer queue, etc.). + void OnFailed(int error_code, const std::string& error_text); + // Called right before the instance is recycled to the ResourcePool. + void BeforeRecycled(); + std::string OnDescription() const; + void SetRemoteConsumed(size_t _remote_consumed); - void TriggerOnConnectIfNeed(); void Wait(void (*on_writable)(StreamId, void*, int), void* arg, const timespec* due_time, bool new_thread, bthread_id_t *join_id); void SendFeedback(int64_t _consumed_bytes); @@ -86,17 +96,20 @@ friend struct butil::DefaultDeleter; void StopIdleTimer(); void HandleRpcResponse(butil::IOBuf* response_buffer); void WriteToHostSocket(butil::IOBuf* b); + // Pack `data` into one or more STRM DATA frames (splitting large data into + // segments) and write them into the host socket in a single Write. + int WritePacked(const butil::IOBuf& data, const StreamWriteOptions* options); + // Roll back `_produced` by `data_length` (under `_congestion_control_mutex`) + // when a write fails. No-op when the congestion window is disabled. + void RollbackProduced(size_t data_length); static int Consume(void *meta, bthread::TaskIterator& iter); static int TriggerOnWritable(bthread_id_t id, void *data, int error_code); static void *RunOnWritable(void* arg); - static void* RunOnConnect(void* arg); - struct ConnectMeta { - int (*on_connect)(int, int, void*); - int ec; - void* arg; - }; + static int SetFailedV(StreamId id, int error_code, + const char* reason_fmt, va_list ap); + void CloseV(int error_code, const char* reason_fmt, va_list ap); struct WritableMeta { void (*on_writable)(StreamId, void*, int); @@ -108,21 +121,36 @@ friend struct butil::DefaultDeleter; bthread_timer_t timer; }; - Socket* _host_socket; // Every stream within a Socket holds a reference - Socket* _fake_socket_weak_ref; // Not holding reference - StreamId _id; + struct PendingWrite { + butil::IOBuf data; + StreamWriteOptions options; + + PendingWrite() = default; + explicit PendingWrite(const butil::IOBuf& d, const StreamWriteOptions* opts) + : data(d) { + if (opts != NULL) { + options = *opts; + } + } + }; + + Socket* _host_socket; // Every stream within a Socket holds a reference. StreamOptions _options; - bthread_mutex_t _connect_mutex; - ConnectMeta _connect_meta; + mutable bthread_mutex_t _connect_mutex; butil::atomic _connected; - bool _closed; - int _error_code; - std::string _error_text; + int _error_code; + std::string _error_text; + // Writes buffered before the stream is connected (the remote stream id + // is unknown until then). Flushed in SetConnected(). + std::vector _pending_writes; bthread_mutex_t _congestion_control_mutex; size_t _produced; size_t _remote_consumed; + // Bytes of this Stream currently included in the host Socket's aggregate + // unconsumed counter. Protected by _congestion_control_mutex. + size_t _socket_unconsumed_size; size_t _cur_buf_size; bthread_id_list_t _writable_wait_list; @@ -132,12 +160,13 @@ friend struct butil::DefaultDeleter; bool _parse_rpc_response; bthread::ExecutionQueueId _consumer_queue; - butil::IOBuf *_pending_buf; + butil::IOBuf* _pending_buf; int64_t _start_idle_timer_us; bthread_timer_t _idle_timer; - std::once_flag _set_host_socket_flag; }; +typedef VersionedRefWithIdUniquePtr StreamUniquePtr; + } // namespace brpc diff --git a/src/brpc/versioned_ref_with_id.h b/src/brpc/versioned_ref_with_id.h index f77d5afa83..3793e2185e 100644 --- a/src/brpc/versioned_ref_with_id.h +++ b/src/brpc/versioned_ref_with_id.h @@ -89,35 +89,74 @@ typename std::enable_if::value, Ret>::type ReturnEmpty() { template typename std::enable_if::value, Ret>::type ReturnEmpty() {} -// Call func_name of class_type if class_type implements func_name, -// otherwise call default function. -#define WRAPPER_OF(class_type, func_name, return_type) \ - struct func_name ## Wrapper { \ - template \ +// Detect whether a type implements the member function `func_name' callable +// with Args..., exposing the result as a compile-time boolean: +// HasMember_::value +// The detector is decoupled from the caller so that it can also be reused in +// standalone static_assert to enforce interface contracts. +#define BRPC_DEFINE_MEMBER_DETECTOR(func_name) \ + template \ + struct HasMember##func_name { \ + template \ static auto Test(int) -> decltype( \ - std::declval().func_name(std::declval()...), std::true_type()); \ - template \ + std::declval().func_name(std::declval()...), std::true_type()); \ + template \ static auto Test(...) -> std::false_type; \ - \ - template \ - typename std::enable_if(0))::value, return_type>::type \ - Call(class_type* obj, Args&&... args) { \ + static constexpr bool value = decltype(Test(0))::value; \ + } + +// Define a static caller `Call' that invokes `obj->func_name(...)' +// if the type implements it, otherwise returns a default-constructed value. +// Requires the detector defined by BRPC_DEFINE_MEMBER_DETECTOR(func_name). +// On C++20, an inline `requires' expression is used directly (no separate +// detector needed); +// On C++17, a single `if constexpr' branch with the detector; +// on C++11/14, two SFINAE overloads so that the body referencing a +// possibly-missing member is never instantiated. +#if __cplusplus >= 202002L +#define BRPC_DEFINE_OPTIONAL_CALLER(func_name, return_type) \ + template \ + static return_type Call##func_name(U* obj, Args&&... args) { \ + if constexpr (requires { obj->func_name(std::forward(args)...); }) { \ + BAIDU_CASSERT((butil::is_result_same< \ + return_type, decltype(&U::func_name), U, Args...>::value), \ + "Params or return type mismatch"); \ + return obj->func_name(std::forward(args)...); \ + } else { \ + return ReturnEmpty(); \ + } \ + } +#elif __cplusplus >= 201703L +#define BRPC_DEFINE_OPTIONAL_CALLER(func_name, return_type) \ + template \ + static return_type Call##func_name(U* obj, Args&&... args) { \ + if constexpr (HasMember##func_name::value) { \ BAIDU_CASSERT((butil::is_result_same< \ - return_type, decltype(&T::func_name), T, Args...>::value), \ + return_type, decltype(&U::func_name), U, Args...>::value), \ "Params or return type mismatch"); \ - return obj->func_name(std::forward(args)...); \ - } \ - \ - template \ - typename std::enable_if(0))::value, return_type>::type \ - Call(class_type* obj, Args&&...) { \ + return obj->func_name(std::forward(args)...); \ + } else { \ return ReturnEmpty(); \ } \ } - -#define WRAPPER_CALL(func_name, obj, ...) func_name ## Wrapper().Call(obj, ## __VA_ARGS__) +#else +#define BRPC_DEFINE_OPTIONAL_CALLER(func_name, return_type) \ + template \ + static typename std::enable_if< \ + HasMember##func_name::value, return_type>::type \ + Call##func_name(U* obj, Args&&... args) { \ + BAIDU_CASSERT((butil::is_result_same< \ + return_type, decltype(&U::func_name), U, Args...>::value), \ + "Params or return type mismatch"); \ + return obj->func_name(std::forward(args)...); \ + } \ + template \ + static typename std::enable_if< \ + !HasMember##func_name::value, return_type>::type \ + Call##func_name(U*, Args&&...) { \ + return ReturnEmpty(); \ + } +#endif // VersionedRefWithId is an efficient data structure, which can be find // in O(1)-time by VRefId. @@ -205,7 +244,11 @@ class VersionedRefWithId { , _this_id(0) , _additional_ref_status(ADDITIONAL_REF_USING) {} - virtual ~VersionedRefWithId() = default; + // Non-virtual on purpose: CRTP static polymorphism needs no vtable, and + // instances are always recycled via return_resource() (never deleted + // through a base pointer), so a virtual destructor would only add a + // useless vptr and hurt cacheline layout. + ~VersionedRefWithId() = default; DISALLOW_COPY_AND_ASSIGN(VersionedRefWithId); // Create a VersionedRefWithId, put the identifier into `id'. @@ -219,14 +262,14 @@ class VersionedRefWithId { // of scope (w/o explicit std::move). User can still access `ptr' // after calling ptr->SetFailed() before release of `ptr'. // This function is wait-free. - // Returns 0 on success, -1 when the Socket was SetFailed(). + // Returns 0 on success, -1 when the object was SetFailed(). static int Address(VRefId id, VersionedRefWithIdUniquePtr* ptr); - // Returns 0 on success, 1 on failed socket, -1 on recycled. + // Returns 0 on success, 1 on failed object, -1 on recycled. static int AddressFailedAsWell(VRefId id, VersionedRefWithIdUniquePtr* ptr); // Re-address current VersionedRefWithId into `ptr'. - // Always succeed even if this socket is failed. + // Always succeed even if this object is failed. void ReAddress(VersionedRefWithIdUniquePtr* ptr); // Returns signed 32-bit referenced-count. @@ -239,12 +282,12 @@ class VersionedRefWithId { // Any later Address() of the identifier shall return NULL. The // VersionedRefWithId is NOT recycled after calling this function, // instead it will be recycled when no one references it. Internal - // fields of the Socket are still accessible after calling this + // fields of the object are still accessible after calling this // function. Calling SetFailed() of a VersionedRefWithId more than // once is OK. // T::OnFailed() will be called when SetFailed() successfully. // This function is lock-free. - // Returns -1 when the Socket was already SetFailed(), 0 otherwise. + // Returns -1 when the object was already SetFailed(), 0 otherwise. template static int SetFailedById(VRefId id, Args&&... args); @@ -301,7 +344,7 @@ friend void DereferenceVersionedRefWithId<>(T* r); _versioned_ref.fetch_add(1, butil::memory_order_release); } - // Make this socket addressable again. + // Make this object addressable again. // If nref is less than `at_least_nref', VersionedRefWithId was // abandoned during revival and cannot be revived. void Revive(int32_t at_least_nref); @@ -310,17 +353,21 @@ friend void DereferenceVersionedRefWithId<>(T* r); typedef butil::ResourceId resource_id_t; // 1. When `failed_as_well=true', returns 0 on success, - // 1 on failed socket, -1 on recycled. + // 1 on failed object, -1 on recycled. // 2. When `failed_as_well=true', returns 0 on success, - // -1 when the Socket was SetFailed(). + // -1 when the object was SetFailed(). static int AddressImpl(VRefId id, bool failed_as_well, VersionedRefWithIdUniquePtr* ptr); - // Callback wrapper of Derived classes. - WRAPPER_OF(T, OnFailed, void); - WRAPPER_OF(T, BeforeAdditionalRefReleased, void); - WRAPPER_OF(T, AfterRevived, void); - WRAPPER_OF(T, OnDescription, std::string); + // Detectors + static callers for optional Derived-class callbacks. + BRPC_DEFINE_MEMBER_DETECTOR(OnFailed); + BRPC_DEFINE_OPTIONAL_CALLER(OnFailed, void); + BRPC_DEFINE_MEMBER_DETECTOR(BeforeAdditionalRefReleased); + BRPC_DEFINE_OPTIONAL_CALLER(BeforeAdditionalRefReleased, void); + BRPC_DEFINE_MEMBER_DETECTOR(AfterRevived); + BRPC_DEFINE_OPTIONAL_CALLER(AfterRevived, void); + BRPC_DEFINE_MEMBER_DETECTOR(OnDescription); + BRPC_DEFINE_OPTIONAL_CALLER(OnDescription, std::string); // unsigned 32-bit version + signed 32-bit referenced-count. // Meaning of version: @@ -329,14 +376,14 @@ friend void DereferenceVersionedRefWithId<>(T* r); // of a VersionedRefWithId on the slot, the version is added with 1 twice. // This is also the version encoded in VRefId. // * Failed version: = created version + 1, SetFailed()-ed but returned. - // * Other versions: the socket is already recycled. + // * Other versions: the object is already recycled. butil::atomic BAIDU_CACHELINE_ALIGNMENT _versioned_ref; // The unique identifier. VRefId _this_id; // Indicates whether additional reference has increased, // decreased, or is increasing. // additional ref status: - // `Socket'、`Create': REF_USING + // constructor / `Create': REF_USING // `SetFailed': REF_USING -> REF_RECYCLED // `Revive' REF_RECYCLED -> REF_REVIVING -> REF_USING butil::atomic _additional_ref_status; @@ -444,7 +491,7 @@ int VersionedRefWithId::AddressImpl( // Addressed a free slot. } } else { - CHECK(false) << "Over dereferenced SocketId=" << id; + CHECK(false) << "Over dereferenced VRefId=" << id; } } return -1; @@ -488,7 +535,7 @@ int VersionedRefWithId::SetFailedImpl(Args&&... args) { butil::memory_order_release, butil::memory_order_relaxed)) { // Call T::OnFailed() to notify the failure of T. - WRAPPER_CALL(OnFailed, static_cast(this), std::forward(args)...); + CallOnFailed(static_cast(this), std::forward(args)...); // Deref additionally which is added at creation so that this // queue's reference will hit 0(recycle) when no one addresses it. ReleaseAdditionalReference(); @@ -507,7 +554,7 @@ int VersionedRefWithId::ReleaseAdditionalReference() { expect, ADDITIONAL_REF_RECYCLED, butil::memory_order_relaxed, butil::memory_order_relaxed)) { - WRAPPER_CALL(BeforeAdditionalRefReleased, static_cast(this)); + CallBeforeAdditionalRefReleased(static_cast(this)); return Dereference(); } @@ -529,7 +576,7 @@ int VersionedRefWithId::Dereference() { if (nref > 1) { return 0; } - if (__builtin_expect(nref == 1, 1)) { + if (BAIDU_LIKELY(nref == 1)) { const uint32_t ver = VersionOfVRef(vref); const uint32_t id_ver = VersionOfVRefId(id); // Besides first successful SetFailed() adds 1 to version, one of @@ -541,9 +588,9 @@ int VersionedRefWithId::Dereference() { // // Note: `ver == id_ver' means this VersionedRefWithId has been `SetRecycle' // before rather than `SetFailed'; `ver == ide_ver+1' means we - // had `SetFailed' this socket before. We should destroy the - // socket under both situation - if (__builtin_expect(ver == id_ver || ver == id_ver + 1, 1)) { + // had `SetFailed' this object before. We should destroy the + // object under both situation + if (BAIDU_LIKELY(ver == id_ver || ver == id_ver + 1)) { // sees nref:1->0, try to set version=id_ver+2,--nref. // No retry: if version changes, the slot is already returned by // another one who sees nref:1->0 concurrently; if nref changes, @@ -590,7 +637,7 @@ void VersionedRefWithId::Revive(int32_t at_least_nref) { int32_t nref = NRefOfVRef(vref); if (nref < at_least_nref) { - // Set the status to REF_RECYCLED since no one uses this socket + // Set the status to REF_RECYCLED since no one uses this object _additional_ref_status.store( ADDITIONAL_REF_RECYCLED, butil::memory_order_relaxed); CHECK_EQ(1, nref); @@ -606,7 +653,7 @@ void VersionedRefWithId::Revive(int32_t at_least_nref) { // Set the status to REF_USING since we add additional ref again _additional_ref_status.store( ADDITIONAL_REF_USING, butil::memory_order_relaxed); - WRAPPER_CALL(AfterRevived, static_cast(this)); + CallAfterRevived(static_cast(this)); return; } } @@ -617,8 +664,7 @@ std::string VersionedRefWithId::description() const { std::string result; result.reserve(128); butil::string_appendf(&result, "%s{id=%" PRIu64 " ", butil::class_name(), id()); - result.append(WRAPPER_CALL( - OnDescription, const_cast(static_cast(this)))); + result.append(CallOnDescription(const_cast(static_cast(this)))); butil::string_appendf(&result, "} (%p)", this); return result; } diff --git a/src/bthread/execution_queue_inl.h b/src/bthread/execution_queue_inl.h index ddf7bc6ba2..9c12e19256 100644 --- a/src/bthread/execution_queue_inl.h +++ b/src/bthread/execution_queue_inl.h @@ -348,11 +348,10 @@ inline ExecutionQueueOptions::ExecutionQueueOptions() {} template -inline int execution_queue_start( - ExecutionQueueId* id, - const ExecutionQueueOptions* options, - int (*execute)(void* meta, TaskIterator&), - void* meta) { +inline int execution_queue_start(ExecutionQueueId* id, + const ExecutionQueueOptions* options, + int (*execute)(void* meta, TaskIterator&), + void* meta) { return ExecutionQueue::create(id, options, execute, meta); } @@ -364,7 +363,7 @@ execution_queue_address(ExecutionQueueId id) { template inline int execution_queue_execute(ExecutionQueueId id, - typename butil::add_const_reference::type task) { + typename butil::add_const_reference::type task) { return execution_queue_execute(id, task, NULL); } @@ -377,9 +376,8 @@ inline int execution_queue_execute(ExecutionQueueId id, template inline int execution_queue_execute(ExecutionQueueId id, - typename butil::add_const_reference::type task, - const TaskOptions* options, - TaskHandle* handle) { + typename butil::add_const_reference::type task, + const TaskOptions* options, TaskHandle* handle) { typename ExecutionQueue::scoped_ptr_t ptr = ExecutionQueue::address(id); if (ptr != NULL) { @@ -390,21 +388,18 @@ inline int execution_queue_execute(ExecutionQueueId id, } template -inline int execution_queue_execute(ExecutionQueueId id, - T&& task) { +inline int execution_queue_execute(ExecutionQueueId id, T&& task) { return execution_queue_execute(id, std::forward(task), NULL); } template -inline int execution_queue_execute(ExecutionQueueId id, - T&& task, +inline int execution_queue_execute(ExecutionQueueId id, T&& task, const TaskOptions* options) { return execution_queue_execute(id, std::forward(task), options, NULL); } template -inline int execution_queue_execute(ExecutionQueueId id, - T&& task, +inline int execution_queue_execute(ExecutionQueueId id, T&& task, const TaskOptions* options, TaskHandle* handle) { typename ExecutionQueue::scoped_ptr_t diff --git a/test/brpc_streaming_rpc_unittest.cpp b/test/brpc_streaming_rpc_unittest.cpp index d6ad16949d..a759bae560 100644 --- a/test/brpc_streaming_rpc_unittest.cpp +++ b/test/brpc_streaming_rpc_unittest.cpp @@ -143,10 +143,14 @@ static void* SendTwoMessagesOnServerExtraStream(void* arg) { const int64_t connect_deadline_us = butil::gettimeofday_us() + 2 * 1000 * 1000L; bool connected = false; while (butil::gettimeofday_us() < connect_deadline_us) { - brpc::SocketUniquePtr ptr; - if (brpc::Socket::Address(sid, &ptr) == 0) { - brpc::Stream* s = static_cast(ptr->conn()); - if (s->_host_socket != NULL && s->_connected) { + brpc::StreamUniquePtr ptr; + if (brpc::Stream::Address(sid, &ptr) == 0) { + brpc::Stream* s = ptr.get(); + // SetConnected() publishes _connected only after _host_socket and + // the remote settings are ready. Check the acquire flag first + // before reading the non-atomic host pointer. + if (s->_connected.load(butil::memory_order_acquire) && + s->_host_socket != NULL) { connected = true; break; } @@ -286,20 +290,6 @@ TEST_F(StreamingRpcTest, batch_create_stream_feedback_race) { ASSERT_EQ(2u, request_streams.size()); state.client_extra_stream_id = request_streams[1]; - // Block SetConnected() on the extra stream to enlarge the race window. - brpc::SocketUniquePtr client_extra_ptr; - ASSERT_EQ(0, brpc::Socket::Address(state.client_extra_stream_id, &client_extra_ptr)); - brpc::Stream* client_extra_stream = static_cast(client_extra_ptr->conn()); - bthread_mutex_lock(&client_extra_stream->_connect_mutex); - struct UnlockGuard { - bthread_mutex_t* m; - ~UnlockGuard() { - if (m) { - bthread_mutex_unlock(m); - } - } - } unlock_guard{&client_extra_stream->_connect_mutex}; - BRPC_SCOPE_EXIT { if (state.server_extra_stream_id != brpc::INVALID_STREAM_ID) { brpc::StreamClose(state.server_extra_stream_id); @@ -317,13 +307,6 @@ TEST_F(StreamingRpcTest, batch_create_stream_feedback_race) { server.Stop(0); server.Join(); - // Release the SocketUniquePtr held above so the fake socket can be - // recycled. Otherwise BeforeRecycle / on_closed for the extra stream - // is deferred until `client_extra_ptr` destructs at scope exit, which - // happens *after* `client_handler` and `state` are destroyed -> UAF - // inside Stream::Consume on Linux. - client_extra_ptr.reset(); - // on_closed() runs asynchronously on each client stream's consumer // bthread. Wait for both before letting handler/state go out of // scope, otherwise Stream::Consume will dereference freed memory. @@ -338,13 +321,11 @@ TEST_F(StreamingRpcTest, batch_create_stream_feedback_race) { stub.Echo(&cntl, &request, &response, brpc::NewCallback(SetAtomicTrue, &state.rpc_done)); // Wait until client consumes the first 64B payload on extra stream. + // This increases the chance that Consume() runs before SetConnected() + // finishes on the extra stream, exercising the SetConnected()/Consume() + // ordering relevant to FEEDBACK sending via the atomic _local_consumed. ASSERT_TRUE(WaitForTrue(state.client_got_first_msg, 2000)); - // Unblock SetConnected(); the fix in PR 3215 should send the first FEEDBACK - // with consumed_size=64 here, making server-side stream writable again. - bthread_mutex_unlock(&client_extra_stream->_connect_mutex); - unlock_guard.m = NULL; - ASSERT_TRUE(WaitForTrue(state.rpc_done, 2000)); ASSERT_FALSE(cntl.Failed()) << cntl.ErrorText(); @@ -587,11 +568,13 @@ TEST_F(StreamingRpcTest, auto_close_if_host_socket_closed) { ASSERT_FALSE(cntl.Failed()) << cntl.ErrorText() << " request_stream=" << request_stream; { - brpc::SocketUniquePtr ptr; - ASSERT_EQ(0, brpc::Socket::Address(request_stream, &ptr)); - brpc::Stream* s = (brpc::Stream*)ptr->conn(); - ASSERT_TRUE(s->_host_socket != NULL); - s->_host_socket->SetFailed(); + brpc::StreamUniquePtr ptr; + ASSERT_EQ(0, brpc::Stream::Address(request_stream, &ptr)); + brpc::Stream* s = ptr.get(); + ASSERT_TRUE(s->_connected.load(butil::memory_order_acquire)); + brpc::Socket* host_socket = s->_host_socket; + ASSERT_TRUE(host_socket != NULL); + host_socket->SetFailed(); } usleep(100); @@ -638,9 +621,10 @@ TEST_F(StreamingRpcTest, failed_when_rst) { usleep(100); } { - brpc::SocketUniquePtr ptr; - ASSERT_EQ(0, brpc::Socket::Address(request_stream, &ptr)); - brpc::Stream* s = (brpc::Stream*)ptr->conn(); + brpc::StreamUniquePtr ptr; + ASSERT_EQ(0, brpc::Stream::Address(request_stream, &ptr)); + brpc::Stream* s = ptr.get(); + ASSERT_TRUE(s->_connected.load(butil::memory_order_acquire)); ASSERT_TRUE(s->_host_socket != NULL); brpc::policy::SendStreamRst(s->_host_socket, s->_remote_settings.stream_id()); @@ -866,11 +850,12 @@ TEST_F(StreamingRpcTest, segment_stream_data_automatically) { brpc::SocketUniquePtr host_socket_ptr; { - brpc::SocketUniquePtr ptr; - ASSERT_EQ(0, brpc::Socket::Address(request_stream, &ptr)); - brpc::Stream *s = (brpc::Stream *)ptr->conn(); - ASSERT_TRUE(s->_host_socket != NULL); - s->_host_socket->ReAddress(&host_socket_ptr); + brpc::StreamUniquePtr ptr; + ASSERT_EQ(0, brpc::Stream::Address(request_stream, &ptr)); + ASSERT_TRUE(ptr->_connected.load(butil::memory_order_acquire)); + brpc::Socket* host_socket = ptr->_host_socket; + ASSERT_TRUE(host_socket != NULL); + host_socket->ReAddress(&host_socket_ptr); } ASSERT_EQ(0, brpc::StreamClose(request_stream)); @@ -883,7 +868,12 @@ TEST_F(StreamingRpcTest, segment_stream_data_automatically) { host_socket_ptr->UpdateStatsEverySecond(now_ms); brpc::SocketStat stat; host_socket_ptr->GetStat(&stat); - ASSERT_LT(N * sizeof(N), stat.out_num_messages_m); + // A whole message (with all its segments) is now written to the host socket + // in a single wait-free Socket::Write, so the number of host-socket messages + // no longer reflects the number of stream frames. Heavy segmentation still + // shows up as extra on-wire bytes: each 1-byte segment carries a full STRM + // frame header + meta, so out_size_m is far larger than the raw payload. + ASSERT_LT(N * sizeof(N), stat.out_size_m); ASSERT_FALSE(handler.failed()); ASSERT_EQ(0, handler.idle_times()); ASSERT_EQ(N, handler._expected_next_value); @@ -1051,11 +1041,11 @@ TEST_F(StreamingRpcTest, batch_create_extra_stream) { for (size_t i = 0; i < request_streams.size(); ++i) { const brpc::StreamId sid = request_streams[i]; ASSERT_TRUE(WaitForTrue([sid]() { - brpc::SocketUniquePtr ptr; - if (brpc::Socket::Address(sid, &ptr) != 0) { + brpc::StreamUniquePtr ptr; + if (brpc::Stream::Address(sid, &ptr) != 0) { return false; } - brpc::Stream* s = static_cast(ptr->conn()); + brpc::Stream* s = ptr.get(); return s->_host_socket != NULL && s->_connected.load(butil::memory_order_acquire); }, 5000)) << "stream_index=" << i; @@ -1136,11 +1126,11 @@ TEST_F(StreamingRpcTest, batch_create_extra_stream_upstream_only) { for (size_t i = 0; i < request_streams.size(); ++i) { const brpc::StreamId sid = request_streams[i]; ASSERT_TRUE(WaitForTrue([sid]() { - brpc::SocketUniquePtr ptr; - if (brpc::Socket::Address(sid, &ptr) != 0) { + brpc::StreamUniquePtr ptr; + if (brpc::Stream::Address(sid, &ptr) != 0) { return false; } - brpc::Stream* s = static_cast(ptr->conn()); + brpc::Stream* s = ptr.get(); return s->_host_socket != NULL && s->_connected.load(butil::memory_order_acquire); }, 5000)) << "stream_index=" << i; @@ -1175,3 +1165,82 @@ TEST_F(StreamingRpcTest, batch_create_extra_stream_upstream_only) { server.Stop(0); server.Join(); } + +TEST_F(StreamingRpcTest, unconsumed_bytes_reclaimed_on_stream_close) { + GFLAGS_NAMESPACE::SetCommandLineOption( + "socket_max_streams_unconsumed_bytes", "10485760"); + BRPC_SCOPE_EXIT { + GFLAGS_NAMESPACE::SetCommandLineOption( + "socket_max_streams_unconsumed_bytes", "0"); + }; + + class BlockingHandler : public brpc::StreamInputHandler { + public: + BlockingHandler() : blocked(true) {} + + int on_received_messages(brpc::StreamId, + butil::IOBuf* const[], size_t) override { + while (blocked.load(std::memory_order_acquire)) { + usleep(100); + } + return 0; + } + void on_idle_timeout(brpc::StreamId) override {} + void on_closed(brpc::StreamId) override {} + + std::atomic blocked; + } handler; + + brpc::StreamOptions opt; + opt.handler = &handler; + opt.max_buf_size = 1024 * 1024; + + brpc::Server server; + MyServiceWithStream service(opt); + ASSERT_EQ(0, server.AddService(&service, brpc::SERVER_DOESNT_OWN_SERVICE)); + ASSERT_EQ(0, server.Start(9007, NULL)); + + brpc::Channel channel; + ASSERT_EQ(0, channel.Init("127.0.0.1:9007", NULL)); + + brpc::Controller cntl; + brpc::StreamId request_stream; + brpc::StreamOptions request_stream_options; + request_stream_options.max_buf_size = 1024 * 1024; + ASSERT_EQ(0, StreamCreate(&request_stream, cntl, &request_stream_options)); + brpc::ScopedStream stream_guard(request_stream); + + test::EchoService_Stub stub(&channel); + stub.Echo(&cntl, &request, &response, NULL); + ASSERT_FALSE(cntl.Failed()) << cntl.ErrorText(); + + brpc::SocketUniquePtr host_socket; + { + brpc::StreamUniquePtr ptr; + ASSERT_EQ(0, brpc::Stream::Address(request_stream, &ptr)); + ASSERT_TRUE(ptr->_connected.load(butil::memory_order_acquire)); + ASSERT_TRUE(ptr->_host_socket != NULL); + ptr->_host_socket->ReAddress(&host_socket); + } + int64_t baseline = host_socket->_total_streams_unconsumed_size.load( + butil::memory_order_relaxed); + + size_t write_size = 100 * 1024; + butil::IOBuf out; + out.append(std::string(write_size, 'x')); + ASSERT_EQ(0, brpc::StreamWrite(request_stream, out)); + ASSERT_TRUE(WaitForTrue([&]() { + return host_socket->_total_streams_unconsumed_size.load( + butil::memory_order_relaxed) >= baseline + static_cast(write_size); + }, 2000)); + + ASSERT_EQ(0, brpc::StreamClose(request_stream)); + ASSERT_TRUE(WaitForTrue([&]() { + return host_socket->_total_streams_unconsumed_size.load( + butil::memory_order_relaxed) == baseline; + }, 2000)); + + handler.blocked.store(false, std::memory_order_release); + server.Stop(0); + server.Join(); +}