本文目录导读:

我来详细介绍SpringBoot集成WebSocket实现推送的完整方案。
基础依赖配置
Maven依赖
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-websocket</artifactId>
</dependency>
WebSocket配置类
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.web.socket.config.annotation.EnableWebSocket;
import org.springframework.web.socket.config.annotation.WebSocketConfigurer;
import org.springframework.web.socket.config.annotation.WebSocketHandlerRegistry;
import org.springframework.web.socket.server.standard.ServerEndpointExporter;
@Configuration
@EnableWebSocket
public class WebSocketConfig implements WebSocketConfigurer {
@Override
public void registerWebSocketHandlers(WebSocketHandlerRegistry registry) {
registry.addHandler(webSocketHandler(), "/ws/push")
.setAllowedOrigins("*")
.addInterceptors(new WebSocketInterceptor());
}
@Bean
public WebSocketHandler webSocketHandler() {
return new MyWebSocketHandler();
}
@Bean
public ServerEndpointExporter serverEndpointExporter() {
return new ServerEndpointExporter();
}
}
WebSocket拦截器
import org.springframework.http.server.ServerHttpRequest;
import org.springframework.http.server.ServerHttpResponse;
import org.springframework.web.socket.WebSocketHandler;
import org.springframework.web.socket.server.HandshakeInterceptor;
import java.util.Map;
public class WebSocketInterceptor implements HandshakeInterceptor {
@Override
public boolean beforeHandshake(ServerHttpRequest request, ServerHttpResponse response,
WebSocketHandler wsHandler, Map<String, Object> attributes) throws Exception {
// 获取请求参数中的用户ID
String userId = request.getURI().getQuery();
if (userId != null && userId.startsWith("userId=")) {
userId = userId.substring(7);
attributes.put("userId", userId);
return true;
}
return false;
}
@Override
public void afterHandshake(ServerHttpRequest request, ServerHttpResponse response,
WebSocketHandler wsHandler, Exception exception) {
// 握手后的处理
}
}
WebSocket处理器
import org.springframework.web.socket.*;
import org.springframework.web.socket.handler.TextWebSocketHandler;
import java.io.IOException;
import java.util.concurrent.ConcurrentHashMap;
public class MyWebSocketHandler extends TextWebSocketHandler {
// 存储所有连接的客户端
private static final ConcurrentHashMap<String, WebSocketSession> sessions = new ConcurrentHashMap<>();
@Override
public void afterConnectionEstablished(WebSocketSession session) throws Exception {
String userId = (String) session.getAttributes().get("userId");
if (userId != null) {
sessions.put(userId, session);
System.out.println("用户 " + userId + " 已连接");
}
}
@Override
protected void handleTextMessage(WebSocketSession session, TextMessage message) throws Exception {
// 处理收到的消息
String payload = message.getPayload();
System.out.println("收到消息: " + payload);
// 可以回复消息
session.sendMessage(new TextMessage("服务器已收到: " + payload));
}
@Override
public void afterConnectionClosed(WebSocketSession session, CloseStatus status) throws Exception {
String userId = (String) session.getAttributes().get("userId");
if (userId != null) {
sessions.remove(userId);
System.out.println("用户 " + userId + " 已断开连接");
}
}
@Override
public void handleTransportError(WebSocketSession session, Throwable exception) throws Exception {
String userId = (String) session.getAttributes().get("userId");
if (userId != null) {
sessions.remove(userId);
}
session.close(CloseStatus.SERVER_ERROR);
}
// 向指定用户发送消息
public void sendToUser(String userId, String message) throws IOException {
WebSocketSession session = sessions.get(userId);
if (session != null && session.isOpen()) {
session.sendMessage(new TextMessage(message));
}
}
// 向所有用户广播消息
public void broadcast(String message) throws IOException {
for (WebSocketSession session : sessions.values()) {
if (session.isOpen()) {
session.sendMessage(new TextMessage(message));
}
}
}
// 获取在线用户数量
public int getOnlineCount() {
return sessions.size();
}
}
推送服务类
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.stereotype.Service;
import java.io.IOException;
@Service
public class WebSocketPushService {
@Autowired
private MyWebSocketHandler webSocketHandler;
/**
* 向指定用户推送消息
*/
public void pushToUser(String userId, String message) {
try {
webSocketHandler.sendToUser(userId, message);
} catch (IOException e) {
e.printStackTrace();
}
}
/**
* 向所有用户广播消息
*/
public void pushToAll(String message) {
try {
webSocketHandler.broadcast(message);
} catch (IOException e) {
e.printStackTrace();
}
}
/**
* 获取在线用户数量
*/
public int getOnlineCount() {
return webSocketHandler.getOnlineCount();
}
}
控制器示例
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.web.bind.annotation.*;
@RestController
@RequestMapping("/api/push")
public class PushController {
@Autowired
private WebSocketPushService pushService;
@PostMapping("/user/{userId}")
public String pushToUser(@PathVariable String userId, @RequestBody String message) {
pushService.pushToUser(userId, message);
return "消息已发送给用户: " + userId;
}
@PostMapping("/all")
public String pushToAll(@RequestBody String message) {
pushService.pushToAll(message);
return "消息已广播给所有用户";
}
@GetMapping("/online")
public int getOnlineCount() {
return pushService.getOnlineCount();
}
}
前端HTML示例
<!DOCTYPE html>
<html>
<head>
<meta charset="UTF-8">WebSocket示例</title>
</head>
<body>
<div>
<h2>WebSocket测试</h2>
<div>
<input type="text" id="userId" placeholder="输入用户ID" value="user001">
<button onclick="connect()">连接</button>
<button onclick="disconnect()">断开</button>
</div>
<div style="margin-top: 20px;">
<div id="messages" style="border: 1px solid #ccc; height: 200px; overflow-y: auto;"></div>
</div>
</div>
<script>
let websocket = null;
let userId = '';
function connect() {
userId = document.getElementById('userId').value;
if (!userId) {
alert('请输入用户ID');
return;
}
// 创建WebSocket连接
websocket = new WebSocket('ws://localhost:8080/ws/push?userId=' + userId);
// 连接成功
websocket.onopen = function(event) {
addMessage('连接成功,用户ID: ' + userId);
};
// 收到消息
websocket.onmessage = function(event) {
addMessage('收到消息: ' + event.data);
};
// 连接关闭
websocket.onclose = function(event) {
addMessage('连接已关闭');
websocket = null;
};
// 连接错误
websocket.onerror = function(event) {
addMessage('连接错误');
};
}
function disconnect() {
if (websocket) {
websocket.close();
websocket = null;
}
}
function addMessage(message) {
const div = document.getElementById('messages');
div.innerHTML += '<div>' + new Date().toLocaleTimeString() + ' - ' + message + '</div>';
div.scrollTop = div.scrollHeight;
}
</script>
</body>
</html>
STOMP协议方式(推荐)
添加依赖
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-websocket</artifactId>
</dependency>
<dependency>
<groupId>org.webjars</groupId>
<artifactId>webjars-locator-core</artifactId>
</dependency>
<dependency>
<groupId>org.webjars</groupId>
<artifactId>sockjs-client</artifactId>
<version>1.1.2</version>
</dependency>
<dependency>
<groupId>org.webjars</groupId>
<artifactId>stomp-websocket</artifactId>
<version>2.3.3</version>
</dependency>
STOMP配置
import org.springframework.context.annotation.Configuration;
import org.springframework.messaging.simp.config.MessageBrokerRegistry;
import org.springframework.web.socket.config.annotation.EnableWebSocketMessageBroker;
import org.springframework.web.socket.config.annotation.StompEndpointRegistry;
import org.springframework.web.socket.config.annotation.WebSocketMessageBrokerConfigurer;
@Configuration
@EnableWebSocketMessageBroker
public class StompWebSocketConfig implements WebSocketMessageBrokerConfigurer {
@Override
public void configureMessageBroker(MessageBrokerRegistry config) {
// 订阅前缀
config.enableSimpleBroker("/topic", "/queue");
// 应用前缀
config.setApplicationDestinationPrefixes("/app");
// 用户前缀
config.setUserDestinationPrefix("/user");
}
@Override
public void registerStompEndpoints(StompEndpointRegistry registry) {
registry.addEndpoint("/stomp-endpoint")
.setAllowedOriginPatterns("*")
.withSockJS();
}
}
STOMP控制器
import org.springframework.messaging.handler.annotation.MessageMapping;
import org.springframework.messaging.handler.annotation.SendTo;
import org.springframework.messaging.simp.SimpMessagingTemplate;
import org.springframework.stereotype.Controller;
@Controller
public class StompController {
private final SimpMessagingTemplate messagingTemplate;
public StompController(SimpMessagingTemplate messagingTemplate) {
this.messagingTemplate = messagingTemplate;
}
@MessageMapping("/send")
@SendTo("/topic/messages")
public String sendMessage(String message) {
return "服务器: " + message;
}
// 向指定用户发送消息
public void sendToUser(String userId, String message) {
messagingTemplate.convertAndSendToUser(userId, "/queue/messages", message);
}
// 广播消息
public void broadcast(String message) {
messagingTemplate.convertAndSend("/topic/messages", message);
}
}
注意事项
- 安全性:建议添加身份验证,防止未授权访问
- 心跳检测:WebSocket需要心跳机制维持连接
- 异常处理:完善异常处理机制,避免连接断开后资源泄漏
- 集群部署:考虑使用Redis或MQ实现跨实例消息推送
- 限流保护:避免消息推送过快导致服务压力
集群部署方案
对于集群环境,推荐使用Redis发布订阅:
@Service
public class RedisWebSocketService {
@Autowired
private RedisTemplate<String, String> redisTemplate;
@Autowired
private WebSocketPushService pushService;
// 发布消息
public void publishMessage(String channel, String message) {
redisTemplate.convertAndSend(channel, message);
}
// 订阅消息
@Bean
public MessageListenerAdapter messageListener() {
return new MessageListenerAdapter(new RedisMessageSubscriber(pushService));
}
}
class RedisMessageSubscriber implements MessageListener {
private final WebSocketPushService pushService;
public RedisMessageSubscriber(WebSocketPushService pushService) {
this.pushService = pushService;
}
@Override
public void onMessage(Message message, byte[] pattern) {
String body = new String(message.getBody());
pushService.pushToAll(body);
}
}
这样完整的SpringBoot WebSocket推送方案就实现了,根据实际需求选择普通WebSocket或STOMP方式。