This is an automated email from the ASF dual-hosted git repository.
Aias00 pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/shenyu.git
The following commit(s) were added to refs/heads/master by this push:
new 1b2a3e46f8 fix(ai-proxy): propagate stream cancellation (#7338)
1b2a3e46f8 is described below
commit 1b2a3e46f8987eb4a04b3d0c572e364eb7780d62
Author: BobSong <[email protected]>
AuthorDate: Wed Sep 30 10:00:09 2026 +0800
fix(ai-proxy): propagate stream cancellation (#7338)
Co-authored-by: BobSong-dev <[email protected]>
---
.../plugin/ai/proxy/enhanced/AiProxyPlugin.java | 3 +
.../enhanced/service/AiProxyExecutorService.java | 6 +-
.../enhanced/service/AiStreamCancellation.java | 70 ++++++++++
.../service/AiProxyStreamCancellationTest.java | 149 +++++++++++++++++++++
4 files changed, 225 insertions(+), 3 deletions(-)
diff --git
a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/AiProxyPlugin.java
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/AiProxyPlugin.java
index 9f5c49869e..1f12cff50f 100644
---
a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/AiProxyPlugin.java
+++
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/AiProxyPlugin.java
@@ -32,6 +32,8 @@ import
org.apache.shenyu.plugin.ai.proxy.enhanced.service.AiProxyConfigService;
import
org.apache.shenyu.plugin.ai.proxy.enhanced.service.AiProxyExecutorService;
import
org.apache.shenyu.plugin.ai.proxy.enhanced.service.AiProxyExecutorService.FallbackContext;
import org.apache.shenyu.plugin.ai.proxy.enhanced.service.UpstreamErrorLogger;
+import org.apache.shenyu.plugin.ai.proxy.enhanced.service.AiStreamCancellation;
+import org.springframework.web.reactive.function.client.WebClient;
import org.apache.shenyu.plugin.api.ShenyuPluginChain;
import org.apache.shenyu.plugin.api.utils.WebFluxResultUtils;
import org.apache.shenyu.plugin.base.AbstractShenyuPlugin;
@@ -248,6 +250,7 @@ public class AiProxyPlugin extends AbstractShenyuPlugin {
throw new IllegalArgumentException("apiKey must not be empty");
}
return OpenAiApi.builder()
+
.webClientBuilder(WebClient.builder().filter(AiStreamCancellation.responseFilter()))
.baseUrl(config.getBaseUrl())
.apiKey(config.getApiKey())
.build();
diff --git
a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyExecutorService.java
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyExecutorService.java
index 1898703725..6e134c50b7 100644
---
a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyExecutorService.java
+++
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyExecutorService.java
@@ -61,9 +61,9 @@ public class AiProxyExecutorService {
public Flux<ChatCompletionChunk> executeDirectStream(final OpenAiApi
mainApi,
final Optional<FallbackContext> fallbackCtxOpt, final
ChatCompletionRequest request,
final String requestBody, final boolean stream) {
- return Flux.defer(() -> {
+ return AiStreamCancellation.propagate(Flux.defer(() -> {
AtomicBoolean emitted = new AtomicBoolean();
- return mainApi.chatCompletionStream(request)
+ return Flux.defer(() -> mainApi.chatCompletionStream(request))
.doOnNext(chunk -> emitted.set(true))
.doOnError(e -> UpstreamErrorLogger.logUpstreamError(LOG,
e, "direct stream"))
.retryWhen(Retry.max(1)
@@ -78,7 +78,7 @@ public class AiProxyExecutorService {
.onErrorResume(error -> emitted.get()
? Flux.error(error)
: handleDirectFallbackStream(error,
fallbackCtxOpt, requestBody, stream));
- });
+ }));
}
/**
diff --git
a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiStreamCancellation.java
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiStreamCancellation.java
new file mode 100644
index 0000000000..c254397f7f
--- /dev/null
+++
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/main/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiStreamCancellation.java
@@ -0,0 +1,70 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.shenyu.plugin.ai.proxy.enhanced.service;
+
+import org.springframework.web.reactive.function.client.ExchangeFilterFunction;
+import reactor.core.publisher.Flux;
+import reactor.core.publisher.SignalType;
+import reactor.core.publisher.Sinks;
+
+/**
+ * Connects downstream cancellation to the raw response body across SDK
windows.
+ * Each subscription owns its signal; cached clients never hold request state.
+ */
+public final class AiStreamCancellation {
+
+ private static final Object CONTEXT_KEY = new Object();
+
+ private AiStreamCancellation() {
+ }
+
+ /**
+ * Creates a filter that cancels the raw HTTP body when its caller
disconnects.
+ *
+ * @return the request-context-aware response filter
+ */
+ public static ExchangeFilterFunction responseFilter() {
+ return (request, next) ->
reactor.core.publisher.Mono.deferContextual(context -> {
+ if (!context.hasKey(CONTEXT_KEY)) {
+ return next.exchange(request);
+ }
+ final Sinks.Empty<Void> cancellation = context.get(CONTEXT_KEY);
+ return next.exchange(request).map(response -> response.mutate()
+ .body(body ->
body.takeUntilOther(cancellation.asMono())).build());
+ });
+ }
+
+ /**
+ * Gives each subscription an independent signal including retries and
fallback.
+ *
+ * @param source the SDK response stream
+ * @param <T> the response element type
+ * @return the stream with request-local cancellation propagation
+ */
+ public static <T> Flux<T> propagate(final Flux<T> source) {
+ return Flux.defer(() -> {
+ final Sinks.Empty<Void> cancellation = Sinks.empty();
+ return source.contextWrite(context -> context.put(CONTEXT_KEY,
cancellation))
+ .doFinally(signal -> {
+ if (signal == SignalType.CANCEL) {
+ cancellation.tryEmitEmpty();
+ }
+ });
+ });
+ }
+}
diff --git
a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/test/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyStreamCancellationTest.java
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/test/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyStreamCancellationTest.java
new file mode 100644
index 0000000000..17d0f6505a
--- /dev/null
+++
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/test/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyStreamCancellationTest.java
@@ -0,0 +1,149 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.shenyu.plugin.ai.proxy.enhanced.service;
+
+import org.junit.jupiter.api.Test;
+import org.apache.shenyu.plugin.ai.common.config.AiCommonConfig;
+import org.springframework.ai.openai.api.OpenAiApi;
+import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest;
+import org.springframework.core.io.buffer.DefaultDataBufferFactory;
+import org.springframework.http.HttpHeaders;
+import org.springframework.http.HttpStatus;
+import org.springframework.http.MediaType;
+import org.springframework.web.reactive.function.client.ClientResponse;
+import org.springframework.web.reactive.function.client.WebClient;
+import reactor.core.publisher.Flux;
+import reactor.core.publisher.Mono;
+import reactor.test.StepVerifier;
+import reactor.core.Disposable;
+
+import java.nio.charset.StandardCharsets;
+import java.time.Duration;
+import java.util.Optional;
+import java.util.ArrayList;
+import java.util.List;
+import java.util.concurrent.atomic.AtomicInteger;
+import java.util.concurrent.atomic.AtomicBoolean;
+
+import static org.junit.jupiter.api.Assertions.assertTrue;
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.when;
+
+/**
+ * Exercises the real Spring AI stream operators without opening sockets.
+ */
+class AiProxyStreamCancellationTest {
+
+ private static final String EVENT = """
+ data:
{"id":"stream-1","object":"chat.completion.chunk","created":0,"model":"fixture","choices":[{"index":0,"delta":{"content":"hello"}}]}
+
+ """;
+
+ @Test
+ void testCancelAfterFirstEventReachesRawBody() {
+ verifyCancellation(Duration.ZERO);
+ }
+
+ @Test
+ void testAsynchronousCancelAfterFirstEventReachesRawBody() {
+ verifyCancellation(Duration.ofMillis(100));
+ }
+
+ private void verifyCancellation(final Duration cancellationDelay) {
+ final AtomicBoolean cancelled = new AtomicBoolean();
+ final Flux<String> events =
Flux.just(EVENT).concatWith(Flux.never()).doOnCancel(() -> cancelled.set(true));
+ final OpenAiApi api = createApi(events);
+ StepVerifier.create(stream(api))
+ .expectNextCount(1)
+ .thenAwait(cancellationDelay)
+ .thenCancel()
+ .verify(Duration.ofSeconds(3));
+ assertTrue(cancelled.get(), "Cancellation must reach the raw WebClient
response, not only the SDK output");
+ }
+
+ @Test
+ void testCancellationBeforeFirstEvent() {
+ final AtomicBoolean cancelled = new AtomicBoolean();
+ final OpenAiApi api = createApi(Flux.<String>never().doOnCancel(() ->
cancelled.set(true)));
+
StepVerifier.create(stream(api)).thenAwait(Duration.ofMillis(100)).thenCancel().verify(Duration.ofSeconds(3));
+ assertTrue(cancelled.get());
+ }
+
+ @Test
+ void testSharedClientSubscriptionsHaveIndependentCancellation() {
+ final List<AtomicBoolean> cancellations = new ArrayList<>();
+ final OpenAiApi api = createApi(Flux.defer(() -> {
+ final AtomicBoolean cancelled = new AtomicBoolean();
+ cancellations.add(cancelled);
+ return Flux.just(EVENT).concatWith(Flux.never()).doOnCancel(() ->
cancelled.set(true));
+ }));
+ final Flux<OpenAiApi.ChatCompletionChunk> shared = stream(api);
+ final AtomicInteger received = new AtomicInteger();
+ final Disposable first = shared.subscribe(chunk ->
received.incrementAndGet());
+ final Disposable second = shared.subscribe(chunk ->
received.incrementAndGet());
+ try {
+ assertEquals(2, received.get());
+ first.dispose();
+ assertTrue(cancellations.get(0).get());
+ assertFalse(cancellations.get(1).get(), "Cancelling one subscriber
must not cancel another on the same cached client");
+ second.dispose();
+ assertTrue(cancellations.get(1).get());
+ } finally {
+ first.dispose();
+ second.dispose();
+ }
+ }
+
+ @Test
+ void testFallbackCancellationReachesRawBody() {
+ final OpenAiApi failing = mock(OpenAiApi.class);
+ final ChatCompletionRequest request =
mock(ChatCompletionRequest.class);
+ when(failing.chatCompletionStream(request)).thenReturn(Flux.error(new
IllegalStateException("fixture failure")));
+ final AtomicBoolean cancelled = new AtomicBoolean();
+ final OpenAiApi fallback =
createApi(Flux.just(EVENT).concatWith(Flux.never()).doOnCancel(() ->
cancelled.set(true)));
+ final AiCommonConfig config = new AiCommonConfig();
+ config.setModel("fixture");
+ final AiProxyExecutorService.FallbackContext context = new
AiProxyExecutorService.FallbackContext(fallback, config);
+ StepVerifier.create(new
AiProxyExecutorService().executeDirectStream(failing, Optional.of(context),
request,
+ "{\"messages\":[{\"role\":\"user\",\"content\":\"test\"}]}",
true))
+
.expectNextCount(1).thenAwait(Duration.ofMillis(100)).thenCancel().verify(Duration.ofSeconds(3));
+ assertTrue(cancelled.get());
+ }
+
+ @Test
+ void testNormalStreamStillCompletes() {
+
StepVerifier.create(stream(createApi(Flux.just(EVENT)))).expectNextCount(1).verifyComplete();
+ }
+
+ private OpenAiApi createApi(final Flux<String> events) {
+ return OpenAiApi.builder().apiKey("fixture-key")
+
.webClientBuilder(WebClient.builder().filter(AiStreamCancellation.responseFilter()).exchangeFunction(request
-> Mono.just(ClientResponse.create(HttpStatus.OK)
+ .header(HttpHeaders.CONTENT_TYPE,
MediaType.TEXT_EVENT_STREAM_VALUE)
+ .body(events.map(event ->
DefaultDataBufferFactory.sharedInstance.wrap(event.getBytes(StandardCharsets.UTF_8))))
+ .build())))
+ .build();
+ }
+
+ private Flux<OpenAiApi.ChatCompletionChunk> stream(final OpenAiApi api) {
+ final ChatCompletionRequest request =
mock(ChatCompletionRequest.class);
+ when(request.stream()).thenReturn(true);
+ return new AiProxyExecutorService().executeDirectStream(api,
Optional.empty(), request, "{}", true);
+ }
+}