Fixup: Add in-flight cancellation tests and ensure super.cancel runs before authzContext cancel
diff --git a/xds/src/main/java/io/grpc/xds/internal/extauthz/ExtAuthzClientCall.java b/xds/src/main/java/io/grpc/xds/internal/extauthz/ExtAuthzClientCall.java index 4242847..eb5ddd7 100644 --- a/xds/src/main/java/io/grpc/xds/internal/extauthz/ExtAuthzClientCall.java +++ b/xds/src/main/java/io/grpc/xds/internal/extauthz/ExtAuthzClientCall.java
@@ -98,8 +98,8 @@ @Override public void cancel( @Nullable String message, @Nullable Throwable cause) { - authzContext.cancel(cause); super.cancel(message, cause); + authzContext.cancel(cause); } /**
diff --git a/xds/src/test/java/io/grpc/xds/internal/extauthz/AuthzCallbackObserverTest.java b/xds/src/test/java/io/grpc/xds/internal/extauthz/AuthzCallbackObserverTest.java index 3ff40fc..88c9d40 100644 --- a/xds/src/test/java/io/grpc/xds/internal/extauthz/AuthzCallbackObserverTest.java +++ b/xds/src/test/java/io/grpc/xds/internal/extauthz/AuthzCallbackObserverTest.java
@@ -625,6 +625,73 @@ observer.onNext(CheckResponse.getDefaultInstance()); } + @Test + public void allow_whenDelayedCallCancelledInFlight_setCallReturnsNull() { + doAnswer(invocation -> { + StreamObserver<CheckResponse> obs = invocation.getArgument(1); + obs.onNext(CheckResponse.newBuilder() + .setStatus(com.google.rpc.Status.newBuilder().setCode(0).build()) + .setOkResponse(OkHttpResponse.getDefaultInstance()) + .build()); + obs.onCompleted(); + return null; + }).when(authzService).check(any(), any()); + + TestDelayedCall<SimpleRequest, SimpleResponse> delayedCall = + new TestDelayedCall<>(MoreExecutors.directExecutor(), scheduler, null); + CapturingListener<SimpleResponse> listener = new CapturingListener<>(); + delayedCall.start(listener, new Metadata()); + delayedCall.cancel("cancelled in flight", null); + + Context.CancellableContext authzCtx = Context.current().withCancellation(); + AuthzCallbackObserver<SimpleRequest, SimpleResponse> observer = + new AuthzCallbackObserver<>( + delayedCall, channel, + SimpleServiceGrpc.getUnaryRpcMethod(), + CallOptions.DEFAULT, + MoreExecutors.directExecutor(), + responseHandler, failClosedConfig(), authzCtx); + + authzCtx.run(() -> { + AuthorizationGrpc.newStub(channel) + .check(CheckRequest.getDefaultInstance(), observer); + }); + + assertThat(capturedBackendHeaders).isNull(); + assertThat(listener.getCloseStatus().getCode()).isEqualTo(Status.Code.CANCELLED); + } + + @Test + public void allow_whenDelayedCallNotStarted_setCallReturnsNull() { + doAnswer(invocation -> { + StreamObserver<CheckResponse> obs = invocation.getArgument(1); + obs.onNext(CheckResponse.newBuilder() + .setStatus(com.google.rpc.Status.newBuilder().setCode(0).build()) + .setOkResponse(OkHttpResponse.getDefaultInstance()) + .build()); + obs.onCompleted(); + return null; + }).when(authzService).check(any(), any()); + + TestDelayedCall<SimpleRequest, SimpleResponse> delayedCall = + new TestDelayedCall<>(MoreExecutors.directExecutor(), scheduler, null); + Context.CancellableContext authzCtx = Context.current().withCancellation(); + AuthzCallbackObserver<SimpleRequest, SimpleResponse> observer = + new AuthzCallbackObserver<>( + delayedCall, channel, + SimpleServiceGrpc.getUnaryRpcMethod(), + CallOptions.DEFAULT, + MoreExecutors.directExecutor(), + responseHandler, failClosedConfig(), authzCtx); + + authzCtx.run(() -> { + AuthorizationGrpc.newStub(channel) + .check(CheckRequest.getDefaultInstance(), observer); + }); + + assertThat(capturedBackendHeaders).isNull(); + } + private static final class TestDelayedCall<ReqT, RespT> extends DelayedClientCall<ReqT, RespT> {
diff --git a/xds/src/test/java/io/grpc/xds/internal/extauthz/ExtAuthzClientCallTest.java b/xds/src/test/java/io/grpc/xds/internal/extauthz/ExtAuthzClientCallTest.java index 69cdbe5..762011d 100644 --- a/xds/src/test/java/io/grpc/xds/internal/extauthz/ExtAuthzClientCallTest.java +++ b/xds/src/test/java/io/grpc/xds/internal/extauthz/ExtAuthzClientCallTest.java
@@ -195,6 +195,47 @@ } @Test + public void cancel_whileCheckInFlight_responseAfterCancelDoesNotForwardToBackend() + throws Exception { + AtomicReference<StreamObserver<CheckResponse>> capturedObserver = new AtomicReference<>(); + CountDownLatch checkCalled = new CountDownLatch(1); + doAnswer(invocation -> { + capturedObserver.set(invocation.getArgument(1)); + checkCalled.countDown(); + return null; + }).when(authzService).check(any(CheckRequest.class), any()); + + HeaderMutations emptyMutations = HeaderMutations.create( + ImmutableList.of(), ImmutableList.of()); + AuthzResponse authzResponse = + AuthzResponse.allow(emptyMutations) + .setResponseHeaderMutations(emptyMutations) + .build(); + when(mockResponseHandler.handleResponse( + any(CheckResponse.class))).thenReturn(authzResponse); + + ExtAuthzClientCall<SimpleRequest, SimpleResponse> call = createCall( + com.google.common.util.concurrent.MoreExecutors.directExecutor(), channel, config); + CapturingListener<SimpleResponse> listener = new CapturingListener<>(); + call.start(listener, new Metadata()); + + assertThat(checkCalled.await(5, TimeUnit.SECONDS)).isTrue(); + + // Cancel while check is pending + call.cancel("client cancelled", null); + assertThat(listener.getCloseStatus().getCode()).isEqualTo(Status.Code.CANCELLED); + + lastBackendHeaders = null; + + // Authz responds after cancellation + capturedObserver.get().onNext(CheckResponse.getDefaultInstance()); + capturedObserver.get().onCompleted(); + + // Verify backend received nothing + assertThat(lastBackendHeaders).isNull(); + } + + @Test public void start_dispatchesAuthzCheckUnderChildContext() throws Exception { Context.Key<String> testKey = Context.key("test-key"); Context dataPlaneContext = Context.current().withValue(testKey, "data-plane-val");