增加 spring websocket 示例

This commit is contained in:
YunaiV
2019-11-21 00:54:57 +08:00
parent 611be14faa
commit b4bd2067d0
12 changed files with 130 additions and 23 deletions

View File

@@ -5,6 +5,7 @@ import cn.iocoder.springboot.lab25.springwebsocket.message.AuthResponse;
import cn.iocoder.springboot.lab25.springwebsocket.message.UserJoinNoticeRequest;
import cn.iocoder.springboot.lab25.springwebsocket.util.WebSocketUtil;
import org.springframework.stereotype.Component;
import org.springframework.util.StringUtils;
import javax.websocket.Session;
@@ -13,12 +14,18 @@ public class AuthMessageHandler implements MessageHandler<AuthRequest> {
@Override
public void execute(Session session, AuthRequest message) {
// 如果未传递 accessToken
if (StringUtils.isEmpty(message.getAccessToken())) {
WebSocketUtil.send(session, AuthResponse.TYPE,
new AuthResponse().setCode(1).setMessage("认证 accessToken 未传入"));
return;
}
// 添加到 WebSocketUtil 中
WebSocketUtil.addUser(session, message.getAccessToken()); // 考虑到代码简化,我们先直接使用 accessToken 作为 User
WebSocketUtil.addSession(session, message.getAccessToken()); // 考虑到代码简化,我们先直接使用 accessToken 作为 User
// 判断是否认证成功。这里,假装直接成功
WebSocketUtil.send(session, AuthResponse.TYPE,
new AuthResponse().setCode(0));
WebSocketUtil.send(session, AuthResponse.TYPE, new AuthResponse().setCode(0));
// 通知所有人,某个人加入了。这个是可选逻辑,仅仅是为了演示
WebSocketUtil.broadcast(UserJoinNoticeRequest.TYPE,

View File

@@ -5,12 +5,21 @@ import cn.iocoder.springboot.lab25.springwebsocket.message.Message;
import javax.websocket.Session;
/**
* 消息处理器
* 消息处理器接口
*/
public interface MessageHandler<T extends Message> {
/**
* 执行处理消息
*
* @param session 会话
* @param message 消息
*/
void execute(Session session, T message);
/**
* @return 消息类型,即每个 Message 实现类上的 TYPE 静态字段
*/
String getType();
}

View File

@@ -1,7 +1,7 @@
package cn.iocoder.springboot.lab25.springwebsocket.message;
/**
* 认证 Message
* 用户认证请求
*/
public class AuthRequest implements Message {

View File

@@ -1,7 +1,7 @@
package cn.iocoder.springboot.lab25.springwebsocket.message;
/**
* 认证结果 Message
* 用户认证响应
*/
public class AuthResponse implements Message {
@@ -11,6 +11,10 @@ public class AuthResponse implements Message {
* 响应状态码
*/
private Integer code;
/**
* 响应提示
*/
private String message;
public Integer getCode() {
return code;
@@ -21,4 +25,13 @@ public class AuthResponse implements Message {
return this;
}
public String getMessage() {
return message;
}
public AuthResponse setMessage(String message) {
this.message = message;
return this;
}
}

View File

@@ -1,7 +1,7 @@
package cn.iocoder.springboot.lab25.springwebsocket.message;
/**
* 基础消息
* 基础消息
*/
public interface Message {
}

View File

@@ -1,7 +1,7 @@
package cn.iocoder.springboot.lab25.springwebsocket.message;
/**
* 发送给单个人消息的成功结果的 Message
* 发送消息响应结果的 Message
*/
public class SendResponse implements Message {
@@ -15,6 +15,10 @@ public class SendResponse implements Message {
* 响应状态码
*/
private Integer code;
/**
* 响应提示
*/
private String message;
public String getMsgId() {
return msgId;
@@ -34,4 +38,13 @@ public class SendResponse implements Message {
return this;
}
public String getMessage() {
return message;
}
public SendResponse setMessage(String message) {
this.message = message;
return this;
}
}

View File

@@ -1,7 +1,7 @@
package cn.iocoder.springboot.lab25.springwebsocket.message;
/**
* 发送给所有人消息的 Message
* 发送给所有人的群聊消息的 Message
*/
public class SendToAllRequest implements Message {

View File

@@ -1,7 +1,7 @@
package cn.iocoder.springboot.lab25.springwebsocket.message;
/**
* 发送给单个人消息的 Message
* 发送给指定人的私聊消息的 Message
*/
public class SendToOneRequest implements Message {

View File

@@ -1,5 +1,8 @@
package cn.iocoder.springboot.lab25.springwebsocket.message;
/**
* 发送消息给一个用户的 Message
*/
public class SendToUserRequest implements Message {
public static final String TYPE = "SEND_TO_USER_REQUEST";

View File

@@ -1,7 +1,7 @@
package cn.iocoder.springboot.lab25.springwebsocket.message;
/**
* 用户加入通知 Message
* 用户加入群聊的通知 Message
*/
public class UserJoinNoticeRequest implements Message {

View File

@@ -18,6 +18,8 @@ public class WebSocketUtil {
private static final Logger LOGGER = LoggerFactory.getLogger(WebSocketUtil.class);
// ========== 会话相关 ==========
/**
* Session 与用户的映射
*/
@@ -27,10 +29,24 @@ public class WebSocketUtil {
*/
private static final Map<String, Session> USER_SESSION_MAP = new ConcurrentHashMap<>();
public static void addSession(Session session) {
SESSION_USER_MAP.put(session, ""); // 使用 "" 占位,因为 ConcurrentHashMap 不允许 value 为空
/**
* 添加 Session 。在这个方法中,会添加用户和 Session 之间的映射
*
* @param session Session
* @param user 用户
*/
public static void addSession(Session session, String user) {
// 更新 USER_SESSION_MAP
USER_SESSION_MAP.put(user, session);
// 更新 SESSION_USER_MAP
SESSION_USER_MAP.put(session, user);
}
/**
* 移除 Session 。
*
* @param session Session
*/
public static void removeSession(Session session) {
// 从 SESSION_USER_MAP 中移除
String user = SESSION_USER_MAP.remove(session);
@@ -40,13 +56,15 @@ public class WebSocketUtil {
}
}
public static void addUser(Session session, String user) {
// 更新 USER_SESSION_MAP
USER_SESSION_MAP.put(user, session);
// 更新 SESSION_USER_MAP
SESSION_USER_MAP.put(session, user);
}
// ========== 消息相关 ==========
/**
* 广播发送消息给所有在线用户
*
* @param type 消息类型
* @param message 消息体
* @param <T> 消息类型
*/
public static <T extends Message> void broadcast(String type, T message) {
// 创建消息
String messageText = buildTextMessage(type, message);
@@ -56,6 +74,14 @@ public class WebSocketUtil {
}
}
/**
* 发送消息给单个用户的 Session
*
* @param session Session
* @param type 消息类型
* @param message 消息体
* @param <T> 消息类型
*/
public static <T extends Message> void send(Session session, String type, T message) {
// 创建消息
String messageText = buildTextMessage(type, message);
@@ -63,6 +89,15 @@ public class WebSocketUtil {
sendTextMessage(session, messageText);
}
/**
* 发送消息给指定用户
*
* @param user 指定用户
* @param type 消息类型
* @param message 消息体
* @param <T> 消息类型
* @return 发送是否成功你那个
*/
public static <T extends Message> boolean send(String user, String type, T message) {
// 获得用户对应的 Session
Session session = USER_SESSION_MAP.get(user);
@@ -75,6 +110,14 @@ public class WebSocketUtil {
return true;
}
/**
* 构建完整的消息
*
* @param type 消息类型
* @param message 消息体
* @param <T> 消息类型
* @return 消息
*/
private static <T extends Message> String buildTextMessage(String type, T message) {
JSONObject messageObject = new JSONObject();
messageObject.put("type", type);
@@ -82,6 +125,12 @@ public class WebSocketUtil {
return messageObject.toString();
}
/**
* 真正发送消息
*
* @param session Session
* @param messageText 消息
*/
private static void sendTextMessage(Session session, String messageText) {
if (session == null) {
LOGGER.error("[sendTextMessage][session 为 null]");

View File

@@ -1,6 +1,7 @@
package cn.iocoder.springboot.lab25.springwebsocket.websocket;
import cn.iocoder.springboot.lab25.springwebsocket.handler.MessageHandler;
import cn.iocoder.springboot.lab25.springwebsocket.message.AuthRequest;
import cn.iocoder.springboot.lab25.springwebsocket.message.Message;
import cn.iocoder.springboot.lab25.springwebsocket.util.WebSocketUtil;
import com.alibaba.fastjson.JSON;
@@ -12,12 +13,14 @@ import org.springframework.beans.factory.InitializingBean;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.context.ApplicationContext;
import org.springframework.stereotype.Controller;
import org.springframework.util.CollectionUtils;
import javax.websocket.*;
import javax.websocket.server.ServerEndpoint;
import java.lang.reflect.ParameterizedType;
import java.lang.reflect.Type;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
@@ -28,9 +31,9 @@ public class WebsocketServerEndpoint implements InitializingBean {
private Logger logger = LoggerFactory.getLogger(getClass());
/**
* 设置成静态变量
* 消息类型与 MessageHandler 的映射
*
* 虽然说 WebsocketServerEndpoint 是单例,但是 Spring Boot 还是会为每个 WebSocket 创建一个 WebsocketServerEndpoint Bean 。
* 注意,这里设置成静态变量。虽然说 WebsocketServerEndpoint 是单例,但是 Spring Boot 还是会为每个 WebSocket 创建一个 WebsocketServerEndpoint Bean 。
*/
private static final Map<String, MessageHandler> HANDLERS = new HashMap<>();
@@ -40,8 +43,18 @@ public class WebsocketServerEndpoint implements InitializingBean {
@OnOpen
public void onOpen(Session session, EndpointConfig config) {
logger.info("[onOpen][session({}) 接入]", session);
// 添加到在线缓存
WebSocketUtil.addSession(session);
// 解析 accessToken
List<String> accessTokenValues = session.getRequestParameterMap().get("accessToken");
String accessToken = !CollectionUtils.isEmpty(accessTokenValues) ? accessTokenValues.get(0) : null;
// 创建 AuthRequest 消息类型
AuthRequest authRequest = new AuthRequest().setAccessToken(accessToken);
// 获得消息处理器
MessageHandler<AuthRequest> messageHandler = HANDLERS.get(AuthRequest.TYPE);
if (messageHandler == null) {
logger.error("[onOpen][认证消息类型,不存在消息处理器]");
return;
}
messageHandler.execute(session, authRequest);
}
@OnMessage