diff --git a/xds/src/main/java/io/grpc/xds/ExternalProcessorClientInterceptor.java b/xds/src/main/java/io/grpc/xds/ExternalProcessorClientInterceptor.java index f61e1002926..7a4585ee240 100644 --- a/xds/src/main/java/io/grpc/xds/ExternalProcessorClientInterceptor.java +++ b/xds/src/main/java/io/grpc/xds/ExternalProcessorClientInterceptor.java @@ -768,6 +768,11 @@ public void request(int numMessages) { super.request(numMessages); return; } + if (!config.getObservabilityMode() + && currentProcessingMode.getResponseBodyMode() != ProcessingMode.BodySendMode.GRPC) { + super.request(numMessages); + return; + } if (!isSidecarReady()) { pendingRequests.addAndGet(numMessages); return; @@ -795,6 +800,10 @@ public void sendMessage(InputStream message) { ExtProcStreamState state = extProcStreamState.get(); if (state.isDraining() || state.isCompleted()) { + if (currentProcessingMode.getRequestBodyMode() == ProcessingMode.BodySendMode.NONE) { + super.sendMessage(message); + return; + } try { ByteString copiedBody = ByteString.readFrom(message); pendingDrainingMessages.add(new KnownLengthInputStream(copiedBody)); @@ -1072,16 +1081,17 @@ public void onReady() { public void onHeaders(Metadata headers) { dataPlaneClientCall.setServerHeadersStartNanos(System.nanoTime()); responseHeadersSent.set(true); - if (dataPlaneClientCall.getExtProcStreamState().get().isDraining()) { - this.savedHeaders = headers; - return; - } boolean sendResponseHeaders = dataPlaneClientCall.getCurrentProcessingMode().getResponseHeaderMode() == ProcessingMode.HeaderSendMode.SEND || dataPlaneClientCall.getCurrentProcessingMode().getResponseHeaderMode() == ProcessingMode.HeaderSendMode.DEFAULT; + if (dataPlaneClientCall.getExtProcStreamState().get().isDraining() && sendResponseHeaders) { + this.savedHeaders = headers; + return; + } + if (dataPlaneClientCall.getPassThroughMode().get() || dataPlaneClientCall.getExtProcStreamState().get().isCompleted() || !sendResponseHeaders) { @@ -1110,8 +1120,11 @@ public void onMessage(InputStream message) { return; } - if (savedHeaders != null - || dataPlaneClientCall.getExtProcStreamState().get().isDraining()) { + boolean checkDrain = dataPlaneClientCall.getExtProcStreamState().get().isDraining() + && dataPlaneClientCall.getCurrentProcessingMode().getResponseBodyMode() + == ProcessingMode.BodySendMode.GRPC; + + if (savedHeaders != null || checkDrain) { try { ByteString copiedBody = ByteString.readFrom(message); savedMessages.add(new KnownLengthInputStream(copiedBody)); @@ -1184,7 +1197,11 @@ public void onClose(Status status, Metadata trailers) { return; } - if (dataPlaneClientCall.getExtProcStreamState().get().isDraining()) { + boolean sendResponseTrailers = + dataPlaneClientCall.getCurrentProcessingMode().getResponseTrailerMode() + == ProcessingMode.HeaderSendMode.SEND; + + if (dataPlaneClientCall.getExtProcStreamState().get().isDraining() && sendResponseTrailers) { return; } @@ -1193,15 +1210,6 @@ public void onClose(Status status, Metadata trailers) { } triggerCloseHandshake(); - - if (dataPlaneClientCall.getConfig().getObservabilityMode()) { - proceedWithClose(); - @SuppressWarnings("unused") - ScheduledFuture unused = dataPlaneClientCall.getScheduler().schedule( - dataPlaneClientCall::closeExtProcStream, - dataPlaneClientCall.getConfig().getDeferredCloseTimeoutNanos(), - TimeUnit.NANOSECONDS); - } } void onReadyNotify() { @@ -1330,6 +1338,15 @@ private void triggerCloseHandshake() { dataPlaneClientCall.closeExtProcStream(); } } + + if (dataPlaneClientCall.getConfig().getObservabilityMode()) { + proceedWithClose(); + @SuppressWarnings("unused") + ScheduledFuture unused = dataPlaneClientCall.getScheduler().schedule( + dataPlaneClientCall::closeExtProcStream, + dataPlaneClientCall.getConfig().getDeferredCloseTimeoutNanos(), + TimeUnit.NANOSECONDS); + } } private void sendResponseBodyToExtProc( diff --git a/xds/src/test/java/io/grpc/xds/ExternalProcessorClientInterceptorTest.java b/xds/src/test/java/io/grpc/xds/ExternalProcessorClientInterceptorTest.java index 9b07cae3477..0ce5ff3ff00 100644 --- a/xds/src/test/java/io/grpc/xds/ExternalProcessorClientInterceptorTest.java +++ b/xds/src/test/java/io/grpc/xds/ExternalProcessorClientInterceptorTest.java @@ -6104,6 +6104,494 @@ public void givenDataPlaneCallIdle_whenIsReadyCalled_thenReturnsFalse() throws E // --- Category 14: Ext-proc request draining --- + @Test + @SuppressWarnings("unchecked") + public void testRequestBodyDrainingBypassedWhenRequestBodyModeNone() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = ExternalProcessor.newBuilder() + .setGrpcService(GrpcService.newBuilder() + .setGoogleGrpc(GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin(Any.newBuilder() + .setTypeUrl("type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode(ProcessingMode.newBuilder() + .setRequestBodyMode(ProcessingMode.BodySendMode.NONE) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarActionLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .setRequestDrain(true) // Trigger Request Drain + .build()); + sidecarActionLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register(InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build().start()); + + CachedChannelManager channelManager = new CachedChannelManager(config -> { + return grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName).directExecutor().build()); + }); + + ExternalProcessorClientInterceptor interceptor = new ExternalProcessorClientInterceptor( + filterConfig, channelManager, scheduler, FAKE_CONTEXT); + + dataPlaneServiceRegistry.addService(ServerServiceDefinition.builder("test.TestService") + .addMethod(METHOD_SAY_HELLO, ServerCalls.asyncUnaryCall( + (request, responseObserver) -> { + responseObserver.onNext("Hello " + request); + responseObserver.onCompleted(); + })) + .build()); + + final List dataPlaneSentMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + ManagedChannel dataPlaneChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(dataPlaneServerName) + .directExecutor() + .intercept(new ClientInterceptor() { + @Override + public ClientCall interceptCall( + MethodDescriptor method, CallOptions callOptions, Channel next) { + return new io.grpc.ForwardingClientCall.SimpleForwardingClientCall( + next.newCall(method, callOptions)) { + @Override + public void sendMessage(ReqT message) { + try { + InputStream stream = (InputStream) message; + byte[] bytes = com.google.common.io.ByteStreams.toByteArray(stream); + dataPlaneSentMessages.add( + new String(bytes, java.nio.charset.StandardCharsets.UTF_8)); + super.sendMessage((ReqT) new java.io.ByteArrayInputStream(bytes)); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + }; + } + }) + .build()); + + CallOptions callOptions = DEFAULT_CALL_OPTIONS.withExecutor(MoreExecutors.directExecutor()); + ClientCall proxyCall = + interceptCall(interceptor, METHOD_SAY_HELLO, callOptions, dataPlaneChannel); + proxyCall.start(new ClientCall.Listener() {}, new Metadata()); + + assertThat(sidecarActionLatch.await(5, TimeUnit.SECONDS)).isTrue(); + // Wait for the drain signal to be received and processed by client call + Thread.sleep(100); + + // Call is now in DRAINING state. + // Send a message. Since request_body_mode is NONE, it should go directly to data plane. + proxyCall.sendMessage("Hello ExtProc"); + + assertThat(dataPlaneSentMessages).containsExactly("Hello ExtProc"); + + proxyCall.cancel("Cleanup", null); + channelManager.close(); + } + + @Test + @SuppressWarnings("unchecked") + public void testResponseBodyDrainingBypassedWhenResponseBodyModeNone() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = ExternalProcessor.newBuilder() + .setGrpcService(GrpcService.newBuilder() + .setGoogleGrpc(GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin(Any.newBuilder() + .setTypeUrl("type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode(ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseBodyMode(ProcessingMode.BodySendMode.NONE) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarActionLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .setRequestDrain(true) // Trigger Request Drain + .build()); + sidecarActionLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register(InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build().start()); + + CachedChannelManager channelManager = new CachedChannelManager(config -> { + return grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName).directExecutor().build()); + }); + + ExternalProcessorClientInterceptor interceptor = new ExternalProcessorClientInterceptor( + filterConfig, channelManager, scheduler, FAKE_CONTEXT); + + final AtomicReference> dataPlaneResponseObserverRef = + new AtomicReference<>(); + dataPlaneServiceRegistry.addService(ServerServiceDefinition.builder("test.TestService") + .addMethod(METHOD_BIDI_STREAMING, ServerCalls.asyncBidiStreamingCall( + new ServerCalls.BidiStreamingMethod() { + @Override + public StreamObserver invoke(StreamObserver responseObserver) { + dataPlaneResponseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(String value) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + })) + .build()); + + final List appReceivedMessages = new java.util.concurrent.CopyOnWriteArrayList<>(); + final CountDownLatch appMessageLatch = new CountDownLatch(1); + ClientCall.Listener appListener = new ClientCall.Listener() { + @Override + public void onMessage(String message) { + appReceivedMessages.add(message); + appMessageLatch.countDown(); + } + }; + + ManagedChannel dataPlaneChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(dataPlaneServerName).directExecutor().build()); + + CallOptions callOptions = DEFAULT_CALL_OPTIONS.withExecutor(MoreExecutors.directExecutor()); + ClientCall proxyCall = + interceptCall(interceptor, METHOD_BIDI_STREAMING, callOptions, dataPlaneChannel); + proxyCall.start(appListener, new Metadata()); + proxyCall.request(10); + + assertThat(sidecarActionLatch.await(5, TimeUnit.SECONDS)).isTrue(); + // Wait for the drain signal to be received and processed by client call + Thread.sleep(100); + + // Send response headers first (they bypass ext_proc because send mode is default SKIP, so + // they proceed immediately) + StreamObserver upstreamResponseObserver = dataPlaneResponseObserverRef.get(); + upstreamResponseObserver.onNext("Dummy for headers"); + + // Now call is in DRAINING state, and savedHeaders is null. + // Send response body message. Since response_body_mode is NONE, it should go directly + // downstream. + upstreamResponseObserver.onNext("Hello Downstream"); + + assertThat(appMessageLatch.await(5, TimeUnit.SECONDS)).isTrue(); + assertThat(appReceivedMessages).contains("Hello Downstream"); + + proxyCall.cancel("Cleanup", null); + channelManager.close(); + } + + @Test + @SuppressWarnings("unchecked") + public void testResponseHeadersDrainingBypassedWhenResponseHeadersSkip() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = ExternalProcessor.newBuilder() + .setGrpcService(GrpcService.newBuilder() + .setGoogleGrpc(GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin(Any.newBuilder() + .setTypeUrl("type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode(ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarActionLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .setRequestDrain(true) // Trigger Request Drain + .build()); + sidecarActionLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register(InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build().start()); + + CachedChannelManager channelManager = new CachedChannelManager(config -> { + return grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName).directExecutor().build()); + }); + + ExternalProcessorClientInterceptor interceptor = new ExternalProcessorClientInterceptor( + filterConfig, channelManager, scheduler, FAKE_CONTEXT); + + final AtomicReference> dataPlaneResponseObserverRef = + new AtomicReference<>(); + dataPlaneServiceRegistry.addService(ServerServiceDefinition.builder("test.TestService") + .addMethod(METHOD_BIDI_STREAMING, ServerCalls.asyncBidiStreamingCall( + new ServerCalls.BidiStreamingMethod() { + @Override + public StreamObserver invoke(StreamObserver responseObserver) { + dataPlaneResponseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(String value) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + })) + .build()); + + final CountDownLatch headersLatch = new CountDownLatch(1); + ClientCall.Listener appListener = new ClientCall.Listener() { + @Override + public void onHeaders(Metadata headers) { + headersLatch.countDown(); + } + }; + + ManagedChannel dataPlaneChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(dataPlaneServerName).directExecutor().build()); + + CallOptions callOptions = DEFAULT_CALL_OPTIONS.withExecutor(MoreExecutors.directExecutor()); + ClientCall proxyCall = + interceptCall(interceptor, METHOD_BIDI_STREAMING, callOptions, dataPlaneChannel); + proxyCall.start(appListener, new Metadata()); + proxyCall.request(10); + + assertThat(sidecarActionLatch.await(5, TimeUnit.SECONDS)).isTrue(); + // Wait for the drain signal to be received and processed by client call + Thread.sleep(100); + + // Call is in DRAINING state. + // Send response headers from server. Since response_header_mode is SKIP, they should go + // directly downstream. + StreamObserver upstreamResponseObserver = dataPlaneResponseObserverRef.get(); + upstreamResponseObserver.onNext("Dummy for headers"); + + assertThat(headersLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + proxyCall.cancel("Cleanup", null); + channelManager.close(); + } + + @Test + @SuppressWarnings("unchecked") + public void testResponseTrailersDrainingBypassedWhenResponseTrailersSkip() throws Exception { + String uniqueExtProcServerName = InProcessServerBuilder.generateName(); + ExternalProcessor proto = ExternalProcessor.newBuilder() + .setGrpcService(GrpcService.newBuilder() + .setGoogleGrpc(GrpcService.GoogleGrpc.newBuilder() + .setTargetUri("in-process:///" + uniqueExtProcServerName) + .addChannelCredentialsPlugin(Any.newBuilder() + .setTypeUrl("type.googleapis.com/envoy.extensions.grpc_service." + + "channel_credentials.insecure.v3.InsecureCredentials") + .build()) + .build()) + .build()) + .setProcessingMode(ProcessingMode.newBuilder() + .setResponseHeaderMode(ProcessingMode.HeaderSendMode.SKIP) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SKIP) + .build()) + .build(); + ConfigOrError configOrError = + provider.parseFilterConfig(Any.pack(proto), filterContext); + assertThat(configOrError.errorDetail).isNull(); + ExternalProcessorFilterConfig filterConfig = configOrError.config; + + final CountDownLatch sidecarActionLatch = new CountDownLatch(1); + ExternalProcessorGrpc.ExternalProcessorImplBase extProcImpl = + new ExternalProcessorGrpc.ExternalProcessorImplBase() { + @Override + public StreamObserver process( + final StreamObserver responseObserver) { + ((ServerCallStreamObserver) responseObserver).request(100); + return new StreamObserver() { + @Override + public void onNext(ProcessingRequest request) { + if (request.hasRequestHeaders()) { + responseObserver.onNext(ProcessingResponse.newBuilder() + .setRequestHeaders(HeadersResponse.newBuilder().build()) + .setRequestDrain(true) // Trigger Request Drain + .build()); + sidecarActionLatch.countDown(); + } + } + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() { + responseObserver.onCompleted(); + } + }; + } + }; + grpcCleanup.register(InProcessServerBuilder.forName(uniqueExtProcServerName) + .addService(extProcImpl) + .directExecutor() + .build().start()); + + CachedChannelManager channelManager = new CachedChannelManager(config -> { + return grpcCleanup.register( + InProcessChannelBuilder.forName(uniqueExtProcServerName).directExecutor().build()); + }); + + ExternalProcessorClientInterceptor interceptor = new ExternalProcessorClientInterceptor( + filterConfig, channelManager, scheduler, FAKE_CONTEXT); + + final AtomicReference> dataPlaneResponseObserverRef = + new AtomicReference<>(); + dataPlaneServiceRegistry.addService(ServerServiceDefinition.builder("test.TestService") + .addMethod(METHOD_BIDI_STREAMING, ServerCalls.asyncBidiStreamingCall( + new ServerCalls.BidiStreamingMethod() { + @Override + public StreamObserver invoke(StreamObserver responseObserver) { + dataPlaneResponseObserverRef.set(responseObserver); + return new StreamObserver() { + @Override + public void onNext(String value) {} + + @Override + public void onError(Throwable t) {} + + @Override + public void onCompleted() {} + }; + } + })) + .build()); + + final CountDownLatch closeLatch = new CountDownLatch(1); + ClientCall.Listener appListener = new ClientCall.Listener() { + @Override + public void onClose(Status status, Metadata trailers) { + closeLatch.countDown(); + } + }; + + ManagedChannel dataPlaneChannel = grpcCleanup.register( + InProcessChannelBuilder.forName(dataPlaneServerName).directExecutor().build()); + + CallOptions callOptions = DEFAULT_CALL_OPTIONS.withExecutor(MoreExecutors.directExecutor()); + ClientCall proxyCall = + interceptCall(interceptor, METHOD_BIDI_STREAMING, callOptions, dataPlaneChannel); + proxyCall.start(appListener, new Metadata()); + proxyCall.request(10); + + assertThat(sidecarActionLatch.await(5, TimeUnit.SECONDS)).isTrue(); + // Wait for the drain signal to be received and processed by client call + Thread.sleep(100); + + // Call is in DRAINING state. + // Complete the server call. Since response_trailer_mode is SKIP, onClose should trigger + // immediately. + StreamObserver upstreamResponseObserver = dataPlaneResponseObserverRef.get(); + upstreamResponseObserver.onCompleted(); + + assertThat(closeLatch.await(5, TimeUnit.SECONDS)).isTrue(); + + proxyCall.cancel("Cleanup", null); + channelManager.close(); + } + @Test @SuppressWarnings("unchecked") public void givenRequestDrainActive_whenIsReadyCalled_thenReturnsFalse() throws Exception { @@ -7592,6 +8080,10 @@ public void givenRequestDrainActive_whenAppRequestsMessages_thenRequestsBuffered .build()) .build()) .build()) + .setProcessingMode(ProcessingMode.newBuilder() + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) .build(); ConfigOrError configOrError = provider.parseFilterConfig(Any.pack(proto), filterContext); @@ -9980,6 +10472,10 @@ public void givenObservabilityModeFalse_whenExtProcBusy_thenAppRequestsAreBuffer .build()) .build()) .build()) + .setProcessingMode(ProcessingMode.newBuilder() + .setResponseBodyMode(ProcessingMode.BodySendMode.GRPC) + .setResponseTrailerMode(ProcessingMode.HeaderSendMode.SEND) + .build()) .setObservabilityMode(false) .build(); ConfigOrError configOrError =