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);
+    }
+}

Reply via email to