Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
import lombok.extern.slf4j.Slf4j;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;

import java.io.IOException;
import java.util.concurrent.atomic.AtomicBoolean;

/**
Expand Down Expand Up @@ -74,8 +75,9 @@ public void sendEvent(String eventName, Object data) {
return;
}
emitter.send(SseEmitter.event().name(eventName).data(data));
} catch (Exception e) {
fail(e);
} catch (IOException e) {
closed.set(true);
log.warn("SSE connection closed while sending event", e);
}
}

Expand All @@ -92,34 +94,4 @@ public void complete() {
}
}

/**
* 异常结束并关闭 SSE 连接
*
* <p>当发生异常时调用此方法,会执行以下操作:</p>
* <ol>
* <li>关闭 SSE 连接并通知客户端异常信息</li>
* <li>不再抛出异常,避免在流式响应已开始后触发全局异常处理器导致响应冲突</li>
* </ol>
*
* @param throwable 导致失败的异常对象
*/
public void fail(Throwable throwable) {
closeWithError(throwable);
log.warn("SSE send failed", throwable);
}

/**
* 内部方法:以异常方式关闭连接
* <p>
* 使用 CAS 操作确保连接只被关闭一次
* 调用 SseEmitter 的 completeWithError 方法,通知客户端连接异常终止
*
* @param throwable 导致连接关闭的异常对象
*/
private void closeWithError(Throwable throwable) {
// 使用 CAS 原子操作,确保只关闭一次
if (closed.compareAndSet(false, true)) {
emitter.completeWithError(throwable);
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
/*
* 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 com.nageoffer.ai.ragent.framework.web;

import org.junit.jupiter.api.Test;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;

import java.io.IOException;

import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.doThrow;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.verify;

class SseEmitterSenderTest {

@Test
void sendIOExceptionShouldNotTriggerAsyncErrorDispatch() throws Exception {
SseEmitter emitter = mock(SseEmitter.class);
doThrow(new IOException("client disconnected"))
.when(emitter).send(any(SseEmitter.SseEventBuilder.class));
SseEmitterSender sender = new SseEmitterSender(emitter);

sender.sendEvent("message", "content");
sender.sendEvent("message", "ignored");

verify(emitter).send(any(SseEmitter.SseEventBuilder.class));
verify(emitter, never()).complete();
verify(emitter, never()).completeWithError(any());
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
/*
* 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 com.nageoffer.ai.ragent.rag.dto;

/**
* 流式对话错误事件载荷
*
* @param error 面向用户的错误提示
*/
public record StreamErrorPayload(String error) {
}
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,11 @@ public enum SSEEventType {
*/
CANCEL("cancel"),

/**
* 错误事件
*/
ERROR("error"),

/**
* 拒绝事件
*/
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
import com.nageoffer.ai.ragent.rag.dto.CompletionPayload;
import com.nageoffer.ai.ragent.rag.dto.MessageDelta;
import com.nageoffer.ai.ragent.rag.dto.MetaPayload;
import com.nageoffer.ai.ragent.rag.dto.StreamErrorPayload;
import com.nageoffer.ai.ragent.rag.enums.SSEEventType;
import com.nageoffer.ai.ragent.framework.context.UserContext;
import com.nageoffer.ai.ragent.framework.convention.ChatMessage;
Expand All @@ -44,6 +45,7 @@ public class StreamChatEventHandler implements StreamCallback {

private static final String TYPE_THINK = "think";
private static final String TYPE_RESPONSE = "response";
private static final String ERROR_MESSAGE = "生成失败,请稍后重试";

private final int messageChunkSize;
private final SseEmitterSender sender;
Expand Down Expand Up @@ -237,8 +239,10 @@ public void onError(Throwable t) {
if (taskManager.isCancelled(taskId)) {
return;
}
log.error("流式对话失败,conversationId:{},taskId:{}", conversationId, taskId, t);
taskManager.unregister(taskId);
sender.fail(t);
sender.sendEvent(SSEEventType.ERROR.value(), new StreamErrorPayload(ERROR_MESSAGE));
sender.complete();
}

private void sendChunked(String type, String content) {
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
/*
* 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 com.nageoffer.ai.ragent.rag.service.handler;

import com.nageoffer.ai.ragent.framework.web.StreamTaskManager;
import com.nageoffer.ai.ragent.infra.config.AIModelProperties;
import com.nageoffer.ai.ragent.rag.core.memory.ConversationMemoryService;
import com.nageoffer.ai.ragent.rag.service.ConversationGroupService;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.ArgumentCaptor;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;
import org.springframework.test.util.ReflectionTestUtils;
import org.springframework.web.servlet.mvc.method.annotation.ResponseBodyEmitter;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;

import java.util.List;

import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;

@ExtendWith(MockitoExtension.class)
class StreamChatEventHandlerTest {

private static final String PUBLIC_ERROR_MESSAGE = "生成失败,请稍后重试";

@Mock
private SseEmitter emitter;

@Mock
private AIModelProperties modelProperties;

@Mock
private ConversationMemoryService memoryService;

@Mock
private ConversationGroupService conversationGroupService;

@Mock
private StreamTaskManager taskManager;

@Test
void businessFailureShouldSendSafeErrorEventAndCompleteNormally() throws Exception {
StreamChatEventHandler handler = new StreamChatEventHandler(StreamChatHandlerParams.builder()
.emitter(emitter)
.conversationId("conversation-1")
.taskId("task-1")
.modelProperties(modelProperties)
.memoryService(memoryService)
.conversationGroupService(conversationGroupService)
.taskManager(taskManager)
.build());

handler.onError(new IllegalStateException("database password leaked"));

ArgumentCaptor<SseEmitter.SseEventBuilder> eventCaptor =
ArgumentCaptor.forClass(SseEmitter.SseEventBuilder.class);
verify(emitter, times(2)).send(eventCaptor.capture());
List<Object> errorEventData = eventCaptor.getAllValues().get(1).build().stream()
.map(ResponseBodyEmitter.DataWithMediaType::getData)
.toList();
assertTrue(errorEventData.stream()
.filter(String.class::isInstance)
.map(String.class::cast)
.anyMatch(data -> data.startsWith("event:error\n")));
Object payload = errorEventData.stream()
.filter(data -> !(data instanceof String))
.findFirst()
.orElseThrow();
assertEquals(PUBLIC_ERROR_MESSAGE, ReflectionTestUtils.getField(payload, "error"));
assertFalse(payload.toString().contains("database password leaked"));
verify(taskManager).unregister("task-1");
verify(emitter).complete();
verify(emitter, never()).completeWithError(any());
}
}