1. 项目概述:从零构建一个健壮的WebSocket消息推送服务
最近在做一个需要实时消息推送的后台管理系统,比如订单状态变更、系统告警、客服聊天这些场景,前端得立刻知道。用HTTP轮询太笨重,长轮询也麻烦,最后选了WebSocket。但真上手才发现,光把连接建起来只是第一步,离“能用”和“好用”还差得远。比如,用户怎么认证?连接断了怎么知道?怎么给特定的一群人发消息而不是广播?这些才是工程实践里的硬骨头。
这个项目就是用SpringBoot搭一个完整的WebSocket服务端,重点解决四个核心问题:连接建立时的身份校验、维持连接可用的心跳机制、服务端主动推送消息,以及按业务逻辑对用户进行分组管理。网上很多教程只讲怎么握手成功,但生产环境里,没心跳的连接说断就断,没校验的服务谁都能连,分组广播实现不好性能就崩。我会结合我趟过的坑,把每个环节的原理、代码和配置细节掰开揉碎了讲,目标是让你看完就能搭出一个稳定、安全、可维护的WebSocket推送服务。
2. 核心组件选型与环境搭建
在SpringBoot里集成WebSocket,主流有两种方式:一是直接用Spring提供的spring-boot-starter-websocket,它底层封装了标准的Java WebSocket API(JSR-356),并提供了更Spring风格的编程模型;二是通过STOMP子协议,它更像一个消息代理,适合复杂的消息路由场景。对于我们的需求——点对点、分组推送、心跳保活——直接使用spring-boot-starter-websocket更轻量、更可控,也更容易理解底层机制。
首先,在pom.xml里引入依赖。这里注意,我们通常还会引入spring-boot-starter-security来做连接时的安全校验,但为了聚焦WebSocket核心,我们先手动实现一个简单的token校验。
<dependency> <groupId>org.springframework.boot</groupId> <artifactId>spring-boot-starter-websocket</artifactId> </dependency> <!-- 用于JSON消息序列化 --> <dependency> <groupId>com.fasterxml.jackson.core</groupId> <artifactId>jackson-databind</artifactId> </dependency>接着,需要一个核心配置类来开启WebSocket支持。这里我直接给出一个增强版的配置,它做了三件事:注册我们的WebSocket处理器、配置允许跨域(前端独立部署时必需)、以及设置消息缓冲区大小。
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; @Configuration @EnableWebSocket public class WebSocketConfig implements WebSocketConfigurer { @Override public void registerWebSocketHandlers(WebSocketHandlerRegistry registry) { registry.addHandler(new MyWebSocketHandler(), "/ws") .addInterceptors(new AuthHandshakeInterceptor()) // 添加握手拦截器 .setAllowedOrigins("*"); // 生产环境应指定具体域名 } }这个配置类里,MyWebSocketHandler是我们处理消息的核心类,AuthHandshakeInterceptor则负责在握手阶段拦截请求,进行身份校验。把校验放在握手阶段,比连接建立后再处理要安全得多,无效的连接请求在握手时就会被拒绝,不会占用服务器资源。
注意:
setAllowedOrigins("*")在开发时图个方便,上线前必须改为具体的前端域名(如"https://yourdomain.com"),这是防止跨站WebSocket劫持(Cross-Site WebSocket Hijacking)最基本的一步。从相关热词里看到bp靶场cross-site websocket hijacking,指的就是这种攻击,攻击者可以在用户浏览器中构造恶意网站,利用用户已有的身份(比如Cookie)向你的WebSocket服务发起连接并窃听或篡改消息。严格限制Origin是首要防线。
3. 握手拦截器:实现连接身份校验
WebSocket协议本身不处理身份认证,我们需要在HTTP升级为WebSocket协议的那个握手请求(Handshake Request)里做文章。Spring WebSocket提供了HandshakeInterceptor接口,允许我们在握手前和握手后插入逻辑。
我实现了一个简单的基于Token的校验拦截器。思路是:前端在建立WebSocket连接时,不能像普通HTTP请求那样在Header里带Authorization,但可以将Token作为一个查询参数(Query Parameter)附在连接URL上,例如:ws://localhost:8080/ws?token=eyJhbGciOiJ...。我们在拦截器里解析这个token并进行验证。
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 AuthHandshakeInterceptor implements HandshakeInterceptor { @Override public boolean beforeHandshake(ServerHttpRequest request, ServerHttpResponse response, WebSocketHandler wsHandler, Map<String, Object> attributes) throws Exception { // 从请求URI中获取token参数 String query = request.getURI().getQuery(); if (query == null || !query.contains("token=")) { // 可以返回false拒绝握手,也可以返回401状态码 response.setStatusCode(HttpStatus.UNAUTHORIZED); return false; } // 简单解析token,实际项目应使用JWT等规范方案 String token = query.substring(query.indexOf("token=") + 6); // 这里模拟一个简单的校验逻辑 if (!isValidToken(token)) { response.setStatusCode(HttpStatus.FORBIDDEN); return false; } // 校验通过,可以将用户信息放入attributes,后续在Handler中可取用 String userId = extractUserIdFromToken(token); attributes.put("userId", userId); return true; } @Override public void afterHandshake(ServerHttpRequest request, ServerHttpResponse response, WebSocketHandler wsHandler, Exception exception) { // 握手成功后调用,可用于记录日志等 } private boolean isValidToken(String token) { // 实现你的token验证逻辑,例如调用认证服务或校验JWT签名 return token != null && token.startsWith("valid_"); } private String extractUserIdFromToken(String token) { // 从token中解析出用户ID,这里简单演示 return token.replace("valid_", ""); } }这里有个关键点:attributes参数。它是一个Map,数据在握手阶段被放入后,会在WebSocket Session建立后,传递给我们自定义的WebSocketHandler。这样,我们就成功地将用户身份信息从HTTP握手请求传递到了WebSocket会话中,为后续的用户分组和定向推送打下了基础。
实操心得:把Token放在URL参数里是一种常见做法,但要注意其可能被浏览器历史记录、日志服务器记录,存在泄漏风险。对于安全性要求极高的场景,可以考虑在握手前,先通过一个普通的HTTP接口认证,服务端返回一个一次性的、短效的
connectTicket,前端再用这个ticket来建立WebSocket连接。不过,大多数内部管理系统,使用HTTPS + Token的方式已经足够。
4. 消息处理器:连接、消息与异常的生命周期管理
WebSocketHandler是处理所有WebSocket事件的核心。Spring提供了一个方便的适配器类TextWebSocketHandler,我们继承它并重写关键方法。这里我设计了一个不仅能处理消息,还能管理用户会话和心跳的增强处理器。
import org.springframework.web.socket.handler.TextWebSocketHandler; import org.springframework.web.socket.WebSocketSession; import org.springframework.web.socket.TextMessage; import java.util.concurrent.ConcurrentHashMap; public class MyWebSocketHandler extends TextWebSocketHandler { // 存储用户ID与WebSocketSession的映射 private static final ConcurrentHashMap<String, WebSocketSession> userSessionMap = new ConcurrentHashMap<>(); // 存储Session与最后活跃时间戳,用于心跳检测 private static final ConcurrentHashMap<String, Long> sessionLastActiveTime = new ConcurrentHashMap<>(); @Override public void afterConnectionEstablished(WebSocketSession session) throws Exception { // 连接建立成功 String userId = (String) session.getAttributes().get("userId"); if (userId != null) { userSessionMap.put(userId, session); sessionLastActiveTime.put(session.getId(), System.currentTimeMillis()); System.out.println("用户 " + userId + " 连接成功,Session ID: " + session.getId()); // 可以在这里向该用户发送一条欢迎消息或连接成功通知 session.sendMessage(new TextMessage("{\"type\":\"system\", \"msg\":\"WebSocket连接已建立\"}")); } else { // 理论上不会走到这里,因为拦截器已校验 session.close(CloseStatus.NOT_ACCEPTABLE.withReason("未识别的用户")); } } @Override protected void handleTextMessage(WebSocketSession session, TextMessage message) throws Exception { // 处理客户端发来的文本消息 String payload = message.getPayload(); String userId = (String) session.getAttributes().get("userId"); System.out.println("收到来自用户 " + userId + " 的消息: " + payload); // 更新该会话的最后活跃时间 sessionLastActiveTime.put(session.getId(), System.currentTimeMillis()); // 这里可以解析消息内容,根据不同的type执行不同逻辑 // 例如:{"type": "ping"} 表示心跳,{"type": "chat", "to": "user2", "content": "hello"} 表示私聊 // 具体业务逻辑解析... } @Override public void handleTransportError(WebSocketSession session, Throwable exception) throws Exception { // 传输过程出错,比如解码失败 System.err.println("WebSocket传输错误,Session ID: " + session.getId() + ", 错误: " + exception.getMessage()); } @Override public void afterConnectionClosed(WebSocketSession session, CloseStatus status) throws Exception { // 连接关闭 String userId = (String) session.getAttributes().get("userId"); if (userId != null) { userSessionMap.remove(userId); sessionLastActiveTime.remove(session.getId()); System.out.println("用户 " + userId + " 连接关闭,状态码: " + status.getCode() + ", 原因: " + status.getReason()); } } }这个处理器框架已经具备了会话管理的基础。userSessionMap让我们能通过用户ID找到对应的WebSocketSession,这是实现点对点推送的关键。sessionLastActiveTime则为后续实现心跳超时断开提供了数据支持。
5. 心跳机制:PING-PONG保活与连接健康度检测
WebSocket连接可能因为网络波动、代理超时、客户端崩溃等原因无声无息地断开(即“死连接”)。服务器和客户端都无法立即感知。心跳机制(Heartbeat)就是双方定期发送一个小数据包(PING/PONG帧)来确认对方是否还在线。
WebSocket协议层面其实定义了PING和PONG控制帧,但Java WebSocket API(JSR-356)并没有直接向应用层暴露发送PING帧的接口。通常,我们在应用层自己实现一个基于文本或二进制消息的“伪心跳”。
我的实现方案是双重的:
- 服务端主动探测:用一个定时任务,定期检查所有会话的最后活跃时间,如果超过阈值(比如30秒),则主动向客户端发送一个PING消息,并等待PONG回应。如果一定时间内没收到PONG,则认为连接已死,主动关闭它。
- 客户端主动上报:要求前端每隔一段时间(比如25秒)主动向服务端发送一个
{"type":"ping"}的消息。服务端收到后,更新该会话的“最后活跃时间”。这种方式更简单,但依赖于客户端的配合。
这里重点讲服务端主动探测的实现。我们需要一个Spring的定时任务组件。
首先,在处理器里增加发送PING和接收PONG的逻辑:
public class MyWebSocketHandler extends TextWebSocketHandler { // ... 之前的成员变量和方法 ... @Override protected void handleTextMessage(WebSocketSession session, TextMessage message) throws Exception { String payload = message.getPayload(); // 更新活跃时间 sessionLastActiveTime.put(session.getId(), System.currentTimeMillis()); // 解析消息 ObjectMapper mapper = new ObjectMapper(); try { JsonNode node = mapper.readTree(payload); String type = node.get("type").asText(); if ("pong".equals(type)) { // 收到客户端对服务端PING的回应 System.out.println("收到来自会话 " + session.getId() + " 的PONG回应"); return; // 心跳回应,不处理其他业务 } else if ("ping".equals(type)) { // 收到客户端主动发来的PING,立即回复PONG session.sendMessage(new TextMessage("{\"type\":\"pong\"}")); return; } // ... 其他业务消息处理 ... } catch (Exception e) { // 消息格式错误处理 } } }然后,创建一个定时任务类,定期扫描并清理死连接:
import org.springframework.scheduling.annotation.EnableScheduling; import org.springframework.scheduling.annotation.Scheduled; import org.springframework.stereotype.Component; import org.springframework.web.socket.WebSocketSession; import java.io.IOException; import java.util.Iterator; import java.util.Map; @Component @EnableScheduling public class WebSocketHeartbeatTask { // 心跳超时时间,单位毫秒 private static final long HEARTBEAT_TIMEOUT = 35000; // 35秒 // 发送PING的间隔,应小于超时时间 private static final long PING_INTERVAL = 30000; // 30秒 @Scheduled(fixedRate = 15000) // 每15秒执行一次检查 public void checkAlive() { long currentTime = System.currentTimeMillis(); Iterator<Map.Entry<String, Long>> iterator = MyWebSocketHandler.sessionLastActiveTime.entrySet().iterator(); while (iterator.hasNext()) { Map.Entry<String, Long> entry = iterator.next(); String sessionId = entry.getKey(); Long lastActiveTime = entry.getValue(); WebSocketSession session = findSessionById(sessionId); // 需要根据sessionId找到session对象 if (session == null || !session.isOpen()) { // 会话已不存在或已关闭,清理记录 iterator.remove(); MyWebSocketHandler.userSessionMap.values().removeIf(s -> s.getId().equals(sessionId)); continue; } long inactiveDuration = currentTime - lastActiveTime; if (inactiveDuration > HEARTBEAT_TIMEOUT) { // 超过超时时间,直接关闭连接 try { session.close(CloseStatus.SESSION_NOT_RELIABLE); System.out.println("因心跳超时关闭会话: " + sessionId); } catch (IOException e) { e.printStackTrace(); } iterator.remove(); MyWebSocketHandler.userSessionMap.values().removeIf(s -> s.getId().equals(sessionId)); } else if (inactiveDuration > PING_INTERVAL) { // 超过PING间隔但未超时,发送一个PING探测 try { session.sendMessage(new TextMessage("{\"type\":\"ping\"}")); System.out.println("向会话 " + sessionId + " 发送PING探测"); } catch (IOException e) { // 发送失败,可能连接已失效,下次检查会处理 e.printStackTrace(); } } } } // 一个辅助方法,根据sessionId从userSessionMap中查找session // 注意:userSessionMap存储的是userId->session,这里需要遍历查找,实际可优化数据结构 private WebSocketSession findSessionById(String targetSessionId) { for (WebSocketSession session : MyWebSocketHandler.userSessionMap.values()) { if (session.getId().equals(targetSessionId)) { return session; } } return null; } }踩坑实录:心跳超时时间
HEARTBEAT_TIMEOUT和PING间隔PING_INTERVAL的设定需要谨慎。间隔太短,会产生大量无用心跳包,增加服务器和网络负担;间隔太长,则无法及时发现死连接。一般建议PING间隔为25-30秒,超时时间比间隔多5-10秒,给网络延迟和客户端处理留出余量。另外,像Nginx这样的反向代理默认会对WebSocket连接有一个60秒的超时(proxy_read_timeout),你的服务端心跳超时必须小于这个值,否则连接会被代理服务器先掐断。
6. 用户分组与定向消息推送
广播消息很简单,遍历userSessionMap的所有session发送即可。但实际业务中,更多是需要按条件推送:比如推送给某个在线的特定用户(点对点)、推送给某个部门的所有人、推送给具有某个角色标签的所有用户。这就是用户分组。
我的实现思路是,在用户连接建立时,不仅记录userId->session的映射,还根据业务规则,将用户加入到不同的“逻辑组”中。这个组可以用一个ConcurrentHashMap<String, Set<String>>来维护,key是组名(如"dept:finance","role:admin"),value是该组内所有用户的ID集合。
首先,在Handler中增加分组管理容器:
public class MyWebSocketHandler extends TextWebSocketHandler { // ... 已有成员变量 ... // 分组存储:组名 -> 用户ID集合 private static final ConcurrentHashMap<String, Set<String>> userGroupMap = new ConcurrentHashMap<>(); @Override public void afterConnectionEstablished(WebSocketSession session) throws Exception { String userId = (String) session.getAttributes().get("userId"); if (userId != null) { userSessionMap.put(userId, session); sessionLastActiveTime.put(session.getId(), System.currentTimeMillis()); // --- 关键:根据业务逻辑将用户加入分组 --- // 假设我们从数据库或上下文中获取用户的部门、角色等信息 List<String> userGroups = getUserGroupsFromDatabase(userId); // 伪方法 for (String group : userGroups) { userGroupMap.computeIfAbsent(group, k -> ConcurrentHashMap.newKeySet()).add(userId); } System.out.println("用户 " + userId + " 加入分组: " + userGroups); // --- 分组结束 --- } } @Override public void afterConnectionClosed(WebSocketSession session, CloseStatus status) throws Exception { String userId = (String) session.getAttributes().get("userId"); if (userId != null) { userSessionMap.remove(userId); sessionLastActiveTime.remove(session.getId()); // --- 关键:用户断开时从所有分组中移除 --- for (Set<String> groupUsers : userGroupMap.values()) { groupUsers.remove(userId); } // 注意:这里不移除空的组,如果需要可以定期清理 // --- 分组结束 --- } } // ... 其他方法 ... }然后,提供几个核心的推送方法:
public class MyWebSocketHandler extends TextWebSocketHandler { // ... 其他代码 ... /** * 向单个用户发送消息 */ public static void sendMessageToUser(String userId, String message) throws IOException { WebSocketSession session = userSessionMap.get(userId); if (session != null && session.isOpen()) { synchronized (session) { // 发送消息需要同步,防止多线程并发写冲突 session.sendMessage(new TextMessage(message)); } } else { // 用户不在线,可以存入消息队列或数据库,待其上线后推送 System.out.println("用户 " + userId + " 不在线,消息无法实时送达"); } } /** * 向特定分组的所有在线用户广播消息 */ public static void sendMessageToGroup(String groupName, String message) { Set<String> userIds = userGroupMap.get(groupName); if (userIds != null && !userIds.isEmpty()) { for (String userId : userIds) { try { sendMessageToUser(userId, message); } catch (IOException e) { System.err.println("向分组用户 " + userId + " 发送消息失败: " + e.getMessage()); } } } } /** * 全局广播(给所有在线用户) */ public static void broadcastMessage(String message) { for (WebSocketSession session : userSessionMap.values()) { if (session.isOpen()) { try { synchronized (session) { session.sendMessage(new TextMessage(message)); } } catch (IOException e) { System.err.println("广播消息失败,Session ID: " + session.getId()); } } } } }这样,在任何一个Spring管理的Bean中(比如Service层),你都可以通过MyWebSocketHandler.sendMessageToUser(userId, msg)或MyWebSocketHandler.sendMessageToGroup("dept:sales", msg)来触发消息推送了。
性能优化提示:当分组内用户数量巨大(比如上万人)时,遍历发送可能会阻塞业务线程。可以考虑将消息推送任务提交给一个专门的线程池来异步执行。另外,
userGroupMap的数据结构可以进一步优化,例如使用Guava的Multimap或者维护一个userId->groups的反向索引,方便用户下线时快速从所有组中移除。
7. 消息格式设计与业务处理
前面我们一直用简单的JSON字符串{"type":"ping"}作为消息。在实际项目中,需要设计一个统一的消息格式,方便前端和后端解析。一个常见的格式如下:
{ "type": "chat/notification/system/...", "sender": "user123", "receiver": "user456 / group:dept1 / all", "timestamp": 1640995200000, "payload": { // 实际的消息内容,结构根据type不同而变化 "title": "新订单", "content": "您有一笔新的订单待处理,订单号:202412310001", "url": "/order/detail/202412310001" } }在Handler的handleTextMessage方法中,我们需要根据type字段来路由到不同的业务处理器:
@Override protected void handleTextMessage(WebSocketSession session, TextMessage message) throws Exception { String payload = message.getPayload(); sessionLastActiveTime.put(session.getId(), System.currentTimeMillis()); // 更新心跳时间 ObjectMapper mapper = new ObjectMapper(); try { JsonNode rootNode = mapper.readTree(payload); String type = rootNode.path("type").asText(null); // 安全获取,避免NPE String sender = (String) session.getAttributes().get("userId"); if ("ping".equals(type)) { // 心跳 session.sendMessage(new TextMessage("{\"type\":\"pong\"}")); return; } else if ("chat".equals(type)) { // 私聊 String receiver = rootNode.path("receiver").asText(); JsonNode msgPayload = rootNode.path("payload"); handlePrivateChat(sender, receiver, msgPayload); } else if ("join_group".equals(type)) { // 动态加入分组 String groupName = rootNode.path("group").asText(); joinGroup(sender, groupName); } else if ("leave_group".equals(type)) { // 动态离开分组 String groupName = rootNode.path("group").asText(); leaveGroup(sender, groupName); } // ... 其他业务类型 ... } catch (JsonProcessingException e) { session.sendMessage(new TextMessage("{\"type\":\"error\", \"msg\":\"消息格式错误\"}")); } catch (Exception e) { session.sendMessage(new TextMessage("{\"type\":\"error\", \"msg\":\"服务器处理异常\"}")); } } private void handlePrivateChat(String sender, String receiver, JsonNode payload) throws IOException { String content = payload.path("content").asText(); String formattedMsg = String.format("{\"type\":\"chat\", \"from\":\"%s\", \"content\":\"%s\"}", sender, content); sendMessageToUser(receiver, formattedMsg); // 可选:也发一份给发送者自己,作为发送成功回执 sendMessageToUser(sender, "{\"type\":\"chat_status\", \"status\":\"sent\", \"to\":\"" + receiver + "\"}"); } private void joinGroup(String userId, String groupName) { userGroupMap.computeIfAbsent(groupName, k -> ConcurrentHashMap.newKeySet()).add(userId); // 通知该用户已加入组 try { sendMessageToUser(userId, "{\"type\":\"system\", \"msg\":\"你已成功加入分组: " + groupName + "\"}"); } catch (IOException e) { e.printStackTrace(); } }这种设计使得消息处理逻辑清晰,易于扩展新的消息类型。
8. 生产环境部署与进阶考量
把代码跑起来只是开始,要上线还得过好几关。
连接数限制与资源管理:每个WebSocket连接都会占用一个线程(取决于容器实现,如Tomcat的NIO模式)和内存。默认的Tomcat配置可能只能处理几千个并发连接。你需要调整application.properties:
# 增大最大连接数 server.tomcat.max-connections=10000 # 增大工作线程数 server.tomcat.threads.max=200 # 调整WebSocket相关的缓冲区大小 server.tomcat.max-swallow-size=2MB同时,必须在代码中做好连接管理,及时清理无效会话(我们的心跳机制就在做这件事),防止内存泄漏。
集群部署与会话共享:上面的代码把所有会话信息(userSessionMap,userGroupMap)都存在单个应用实例的内存里。一旦部署多台实例,用户可能连到A实例,但推送消息的请求发到了B实例,B实例上根本没有这个用户的session,导致推送失败。
解决方案是引入外部存储来共享会话信息,例如Redis。
- 连接建立时,将
userId、instanceId(实例标识,如IP:PORT)和group信息存入Redis,并设置过期时间(略大于心跳超时时间)。 - 发送消息时,先根据
userId或group从Redis查出所有在线的用户及其所在的instanceId。 - 如果目标用户就在当前实例,直接发送;如果在其他实例,则需要通过消息队列(如RabbitMQ、Kafka)或者HTTP调用(需要实例间有通信能力)将消息转发到对应实例去发送。
这是一个架构上的重大变化,通常会引入Spring Cloud、WebSocket Stomp Broker Relay或者自研一个轻量的消息转发层。
前端连接与重连策略:前端不能只连一次就完事。需要监听WebSocket的onclose和onerror事件,实现自动重连,并采用指数退避策略(比如第一次断线等1秒重连,第二次等2秒,第三次等4秒...)。重连时同样要带上认证Token。
监控与日志:务必记录连接建立、关闭、消息收发、异常断开等关键事件,并监控活跃连接数、消息吞吐量等指标,这对排查线上问题至关重要。
我自己在项目上线后,就遇到过因为Nginx的proxy_read_timeout设置比服务端心跳间隔短,导致连接频繁被代理断开的问题。后来统一了超时配置,并在服务端增加了更细致的连接状态日志,才稳定下来。WebSocket看似简单,真想在生产环境扛住流量,每一个细节都得抠。