This is an automated email from the ASF dual-hosted git repository.
wasphin pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/brpc.git
The following commit(s) were added to refs/heads/master by this push:
new 4a429cdc Reject invalid stream feedback frames (#3494)
4a429cdc is described below
commit 4a429cdc2153eff616902edfd2349823d6d7c5b5
Author: Xiaofeng Wang <[email protected]>
AuthorDate: Fri Aug 28 13:41:24 2026 +0800
Reject invalid stream feedback frames (#3494)
* Reject invalid stream feedback frames
Close only the affected stream when feedback carries a payload.
Apply the same behavior when flow control is disabled.
* Keep stream feedback window positive
Preserve a one-byte congestion window so valid feedback remains enabled.
* Wait for stream close in feedback test
Keep the handler and host socket alive until asynchronous cleanup finishes.
---
src/brpc/stream.cpp | 109 +++++++++++++++++------------
src/brpc/stream_impl.h | 2 +-
test/brpc_streaming_rpc_unittest.cpp | 128 +++++++++++++++++++++++++++++++++++
3 files changed, 196 insertions(+), 43 deletions(-)
diff --git a/src/brpc/stream.cpp b/src/brpc/stream.cpp
index 42cb0eab..ea4ef852 100644
--- a/src/brpc/stream.cpp
+++ b/src/brpc/stream.cpp
@@ -338,57 +338,71 @@ int Stream::AppendIfNotFull(const butil::IOBuf &data,
return 0;
}
-void Stream::SetRemoteConsumed(size_t new_remote_consumed) {
- CHECK(_cur_buf_size > 0);
+int Stream::SetRemoteConsumed(size_t new_remote_consumed) {
bthread_id_list_t tmplist;
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);
- return;
- }
- const bool was_full = _produced >= _remote_consumed + _cur_buf_size;
-
- if (FLAGS_socket_max_streams_unconsumed_bytes > 0 && _host_socket !=
nullptr) {
- 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;
+ BRPC_SCOPE_EXIT { bthread_id_list_destroy(&tmplist); };
+ {
+ BAIDU_SCOPED_LOCK(_congestion_control_mutex);
+ if (_cur_buf_size == 0) {
+ return -1;
}
- 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;
+ if (_remote_consumed >= new_remote_consumed) {
+ return 0;
+ }
+ const bool was_full = _produced >= _remote_consumed + _cur_buf_size;
+
+ if (FLAGS_socket_max_streams_unconsumed_bytes > 0 &&
+ _host_socket != nullptr) {
+ 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;
}
- 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 {
- _cur_buf_size *= 2;
+ 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 if (_cur_buf_size > 1) {
+ _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)) {
+ if (_options.max_buf_size > 0 &&
+ _cur_buf_size * 2 >
+ (size_t)_options.max_buf_size) {
+ _cur_buf_size = _options.max_buf_size;
+ } else {
+ _cur_buf_size *= 2;
+ }
}
}
- }
- _remote_consumed = new_remote_consumed;
- const bool is_full = _produced >= _remote_consumed + _cur_buf_size;
- if (was_full && !is_full) {
- bthread_id_list_swap(&tmplist, &_writable_wait_list);
+ _remote_consumed = new_remote_consumed;
+ const bool is_full =
+ _produced >= _remote_consumed + _cur_buf_size;
+ if (was_full && !is_full) {
+ bthread_id_list_swap(&tmplist, &_writable_wait_list);
+ }
}
- bthread_mutex_unlock(&_congestion_control_mutex);
// broadcast
bthread_id_list_reset(&tmplist, 0);
- bthread_id_list_destroy(&tmplist);
+ return 0;
}
void* Stream::RunOnWritable(void* arg) {
@@ -610,10 +624,21 @@ int Stream::OnReceived(const StreamFrameMeta& fm,
butil::IOBuf *buf, Socket* soc
switch (fm.frame_type()) {
case FRAME_TYPE_FEEDBACK:
+ if (!buf->empty()) {
+ LOG(WARNING) << "Close stream=" << id()
+ << " whose feedback frame has payload_size="
+ << buf->size();
+ Close(EPROTO, "Feedback frame must not contain a payload");
+ return -1;
+ }
if (_connected.load(butil::memory_order_acquire)) {
- SetRemoteConsumed(fm.feedback().consumed_size());
+ if (SetRemoteConsumed(fm.feedback().consumed_size()) != 0) {
+ LOG(WARNING) << "Close stream=" << id()
+ << " that received unexpected feedback";
+ Close(EPROTO, "Feedback is disabled for this stream");
+ return -1;
+ }
}
- CHECK(buf->empty());
break;
case FRAME_TYPE_DATA:
if (buf->length() > FLAGS_max_body_size ||
diff --git a/src/brpc/stream_impl.h b/src/brpc/stream_impl.h
index 4d3dd0c5..e64c14c0 100644
--- a/src/brpc/stream_impl.h
+++ b/src/brpc/stream_impl.h
@@ -88,7 +88,7 @@ friend class VersionedRefWithId<Stream>;
void BeforeRecycled();
std::string OnDescription() const;
- void SetRemoteConsumed(size_t _remote_consumed);
+ int SetRemoteConsumed(size_t _remote_consumed);
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);
diff --git a/test/brpc_streaming_rpc_unittest.cpp
b/test/brpc_streaming_rpc_unittest.cpp
index 7d1dc366..3e3d0f7c 100644
--- a/test/brpc_streaming_rpc_unittest.cpp
+++ b/test/brpc_streaming_rpc_unittest.cpp
@@ -393,6 +393,134 @@ public:
std::atomic<int> failure_code{0};
};
+class FeedbackValidationHandler : public brpc::StreamInputHandler {
+public:
+ int on_received_messages(brpc::StreamId,
+ butil::IOBuf* const[],
+ size_t) override {
+ return 0;
+ }
+
+ void on_idle_timeout(brpc::StreamId) override {}
+ void on_closed(brpc::StreamId) override {
+ closed.store(true, std::memory_order_release);
+ }
+ void on_failed(brpc::StreamId, int error_code,
+ const std::string&) override {
+ failure_code.store(error_code, std::memory_order_release);
+ }
+
+ std::atomic<int> failure_code{0};
+ std::atomic<bool> closed{false};
+};
+
+TEST_F(StreamingRpcTest, reject_malformed_feedback_frames) {
+ brpc::Server server;
+ MyServiceWithStream service;
+ ASSERT_EQ(0, server.AddService(
+ &service, brpc::SERVER_DOESNT_OWN_SERVICE));
+ ASSERT_EQ(0, server.Start(0, nullptr));
+
+ brpc::Channel channel;
+ ASSERT_EQ(0, channel.Init(server.listen_address(), nullptr));
+ test::EchoService_Stub stub(&channel);
+
+ auto check_feedback = [&](int max_buf_size, bool add_payload,
+ FeedbackValidationHandler* handler) {
+ brpc::Controller cntl;
+ brpc::StreamOptions options;
+ options.handler = handler;
+ options.min_buf_size = max_buf_size;
+ options.max_buf_size = max_buf_size;
+ brpc::StreamId request_stream;
+ ASSERT_EQ(0, brpc::StreamCreate(&request_stream, cntl, &options));
+ brpc::ScopedStream stream_guard(request_stream);
+ test::EchoResponse rpc_response;
+ stub.Echo(&cntl, &request, &rpc_response, nullptr);
+ ASSERT_FALSE(cntl.Failed()) << cntl.ErrorText();
+
+ brpc::StreamUniquePtr stream;
+ ASSERT_EQ(0, brpc::Stream::Address(request_stream, &stream));
+ ASSERT_TRUE(stream->_connected.load(butil::memory_order_acquire));
+ ASSERT_NE(nullptr, stream->_host_socket);
+
+ brpc::StreamFrameMeta feedback;
+ feedback.set_stream_id(request_stream);
+ feedback.set_frame_type(brpc::FRAME_TYPE_FEEDBACK);
+ feedback.mutable_feedback()->set_consumed_size(1);
+ butil::IOBuf payload;
+ if (add_payload) {
+ payload.append("unexpected payload");
+ }
+ ASSERT_EQ(-1, stream->OnReceived(
+ feedback, &payload, stream->_host_socket));
+ ASSERT_TRUE(WaitForTrue([handler]() {
+ return handler->failure_code.load(std::memory_order_acquire) != 0;
+ }, 2000));
+ ASSERT_EQ(EPROTO,
+ handler->failure_code.load(std::memory_order_relaxed));
+ };
+
+ FeedbackValidationHandler payload_handler;
+ check_feedback(1024, true, &payload_handler);
+
+ FeedbackValidationHandler disabled_handler;
+ check_feedback(0, false, &disabled_handler);
+
+ std::string old_socket_limit;
+ ASSERT_TRUE(GFLAGS_NAMESPACE::GetCommandLineOption(
+ "socket_max_streams_unconsumed_bytes", &old_socket_limit));
+ ASSERT_FALSE(GFLAGS_NAMESPACE::SetCommandLineOption(
+ "socket_max_streams_unconsumed_bytes", "1").empty());
+ BRPC_SCOPE_EXIT {
+ GFLAGS_NAMESPACE::SetCommandLineOption(
+ "socket_max_streams_unconsumed_bytes",
+ old_socket_limit.c_str());
+ };
+
+ FeedbackValidationHandler valid_handler;
+ brpc::Controller cntl;
+ brpc::StreamOptions options;
+ options.handler = &valid_handler;
+ options.min_buf_size = 0;
+ options.max_buf_size = 1;
+ brpc::StreamId request_stream;
+ ASSERT_EQ(0, brpc::StreamCreate(&request_stream, cntl, &options));
+ brpc::ScopedStream stream_guard(request_stream);
+ test::EchoResponse rpc_response;
+ stub.Echo(&cntl, &request, &rpc_response, nullptr);
+ ASSERT_FALSE(cntl.Failed()) << cntl.ErrorText();
+
+ brpc::StreamUniquePtr stream;
+ ASSERT_EQ(0, brpc::Stream::Address(request_stream, &stream));
+ ASSERT_TRUE(stream->_connected.load(butil::memory_order_acquire));
+ ASSERT_NE(nullptr, stream->_host_socket);
+ brpc::SocketUniquePtr host_socket;
+ stream->_host_socket->ReAddress(&host_socket);
+ const int64_t old_unconsumed =
+ host_socket->_total_streams_unconsumed_size.exchange(
+ 2, butil::memory_order_relaxed);
+ BRPC_SCOPE_EXIT {
+ host_socket->_total_streams_unconsumed_size.store(
+ old_unconsumed, butil::memory_order_relaxed);
+ };
+
+ brpc::StreamFrameMeta feedback;
+ feedback.set_stream_id(request_stream);
+ feedback.set_frame_type(brpc::FRAME_TYPE_FEEDBACK);
+ feedback.mutable_feedback()->set_consumed_size(1);
+ butil::IOBuf payload;
+ ASSERT_EQ(0, stream->OnReceived(
+ feedback, &payload, host_socket.get()));
+ ASSERT_EQ(1u, stream->_cur_buf_size);
+ ASSERT_EQ(0, valid_handler.failure_code.load(
+ std::memory_order_acquire));
+ stream.reset();
+ ASSERT_EQ(0, brpc::StreamClose(request_stream));
+ ASSERT_EQ(request_stream, stream_guard.release());
+ ASSERT_TRUE(WaitForTrue(valid_handler.closed, 2000));
+}
+
TEST_F(StreamingRpcTest, limit_reassembled_message_size) {
std::string old_max_body_size;
std::string old_segment_size;
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]