Skip to content

Commit 350dc40

Browse files
committed
xds: Fix TSAN data race on ClientCall cancellation in ExternalProcessorClientInterceptor
Prevent concurrent cancellations of the underlying ClientCall in ExternalProcessorClientInterceptor: - Wrap rawCall with SimpleForwardingClientCall using an AtomicBoolean to ensure the underlying ClientCall.cancel() is executed at most once, even if invoked concurrently across threads or from DelayedListener. - In DataPlaneClientCall.cancel(), atomically transition extProcStreamState to FAILED, catch exceptions during onError(), and clear extProcClientCallRequestObserver. - In sendToExtProc(), return early if the ext-proc stream is already completed or the observer is null, and catch unexpected onNext() exceptions to trigger internalOnError() rather than letting exceptions escape into listener callbacks. - Safely complete and clear extProcClientCallRequestObserver in closeExtProcStream() and halfCloseExtProcStream(). - Route all rawCall.cancel() calls in sendMessage(), handleImmediateResponse(), and DataPlaneListener through cancelDownstream().
1 parent 2858db2 commit 350dc40

1 file changed

Lines changed: 62 additions & 22 deletions

File tree

‎xds/src/main/java/io/grpc/xds/ExternalProcessorClientInterceptor.java‎

Lines changed: 62 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -238,8 +238,18 @@ public <ReqT, RespT> ClientCall<ReqT, RespT> interceptCall(
238238
MethodDescriptor<InputStream, InputStream> rawMethod =
239239
(MethodDescriptor<InputStream, InputStream>) (MethodDescriptor<?, ?>) method;
240240
ClientCall<InputStream, InputStream> rawCall =
241-
(ClientCall<InputStream, InputStream>) (ClientCall<?, ?>)
242-
next.newCall(method, callOptions);
241+
new SimpleForwardingClientCall<InputStream, InputStream>(
242+
(ClientCall<InputStream, InputStream>) (ClientCall<?, ?>)
243+
next.newCall(method, callOptions)) {
244+
private final AtomicBoolean cancelled = new AtomicBoolean(false);
245+
246+
@Override
247+
public void cancel(@Nullable String message, @Nullable Throwable cause) {
248+
if (cancelled.compareAndSet(false, true)) {
249+
super.cancel(message, cause);
250+
}
251+
}
252+
};
243253

244254
// Create a local subclass instance to buffer outbound actions
245255
DataPlaneDelayedCall<InputStream, InputStream> delayedCall =
@@ -397,13 +407,18 @@ private boolean validateCompressionSupport(BodyResponse bodyResponse) {
397407
.withDescription("gRPC message compression not supported in ext_proc")
398408
.asRuntimeException();
399409
synchronized (streamLock) {
400-
if (!extProcStreamState.get().isCompleted()
401-
&& extProcClientCallRequestObserver != null) {
402-
extProcClientCallRequestObserver.onError(ex);
410+
if (markExtProcStreamFailed(extProcStreamState)) {
411+
if (extProcClientCallRequestObserver != null) {
412+
try {
413+
extProcClientCallRequestObserver.onError(ex);
414+
} catch (Throwable ignored) {
415+
// Ignore exceptions during cancel/onError propagation
416+
}
417+
extProcClientCallRequestObserver = null;
418+
}
403419
}
404420
}
405421
activateCall();
406-
markExtProcStreamFailed(extProcStreamState);
407422
cancelDownstream("gRPC message compression not supported in ext_proc", ex);
408423
closeExtProcStream();
409424
return false;
@@ -600,6 +615,9 @@ public void onError(Throwable t) {
600615
@Override
601616
public void onCompleted() {
602617
if (markExtProcStreamCompleted(extProcStreamState)) {
618+
synchronized (streamLock) {
619+
extProcClientCallRequestObserver = null;
620+
}
603621
handleFailOpen(wrappedListener);
604622
}
605623
}
@@ -628,8 +646,9 @@ public void onCompleted() {
628646
}
629647

630648
private void sendToExtProc(ProcessingRequest request) {
649+
Throwable sendError = null;
631650
synchronized (streamLock) {
632-
if (extProcStreamState.get().isCompleted()) {
651+
if (extProcStreamState.get().isCompleted() || extProcClientCallRequestObserver == null) {
633652
return;
634653
}
635654

@@ -670,7 +689,14 @@ private void sendToExtProc(ProcessingRequest request) {
670689
.build();
671690
}
672691

673-
extProcClientCallRequestObserver.onNext(requestToSend);
692+
try {
693+
extProcClientCallRequestObserver.onNext(requestToSend);
694+
} catch (Throwable t) {
695+
sendError = t;
696+
}
697+
}
698+
if (sendError != null) {
699+
internalOnError(sendError);
674700
}
675701
}
676702

@@ -690,7 +716,11 @@ private void closeExtProcStream() {
690716
synchronized (streamLock) {
691717
if (markExtProcStreamCompleted(extProcStreamState)) {
692718
if (extProcClientCallRequestObserver != null) {
693-
extProcClientCallRequestObserver.onCompleted();
719+
try {
720+
extProcClientCallRequestObserver.onCompleted();
721+
} catch (Throwable ignored) {
722+
}
723+
extProcClientCallRequestObserver = null;
694724
}
695725
}
696726
}
@@ -722,7 +752,10 @@ private void internalOnError(Throwable t) {
722752
private void halfCloseExtProcStream() {
723753
synchronized (streamLock) {
724754
if (!extProcStreamState.get().isCompleted() && extProcClientCallRequestObserver != null) {
725-
extProcClientCallRequestObserver.onCompleted();
755+
try {
756+
extProcClientCallRequestObserver.onCompleted();
757+
} catch (Throwable ignored) {
758+
}
726759
}
727760
}
728761
}
@@ -809,7 +842,7 @@ public void sendMessage(InputStream message) {
809842
ByteString copiedBody = ByteString.readFrom(message);
810843
pendingDrainingMessages.add(new KnownLengthInputStream(copiedBody));
811844
} catch (IOException e) {
812-
rawCall.cancel("Failed to copy outbound message for buffering", e);
845+
cancelDownstream("Failed to copy outbound message for buffering", e);
813846
}
814847
return;
815848
}
@@ -835,7 +868,7 @@ public void sendMessage(InputStream message) {
835868
super.sendMessage(new KnownLengthInputStream(bodyByteString));
836869
}
837870
} catch (IOException e) {
838-
rawCall.cancel("Failed to serialize message for External Processor", e);
871+
cancelDownstream("Failed to serialize message for External Processor", e);
839872
}
840873
}
841874

@@ -900,7 +933,7 @@ public void halfClose() {
900933
.build());
901934
}
902935

903-
private void cancelDownstream(@Nullable String message, @Nullable Throwable cause) {
936+
void cancelDownstream(@Nullable String message, @Nullable Throwable cause) {
904937
if (downstreamCancelled.compareAndSet(false, true)) {
905938
delayedCall.cancel(message, cause);
906939
}
@@ -909,12 +942,19 @@ private void cancelDownstream(@Nullable String message, @Nullable Throwable caus
909942
@Override
910943
public void cancel(@Nullable String message, @Nullable Throwable cause) {
911944
synchronized (streamLock) {
912-
if (!extProcStreamState.get().isCompleted() && extProcClientCallRequestObserver != null) {
913-
extProcClientCallRequestObserver.onError(
914-
Status.CANCELLED
915-
.withDescription(message)
916-
.withCause(cause)
917-
.asRuntimeException());
945+
if (markExtProcStreamFailed(extProcStreamState)) {
946+
if (extProcClientCallRequestObserver != null) {
947+
try {
948+
extProcClientCallRequestObserver.onError(
949+
Status.CANCELLED
950+
.withDescription(message)
951+
.withCause(cause)
952+
.asRuntimeException());
953+
} catch (Throwable ignored) {
954+
// Ignore exceptions during cancel/onError propagation
955+
}
956+
extProcClientCallRequestObserver = null;
957+
}
918958
}
919959
}
920960
cancelDownstream(message, cause);
@@ -970,7 +1010,7 @@ private void handleImmediateResponse(ImmediateResponse immediate, DataPlaneListe
9701010
// If sent in response to any other event, it will cause the data plane RPC to
9711011
// immediately fail with the specified status as if it were an out-of-band
9721012
// cancellation.
973-
rawCall.cancel(status.getDescription(), null);
1013+
cancelDownstream(status.getDescription(), null);
9741014
listener.unblockAfterStreamComplete();
9751015
}
9761016
closeExtProcStream();
@@ -1153,7 +1193,7 @@ public void onMessage(InputStream message) {
11531193
ByteString copiedBody = ByteString.readFrom(message);
11541194
savedMessages.add(new KnownLengthInputStream(copiedBody));
11551195
} catch (IOException e) {
1156-
rawCall.cancel("Failed to copy inbound message for buffering", e);
1196+
dataPlaneClientCall.cancelDownstream("Failed to copy inbound message for buffering", e);
11571197
}
11581198
return;
11591199
}
@@ -1184,7 +1224,7 @@ public void onMessage(InputStream message) {
11841224
() -> delegate().onMessage(bodyByteString.newInput()));
11851225
}
11861226
} catch (IOException e) {
1187-
rawCall.cancel("Failed to read server response", e);
1227+
dataPlaneClientCall.cancelDownstream("Failed to read server response", e);
11881228
}
11891229
}
11901230

0 commit comments

Comments
 (0)