diff --git a/src/brpc/details/controller_private_accessor.h b/src/brpc/details/controller_private_accessor.h index 1aad5b2b4e..d752d7fe5e 100644 --- a/src/brpc/details/controller_private_accessor.h +++ b/src/brpc/details/controller_private_accessor.h @@ -58,6 +58,26 @@ class ControllerPrivateAccessor { return _cntl->_current_call.sending_sock.get(); } + bool does_response_match_request_socket( + CallId cid, SocketId response_socket_id) const { + const Controller::Call* call = nullptr; + if (cid == _cntl->current_id() || + (cid == _cntl->_correlation_id && + _cntl->_current_call.sending_sock == nullptr)) { + call = &_cntl->_current_call; + } else if (_cntl->_unfinished_call != nullptr && + cid == _cntl->get_id(_cntl->_unfinished_call->nretry)) { + call = _cntl->_unfinished_call; + } else { + return false; + } + // A regular client request always records the socket used to send it. + // Keep direct response injection available to internal tools and tests + // that do not issue a request through Controller::IssueRPC(). + return call->sending_sock == nullptr || + call->sending_sock->id() == response_socket_id; + } + int64_t real_timeout_ms() { return _cntl->_real_timeout_ms; } diff --git a/src/brpc/policy/baidu_rpc_protocol.cpp b/src/brpc/policy/baidu_rpc_protocol.cpp index baa691e029..54e611861c 100644 --- a/src/brpc/policy/baidu_rpc_protocol.cpp +++ b/src/brpc/policy/baidu_rpc_protocol.cpp @@ -968,7 +968,23 @@ void ProcessRpcResponse(InputMessageBase* msg_base) { } return; } - + + if (cntl == nullptr || !ControllerPrivateAccessor(cntl) + .does_response_match_request_socket(cid, msg->socket()->id())) { + LOG(WARNING) << "correlation_id=" << cid.value + << " of the response from " << *msg->socket() + << " does not match a request sent over it, drop it"; + CHECK_EQ(0, bthread_id_unlock(cid)); + if (remote_stream_id != INVALID_STREAM_ID) { + SendStreamRst(msg->socket(), remote_stream_id); + const auto& extra_stream_ids = meta.stream_settings().extra_stream_ids(); + for (int i = 0; i < extra_stream_ids.size(); ++i) { + SendStreamRst(msg->socket(), extra_stream_ids[i]); + } + } + return; + } + ControllerPrivateAccessor accessor(cntl); if (remote_stream_id != INVALID_STREAM_ID) { accessor.set_remote_stream_settings( diff --git a/src/brpc/policy/hulu_pbrpc_protocol.cpp b/src/brpc/policy/hulu_pbrpc_protocol.cpp index 146f88585a..1e2132fecb 100644 --- a/src/brpc/policy/hulu_pbrpc_protocol.cpp +++ b/src/brpc/policy/hulu_pbrpc_protocol.cpp @@ -610,7 +610,16 @@ void ProcessHuluResponse(InputMessageBase* msg_base) { << "Fail to lock correlation_id=" << cid << ": " << berror(rc); return; } - + + if (cntl == nullptr || !ControllerPrivateAccessor(cntl) + .does_response_match_request_socket(cid, msg->socket()->id())) { + LOG(WARNING) << "correlation_id=" << cid.value + << " of the response from " << *msg->socket() + << " does not match a request sent over it, drop it"; + CHECK_EQ(0, bthread_id_unlock(cid)); + return; + } + ControllerPrivateAccessor accessor(cntl); if (auto span = accessor.span()) { span->set_base_real_us(msg->base_real_us()); @@ -730,4 +739,3 @@ void PackHuluRequest(butil::IOBuf* req_buf, } // namespace policy } // namespace brpc - diff --git a/src/brpc/policy/public_pbrpc_protocol.cpp b/src/brpc/policy/public_pbrpc_protocol.cpp index 111a863e91..811e072d71 100644 --- a/src/brpc/policy/public_pbrpc_protocol.cpp +++ b/src/brpc/policy/public_pbrpc_protocol.cpp @@ -173,6 +173,15 @@ void ProcessPublicPbrpcResponse(InputMessageBase* msg_base) { return; } + if (cntl == nullptr || !ControllerPrivateAccessor(cntl) + .does_response_match_request_socket(cid, msg->socket()->id())) { + LOG(WARNING) << "correlation_id=" << cid.value + << " of the response from " << *msg->socket() + << " does not match a request sent over it, drop it"; + CHECK_EQ(0, bthread_id_unlock(cid)); + return; + } + ControllerPrivateAccessor accessor(cntl); if (auto span = accessor.span()) { span->set_base_real_us(msg->base_real_us()); @@ -284,4 +293,3 @@ void PackPublicPbrpcRequest(butil::IOBuf* buf, } // namespace policy } // namespace brpc - diff --git a/src/brpc/policy/sofa_pbrpc_protocol.cpp b/src/brpc/policy/sofa_pbrpc_protocol.cpp index fa51259050..3a8fdd7856 100644 --- a/src/brpc/policy/sofa_pbrpc_protocol.cpp +++ b/src/brpc/policy/sofa_pbrpc_protocol.cpp @@ -516,7 +516,16 @@ void ProcessSofaResponse(InputMessageBase* msg_base) { << "Fail to lock correlation_id=" << cid << ": " << berror(rc); return; } - + + if (cntl == nullptr || !ControllerPrivateAccessor(cntl) + .does_response_match_request_socket(cid, msg->socket()->id())) { + LOG(WARNING) << "correlation_id=" << cid.value + << " of the response from " << *msg->socket() + << " does not match a request sent over it, drop it"; + CHECK_EQ(0, bthread_id_unlock(cid)); + return; + } + ControllerPrivateAccessor accessor(cntl); if (auto span = accessor.span()) { span->set_base_real_us(msg->base_real_us()); diff --git a/test/brpc_channel_unittest.cpp b/test/brpc_channel_unittest.cpp index 1467939aa4..047d978c12 100644 --- a/test/brpc_channel_unittest.cpp +++ b/test/brpc_channel_unittest.cpp @@ -21,6 +21,7 @@ #include #include +#include #include #include #include @@ -28,12 +29,15 @@ #include "butil/macros.h" #include "butil/logging.h" #include "butil/files/temp_file.h" +#include "butil/fd_guard.h" #include "brpc/socket.h" #include "brpc/acceptor.h" #include "brpc/server.h" #include "brpc/policy/baidu_rpc_protocol.h" #include "brpc/policy/baidu_rpc_meta.pb.h" #include "brpc/policy/most_common_message.h" +#include "brpc/policy/public_pbrpc_protocol.h" +#include "brpc/policy/streaming_rpc_protocol.h" #include "brpc/channel.h" #include "brpc/details/load_balancer_with_naming.h" #include "brpc/parallel_channel.h" @@ -2282,6 +2286,140 @@ class MyShared : public brpc::SharedObject { int MyShared::nctor = 0; int MyShared::ndtor = 0; +TEST(ResponseSocketTest, baidu_response_requires_sending_socket) { + brpc::Controller cntl; + test::EchoResponse res; + cntl._response = &res; + ASSERT_EQ(0, bthread_id_lock_and_reset_range( + cntl.call_id(), nullptr, 2)); + ASSERT_EQ(0, bthread_id_unlock(cntl.current_id())); + brpc::SocketId sending_id; + brpc::SocketId foreign_id; + ASSERT_EQ(0, brpc::Socket::Create(brpc::SocketOptions(), &sending_id)); + int fds[2]; + ASSERT_EQ(0, socketpair(AF_UNIX, SOCK_STREAM, 0, fds)); + butil::fd_guard response_fd(fds[0]); + butil::fd_guard peer_fd(fds[1]); + brpc::SocketOptions foreign_options; + foreign_options.fd = response_fd; + ASSERT_EQ(0, brpc::Socket::Create(foreign_options, &foreign_id)); + response_fd.release(); + ASSERT_EQ(0, brpc::Socket::Address( + sending_id, &cntl._current_call.sending_sock)); + brpc::SocketUniquePtr foreign_socket; + ASSERT_EQ(0, brpc::Socket::Address(foreign_id, &foreign_socket)); + + brpc::policy::RpcMeta meta; + meta.set_correlation_id(cntl.current_id().value); + meta.mutable_response()->set_error_code(0); + auto make_response = [&meta](brpc::Socket* socket) { + auto* msg = brpc::policy::MostCommonMessage::Get(); + butil::IOBufAsZeroCopyOutputStream meta_stream(&msg->meta); + EXPECT_TRUE(meta.SerializeToZeroCopyStream(&meta_stream)); + test::EchoResponse response; + response.set_message("matched"); + butil::IOBufAsZeroCopyOutputStream payload_stream(&msg->payload); + EXPECT_TRUE(response.SerializeToZeroCopyStream(&payload_stream)); + socket->ReAddress(&msg->_socket); + socket->PostponeEOF(); + return msg; + }; + const int64_t stream_ids[] = {123, 456, 789}; + meta.mutable_stream_settings()->set_stream_id(stream_ids[0]); + meta.mutable_stream_settings()->add_extra_stream_ids(stream_ids[1]); + meta.mutable_stream_settings()->add_extra_stream_ids(stream_ids[2]); + brpc::policy::ProcessRpcResponse(make_response(foreign_socket.get())); + EXPECT_TRUE(res.message().empty()); + EXPECT_EQ(0, cntl.ErrorCode()); + EXPECT_FALSE(cntl.has_remote_stream()); + + // All advertised streams must be reset on the response's arrival socket. + butil::IOBuf expected; + for (int64_t stream_id : stream_ids) { + brpc::StreamFrameMeta frame_meta; + frame_meta.set_stream_id(stream_id); + frame_meta.set_frame_type(brpc::FRAME_TYPE_RST); + brpc::policy::PackStreamMessage(&expected, frame_meta, nullptr); + } + butil::IOBuf received; + const int64_t deadline = butil::cpuwide_time_ms() + 5000; + while (received.size() < expected.size()) { + const int64_t remaining = deadline - butil::cpuwide_time_ms(); + ASSERT_GT(remaining, 0); + pollfd pfd = {peer_fd, POLLIN, 0}; + ASSERT_EQ(1, poll(&pfd, 1, remaining)); + char buf[1024]; + const ssize_t n = read(peer_fd, buf, sizeof(buf)); + ASSERT_GT(n, 0); + received.append(buf, n); + } + EXPECT_EQ(expected, received); + meta.clear_stream_settings(); + + // A real request uses a versioned ID, not the timeout/cancel base ID. + meta.set_correlation_id(cntl.call_id().value); + brpc::policy::ProcessRpcResponse( + make_response(cntl._current_call.sending_sock.get())); + EXPECT_TRUE(res.message().empty()); + EXPECT_EQ(0, cntl.ErrorCode()); + meta.set_correlation_id(cntl.current_id().value); + brpc::policy::ProcessRpcResponse( + make_response(cntl._current_call.sending_sock.get())); + EXPECT_EQ("matched", res.message()); + EXPECT_EQ(0, cntl.ErrorCode()); +} + +TEST(ResponseSocketTest, real_rpc_responses_match_sending_socket) { + class EchoService : public test::EchoService { + void Echo(google::protobuf::RpcController*, + const test::EchoRequest* request, + test::EchoResponse* response, + google::protobuf::Closure* done) override { + brpc::ClosureGuard done_guard(done); + response->set_message("received " + request->message()); + } + } service; + brpc::Server server; + ASSERT_EQ(0, server.AddService(&service, brpc::SERVER_DOESNT_OWN_SERVICE)); + brpc::ServerOptions server_options; + server_options.nshead_service = new brpc::policy::PublicPbrpcServiceAdaptor; + ASSERT_EQ(0, server.Start(0, &server_options)); + + const char* protocols[] = { + "baidu_std", "hulu_pbrpc", "sofa_pbrpc", "public_pbrpc"}; + const char* connections[] = {"single", "pooled", "short"}; + for (const char* protocol : protocols) { + for (const char* connection : connections) { + // Public pbrpc uses the half-duplex nshead server adaptor. + if (strcmp(protocol, "public_pbrpc") == 0 && + strcmp(connection, "single") == 0) { + continue; + } + SCOPED_TRACE(protocol); + SCOPED_TRACE(connection); + brpc::ChannelOptions options; + options.protocol = protocol; + options.connection_type = connection; + options.timeout_ms = 5000; + options.max_retry = 0; + brpc::Channel channel; + ASSERT_EQ(0, channel.Init(server.listen_address(), &options)); + test::EchoService_Stub stub(&channel); + for (int i = 0; i < 2; ++i) { + brpc::Controller cntl; + test::EchoRequest request; + test::EchoResponse response; + request.set_message("socket binding"); + stub.Echo(&cntl, &request, &response, nullptr); + ASSERT_FALSE(cntl.Failed()) << cntl.ErrorText(); + EXPECT_EQ("received socket binding", response.message()); + } + } + } + server.Stop(0); + server.Join(); +} + TEST_F(ChannelTest, intrusive_ptr_sanity) { MyShared::nctor = 0; MyShared::ndtor = 0; diff --git a/test/brpc_hulu_pbrpc_protocol_unittest.cpp b/test/brpc_hulu_pbrpc_protocol_unittest.cpp index e33d860e72..cec42cca9f 100644 --- a/test/brpc_hulu_pbrpc_protocol_unittest.cpp +++ b/test/brpc_hulu_pbrpc_protocol_unittest.cpp @@ -265,6 +265,114 @@ TEST_F(HuluTest, process_response_after_eof) { ASSERT_TRUE(_socket->Failed()); } +TEST_F(HuluTest, process_response_from_foreign_socket_is_dropped) { + brpc::policy::HuluRpcResponseMeta meta; + test::EchoResponse res; + brpc::Controller cntl; + cntl._response = &res; + ASSERT_EQ(0, bthread_id_lock_and_reset_range( + cntl.call_id(), nullptr, 2)); + ASSERT_EQ(0, bthread_id_unlock(cntl.current_id())); + meta.set_correlation_id(cntl.current_id().value); + + brpc::SocketId sending_id; + brpc::SocketOptions sending_options; + ASSERT_EQ(0, brpc::Socket::Create(sending_options, &sending_id)); + ASSERT_EQ(0, brpc::Socket::Address( + sending_id, &cntl._current_call.sending_sock)); + + brpc::policy::MostCommonMessage* msg = MakeResponseMessage(meta); + ProcessMessage(brpc::policy::ProcessHuluResponse, msg, false); + + EXPECT_TRUE(res.message().empty()); + EXPECT_EQ(0, cntl.ErrorCode()); + + // Dropping the response must leave the call available for its real peer. + msg = MakeResponseMessage(meta); + cntl._current_call.sending_sock->ReAddress(&msg->_socket); + ProcessMessage(brpc::policy::ProcessHuluResponse, msg, false); + EXPECT_EQ(EXP_RESPONSE, res.message()); + EXPECT_EQ(0, cntl.ErrorCode()); +} + +TEST_F(HuluTest, process_response_from_sending_socket_is_accepted) { + brpc::policy::HuluRpcResponseMeta meta; + test::EchoResponse res; + brpc::Controller cntl; + cntl._response = &res; + ASSERT_EQ(0, bthread_id_lock_and_reset_range( + cntl.call_id(), nullptr, 2)); + ASSERT_EQ(0, bthread_id_unlock(cntl.current_id())); + meta.set_correlation_id(cntl.current_id().value); + _socket->ReAddress(&cntl._current_call.sending_sock); + + brpc::policy::MostCommonMessage* msg = MakeResponseMessage(meta); + ProcessMessage(brpc::policy::ProcessHuluResponse, msg, false); + + EXPECT_EQ(EXP_RESPONSE, res.message()); + EXPECT_EQ(0, cntl.ErrorCode()); +} + +TEST_F(HuluTest, process_response_from_previous_retry_is_dropped) { + brpc::Controller cntl; + test::EchoResponse res; + cntl._response = &res; + ASSERT_EQ(0, bthread_id_lock_and_reset_range( + cntl.call_id(), nullptr, 3)); + const brpc::CallId previous_id = cntl.current_id(); + cntl._current_call.nretry = 1; + ASSERT_EQ(0, bthread_id_unlock(cntl.current_id())); + _socket->ReAddress(&cntl._current_call.sending_sock); + + // Even on the right socket, a completed attempt cannot supply the result. + brpc::policy::HuluRpcResponseMeta meta; + meta.set_correlation_id(previous_id.value); + ProcessMessage(brpc::policy::ProcessHuluResponse, + MakeResponseMessage(meta), false); + EXPECT_TRUE(res.message().empty()); + EXPECT_EQ(0, cntl.ErrorCode()); + + meta.set_correlation_id(cntl.current_id().value); + ProcessMessage(brpc::policy::ProcessHuluResponse, + MakeResponseMessage(meta), false); + EXPECT_EQ(EXP_RESPONSE, res.message()); + EXPECT_EQ(0, cntl.ErrorCode()); +} + +TEST_F(HuluTest, process_response_matches_unfinished_backup_socket) { + brpc::Controller cntl; + test::EchoResponse res; + cntl._response = &res; + ASSERT_EQ(0, bthread_id_lock_and_reset_range( + cntl.call_id(), nullptr, 3)); + ASSERT_EQ(0, bthread_id_unlock(cntl.current_id())); + const brpc::CallId previous_id = cntl.current_id(); + + brpc::SocketId previous_socket_id; + ASSERT_EQ(0, brpc::Socket::Create( + brpc::SocketOptions(), &previous_socket_id)); + ASSERT_EQ(0, brpc::Socket::Address( + previous_socket_id, &cntl._current_call.sending_sock)); + // Mirror starting a backup: preserve the first attempt and its socket. + cntl._unfinished_call = new brpc::Controller::Call(&cntl._current_call); + ++cntl._current_call.nretry; + _socket->ReAddress(&cntl._current_call.sending_sock); + + brpc::policy::HuluRpcResponseMeta meta; + meta.set_correlation_id(previous_id.value); + ProcessMessage(brpc::policy::ProcessHuluResponse, + MakeResponseMessage(meta), false); + EXPECT_TRUE(res.message().empty()); + EXPECT_EQ(0, cntl.ErrorCode()); + + // The preserved attempt may still finish, but only on its own socket. + brpc::policy::MostCommonMessage* msg = MakeResponseMessage(meta); + cntl._unfinished_call->sending_sock->ReAddress(&msg->_socket); + ProcessMessage(brpc::policy::ProcessHuluResponse, msg, false); + EXPECT_EQ(EXP_RESPONSE, res.message()); + EXPECT_EQ(0, cntl.ErrorCode()); +} + TEST_F(HuluTest, process_response_error_code) { const int ERROR_CODE = 12345; brpc::policy::HuluRpcResponseMeta meta; diff --git a/test/brpc_public_pbrpc_protocol_unittest.cpp b/test/brpc_public_pbrpc_protocol_unittest.cpp index f89ae4a6dd..31ce7294dc 100644 --- a/test/brpc_public_pbrpc_protocol_unittest.cpp +++ b/test/brpc_public_pbrpc_protocol_unittest.cpp @@ -250,6 +250,32 @@ TEST_F(PublicPbrpcTest, process_response_after_eof) { ASSERT_TRUE(_socket->Failed()); } +TEST_F(PublicPbrpcTest, process_response_requires_sending_socket) { + brpc::Controller cntl; + test::EchoResponse res; + cntl._response = &res; + ASSERT_EQ(0, bthread_id_lock_and_reset_range( + cntl.call_id(), nullptr, 2)); + ASSERT_EQ(0, bthread_id_unlock(cntl.current_id())); + brpc::SocketId sending_id; + ASSERT_EQ(0, brpc::Socket::Create(brpc::SocketOptions(), &sending_id)); + ASSERT_EQ(0, brpc::Socket::Address( + sending_id, &cntl._current_call.sending_sock)); + + brpc::policy::PublicPbrpcResponse meta; + meta.add_responsebody()->set_id(cntl.current_id().value); + meta.mutable_responsehead()->set_code(0); + ProcessMessage(brpc::policy::ProcessPublicPbrpcResponse, MakeResponseMessage(&meta), false); + EXPECT_TRUE(res.message().empty()); + EXPECT_EQ(0, cntl.ErrorCode()); + + brpc::policy::MostCommonMessage* msg = MakeResponseMessage(&meta); + cntl._current_call.sending_sock->ReAddress(&msg->_socket); + ProcessMessage(brpc::policy::ProcessPublicPbrpcResponse, msg, false); + EXPECT_EQ(EXP_RESPONSE, res.message()); + EXPECT_EQ(0, cntl.ErrorCode()); +} + TEST_F(PublicPbrpcTest, process_response_error_code) { const int ERROR_CODE = 12345; brpc::policy::PublicPbrpcResponse meta; diff --git a/test/brpc_sofa_pbrpc_protocol_unittest.cpp b/test/brpc_sofa_pbrpc_protocol_unittest.cpp index e03575e59c..f3cb422b65 100644 --- a/test/brpc_sofa_pbrpc_protocol_unittest.cpp +++ b/test/brpc_sofa_pbrpc_protocol_unittest.cpp @@ -296,6 +296,32 @@ TEST_F(SofaTest, reject_huge_meta_size) { ASSERT_EQ(brpc::PARSE_ERROR_TOO_BIG_DATA, pr.error()); } +TEST_F(SofaTest, process_response_requires_sending_socket) { + brpc::Controller cntl; + test::EchoResponse res; + cntl._response = &res; + ASSERT_EQ(0, bthread_id_lock_and_reset_range( + cntl.call_id(), nullptr, 2)); + ASSERT_EQ(0, bthread_id_unlock(cntl.current_id())); + brpc::SocketId sending_id; + ASSERT_EQ(0, brpc::Socket::Create(brpc::SocketOptions(), &sending_id)); + ASSERT_EQ(0, brpc::Socket::Address( + sending_id, &cntl._current_call.sending_sock)); + + brpc::policy::SofaRpcMeta meta; + meta.set_type(brpc::policy::SofaRpcMeta::RESPONSE); + meta.set_sequence_id(cntl.current_id().value); + ProcessMessage(brpc::policy::ProcessSofaResponse, MakeResponseMessage(meta), false); + EXPECT_TRUE(res.message().empty()); + EXPECT_EQ(0, cntl.ErrorCode()); + + brpc::policy::MostCommonMessage* msg = MakeResponseMessage(meta); + cntl._current_call.sending_sock->ReAddress(&msg->_socket); + ProcessMessage(brpc::policy::ProcessSofaResponse, msg, false); + EXPECT_EQ(EXP_RESPONSE, res.message()); + EXPECT_EQ(0, cntl.ErrorCode()); +} + TEST_F(SofaTest, process_response_error_code) { const int ERROR_CODE = 12345; brpc::policy::SofaRpcMeta meta;