Java推送系统案例

wen java案例 3

本文目录导读:

Java推送系统案例

  1. 系统架构设计
  2. WebSocket实现方案
  3. SSE(Server-Sent Events)实现
  4. 前端代码
  5. 推送控制器
  6. 配置文件
  7. 功能说明

我来为您提供一个完整的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推送系统的核心实现,您可以根据实际需求选择合适的方案进行扩展。

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