Skip to content

Commit eb31fa5

Browse files
authored
Make _state atomic to prevent concurrent _read_buf mutation on TCP fallback (#3347)
RdmaEndpoint::_state was a plain enum, written by the handshake bthread and read concurrently by the event-dispatching thread (OnNewDataFromTcp). This is a data race, and on a weak memory model it can let the two threads concurrently mutate _socket->_read_buf. Make _state a butil::atomic<State>: - Terminal-state stores use release and the matching loads use acquire, so data published before a terminal state (the magic bytes put back into _read_buf, and the RDMA window/resource setup before ESTABLISHED) is visible to the reader. - Non-terminal handshake transitions use relaxed.
1 parent 4dc6ac8 commit eb31fa5

2 files changed

Lines changed: 53 additions & 46 deletions

File tree

‎src/brpc/rdma/rdma_endpoint.cpp‎

Lines changed: 51 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -161,7 +161,7 @@ RdmaEndpoint::~RdmaEndpoint() {
161161
void RdmaEndpoint::Reset() {
162162
DeallocateResources();
163163

164-
_state = UNINIT;
164+
_state.store(UNINIT, butil::memory_order_relaxed);
165165
_resource = NULL;
166166
_send_cq_events = 0;
167167
_recv_cq_events = 0;
@@ -195,7 +195,8 @@ void RdmaConnect::StartConnect(const Socket* socket,
195195
return;
196196
}
197197
if (!IsRdmaAvailable()) {
198-
rdma_transport->_rdma_ep->_state = RdmaEndpoint::FALLBACK_TCP;
198+
rdma_transport->_rdma_ep->_state.store(RdmaEndpoint::FALLBACK_TCP,
199+
butil::memory_order_relaxed);
199200
rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF;
200201
done(0, data);
201202
return;
@@ -206,7 +207,8 @@ void RdmaConnect::StartConnect(const Socket* socket,
206207
bthread_attr_t attr = BTHREAD_ATTR_NORMAL;
207208
bthread_attr_set_name(&attr, "RdmaProcessHandshakeAtClient");
208209
if (bthread_start_background(&tid, &attr,
209-
RdmaEndpoint::ProcessHandshakeAtClient, rdma_transport->_rdma_ep) < 0) {
210+
RdmaEndpoint::ProcessHandshakeAtClient,
211+
rdma_transport->_rdma_ep) < 0) {
210212
LOG(FATAL) << "Fail to start handshake bthread";
211213
Run();
212214
} else {
@@ -230,7 +232,7 @@ static void TryReadOnTcpDuringRdmaEst(Socket* s) {
230232
const int saved_errno = errno;
231233
PLOG(WARNING) << "Fail to read from " << s;
232234
s->SetFailed(saved_errno, "Fail to read from %s: %s",
233-
s->description().c_str(), berror(saved_errno));
235+
s->description().c_str(), berror(saved_errno));
234236
return;
235237
}
236238
if (!s->MoreReadEvents(&progress)) {
@@ -255,22 +257,22 @@ void RdmaEndpoint::OnNewDataFromTcp(Socket* m) {
255257

256258
int progress = Socket::PROGRESS_INIT;
257259
while (true) {
258-
if (ep->_state == UNINIT) {
260+
State state = ep->_state.load(butil::memory_order_acquire);
261+
if (state == UNINIT) {
259262
if (!m->CreatedByConnect()) {
260263
if (!IsRdmaAvailable()) {
261-
ep->_state = FALLBACK_TCP;
262264
rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF;
265+
ep->_state.store(FALLBACK_TCP, butil::memory_order_relaxed);
263266
continue;
264267
}
265268
bthread_t tid;
266-
ep->_state = S_HELLO_WAIT;
269+
ep->_state.store(S_HELLO_WAIT, butil::memory_order_relaxed);
267270
SocketUniquePtr s;
268271
m->ReAddress(&s);
269272
bthread_attr_t attr = BTHREAD_ATTR_NORMAL;
270273
bthread_attr_set_name(&attr, "RdmaProcessHandshakeAtServer");
271-
if (bthread_start_background(&tid, &attr,
272-
ProcessHandshakeAtServer, ep) < 0) {
273-
ep->_state = UNINIT;
274+
if (bthread_start_background(&tid, &attr, ProcessHandshakeAtServer, ep) < 0) {
275+
ep->_state.store(UNINIT, butil::memory_order_relaxed);
274276
LOG(FATAL) << "Fail to start handshake bthread";
275277
} else {
276278
s.release();
@@ -280,13 +282,13 @@ void RdmaEndpoint::OnNewDataFromTcp(Socket* m) {
280282
// starts handshake. This will be handled by client handshake.
281283
// Ignore the exception here.
282284
}
283-
} else if (ep->_state < ESTABLISHED) { // during handshake
285+
} else if (state < ESTABLISHED) { // during handshake
284286
ep->_read_butex->fetch_add(1, butil::memory_order_release);
285287
bthread::butex_wake(ep->_read_butex);
286-
} else if (ep->_state == FALLBACK_TCP){ // handshake finishes
288+
} else if (state == FALLBACK_TCP){ // handshake finishes
287289
InputMessenger::OnNewMessages(m);
288290
return;
289-
} else if (ep->_state == ESTABLISHED) {
291+
} else if (state == ESTABLISHED) {
290292
TryReadOnTcpDuringRdmaEst(ep->_socket);
291293
return;
292294
}
@@ -422,9 +424,10 @@ int RdmaEndpoint::WriteToFd(butil::IOBuf* data) {
422424

423425
inline void RdmaEndpoint::TryReadOnTcp() {
424426
if (_socket->_nevent.fetch_add(1, butil::memory_order_acq_rel) == 0) {
425-
if (_state == FALLBACK_TCP) {
427+
State state = _state.load(butil::memory_order_acquire);
428+
if (state == FALLBACK_TCP) {
426429
InputMessenger::OnNewMessages(_socket);
427-
} else if (_state == ESTABLISHED) {
430+
} else if (state == ESTABLISHED) {
428431
TryReadOnTcpDuringRdmaEst(_socket);
429432
}
430433
}
@@ -475,28 +478,28 @@ void* RdmaEndpoint::ProcessHandshakeAtClient(void* arg) {
475478
ep->_handshake_version = handshake->ProtocolVersion();
476479

477480
// First initialize CQ and QP resources.
478-
ep->_state = C_ALLOC_QPCQ;
481+
ep->_state.store(C_ALLOC_QPCQ, butil::memory_order_relaxed);
479482
if (ep->AllocateResources() < 0) {
480483
LOG(WARNING) << "Fallback to tcp:" << s->description();
481484
rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF;
482-
ep->_state = FALLBACK_TCP;
485+
ep->_state.store(FALLBACK_TCP, butil::memory_order_release);
483486
return NULL;
484487
}
485488

486489
// Send hello message to server
487-
ep->_state = C_HELLO_SEND;
490+
ep->_state.store(C_HELLO_SEND, butil::memory_order_relaxed);
488491
if (handshake->SendLocalHello() < 0) {
489492
int saved_errno = errno;
490493
PLOG(WARNING) << "Fail to send hello message to server:"
491494
<< s->description();
492495
s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s",
493496
s->description().c_str(), berror(saved_errno));
494-
ep->_state = FAILED;
497+
ep->_state.store(FAILED, butil::memory_order_relaxed);
495498
return NULL;
496499
}
497500

498501
// Receive and parse remote hello.
499-
ep->_state = C_HELLO_WAIT;
502+
ep->_state.store(C_HELLO_WAIT, butil::memory_order_relaxed);
500503
ParsedHello remote{};
501504
bool negotiated = false;
502505
if (handshake->ReceiveAndParseRemoteHello(&remote, &negotiated) < 0) {
@@ -505,7 +508,7 @@ void* RdmaEndpoint::ProcessHandshakeAtClient(void* arg) {
505508
<< s->description();
506509
s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s",
507510
s->description().c_str(), berror(saved_errno));
508-
ep->_state = FAILED;
511+
ep->_state.store(FAILED, butil::memory_order_relaxed);
509512
return NULL;
510513
}
511514

@@ -515,7 +518,7 @@ void* RdmaEndpoint::ProcessHandshakeAtClient(void* arg) {
515518
rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF;
516519
} else {
517520
ep->ApplyRemoteHello(remote);
518-
ep->_state = C_BRINGUP_QP;
521+
ep->_state.store(C_BRINGUP_QP, butil::memory_order_relaxed);
519522
if (ep->BringUpQp(remote.lid, remote.gid, remote.qp_num) < 0) {
520523
LOG(WARNING) << "Fail to bringup QP, fallback to tcp:"
521524
<< s->description();
@@ -526,26 +529,27 @@ void* RdmaEndpoint::ProcessHandshakeAtClient(void* arg) {
526529
}
527530

528531
// Send ACK message to server
529-
ep->_state = C_ACK_SEND;
530-
uint32_t flags = rdma_transport->_rdma_state != RdmaTransport::RDMA_OFF ? HELLO_ACK_RDMA_OK : 0;
532+
ep->_state.store(C_ACK_SEND, butil::memory_order_relaxed);
533+
bool rdma_on = rdma_transport->_rdma_state == RdmaTransport::RDMA_ON;
534+
uint32_t flags = rdma_on ? HELLO_ACK_RDMA_OK : 0;
531535
uint32_t flags_be = butil::HostToNet32(flags);
532536
if (ep->WriteToFd(&flags_be, HELLO_ACK_LEN) < 0) {
533537
int saved_errno = errno;
534538
PLOG(WARNING) << "Fail to send Ack Message to server:"
535539
<< s->description();
536540
s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s",
537541
s->description().c_str(), berror(saved_errno));
538-
ep->_state = FAILED;
542+
ep->_state.store(FAILED, butil::memory_order_relaxed);
539543
return NULL;
540544
}
541545

542546
if (rdma_transport->_rdma_state == RdmaTransport::RDMA_ON) {
543-
ep->_state = ESTABLISHED;
547+
ep->_state.store(ESTABLISHED, butil::memory_order_release);
544548
LOG_IF(INFO, FLAGS_rdma_trace_verbose)
545549
<< "Client handshake ends (use rdma v" << ep->_handshake_version
546550
<< ") on " << s->description();
547551
} else {
548-
ep->_state = FALLBACK_TCP;
552+
ep->_state.store(FALLBACK_TCP, butil::memory_order_release);
549553
LOG_IF(INFO, FLAGS_rdma_trace_verbose)
550554
<< "Client handshake ends (use tcp) on " << s->description();
551555
}
@@ -578,15 +582,15 @@ void* RdmaEndpoint::ProcessHandshakeAtServer(void* arg) {
578582
LOG_IF(INFO, FLAGS_rdma_trace_verbose)
579583
<< "Start handshake on " << s->description();
580584

581-
ep->_state = S_HELLO_WAIT;
585+
ep->_state.store(S_HELLO_WAIT, butil::memory_order_relaxed);
582586
uint8_t magic[MAGIC_STR_LEN];
583587
if (ep->ReadFromFd(magic, MAGIC_STR_LEN) < 0) {
584588
int saved_errno = errno;
585589
PLOG(WARNING) << "Fail to read Hello Message from client:"
586590
<< s->description() << " " << s->_remote_side;
587591
s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s",
588592
s->description().c_str(), berror(saved_errno));
589-
ep->_state = FAILED;
593+
ep->_state.store(FAILED, butil::memory_order_relaxed);
590594
return NULL;
591595
}
592596

@@ -598,8 +602,11 @@ void* RdmaEndpoint::ProcessHandshakeAtServer(void* arg) {
598602
<< s->description();
599603
// We need to copy data read back to _socket->_read_buf.
600604
s->_read_buf.append(magic, MAGIC_STR_LEN);
601-
ep->_state = FALLBACK_TCP;
602605
rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF;
606+
// Use release memory order to publish the magic bytes appended
607+
// above to whoever reads `_state == FALLBACK_TCP` (the event
608+
// thread in OnNewDataFromTcp).
609+
ep->_state.store(FALLBACK_TCP, butil::memory_order_release);
603610
ep->TryReadOnTcp();
604611
return NULL;
605612
}
@@ -614,7 +621,7 @@ void* RdmaEndpoint::ProcessHandshakeAtServer(void* arg) {
614621
<< s->description();
615622
s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s",
616623
s->description().c_str(), berror(saved_errno));
617-
ep->_state = FAILED;
624+
ep->_state.store(FAILED, butil::memory_order_relaxed);
618625
return NULL;
619626
}
620627

@@ -624,13 +631,13 @@ void* RdmaEndpoint::ProcessHandshakeAtServer(void* arg) {
624631
rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF;
625632
} else {
626633
ep->ApplyRemoteHello(remote);
627-
ep->_state = S_ALLOC_QPCQ;
634+
ep->_state.store(S_ALLOC_QPCQ, butil::memory_order_relaxed);
628635
if (ep->AllocateResources() < 0) {
629636
LOG(WARNING) << "Fail to allocate rdma resources, fallback to tcp:"
630637
<< s->description();
631638
rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF;
632639
} else {
633-
ep->_state = S_BRINGUP_QP;
640+
ep->_state.store(S_BRINGUP_QP, butil::memory_order_relaxed);
634641
if (ep->BringUpQp(remote.lid, remote.gid, remote.qp_num) < 0) {
635642
LOG(WARNING) << "Fail to bringup QP, fallback to tcp:"
636643
<< s->description();
@@ -639,31 +646,30 @@ void* RdmaEndpoint::ProcessHandshakeAtServer(void* arg) {
639646
}
640647
}
641648

642-
ep->_state = S_HELLO_SEND;
649+
ep->_state.store(S_HELLO_SEND, butil::memory_order_relaxed);
643650
if (handshake->SendLocalHello() < 0) {
644651
int saved_errno = errno;
645652
PLOG(WARNING) << "Fail to send Hello Message to client:"
646653
<< s->description();
647654
s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s",
648655
s->description().c_str(), berror(saved_errno));
649-
ep->_state = FAILED;
656+
ep->_state.store(FAILED, butil::memory_order_relaxed);
650657
return NULL;
651658
}
652659

653-
ep->_state = S_ACK_WAIT;
660+
ep->_state.store(S_ACK_WAIT, butil::memory_order_relaxed);
654661
uint32_t flags_be = 0;
655662
if (ep->ReadFromFd(&flags_be, HELLO_ACK_LEN) < 0) {
656663
int saved_errno = errno;
657664
PLOG(WARNING) << "Fail to read ack message from client:"
658665
<< s->description();
659666
s->SetFailed(saved_errno, "Fail to complete rdma handshake from %s: %s",
660667
s->description().c_str(), berror(saved_errno));
661-
ep->_state = FAILED;
668+
ep->_state.store(FAILED, butil::memory_order_relaxed);
662669
return NULL;
663670
}
664671
uint32_t flags = butil::NetToHost32(flags_be);
665672
bool client_ack_ok = (flags & HELLO_ACK_RDMA_OK) != 0;
666-
667673
if (client_ack_ok) {
668674
if (rdma_transport->_rdma_state == RdmaTransport::RDMA_OFF) {
669675
// Client asked for RDMA but we are falling back: protocol
@@ -673,17 +679,17 @@ void* RdmaEndpoint::ProcessHandshakeAtServer(void* arg) {
673679
<< "RDMA_OFF state: " << s->description();
674680
s->SetFailed(EPROTO, "Fail to complete rdma handshake from %s: %s",
675681
s->description().c_str(), berror(EPROTO));
676-
ep->_state = FAILED;
682+
ep->_state.store(FAILED, butil::memory_order_relaxed);
677683
return NULL;
678684
}
679685
rdma_transport->_rdma_state = RdmaTransport::RDMA_ON;
680-
ep->_state = ESTABLISHED;
686+
ep->_state.store(ESTABLISHED, butil::memory_order_release);
681687
LOG_IF(INFO, FLAGS_rdma_trace_verbose)
682688
<< "Server handshake ends (use rdma v" << ep->_handshake_version
683689
<< ") on " << s->description();
684690
} else {
685691
rdma_transport->_rdma_state = RdmaTransport::RDMA_OFF;
686-
ep->_state = FALLBACK_TCP;
692+
ep->_state.store(FALLBACK_TCP, butil::memory_order_release);
687693
LOG_IF(INFO, FLAGS_rdma_trace_verbose)
688694
<< "Server handshake ends (use tcp) on " << s->description();
689695
}
@@ -712,7 +718,8 @@ friend class RdmaEndpoint;
712718
// blocks or first max_len bytes.
713719
// Return: the bytes included in the sglist, or -1 if failed
714720
ssize_t cut_into_sglist_and_iobuf(ibv_sge* sglist, size_t* sge_index,
715-
butil::IOBuf* to, size_t max_sge, size_t max_len) {
721+
butil::IOBuf* to, size_t max_sge,
722+
size_t max_len) {
716723
size_t len = 0;
717724
while (*sge_index < max_sge) {
718725
if (len == max_len || _ref_num() == 0) {
@@ -967,7 +974,7 @@ ssize_t RdmaEndpoint::HandleCompletion(ibv_wc& wc) {
967974
if (wc.byte_len < (uint32_t)FLAGS_rdma_zerocopy_min_size) {
968975
zerocopy = false;
969976
}
970-
CHECK(_state != FALLBACK_TCP);
977+
CHECK_NE(_state.load(butil::memory_order_acquire), FALLBACK_TCP);
971978
if (zerocopy) {
972979
_rbuf[_rq_received].cutn(&_socket->_read_buf, wc.byte_len);
973980
} else {
@@ -1586,7 +1593,7 @@ void RdmaEndpoint::PollCq(Socket* m) {
15861593
}
15871594

15881595
std::string RdmaEndpoint::GetStateStr() const {
1589-
switch (_state) {
1596+
switch (_state.load(butil::memory_order_relaxed)) {
15901597
case UNINIT: return "UNINIT";
15911598
case C_ALLOC_QPCQ: return "C_ALLOC_QPCQ";
15921599
case C_HELLO_SEND: return "C_HELLO_SEND";

‎src/brpc/rdma/rdma_endpoint.h‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -250,7 +250,7 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&);
250250
std::string GetStateStr() const;
251251

252252
// Try to read data on TCP fd in _socket
253-
inline void TryReadOnTcp();
253+
void TryReadOnTcp();
254254

255255
// Add cq socket id to poller
256256
void PollerAddCqSid();
@@ -262,7 +262,7 @@ friend int v3_wire::WriteV3Hello(RdmaEndpoint*, const RdmaHello&);
262262
Socket* _socket;
263263

264264
// State of Handshake
265-
State _state;
265+
butil::atomic<State> _state;
266266

267267
// Wire-level handshake protocol version (set by dispatch in
268268
// ProcessHandshakeAtClient/Server). Aligned with the protocol code:

0 commit comments

Comments
 (0)