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 41dee74d57 fix(ai-proxy): avoid replaying partial streams (#7205)
41dee74d57 is described below
commit 41dee74d577ff7166588a00fc719ae41d53f272f
Author: Liming Deng <[email protected]>
AuthorDate: Thu Sep 24 14:13:39 2026 +0800
fix(ai-proxy): avoid replaying partial streams (#7205)
Co-authored-by: aias00 <[email protected]>
---
.../enhanced/service/AiProxyExecutorService.java | 31 +++++++++++++---------
.../service/AiProxyExecutorServiceTest.java | 23 ++++++++++++++++
2 files changed, 42 insertions(+), 12 deletions(-)
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 4628fee477..1898703725 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
@@ -38,6 +38,7 @@ import reactor.util.retry.Retry;
import java.time.Duration;
import java.util.Objects;
import java.util.Optional;
+import java.util.concurrent.atomic.AtomicBoolean;
/**
* AI proxy executor service.
@@ -60,18 +61,24 @@ public class AiProxyExecutorService {
public Flux<ChatCompletionChunk> executeDirectStream(final OpenAiApi
mainApi,
final Optional<FallbackContext> fallbackCtxOpt, final
ChatCompletionRequest request,
final String requestBody, final boolean stream) {
- return mainApi.chatCompletionStream(request)
- .doOnError(e -> UpstreamErrorLogger.logUpstreamError(LOG, e,
"direct stream"))
- .retryWhen(Retry.max(1)
- .filter(AiProxyExecutorService::isRetryable)
- .onRetryExhaustedThrow((retryBackoffSpec, retrySignal)
-> {
- LOG.warn("Direct stream retry exhausted.
Triggering fallback.",
- retrySignal.failure());
- return new NonTransientAiException(
- "Direct stream failed after 1 retry.
Triggering fallback.",
- retrySignal.failure());
- }))
- .onErrorResume(e -> handleDirectFallbackStream(e,
fallbackCtxOpt, requestBody, stream));
+ return Flux.defer(() -> {
+ AtomicBoolean emitted = new AtomicBoolean();
+ return mainApi.chatCompletionStream(request)
+ .doOnNext(chunk -> emitted.set(true))
+ .doOnError(e -> UpstreamErrorLogger.logUpstreamError(LOG,
e, "direct stream"))
+ .retryWhen(Retry.max(1)
+ .filter(error -> !emitted.get() &&
isRetryable(error))
+ .onRetryExhaustedThrow((retryBackoffSpec,
retrySignal) -> {
+ LOG.warn("Direct stream retry exhausted.
Triggering fallback.",
+ retrySignal.failure());
+ return new NonTransientAiException(
+ "Direct stream failed after 1 retry.
Triggering fallback.",
+ retrySignal.failure());
+ }))
+ .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/test/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyExecutorServiceTest.java
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/test/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyExecutorServiceTest.java
index 38597d0dbe..76788c2971 100644
---
a/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/test/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyExecutorServiceTest.java
+++
b/shenyu-plugin/shenyu-plugin-ai/shenyu-plugin-ai-proxy/src/test/java/org/apache/shenyu/plugin/ai/proxy/enhanced/service/AiProxyExecutorServiceTest.java
@@ -35,6 +35,7 @@ import java.util.Optional;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.never;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
@@ -97,6 +98,28 @@ public class AiProxyExecutorServiceTest {
verify(fallbackApi,
times(1)).chatCompletionStream(any(ChatCompletionRequest.class));
}
+ @Test
+ void testExecuteDirectStreamDoesNotRetryOrFallbackAfterEmission() {
+ final OpenAiApi mainApi = mock(OpenAiApi.class);
+ final OpenAiApi fallbackApi = mock(OpenAiApi.class);
+ final ChatCompletionRequest request =
mock(ChatCompletionRequest.class);
+ final ChatCompletionChunk firstChunk = mock(ChatCompletionChunk.class);
+ when(mainApi.chatCompletionStream(request)).thenReturn(
+ Flux.concat(Flux.just(firstChunk), Flux.error(new
RuntimeException("mid-stream error"))));
+
+ final AiCommonConfig fallbackConfig = new AiCommonConfig();
+ fallbackConfig.setModel("fallback-model");
+ final AiProxyExecutorService.FallbackContext ctx = new
AiProxyExecutorService.FallbackContext(fallbackApi, fallbackConfig);
+
+ StepVerifier.create(executorService.executeDirectStream(mainApi,
Optional.of(ctx), request, REQUEST_BODY, true))
+ .expectNext(firstChunk)
+ .expectErrorMessage("mid-stream error")
+ .verify();
+
+ verify(mainApi, times(1)).chatCompletionStream(request);
+ verify(fallbackApi,
never()).chatCompletionStream(any(ChatCompletionRequest.class));
+ }
+
@Test
void testExecuteDirectCallSuccess() {
final OpenAiApi mainApi = mock(OpenAiApi.class);