本文目录导读:

我来为您提供一个完整的Java推送系统案例,包含WebSocket和SSE两种主流实现方式。
系统架构设计
技术选型
- 后端:Spring Boot 2.x + WebSocket + SSE
- 消息中间件:Redis Pub/Sub(用于集群广播)
- 前端:原生JavaScript + WebSocket API
系统结构
┌─────────────┐ ┌──────────────┐ ┌─────────────┐
│ 客户端A │◄───┤ │◄───┤ 客户端B │
└─────────────┘ │ 推送服务器 │ └─────────────┘
┌─────────────┐ │ (集群模式) │ ┌─────────────┐
│ 客户端C │◄───┤ │◄───┤ 客户端D │
└─────────────┘ └──────┬───────┘ └─────────────┘
│
┌──────┴───────┐
│ Redis消息 │
│ 广播中心 │
└──────────────┘
WebSocket实现方案
Maven依赖配置
<dependencies>
<!-- Spring Boot WebSocket -->
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-websocket</artifactId>
</dependency>
<!-- Redis -->
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-data-redis</artifactId>
</dependency>
<!-- JSON处理 -->
<dependency>
<groupId>com.alibaba</groupId>
<artifactId>fastjson</artifactId>
<version>2.0.25</version>
</dependency>
<!-- Lombok -->
<dependency>
<groupId>org.projectlombok</groupId>
<artifactId>lombok</artifactId>
<optional>true</optional>
</dependency>
</dependencies>
WebSocket配置类
package com.example.push.config;
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 {
private final PushWebSocketHandler pushWebSocketHandler;
private final WebSocketInterceptor webSocketInterceptor;
public WebSocketConfig(PushWebSocketHandler pushWebSocketHandler,
WebSocketInterceptor webSocketInterceptor) {
this.pushWebSocketHandler = pushWebSocketHandler;
this.webSocketInterceptor = webSocketInterceptor;
}
@Override
public void registerWebSocketHandlers(WebSocketHandlerRegistry registry) {
registry.addHandler(pushWebSocketHandler, "/ws/push")
.addInterceptors(webSocketInterceptor)
.setAllowedOrigins("*");
}
}
WebSocket拦截器
package com.example.push.config;
import org.springframework.http.server.ServerHttpRequest;
import org.springframework.http.server.ServerHttpResponse;
import org.springframework.stereotype.Component;
import org.springframework.web.socket.WebSocketHandler;
import org.springframework.web.socket.server.HandshakeInterceptor;
import java.util.Map;
@Component
public class WebSocketInterceptor implements HandshakeInterceptor {
@Override
public boolean beforeHandshake(ServerHttpRequest request, ServerHttpResponse response,
WebSocketHandler wsHandler, Map<String, Object> attributes) {
// 从URL参数中获取用户ID
String path = request.getURI().getPath();
String query = request.getURI().getQuery();
if (query != null && query.contains("userId=")) {
String userId = query.substring(query.indexOf("userId=") + 7);
if (userId.contains("&")) {
userId = userId.substring(0, userId.indexOf("&"));
}
attributes.put("userId", userId);
return true;
}
return false;
}
@Override
public void afterHandshake(ServerHttpRequest request, ServerHttpResponse response,
WebSocketHandler wsHandler, Exception exception) {
// 握手完成后的处理
}
}
WebSocket处理器
package com.example.push.handler;
import com.alibaba.fastjson.JSON;
import com.example.push.model.PushMessage;
import com.example.push.service.PushService;
import lombok.extern.slf4j.Slf4j;
import org.springframework.stereotype.Component;
import org.springframework.web.socket.*;
import org.springframework.web.socket.handler.TextWebSocketHandler;
import java.io.IOException;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
@Slf4j
@Component
public class PushWebSocketHandler extends TextWebSocketHandler {
// 存储用户连接,key: userId, value: 连接
private static final Map<String, WebSocketSession> SESSIONS = new ConcurrentHashMap<>();
private final PushService pushService;
public PushWebSocketHandler(PushService pushService) {
this.pushService = pushService;
}
@Override
public void afterConnectionEstablished(WebSocketSession session) {
String userId = (String) session.getAttributes().get("userId");
SESSIONS.put(userId, session);
log.info("用户 {} 已连接,当前在线用户数: {}", userId, SESSIONS.size());
// 发送欢迎消息
sendMessage(userId, new PushMessage("system", "连接成功", "欢迎使用推送系统"));
// 推送未消费的消息
pushService.pushPendingMessages(userId);
}
@Override
protected void handleTextMessage(WebSocketSession session, TextMessage message) {
String userId = (String) session.getAttributes().get("userId");
log.info("收到用户 {} 的消息: {}", userId, message.getPayload());
// 处理心跳消息
if ("PING".equals(message.getPayload())) {
sendMessage(userId, new PushMessage("system", "PONG", "心跳响应"));
}
}
@Override
public void afterConnectionClosed(WebSocketSession session, CloseStatus status) {
String userId = (String) session.getAttributes().get("userId");
SESSIONS.remove(userId);
log.info("用户 {} 断开连接,当前在线用户数: {}", userId, SESSIONS.size());
}
@Override
public void handleTransportError(WebSocketSession session, Throwable exception) {
String userId = (String) session.getAttributes().get("userId");
SESSIONS.remove(userId);
log.error("连接异常,用户: {}", userId, exception);
}
/**
* 发送消息给指定用户
*/
public static void sendMessage(String userId, PushMessage message) {
WebSocketSession session = SESSIONS.get(userId);
if (session != null && session.isOpen()) {
try {
synchronized (session) {
session.sendMessage(new TextMessage(JSON.toJSONString(message)));
}
} catch (IOException e) {
log.error("发送消息失败,用户: {}", userId, e);
}
}
}
/**
* 广播消息给所有用户
*/
public static void broadcast(PushMessage message) {
SESSIONS.forEach((userId, session) -> {
if (session.isOpen()) {
try {
synchronized (session) {
session.sendMessage(new TextMessage(JSON.toJSONString(message)));
}
} catch (IOException e) {
log.error("广播消息失败,用户: {}", userId, e);
}
}
});
}
/**
* 获取在线用户数
*/
public static int getOnlineCount() {
return SESSIONS.size();
}
}
推送服务
package com.example.push.service;
import com.alibaba.fastjson.JSON;
import com.example.push.handler.PushWebSocketHandler;
import com.example.push.model.PushMessage;
import lombok.extern.slf4j.Slf4j;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.data.redis.core.RedisTemplate;
import org.springframework.stereotype.Service;
import java.util.concurrent.TimeUnit;
@Slf4j
@Service
public class PushServiceImpl implements PushService {
@Autowired
private RedisTemplate<String, String> redisTemplate;
private static final String MESSAGE_CHANNEL = "push:message";
private static final String OFFLINE_MESSAGE_KEY = "push:offline:";
private static final String PENDING_MESSAGE_KEY = "push:pending:";
@Override
public void sendMessage(String userId, PushMessage message) {
// 先尝试直接发送
PushWebSocketHandler.sendMessage(userId, message);
// 发布到Redis,用于集群其他节点处理
String messageJson = JSON.toJSONString(message);
message.setTargetUserId(userId);
redisTemplate.convertAndSend(MESSAGE_CHANNEL, JSON.toJSONString(message));
}
@Override
public void broadcast(PushMessage message) {
// 本地广播
PushWebSocketHandler.broadcast(message);
// 发布到Redis,用于集群广播
message.setBroadcast(true);
redisTemplate.convertAndSend(MESSAGE_CHANNEL, JSON.toJSONString(message));
}
@Override
public void saveOfflineMessage(String userId, PushMessage message) {
String key = OFFLINE_MESSAGE_KEY + userId;
// 存储离线消息,保留最近100条
redisTemplate.opsForList().rightPush(key, JSON.toJSONString(message));
redisTemplate.opsForList().trim(key, -100, -1);
redisTemplate.expire(key, 7, TimeUnit.DAYS);
}
@Override
public void pushPendingMessages(String userId) {
String key = PENDING_MESSAGE_KEY + userId;
// 获取并清除待消费的消息
List<String> messages = redisTemplate.opsForList().range(key, 0, -1);
if (messages != null && !messages.isEmpty()) {
for (String messageJson : messages) {
PushMessage message = JSON.parseObject(messageJson, PushMessage.class);
PushWebSocketHandler.sendMessage(userId, message);
}
redisTemplate.delete(key);
}
}
}
Redis消息监听器
package com.example.push.listener;
import com.alibaba.fastjson.JSON;
import com.example.push.handler.PushWebSocketHandler;
import com.example.push.model.PushMessage;
import lombok.extern.slf4j.Slf4j;
import org.springframework.data.redis.connection.Message;
import org.springframework.data.redis.connection.MessageListener;
import org.springframework.stereotype.Component;
@Slf4j
@Component
public class PushMessageListener implements MessageListener {
@Override
public void onMessage(Message message, byte[] pattern) {
String messageBody = new String(message.getBody());
PushMessage pushMessage = JSON.parseObject(messageBody, PushMessage.class);
if (pushMessage.isBroadcast()) {
// 广播消息,只发送给本地连接的用户的用户
PushWebSocketHandler.broadcast(pushMessage);
} else {
// 单用户消息
PushWebSocketHandler.sendMessage(pushMessage.getTargetUserId(), pushMessage);
}
}
}
模型类
package com.example.push.model;
import lombok.Data;
import java.time.LocalDateTime;
@Data
public class PushMessage {
private String type; // 消息类型(system, chat, alert等)
private String content; // 消息内容
private String fromUserId; // 发送者
private String targetUserId; // 目标用户(广播时为null)
private boolean broadcast; // 是否广播
private LocalDateTime sendTime; // 发送时间
public PushMessage() {}
public PushMessage(String type, String content, String fromUserId) {
this.type = type;
this.content = content;
this.fromUserId = fromUserId;
this.sendTime = LocalDateTime.now();
}
}
SSE(Server-Sent Events)实现
SSE控制器
package com.example.push.controller;
import com.example.push.model.PushMessage;
import com.example.push.service.SseService;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.MediaType;
import org.springframework.web.bind.annotation.*;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
@RestController
@RequestMapping("/api/sse")
public class SseController {
@Autowired
private SseService sseService;
/**
* 创建SSE连接
*/
@GetMapping(value = "/connect/{userId}", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
public SseEmitter connect(@PathVariable String userId) {
return sseService.createConnection(userId);
}
/**
* 发送消息
*/
@PostMapping("/send")
public void send(@RequestBody PushMessage message) {
sseService.sendMessage(message.getTargetUserId(), message);
}
/**
* 广播消息
*/
@PostMapping("/broadcast")
public void broadcast(@RequestBody PushMessage message) {
sseService.broadcast(message);
}
/**
* 获取在线用户数
*/
@GetMapping("/online-count")
public int getOnlineCount() {
return sseService.getOnlineCount();
}
}
SSE服务
package com.example.push.service;
import com.example.push.model.PushMessage;
import lombok.extern.slf4j.Slf4j;
import org.springframework.stereotype.Service;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
import java.io.IOException;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.Executors;
import java.util.concurrent.ScheduledExecutorService;
import java.util.concurrent.TimeUnit;
@Slf4j
@Service
public class SseService {
// 存储用户连接
private static final Map<String, SseEmitter> CONNECTIONS = new ConcurrentHashMap<>();
// 心跳检测调度器
private final ScheduledExecutorService scheduler = Executors.newScheduledThreadPool(1);
public SseService() {
// 启动心跳检测
startHeartbeatCheck();
}
/**
* 创建连接
*/
public SseEmitter createConnection(String userId) {
SseEmitter emitter = new SseEmitter(0L); // 0表示不超时
CONNECTIONS.put(userId, emitter);
// 设置回调
emitter.onCompletion(() -> {
CONNECTIONS.remove(userId);
log.info("用户 {} 的连接已完成", userId);
});
emitter.onTimeout(() -> {
CONNECTIONS.remove(userId);
log.info("用户 {} 的连接超时", userId);
emitter.complete();
});
emitter.onError(e -> {
CONNECTIONS.remove(userId);
log.error("用户 {} 的连接错误: {}", userId, e.getMessage());
});
// 发送连接成功消息
try {
emitter.send(SseEmitter.event()
.name("connected")
.data("连接成功"));
} catch (IOException e) {
log.error("发送连接成功消息失败", e);
}
log.info("用户 {} 已连接,当前在线用户数: {}", userId, CONNECTIONS.size());
return emitter;
}
/**
* 发送消息给指定用户
*/
public void sendMessage(String userId, PushMessage message) {
SseEmitter emitter = CONNECTIONS.get(userId);
if (emitter != null) {
try {
emitter.send(SseEmitter.event()
.name("message")
.data(message));
} catch (IOException e) {
log.error("发送消息失败,用户: {}", userId, e);
CONNECTIONS.remove(userId);
}
}
}
/**
* 广播消息
*/
public void broadcast(PushMessage message) {
CONNECTIONS.forEach((userId, emitter) -> {
try {
emitter.send(SseEmitter.event()
.name("message")
.data(message));
} catch (IOException e) {
log.error("广播消息失败,用户: {}", userId, e);
CONNECTIONS.remove(userId);
}
});
}
/**
* 获取在线用户数
*/
public int getOnlineCount() {
return CONNECTIONS.size();
}
/**
* 心跳检测,清理死连接
*/
private void startHeartbeatCheck() {
scheduler.scheduleAtFixedRate(() -> {
CONNECTIONS.forEach((userId, emitter) -> {
try {
// 发送心跳消息
emitter.send(SseEmitter.event()
.name("heartbeat")
.data(System.currentTimeMillis()));
} catch (IOException e) {
log.info("用户 {} 的连接已断开", userId);
CONNECTIONS.remove(userId);
emitter.complete();
}
});
}, 30, 30, TimeUnit.SECONDS);
}
}
前端代码
WebSocket客户端
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">WebSocket推送客户端</title>
<style>
body { font-family: Arial, sans-serif; margin: 20px; }
.message { background-color: #f5f5f5; padding: 10px; margin-bottom: 10px; border-radius: 5px; }
.system { color: blue; }
.alert { color: red; }
.chat { color: green; }
#messageArea { max-height: 400px; overflow-y: auto; margin-bottom: 20px; }
input, button { padding: 8px; margin-right: 5px; }
#status { color: orange; margin-bottom: 10px; }
.connected { color: green; }
.disconnected { color: red; }
</style>
</head>
<body>
<h1>实时推送系统 - WebSocket客户端</h1>
<div id="status">未连接</div>
<div>
<input type="text" id="userId" placeholder="用户ID" value="user_001">
<button onclick="connect()">连接</button>
<button onclick="disconnect()">断开</button>
</div>
<div style="margin-top: 20px;">
<input type="text" id="targetUserId" placeholder="目标用户ID">
<input type="text" id="messageContent" placeholder="消息内容">
<button onclick="sendMessage()">发送</button>
<button onclick="broadcast()">广播</button>
</div>
<div id="messageArea" style="margin-top: 20px;"></div>
<script>
var ws = null;
function connect() {
var userId = document.getElementById('userId').value;
if (!userId) {
alert('请输入用户ID');
return;
}
// 创建WebSocket连接
var url = 'ws://localhost:8080/ws/push?userId=' + userId;
ws = new WebSocket(url);
ws.onopen = function() {
document.getElementById('status').textContent = '已连接';
document.getElementById('status').className = 'connected';
addMessage('system', '连接成功');
// 发送心跳
setInterval(function() {
if (ws.readyState === WebSocket.OPEN) {
ws.send('PING');
}
}, 30000);
};
ws.onmessage = function(event) {
var message = JSON.parse(event.data);
addMessage(message.type, message.content);
};
ws.onclose = function() {
document.getElementById('status').textContent = '已断开';
document.getElementById('status').className = 'disconnected';
addMessage('system', '连接已断开');
};
ws.onerror = function(error) {
document.getElementById('status').textContent = '连接错误';
addMessage('system', '连接错误: ' + error);
};
}
function disconnect() {
if (ws) {
ws.close();
ws = null;
}
}
function sendMessage() {
var targetUserId = document.getElementById('targetUserId').value;
var content = document.getElementById('messageContent').value;
if (!targetUserId || !content) {
alert('请填写目标和内容');
return;
}
// 使用HTTP接口发送
fetch('/api/push/send', {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
targetUserId: targetUserId,
content: content,
type: 'chat'
})
});
document.getElementById('messageContent').value = '';
}
function broadcast() {
var content = document.getElementById('messageContent').value;
if (!content) {
alert('请填写内容');
return;
}
fetch('/api/push/broadcast', {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify({
content: content,
type: 'alert'
})
});
document.getElementById('messageContent').value = '';
}
function addMessage(type, content) {
var messageArea = document.getElementById('messageArea');
var messageDiv = document.createElement('div');
messageDiv.className = 'message ' + type;
var time = new Date().toLocaleTimeString();
messageDiv.innerHTML = '[' + time + '] [' + type + '] ' + content;
messageArea.appendChild(messageDiv);
messageArea.scrollTop = messageArea.scrollHeight;
}
</script>
</body>
</html>
SSE客户端
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">SSE推送客户端</title>
<style>
body { font-family: Arial, sans-serif; margin: 20px; }
.message { background-color: #f5f5f5; padding: 10px; margin-bottom: 10px; border-radius: 5px; }
#connectionState { color: orange; }
.connected { color: green; }
.disconnected { color: red; }
</style>
</head>
<body>
<h1>SSE实时推送客户端</h1>
<div id="connectionState">未连接</div>
<div>
<input type="text" id="sseUserId" placeholder="用户ID" value="user_001">
<button onclick="connectSSE()">连接</button>
<button onclick="disconnectSSE()">断开</button>
</div>
<div id="sseMessages" style="margin-top: 20px;"></div>
<script>
var eventSource = null;
function connectSSE() {
var userId = document.getElementById('sseUserId').value;
if (!userId) {
alert('请输入用户ID');
return;
}
// 创建SSE连接
eventSource = new EventSource('/api/sse/connect/' + userId);
// 连接打开
eventSource.onopen = function(event) {
document.getElementById('connectionState').textContent = '已连接';
document.getElementById('connectionState').className = 'connected';
addSSEMessage('system', 'SSE连接成功');
};
// 连接的响应消息
eventSource.addEventListener('connected', function(event) {
addSSEMessage('system', '用户连接成功');
});
// 普通消息
eventSource.addEventListener('message', function(event) {
var message = JSON.parse(event.data);
addSSEMessage(message.type, message.content);
});
// 心跳消息
eventSource.addEventListener('heartbeat', function(event) {
// 可以忽略心跳
});
// 错误处理
eventSource.onerror = function(error) {
document.getElementById('connectionState').textContent = '连接错误';
document.getElementById('connectionState').className = 'disconnected';
addSSEMessage('system', 'SSE连接错误');
};
}
function disconnectSSE() {
if (eventSource) {
eventSource.close();
eventSource = null;
document.getElementById('connectionState').textContent = '已断开';
document.getElementById('connectionState').className = 'disconnected';
}
}
function addSSEMessage(type, content) {
var messageArea = document.getElementById('sseMessages');
var messageDiv = document.createElement('div');
messageDiv.className = 'message';
var time = new Date().toLocaleTimeString();
messageDiv.innerHTML = '[' + time + '] [' + type + '] ' + content;
messageArea.appendChild(messageDiv);
messageArea.scrollTop = messageArea.scrollHeight;
}
</script>
</body>
</html>
推送控制器
package com.example.push.controller;
import com.example.push.model.PushMessage;
import com.example.push.service.PushService;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.web.bind.annotation.*;
@RestController
@RequestMapping("/api/push")
public class PushController {
@Autowired
private PushService pushService;
/**
* 发送消息给用户
*/
@PostMapping("/send")
public void send(@RequestBody PushMessage message) {
pushService.sendMessage(message.getTargetUserId(), message);
}
/**
* 广播消息
*/
@PostMapping("/broadcast")
public void broadcast(@RequestBody PushMessage message) {
pushService.broadcast(message);
}
/**
* 获取在线用户数
*/
@GetMapping("/online-count")
public int getOnlineCount() {
return PushWebSocketHandler.getOnlineCount();
}
}
配置文件
# application.yml
server:
port: 8080
spring:
redis:
host: localhost
port: 6379
password:
database: 0
# 自定义配置
push:
# 推送线程池大小
core-pool-size: 10
max-pool-size: 50
# 队列容量
queue-capacity: 10000
功能说明
核心功能
- 单用户推送:向指定用户发送消息
- 广播推送:向所有在线用户广播消息
- 离线消息:用户离线时保存消息,上线后推送
- 心跳检测:维持连接,清理无效连接
- 集群支持:通过Redis发布订阅实现集群消息广播
WebSocket vs SSE 对比
| 特性 | WebSocket | SSE |
|---|---|---|
| 通信方向 | 双向通信 | 单向(服务器→客户端) |
| 协议 | ws:// | http:// |
| 二进制数据 | 支持 | 不支持 |
| 自动重连 | 不支持 | 支持 |
| 兼容性 | 较好 | 现代浏览器支持 |
使用场景建议
- WebSocket:需要双向通信的场景,如聊天室、游戏
- SSE:服务器单向推送,如新闻通知、股票行情
这个案例涵盖了Java推送系统的核心实现,您可以根据实际需求选择合适的方案进行扩展。