Explorar el Código

fix 优化mj websocket

zwhui hace 1 año
padre
commit
54b18d1401

+ 233 - 0
midjourney-netty-websocket-web/src/main/java/com/yhlxj/netty/websocket/midjounrey/service/SocketOldServer.java

@@ -0,0 +1,233 @@
+package com.yhlxj.netty.websocket.midjounrey.service;
+
+import cn.hutool.core.collection.CollectionUtil;
+import cn.hutool.json.JSONObject;
+import cn.hutool.json.JSONUtil;
+import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
+import com.yhlxj.netty.websocket.midjounrey.dao.model.entity.MidjourneyUser;
+import com.yhlxj.netty.websocket.midjounrey.dao.model.enums.WsMessageTypeEnum;
+import com.yhlxj.netty.websocket.starter.annotations.*;
+import com.yhlxj.netty.websocket.starter.socket.Session;
+import io.netty.handler.codec.http.HttpHeaders;
+import io.netty.handler.timeout.IdleStateEvent;
+import lombok.RequiredArgsConstructor;
+import lombok.extern.slf4j.Slf4j;
+import org.apache.commons.lang3.StringUtils;
+import org.springframework.stereotype.Component;
+
+import java.util.Set;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.ConcurrentMap;
+
+/**
+ * @author chan
+ * @date 2024-09-06 22:39
+ */
+@Component
+@WsServerEndpoint(value = "/ws/midjourney/{userToken}/{taskId}")
+@Slf4j
+@RequiredArgsConstructor
+public class SocketOldServer {
+
+	private static ConcurrentMap<String, Set<Session>> sessionPool = new ConcurrentHashMap<>();
+	private static ConcurrentMap<String,String> sessionIds = new ConcurrentHashMap<>();
+
+	private static final String ADMIN = "yhlxj";
+
+	private final MidjourneyService midjourneyService;
+
+	@HandshakeBefore
+	public void handshakeBefore(HttpHeaders headers,@PathParam(value="userToken") String userToken) {
+		log.info("handshakeBefore userToken: {}  host: {}", userToken, headers.get("HOST"));
+	}
+
+	/**
+	 * 用户连接时触发
+	 */
+	@OnOpen
+	public void open(Session session,@PathParam(value="userToken") String userToken, @PathParam(value = "taskId") String taskId){
+		log.info("[WS建立连接],token:{}, taskId:{}", userToken, taskId);
+		//检查uniodId是否合法
+		if (StringUtils.isBlank(userToken) || StringUtils.isBlank(taskId)) {
+
+			log.error("[WS建立连接失败],[参数非法],{},{}", userToken, taskId);
+
+			session.sendText("入参错误");
+			session.close();
+		}
+		MidjourneyUser user = getUser(userToken);
+		if (user == null) {
+			log.error("[WS建立连接失败],[用户过期],{},{}", userToken, taskId);
+			session.sendText("用户过期");
+			session.close();
+		} else {
+			if (user.getIsBlack()) {
+				log.error("[WS建立连接失败],[用户拉黑],{},{}", userToken, taskId);
+				session.sendText("状态异常");
+				session.close();
+			}
+			MidjourneyUserConversation conversation = midjourneyService.getConversationById(taskId);
+
+			if (conversation == null) {
+				session.sendText(packMsg(WsMessageTypeEnum.SERVER_TASK_NOT_FIND, "任务不存在."));
+			} else {
+				session.sendText(packMsg(WsMessageTypeEnum.SERVER_TASK_INFO, new JSONObject(conversation)));
+			}
+		}
+		Set<Session> sessions = sessionPool.computeIfAbsent(userToken, k -> ConcurrentHashMap.newKeySet());
+		if (sessions.size() >= 3){
+			sessions.stream().findFirst().ifPresent(s -> onClose(s,userToken));
+		}
+		sessions.add(session);
+		sessionPool.put(userToken, sessions);
+		sessionIds.put(session.getId(), userToken);
+	}
+
+	/**
+	 * 收到信息时触发
+	 * @param message
+	 */
+	@OnMessage
+	public void onMessage(Session session, String message, @PathParam(value = "userToken") String userToken, @PathParam(value = "taskId") String taskId){
+		try {
+			log.info("[WS收到消息],{},{},{}", userToken, taskId, message);
+
+			JSONObject body = null;//消息具体内容
+
+			//处理消息类型,是轮询还是消息
+			switch (WsMessageTypeEnum.getMessageType(message.split(":")[0])) {
+				//ping pong
+				case PING:
+					session.sendText("pong");
+					return;
+				case MESSAGE:
+					body = JSONUtil.parseObj(message.substring(4));
+					break;
+				default:
+					log.error("[WS消息类型异常],{},{},{}", userToken, taskId, message);
+					return;
+			}
+
+			//根据子类型进行处理
+			if (WsMessageTypeEnum.getMessageType(body.getStr("type")) == WsMessageTypeEnum.CLIENT_TASK_ID) { //
+				MidjourneyUserConversation conversation = midjourneyService.getConversationById(body.getStr("taskId"));
+				session.sendText(packMsg(WsMessageTypeEnum.SERVER_TASK_INFO, new JSONObject(conversation)));
+			} else {
+				log.error("[WS消息子类型异常],{},{},{}", userToken, taskId, message);
+			}
+
+		} catch (Exception e) {
+			log.error("[WS接收消息失败],{},{},{},{}", userToken, taskId, message, e);
+		}
+	}
+
+	/**
+	 * 连接关闭触发
+	 */
+	@OnClose
+	public void onClose(Session session,@PathParam String userToken){
+		Set<Session> sessions = sessionPool.get(userToken);
+		if (CollectionUtil.isNotEmpty(sessions)) {
+			sessions.removeIf(session1 -> session1.getId().equals(session.getId()));
+			sessionIds.remove(session.getId());
+		}
+		log.info("client【{}】【{}】断开连接",userToken, session.getId());
+	}
+
+	/**
+	 * 发生错误时触发
+	 * @param session
+	 * @param error
+	 */
+	@OnError
+	public void onError(Session session, Throwable error, @PathParam("userToken") String userToken) {
+		log.info("client【{}】onError 【{}】",userToken, error.getMessage());
+		Set<Session> sessions = sessionPool.get(userToken);
+		if (CollectionUtil.isNotEmpty(sessions)) {
+			sessions.removeIf(session1 -> session1.getId().equals(session.getId()));
+			sessionIds.remove(session.getId());
+		}
+		log.info("client【{}】【{}】onError",userToken, session.getId());
+	}
+
+	/**
+	 * 发生事件时触发
+	 * @param session
+	 * @param evt
+	 */
+	@OnEvent
+	public void onEvent(Session session,@PathParam(value="userToken") String userToken, Object evt) {
+		if (evt instanceof IdleStateEvent) {
+			IdleStateEvent idleStateEvent = (IdleStateEvent) evt;
+			switch (idleStateEvent.state()) {
+				case READER_IDLE:
+					log.info("clent : {} heartbeat read timeout event",userToken);
+					break;
+				case WRITER_IDLE:
+					log.info("clent : {} heartbeat write timeout event",userToken);
+					session.close();
+					break;
+				case ALL_IDLE:
+					log.info("clent : {} heartbeat all timeout event",userToken);
+					break;
+				default:
+					break;
+			}
+		}
+	}
+
+	/**
+	 *信息发送的方法
+	 */
+	public static void sendMessage(WsMessageTypeEnum typeEnum,String message,String userToken){
+		sendMessage(packMsg(typeEnum, message), userToken);
+	}
+
+	public static void sendMessage(String message,String userToken){
+		Set<Session> sessions = sessionPool.get(userToken);
+		if (CollectionUtil.isNotEmpty(sessions)) {
+			sessions.forEach((session)->{
+			log.info("[WS发送消息],session:{},message:{}",session, message);
+				session.sendText(packMsg(WsMessageTypeEnum.SERVER_TASK_INFO, new JSONObject(message)));
+			});
+		}
+	}
+
+	/**
+	 * 获取当前连接数
+	 * @return
+	 */
+	public static int getOnlineNum(){
+		if(sessionIds.values().contains(ADMIN)) {
+
+			return sessionPool.size()-1;
+		}
+		return sessionPool.size();
+	}
+
+	/**
+	 * 获取在线用户名以逗号隔开
+	 * @return
+	 */
+	public static String getOnlineUsers(){
+		StringBuffer users = new StringBuffer();
+		for (String key : sessionIds.keySet()) {//ADMIN是服务端自己的连接,不能算在线人数
+			if (!ADMIN.equals(sessionIds.get(key))) {
+				users.append(sessionIds.get(key)+",");
+			}
+		}
+		return users.toString();
+	}
+
+
+	public MidjourneyUser getUser(String userToken) {
+		return midjourneyService.getUser(userToken);
+	}
+
+	public static String packMsg(WsMessageTypeEnum type, Object obj){
+		JSONObject result = new JSONObject();
+		result.putOpt("type", type.type);
+		result.putOpt("content", JSONUtil.toJsonStr(obj));
+		return WsMessageTypeEnum.MESSAGE.type + ":" + result.toString();
+	}
+}

+ 12 - 21
midjourney-netty-websocket-web/src/main/java/com/yhlxj/netty/websocket/midjounrey/service/SocketServer.java

@@ -24,7 +24,7 @@ import java.util.concurrent.ConcurrentMap;
  * @date 2024-09-06 22:39
  */
 @Component
-@WsServerEndpoint(value = "/ws/midjourney/{userToken}/{taskId}")
+@WsServerEndpoint(value = "/ws/midjourney/{userToken}")
 @Slf4j
 @RequiredArgsConstructor
 public class SocketServer {
@@ -45,34 +45,25 @@ public class SocketServer {
 	 * 用户连接时触发
 	 */
 	@OnOpen
-	public void open(Session session,@PathParam(value="userToken") String userToken, @PathParam(value = "taskId") String taskId){
-		log.info("[WS建立连接],token:{}, taskId:{}", userToken, taskId);
+	public void open(Session session,@PathParam(value="userToken") String userToken){
+		log.info("[WS建立连接],token:{}", userToken);
 		//检查uniodId是否合法
-		if (StringUtils.isBlank(userToken) || StringUtils.isBlank(taskId)) {
-
-			log.error("[WS建立连接失败],[参数非法],{},{}", userToken, taskId);
-
+		if (StringUtils.isBlank(userToken)) {
+			log.error("[WS建立连接失败],[参数非法],{}", userToken);
 			session.sendText("入参错误");
 			session.close();
 		}
 		MidjourneyUser user = getUser(userToken);
 		if (user == null) {
-			log.error("[WS建立连接失败],[用户过期],{},{}", userToken, taskId);
+			log.error("[WS建立连接失败],[用户过期],{}", userToken);
 			session.sendText("用户过期");
 			session.close();
 		} else {
 			if (user.getIsBlack()) {
-				log.error("[WS建立连接失败],[用户拉黑],{},{}", userToken, taskId);
+				log.error("[WS建立连接失败],[用户拉黑],{}", userToken);
 				session.sendText("状态异常");
 				session.close();
 			}
-			MidjourneyUserConversation conversation = midjourneyService.getConversationById(taskId);
-
-			if (conversation == null) {
-				session.sendText(packMsg(WsMessageTypeEnum.SERVER_TASK_NOT_FIND, "任务不存在."));
-			} else {
-				session.sendText(packMsg(WsMessageTypeEnum.SERVER_TASK_INFO, new JSONObject(conversation)));
-			}
 		}
 		Set<Session> sessions = sessionPool.computeIfAbsent(userToken, k -> ConcurrentHashMap.newKeySet());
 		if (sessions.size() >= 3){
@@ -88,9 +79,9 @@ public class SocketServer {
 	 * @param message
 	 */
 	@OnMessage
-	public void onMessage(Session session, String message, @PathParam(value = "userToken") String userToken, @PathParam(value = "taskId") String taskId){
+	public void onMessage(Session session, String message, @PathParam(value = "userToken") String userToken){
 		try {
-			log.info("[WS收到消息],{},{},{}", userToken, taskId, message);
+			log.info("[WS收到消息],{},{}", userToken, message);
 
 			JSONObject body = null;//消息具体内容
 
@@ -104,7 +95,7 @@ public class SocketServer {
 					body = JSONUtil.parseObj(message.substring(4));
 					break;
 				default:
-					log.error("[WS消息类型异常],{},{},{}", userToken, taskId, message);
+					log.error("[WS消息类型异常],{},{}", userToken, message);
 					return;
 			}
 
@@ -113,11 +104,11 @@ public class SocketServer {
 				MidjourneyUserConversation conversation = midjourneyService.getConversationById(body.getStr("taskId"));
 				session.sendText(packMsg(WsMessageTypeEnum.SERVER_TASK_INFO, new JSONObject(conversation)));
 			} else {
-				log.error("[WS消息子类型异常],{},{},{}", userToken, taskId, message);
+				log.error("[WS消息子类型异常],{},{}", userToken, message);
 			}
 
 		} catch (Exception e) {
-			log.error("[WS接收消息失败],{},{},{},{}", userToken, taskId, message, e);
+			log.error("[WS接收消息失败],{},{},{}", userToken, message, e);
 		}
 	}