本文目录导读:

我来详细介绍Java中实现WebSocket的几种常见方式。
使用Java EE标准API(JSR 356)
添加依赖(Maven)
<!-- Java EE WebSocket API -->
<dependency>
<groupId>javax.websocket</groupId>
<artifactId>javax.websocket-api</artifactId>
<version>1.1</version>
<scope>provided</scope>
</dependency>
<!-- 如果使用Tomcat,需要添加 -->
<dependency>
<groupId>org.apache.tomcat</groupId>
<artifactId>tomcat-websocket</artifactId>
<version>9.0.65</version>
</dependency>
服务端实现
import javax.websocket.*;
import javax.websocket.server.ServerEndpoint;
import java.io.IOException;
import java.util.concurrent.CopyOnWriteArraySet;
@ServerEndpoint("/websocket/{userId}")
public class WebSocketServer {
// 存储所有连接的客户端
private static CopyOnWriteArraySet<WebSocketServer> clients =
new CopyOnWriteArraySet<>();
// 与客户端的会话
private Session session;
private String userId;
/**
* 连接建立成功调用的方法
*/
@OnOpen
public void onOpen(Session session, @PathParam("userId") String userId) {
this.session = session;
this.userId = userId;
clients.add(this);
System.out.println("新连接加入:" + userId);
// 发送欢迎消息
sendMessage("欢迎连接WebSocket服务器!");
}
/**
* 连接关闭调用的方法
*/
@OnClose
public void onClose() {
clients.remove(this);
System.out.println("连接关闭:" + userId);
}
/**
* 收到客户端消息后调用的方法
*/
@OnMessage
public void onMessage(String message, Session session) {
System.out.println("来自客户端的消息:" + message);
// 处理消息
handleMessage(message);
}
/**
* 发生错误时调用
*/
@OnError
public void onError(Session session, Throwable error) {
System.out.println("发生错误:" + error.getMessage());
error.printStackTrace();
}
/**
* 发送消息到当前客户端
*/
public void sendMessage(String message) {
try {
if (this.session.isOpen()) {
this.session.getBasicRemote().sendText(message);
}
} catch (IOException e) {
e.printStackTrace();
}
}
/**
* 群发消息
*/
public static void broadcast(String message) {
for (WebSocketServer client : clients) {
client.sendMessage(message);
}
}
/**
* 发送消息给指定用户
*/
public static void sendToUser(String userId, String message) {
for (WebSocketServer client : clients) {
if (client.userId.equals(userId)) {
client.sendMessage(message);
break;
}
}
}
/**
* 处理消息的逻辑
*/
private void handleMessage(String message) {
// 简单回显
sendMessage("服务器已收到消息:" + message);
// 或者广播消息
broadcast("用户 " + userId + " 发送:" + message);
}
/**
* 获取在线用户数
*/
public static int getOnlineCount() {
return clients.size();
}
}
客户端实现(JavaScript示例)
// 创建WebSocket连接
let userId = "user_" + Date.now();
let ws = new WebSocket("ws://localhost:8080/websocket/" + userId);
// 连接建立时触发
ws.onopen = function() {
console.log("WebSocket连接已建立");
document.getElementById("status").innerHTML = "已连接";
};
// 收到服务器消息时触发
ws.onmessage = function(event) {
console.log("收到消息:" + event.data);
displayMessage(event.data);
};
// 连接关闭时触发
ws.onclose = function() {
console.log("WebSocket连接已关闭");
document.getElementById("status").innerHTML = "已断开";
};
// 连接出错时触发
ws.onerror = function(error) {
console.log("WebSocket错误:" + error);
};
// 发送消息
function sendMessage() {
let message = document.getElementById("message").value;
ws.send(message);
}
// 关闭连接
function closeConnection() {
ws.close();
}
// 显示消息
function displayMessage(message) {
let container = document.getElementById("messages");
let div = document.createElement("div");
div.textContent = message;
container.appendChild(div);
}
使用Spring Boot + WebSocket
添加Maven依赖
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-websocket</artifactId>
</dependency>
Spring Boot配置类
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.ServletServerContainerFactoryBean;
@Configuration
@EnableWebSocket
public class WebSocketConfig implements WebSocketConfigurer {
@Override
public void registerWebSocketHandlers(WebSocketHandlerRegistry registry) {
registry.addHandler(chatWebSocketHandler(), "/chat/{roomId}")
.setAllowedOrigins("*") // 允许跨域
.addInterceptors(new ChatHandshakeInterceptor());
// 支持SockJS
registry.addHandler(chatWebSocketHandler(), "/chat/sockjs/{roomId}")
.setAllowedOrigins("*")
.addInterceptors(new ChatHandshakeInterceptor())
.withSockJS();
}
@Bean
public ChatWebSocketHandler chatWebSocketHandler() {
return new ChatWebSocketHandler();
}
@Bean
public ServletServerContainerFactoryBean createWebSocketContainer() {
ServletServerContainerFactoryBean container =
new ServletServerContainerFactoryBean();
container.setMaxTextMessageBufferSize(8192);
container.setMaxBinaryMessageBufferSize(8192);
container.setMaxSessionIdleTimeout(600000L); // 10分钟超时
return container;
}
}
WebSocket处理器
import org.springframework.web.socket.*;
import org.springframework.web.socket.handler.TextWebSocketHandler;
import java.util.concurrent.ConcurrentHashMap;
public class ChatWebSocketHandler extends TextWebSocketHandler {
// 存储连接,key为roomId,value为该房间的用户连接列表
private static ConcurrentHashMap<String, ConcurrentHashMap<String, WebSocketSession>>
roomSessions = new ConcurrentHashMap<>();
@Override
public void afterConnectionEstablished(WebSocketSession session) {
String roomId = getRoomId(session);
String userId = getUserId(session);
roomSessions.computeIfAbsent(roomId, k -> new ConcurrentHashMap<>())
.put(userId, session);
System.out.println("用户 " + userId + " 加入房间 " + roomId);
// 广播用户加入消息
broadcastToRoom(roomId,
"{\"type\":\"join\", \"userId\":\"" + userId + "\"}");
}
@Override
protected void handleTextMessage(WebSocketSession session, TextMessage message) {
String payload = message.getPayload();
String roomId = getRoomId(session);
// 广播消息到房间所有人
broadcastToRoom(roomId, payload);
}
@Override
public void afterConnectionClosed(WebSocketSession session, CloseStatus status) {
String roomId = getRoomId(session);
String userId = getUserId(session);
ConcurrentHashMap<String, WebSocketSession> sessions = roomSessions.get(roomId);
if (sessions != null) {
sessions.remove(userId);
if (sessions.isEmpty()) {
roomSessions.remove(roomId);
}
}
// 广播用户离开消息
broadcastToRoom(roomId,
"{\"type\":\"leave\", \"userId\":\"" + userId + "\"}");
}
@Override
public void handleTransportError(WebSocketSession session, Throwable exception) {
System.out.println("传输错误:" + exception.getMessage());
}
/**
* 向房间广播消息
*/
private void broadcastToRoom(String roomId, String message) {
ConcurrentHashMap<String, WebSocketSession> sessions = roomSessions.get(roomId);
if (sessions != null) {
TextMessage textMessage = new TextMessage(message);
sessions.forEach((userId, session) -> {
if (session.isOpen()) {
try {
session.sendMessage(textMessage);
} catch (Exception e) {
e.printStackTrace();
}
}
});
}
}
/**
* 从session中获取房间ID
*/
private String getRoomId(WebSocketSession session) {
return (String) session.getAttributes().get("roomId");
}
/**
* 从session中获取用户ID
*/
private String getUserId(WebSocketSession session) {
return (String) session.getAttributes().get("userId");
}
}
握手拦截器
import org.springframework.http.server.ServerHttpRequest;
import org.springframework.http.server.ServerHttpResponse;
import org.springframework.web.socket.WebSocketHandler;
import org.springframework.web.socket.server.support.HttpSessionHandshakeInterceptor;
import java.util.Map;
public class ChatHandshakeInterceptor extends HttpSessionHandshakeInterceptor {
@Override
public boolean beforeHandshake(ServerHttpRequest request,
ServerHttpResponse response,
WebSocketHandler wsHandler,
Map<String, Object> attributes) {
// 从URL中提取参数
String path = request.getURI().getPath();
String[] parts = path.split("/");
// 假设URL格式为: /chat/{roomId}
if (parts.length >= 3) {
attributes.put("roomId", parts[2]);
}
// 设置用户ID(实际项目可从session或token获取)
attributes.put("userId", "user_" + System.currentTimeMillis());
return super.beforeHandshake(request, response, wsHandler, attributes);
}
@Override
public void afterHandshake(ServerHttpRequest request,
ServerHttpResponse response,
WebSocketHandler wsHandler,
Exception ex) {
super.afterHandshake(request, response, wsHandler, ex);
}
}
心跳检测(可选)
import org.springframework.scheduling.annotation.Scheduled;
import org.springframework.stereotype.Component;
import org.springframework.web.socket.TextMessage;
import org.springframework.web.socket.WebSocketSession;
import java.io.IOException;
import java.util.concurrent.ConcurrentHashMap;
@Component
public class HeartbeatTask {
@Scheduled(fixedRate = 30000) // 30秒执行一次
public void sendHeartbeat() {
TextMessage heartbeat = new TextMessage("{\"type\":\"heartbeat\"}");
// 遍历所有房间的连接
ChatWebSocketHandler.getRoomSessions().forEach((roomId, sessions) -> {
sessions.forEach((userId, session) -> {
try {
if (session.isOpen()) {
session.sendMessage(heartbeat);
}
} catch (IOException e) {
// 连接可能已断开
e.printStackTrace();
}
});
});
}
}
配置application.properties
# WebSocket配置 server.port=8080 # 设置WebSocket缓冲区大小 spring.websocket.max-text-message-size=8192 spring.websocket.max-binary-message-size=8192
客户端示例(使用SockJS)
// 使用SockJS连接
var sock = new SockJS('http://localhost:8080/chat/sockjs/room001');
var ws = new WebSocket(sock);
ws.onopen = function() {
console.log('连接建立');
document.getElementById('status').textContent = '已连接';
};
ws.onmessage = function(event) {
var data = JSON.parse(event.data);
console.log('收到消息:', data);
displayMessage(data);
};
ws.onclose = function() {
console.log('连接关闭');
document.getElementById('status').textContent = '已断开';
};
function displayMessage(data) {
var div = document.createElement('div');
div.textContent = data.type + ': ' + data.userId;
document.getElementById('messages').appendChild(div);
}
function sendMessage() {
var message = document.getElementById('message').value;
ws.send(JSON.stringify({
type: 'message',
content: message
}));
}
实现WebSocket的关键点:
- 选择合适的实现方式:Java EE标准、Spring Boot集成或原生API
- 处理连接生命周期:OnOpen、OnClose、OnMessage、OnError
- 管理会话:存储和管理客户端连接
- 消息广播:向所有或特定客户端发送消息
- 心跳机制:保持连接活性
- 错误处理:处理各种异常情况
Spring Boot集成的WebSocket在企业项目中更常用,它提供了更完善的配置和集成方案。