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");