|
@@ -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();
|
|
|
|
|
+ }
|
|
|
|
|
+}
|