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 21a7b242ad refactor: replace blocking WebSocket basic send with
non-blocking async send in WebsocketCollector (#6980)
21a7b242ad is described below
commit 21a7b242ad640a8bfcebbfe92a087f3e34096a40
Author: Limbo <[email protected]>
AuthorDate: Sat Sep 5 09:46:12 2026 +0800
refactor: replace blocking WebSocket basic send with non-blocking async
send in WebsocketCollector (#6980)
Co-authored-by: aias00 <[email protected]>
---
.../listener/websocket/WebsocketCollector.java | 87 +++++++++++-
.../listener/websocket/WebsocketCollectorTest.java | 150 ++++++++++++++-------
2 files changed, 184 insertions(+), 53 deletions(-)
diff --git
a/shenyu-admin/src/main/java/org/apache/shenyu/admin/listener/websocket/WebsocketCollector.java
b/shenyu-admin/src/main/java/org/apache/shenyu/admin/listener/websocket/WebsocketCollector.java
index 0631b5ddd0..7260b5ac70 100644
---
a/shenyu-admin/src/main/java/org/apache/shenyu/admin/listener/websocket/WebsocketCollector.java
+++
b/shenyu-admin/src/main/java/org/apache/shenyu/admin/listener/websocket/WebsocketCollector.java
@@ -46,10 +46,11 @@ import jakarta.websocket.OnOpen;
import jakarta.websocket.Session;
import jakarta.websocket.server.ServerEndpoint;
-import java.io.IOException;
+import java.util.ArrayDeque;
import java.util.Map;
import java.util.Objects;
import java.util.Optional;
+import java.util.Queue;
import java.util.Set;
import java.util.concurrent.CopyOnWriteArraySet;
@@ -66,6 +67,8 @@ public class WebsocketCollector {
private static final Set<Session> SESSION_SET = new
CopyOnWriteArraySet<>();
private static final Map<String, Set<Session>> NAMESPACE_SESSION_MAP =
Maps.newConcurrentMap();
+
+ private static final Map<Session, SessionSendQueue> SESSION_SEND_QUEUES =
Maps.newConcurrentMap();
private static final String SESSION_KEY = "sessionKey";
@@ -249,6 +252,7 @@ public class WebsocketCollector {
sendMessageBySession(session, message);
} else {
SESSION_SET.remove(session);
+ removeSessionSendQueue(session);
}
}
} else {
@@ -279,6 +283,7 @@ public class WebsocketCollector {
sendMessageBySession(session, message);
} else {
NAMESPACE_SESSION_MAP.getOrDefault(namespaceId,
Sets.newConcurrentHashSet()).remove(session);
+ removeSessionSendQueue(session);
}
}
} else {
@@ -288,16 +293,20 @@ public class WebsocketCollector {
}
- private static synchronized void sendMessageBySession(final Session
session, final String message) {
- try {
- session.getBasicRemote().sendText(message);
- } catch (IOException e) {
- LOG.error("websocket send result is exception: ", e);
+ private static void sendMessageBySession(final Session session, final
String message) {
+ SESSION_SEND_QUEUES.computeIfAbsent(session,
SessionSendQueue::new).send(message);
+ }
+
+ private static void removeSessionSendQueue(final Session session) {
+ SessionSendQueue sendQueue = SESSION_SEND_QUEUES.remove(session);
+ if (Objects.nonNull(sendQueue)) {
+ sendQueue.close();
}
}
private void clearSession(final Session session) {
SESSION_SET.remove(session);
+ removeSessionSendQueue(session);
String namespaceId = getNamespaceId(session);
if (StringUtils.isNotBlank(namespaceId)) {
NAMESPACE_SESSION_MAP.getOrDefault(namespaceId,
Sets.newConcurrentHashSet()).remove(session);
@@ -325,4 +334,70 @@ public class WebsocketCollector {
return json;
}
}
+
+ private static final class SessionSendQueue {
+
+ private final Session session;
+
+ private final Queue<String> messages = new ArrayDeque<>();
+
+ private boolean sending;
+
+ private boolean closed;
+
+ private SessionSendQueue(final Session session) {
+ this.session = session;
+ }
+
+ private void send(final String message) {
+ boolean startSending = false;
+ synchronized (this) {
+ if (closed) {
+ return;
+ }
+ messages.offer(message);
+ if (!sending) {
+ sending = true;
+ startSending = true;
+ }
+ }
+ if (startSending) {
+ sendNext();
+ }
+ }
+
+ private void sendNext() {
+ final String message;
+ synchronized (this) {
+ if (closed) {
+ sending = false;
+ messages.clear();
+ return;
+ }
+ message = messages.poll();
+ if (Objects.isNull(message)) {
+ sending = false;
+ return;
+ }
+ }
+ try {
+ session.getAsyncRemote().sendText(message, result -> {
+ if (!result.isOK()) {
+ LOG.error("websocket send result is exception: ",
result.getException());
+ }
+ sendNext();
+ });
+ } catch (RuntimeException ex) {
+ LOG.error("websocket send result is exception: ", ex);
+ sendNext();
+ }
+ }
+
+ private void close() {
+ synchronized (this) {
+ closed = true;
+ messages.clear();
+ }
+ }
+ }
}
diff --git
a/shenyu-admin/src/test/java/org/apache/shenyu/admin/listener/websocket/WebsocketCollectorTest.java
b/shenyu-admin/src/test/java/org/apache/shenyu/admin/listener/websocket/WebsocketCollectorTest.java
index c009e61076..72e1f1ac7e 100644
---
a/shenyu-admin/src/test/java/org/apache/shenyu/admin/listener/websocket/WebsocketCollectorTest.java
+++
b/shenyu-admin/src/test/java/org/apache/shenyu/admin/listener/websocket/WebsocketCollectorTest.java
@@ -18,6 +18,8 @@
package org.apache.shenyu.admin.listener.websocket;
import jakarta.websocket.RemoteEndpoint;
+import jakarta.websocket.SendHandler;
+import jakarta.websocket.SendResult;
import jakarta.websocket.Session;
import org.apache.shenyu.admin.config.properties.ClusterProperties;
import org.apache.shenyu.admin.mode.cluster.service.ClusterSelectMasterService;
@@ -34,6 +36,7 @@ import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
+import org.mockito.ArgumentCaptor;
import org.mockito.Mock;
import org.mockito.MockedStatic;
import org.mockito.junit.jupiter.MockitoExtension;
@@ -44,7 +47,6 @@ import org.slf4j.LoggerFactory;
import org.springframework.context.ConfigurableApplicationContext;
import org.springframework.test.util.ReflectionTestUtils;
-import java.io.IOException;
import java.util.HashMap;
import java.util.Map;
import java.util.Objects;
@@ -53,8 +55,11 @@ import java.util.Set;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertThrows;
+import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyString;
+import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.ArgumentMatchers.isA;
+import static org.mockito.Mockito.doAnswer;
import static org.mockito.Mockito.doNothing;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.mockStatic;
@@ -122,6 +127,11 @@ public final class WebsocketCollectorTest {
if (Objects.nonNull(namespaceMap)) {
namespaceMap.clear();
}
+ Map<Session, ?> sessionSendQueues =
+ (Map<Session, ?>)
ReflectionTestUtils.getField(WebsocketCollector.class, "SESSION_SEND_QUEUES");
+ if (Objects.nonNull(sessionSendQueues)) {
+ sessionSendQueues.clear();
+ }
}
@Test
@@ -176,24 +186,23 @@ public final class WebsocketCollectorTest {
}
@Test
- void testOnMessageRunningModeStandalone() throws IOException {
+ void testOnMessageRunningModeStandalone() {
ClusterProperties clusterProperties = mock(ClusterProperties.class);
when(clusterProperties.isEnabled()).thenReturn(false);
when(SpringBeanUtils.getInstance().getBean(ClusterProperties.class)).thenReturn(clusterProperties);
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
- when(session.getBasicRemote()).thenReturn(basic);
+ final RemoteEndpoint.Async async = mockSuccessfulAsyncRemote(session);
websocketCollector.onOpen(session);
websocketCollector.onMessage(DataEventTypeEnum.RUNNING_MODE.name(),
session);
- verify(basic, times(1)).sendText(anyString());
+ verify(async, times(1)).sendText(anyString(), any(SendHandler.class));
websocketCollector.onClose(session);
ThreadLocalUtils.remove("sessionKey");
}
@Test
- void testOnMessageRunningModeCluster() throws IOException {
+ void testOnMessageRunningModeCluster() {
ClusterProperties clusterProperties = mock(ClusterProperties.class);
when(clusterProperties.isEnabled()).thenReturn(true);
ClusterSelectMasterService masterService =
mock(ClusterSelectMasterService.class);
@@ -202,13 +211,12 @@ public final class WebsocketCollectorTest {
when(SpringBeanUtils.getInstance().getBean(ClusterProperties.class)).thenReturn(clusterProperties);
when(SpringBeanUtils.getInstance().getBean(ClusterSelectMasterService.class)).thenReturn(masterService);
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
- when(session.getBasicRemote()).thenReturn(basic);
+ final RemoteEndpoint.Async async = mockSuccessfulAsyncRemote(session);
websocketCollector.onOpen(session);
websocketCollector.onMessage(DataEventTypeEnum.RUNNING_MODE.name(),
session);
- verify(basic, times(1)).sendText(anyString());
+ verify(async, times(1)).sendText(anyString(), any(SendHandler.class));
websocketCollector.onClose(session);
ThreadLocalUtils.remove("sessionKey");
}
@@ -271,28 +279,26 @@ public final class WebsocketCollectorTest {
}
@Test
- void testSendOldApi() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
- when(session.getBasicRemote()).thenReturn(basic);
+ void testSendOldApi() {
+ final RemoteEndpoint.Async async = mockSuccessfulAsyncRemote(session);
when(session.isOpen()).thenReturn(true);
websocketCollector.onOpen(session);
assertEquals(1L, getSessionSetSize());
WebsocketCollector.send(null, DataEventTypeEnum.MYSELF);
- verify(basic, times(0)).sendText(null);
+ verify(async, times(0)).sendText(eq(null), any(SendHandler.class));
ThreadLocalUtils.put("sessionKey", session);
WebsocketCollector.send("test_message_1", DataEventTypeEnum.MYSELF);
- verify(basic, times(1)).sendText("test_message_1");
+ verify(async, times(1)).sendText(eq("test_message_1"),
any(SendHandler.class));
WebsocketCollector.send("test_message_2", DataEventTypeEnum.CREATE);
- verify(basic, times(1)).sendText("test_message_2");
+ verify(async, times(1)).sendText(eq("test_message_2"),
any(SendHandler.class));
doNothing().when(loggerSpy).warn(anyString(), anyString());
websocketCollector.onClose(session);
ThreadLocalUtils.remove("sessionKey");
}
@Test
- void testSendOldApiMyselfClosedSession() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
- when(session.getBasicRemote()).thenReturn(basic);
+ void testSendOldApiMyselfClosedSession() {
+ final RemoteEndpoint.Async async = mockSuccessfulAsyncRemote(session);
websocketCollector.onOpen(session);
// Mark session as closed
@@ -300,25 +306,25 @@ public final class WebsocketCollectorTest {
ThreadLocalUtils.put("sessionKey", session);
WebsocketCollector.send("msg", DataEventTypeEnum.MYSELF);
// closed session → removed from SESSION_SET, no sendText
- verify(basic, never()).sendText("msg");
+ verify(async, never()).sendText(eq("msg"), any(SendHandler.class));
assertEquals(0L, getSessionSetSize());
ThreadLocalUtils.remove("sessionKey");
}
@Test
- void testSendOldApiMyselfNullSession() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
+ void testSendOldApiMyselfNullSession() {
+ final RemoteEndpoint.Async async = mock(RemoteEndpoint.Async.class);
// No session in ThreadLocal
ThreadLocalUtils.remove("sessionKey");
WebsocketCollector.send("msg", DataEventTypeEnum.MYSELF);
- verify(basic, never()).sendText(anyString());
+ verify(async, never()).sendText(anyString(), any(SendHandler.class));
}
@Test
- void testSendWithNamespaceIdBlankMessageNoOp() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
+ void testSendWithNamespaceIdBlankMessageNoOp() {
+ final RemoteEndpoint.Async async = mock(RemoteEndpoint.Async.class);
WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID, "",
DataEventTypeEnum.CREATE);
- verify(basic, never()).sendText(anyString());
+ verify(async, never()).sendText(anyString(), any(SendHandler.class));
}
@Test
@@ -328,71 +334,121 @@ public final class WebsocketCollectorTest {
}
@Test
- void testSendWithNamespaceIdMyself() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
- when(session.getBasicRemote()).thenReturn(basic);
+ void testSendWithNamespaceIdMyself() {
+ final RemoteEndpoint.Async async = mockSuccessfulAsyncRemote(session);
when(session.isOpen()).thenReturn(true);
websocketCollector.onOpen(session);
ThreadLocalUtils.put("sessionKey", session);
WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID, "ns-msg",
DataEventTypeEnum.MYSELF);
- verify(basic, times(1)).sendText("ns-msg");
+ verify(async, times(1)).sendText(eq("ns-msg"), any(SendHandler.class));
websocketCollector.onClose(session);
ThreadLocalUtils.remove("sessionKey");
}
@Test
- void testSendWithNamespaceIdMyselfClosedSession() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
- when(session.getBasicRemote()).thenReturn(basic);
+ void testSendWithNamespaceIdMyselfClosedSession() {
+ final RemoteEndpoint.Async async = mockSuccessfulAsyncRemote(session);
websocketCollector.onOpen(session);
when(session.isOpen()).thenReturn(false);
ThreadLocalUtils.put("sessionKey", session);
WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID, "ns-msg",
DataEventTypeEnum.MYSELF);
- verify(basic, never()).sendText("ns-msg");
+ verify(async, never()).sendText(eq("ns-msg"), any(SendHandler.class));
ThreadLocalUtils.remove("sessionKey");
}
@Test
- void testSendWithNamespaceIdMyselfNullSession() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
+ void testSendWithNamespaceIdMyselfNullSession() {
+ final RemoteEndpoint.Async async = mock(RemoteEndpoint.Async.class);
ThreadLocalUtils.remove("sessionKey");
WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID, "ns-msg",
DataEventTypeEnum.MYSELF);
- verify(basic, never()).sendText(anyString());
+ verify(async, never()).sendText(anyString(), any(SendHandler.class));
}
@Test
- void testSendWithNamespaceIdNonMyselfBroadcast() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
- when(session.getBasicRemote()).thenReturn(basic);
+ void testSendWithNamespaceIdNonMyselfBroadcast() {
+ final RemoteEndpoint.Async async = mockSuccessfulAsyncRemote(session);
when(session.isOpen()).thenReturn(true);
websocketCollector.onOpen(session);
WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID,
"broadcast-msg", DataEventTypeEnum.CREATE);
- verify(basic, times(1)).sendText("broadcast-msg");
+ verify(async, times(1)).sendText(eq("broadcast-msg"),
any(SendHandler.class));
websocketCollector.onClose(session);
}
@Test
- void testSendBySessionIOException() throws IOException {
- RemoteEndpoint.Basic basic = mock(RemoteEndpoint.Basic.class);
- when(session.getBasicRemote()).thenReturn(basic);
+ void testSendBySessionFailure() {
+ final RemoteEndpoint.Async async = mock(RemoteEndpoint.Async.class);
+ when(session.getAsyncRemote()).thenReturn(async);
when(session.isOpen()).thenReturn(true);
websocketCollector.onOpen(session);
- // IOException on sendText should be caught internally
- org.mockito.Mockito.doThrow(new IOException("io
error")).when(basic).sendText(anyString());
+ doAnswer(invocation -> {
+ SendHandler handler = invocation.getArgument(1);
+ handler.onResult(new SendResult(new IllegalStateException("send
error")));
+ return null;
+ }).when(async).sendText(anyString(), any(SendHandler.class));
WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID,
"fail-msg", DataEventTypeEnum.CREATE);
- // no exception propagated; verify attempted send
- verify(basic, times(1)).sendText("fail-msg");
+ verify(async, times(1)).sendText(eq("fail-msg"),
any(SendHandler.class));
+
+ websocketCollector.onClose(session);
+ }
+
+ @Test
+ void testSendDoesNotWaitForOtherSession() {
+ Session anotherSession = mock(Session.class);
+ Map<String, Object> userProperties = session.getUserProperties();
+ when(anotherSession.isOpen()).thenReturn(true);
+ when(anotherSession.getUserProperties()).thenReturn(userProperties);
+ RemoteEndpoint.Async firstAsync = mock(RemoteEndpoint.Async.class);
+ RemoteEndpoint.Async secondAsync = mock(RemoteEndpoint.Async.class);
+ when(session.getAsyncRemote()).thenReturn(firstAsync);
+ when(anotherSession.getAsyncRemote()).thenReturn(secondAsync);
+ websocketCollector.onOpen(session);
+ websocketCollector.onOpen(anotherSession);
+
+ WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID,
"broadcast-msg", DataEventTypeEnum.CREATE);
+
+ verify(firstAsync).sendText(eq("broadcast-msg"),
any(SendHandler.class));
+ verify(secondAsync).sendText(eq("broadcast-msg"),
any(SendHandler.class));
+ websocketCollector.onClose(session);
+ websocketCollector.onClose(anotherSession);
+ }
+
+ @Test
+ void testSendMessagesSequentiallyForSameSession() {
+ final RemoteEndpoint.Async async = mock(RemoteEndpoint.Async.class);
+ when(session.getAsyncRemote()).thenReturn(async);
+ websocketCollector.onOpen(session);
+ WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID,
"first-message", DataEventTypeEnum.CREATE);
+ WebsocketCollector.send(Constants.SYS_DEFAULT_NAMESPACE_ID,
"second-message", DataEventTypeEnum.CREATE);
+
+ ArgumentCaptor<SendHandler> handlerCaptor =
ArgumentCaptor.forClass(SendHandler.class);
+ verify(async).sendText(eq("first-message"), handlerCaptor.capture());
+ verify(async, never()).sendText(eq("second-message"),
any(SendHandler.class));
+
+ handlerCaptor.getValue().onResult(new SendResult());
+
+ verify(async).sendText(eq("second-message"), any(SendHandler.class));
websocketCollector.onClose(session);
}
+ private RemoteEndpoint.Async mockSuccessfulAsyncRemote(final Session
targetSession) {
+ final RemoteEndpoint.Async async = mock(RemoteEndpoint.Async.class);
+ when(targetSession.getAsyncRemote()).thenReturn(async);
+ doAnswer(invocation -> {
+ SendHandler handler = invocation.getArgument(1);
+ handler.onResult(new SendResult());
+ return null;
+ }).when(async).sendText(anyString(), any(SendHandler.class));
+ return async;
+ }
+
private long getSessionSetSize() {
Set sessionSet = (Set)
ReflectionTestUtils.getField(WebsocketCollector.class, "SESSION_SET");
return Objects.isNull(sessionSet) ? -1 : sessionSet.size();