SpringBoot集成WebSocket推送

wen java案例 2

本文目录导读:

SpringBoot集成WebSocket推送

  1. 基础依赖配置
  2. WebSocket配置类
  3. WebSocket拦截器
  4. WebSocket处理器
  5. 推送服务类
  6. 控制器示例
  7. 前端HTML示例
  8. STOMP协议方式(推荐)
  9. 注意事项
  10. 集群部署方案

我来详细介绍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);
    }
}

注意事项

  1. 安全性:建议添加身份验证,防止未授权访问
  2. 心跳检测:WebSocket需要心跳机制维持连接
  3. 异常处理:完善异常处理机制,避免连接断开后资源泄漏
  4. 集群部署:考虑使用Redis或MQ实现跨实例消息推送
  5. 限流保护:避免消息推送过快导致服务压力

集群部署方案

对于集群环境,推荐使用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方式。

抱歉,评论功能暂时关闭!