Kaynağa Gözat

fix 修改mj缓存序列化问题

zwhui 1 yıl önce
ebeveyn
işleme
6648339a42

+ 11 - 2
midjourney-netty-websocket-web/src/main/java/com/yhlxj/netty/websocket/midjounrey/dao/model/entity/MidjourneyUserConversation.java → midjourney-netty-websocket-web/src/main/java/com/yhlxj/dao/model/entity/MidjourneyUserConversation.java

@@ -1,4 +1,4 @@
-package com.yhlxj.netty.websocket.midjounrey.dao.model.entity;
+package com.yhlxj.dao.model.entity;
 
 import com.baomidou.mybatisplus.annotation.TableField;
 import com.ejlchina.searcher.bean.DbIgnore;
@@ -6,6 +6,9 @@ import com.ejlchina.searcher.bean.SearchBean;
 import com.fasterxml.jackson.annotation.JsonInclude;
 import com.yhlxj.netty.websocket.midjounrey.dao.model.BaseEntity;
 import lombok.Data;
+import lombok.Getter;
+import lombok.Setter;
+import lombok.ToString;
 import org.apache.commons.lang3.StringUtils;
 
 import java.io.Serializable;
@@ -14,13 +17,19 @@ import java.io.Serializable;
  * @author zwhui
  * @date 2024/4/23 11:10
  */
-@Data
+@Getter
+@Setter
+@ToString
 @SearchBean(tables = "midjourney_user_conversation")
 @JsonInclude(JsonInclude.Include.NON_NULL)
 public class MidjourneyUserConversation extends BaseEntity implements Serializable {
 
     private Long userId;
 
+    @DbIgnore
+    @TableField(exist = false)
+    private String userToken;
+
     /**
      * plus服务任务id
      */

+ 33 - 1
midjourney-netty-websocket-web/src/main/java/com/yhlxj/netty/websocket/midjounrey/Application.java

@@ -1,8 +1,19 @@
 package com.yhlxj.netty.websocket.midjounrey;
 
+import com.fasterxml.jackson.annotation.JsonAutoDetect;
+import com.fasterxml.jackson.annotation.JsonTypeInfo;
+import com.fasterxml.jackson.annotation.PropertyAccessor;
+import com.fasterxml.jackson.databind.ObjectMapper;
+import com.fasterxml.jackson.databind.jsontype.impl.LaissezFaireSubTypeValidator;
 import com.yhlxj.netty.websocket.midjounrey.service.SocketServer;
 import org.springframework.boot.SpringApplication;
 import org.springframework.boot.autoconfigure.SpringBootApplication;
+import org.springframework.context.annotation.Bean;
+import org.springframework.data.redis.connection.RedisConnectionFactory;
+import org.springframework.data.redis.core.RedisTemplate;
+import org.springframework.data.redis.core.StringRedisTemplate;
+import org.springframework.data.redis.serializer.Jackson2JsonRedisSerializer;
+import org.springframework.data.redis.serializer.StringRedisSerializer;
 
 
 /**
@@ -11,7 +22,28 @@ import org.springframework.boot.autoconfigure.SpringBootApplication;
  */
 @SpringBootApplication
 public class Application {
-    public static void main(String[] args)  {
+    public static void main(String[] args) {
         SpringApplication.run(Application.class).getBean(SocketServer.class);
     }
+
+    @Bean(name = "RedisTemp")
+    public RedisTemplate<String, String> redisTemplate(RedisConnectionFactory factory) {
+        StringRedisTemplate template = new StringRedisTemplate(factory);
+        //定义key序列化方式
+        //RedisSerializer<String> redisSerializer = new StringRedisSerializer();//Long类型会出现异常信息;需要我们上面的自定义key生成策略,一般没必要
+        //定义value的序列化方式
+        Jackson2JsonRedisSerializer jackson2JsonRedisSerializer = new Jackson2JsonRedisSerializer(Object.class);
+        ObjectMapper om = new ObjectMapper();
+        om.setVisibility(PropertyAccessor.ALL, JsonAutoDetect.Visibility.ANY);
+        om.activateDefaultTyping(LaissezFaireSubTypeValidator.instance,
+                ObjectMapper.DefaultTyping.NON_FINAL,
+                JsonTypeInfo.As.WRAPPER_ARRAY);
+        jackson2JsonRedisSerializer.setObjectMapper(om);
+
+        // template.setKeySerializer(redisSerializer);
+        template.setValueSerializer(jackson2JsonRedisSerializer);
+        template.setHashValueSerializer(jackson2JsonRedisSerializer);
+        template.afterPropertiesSet();
+        return template;
+    }
 }

+ 1 - 1
midjourney-netty-websocket-web/src/main/java/com/yhlxj/netty/websocket/midjounrey/dao/mapper/MidjourneyUserConversationMapper.java

@@ -1,7 +1,7 @@
 package com.yhlxj.netty.websocket.midjounrey.dao.mapper;
 
 import com.baomidou.mybatisplus.core.mapper.BaseMapper;
-import com.yhlxj.netty.websocket.midjounrey.dao.model.entity.MidjourneyUserConversation;
+import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
 
 /**
  * @author zwhui

+ 1 - 1
midjourney-netty-websocket-web/src/main/java/com/yhlxj/netty/websocket/midjounrey/redis/RedisService.java

@@ -26,7 +26,7 @@ import java.util.concurrent.TimeUnit;
 @SuppressWarnings("all")
 public class RedisService {
 
-    @Resource
+    @Resource(name = "RedisTemp")
     private RedisTemplate<String, Object> redisTemplate;
 
     public static String env = "dev";

+ 5 - 15
midjourney-netty-websocket-web/src/main/java/com/yhlxj/netty/websocket/midjounrey/service/MidjourneyService.java

@@ -1,12 +1,11 @@
 package com.yhlxj.netty.websocket.midjounrey.service;
 
-import cn.hutool.json.JSONUtil;
 import com.baomidou.mybatisplus.core.toolkit.Wrappers;
 import com.cyksj.common.util.QueryWrapperUtils;
+import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
 import com.yhlxj.netty.websocket.midjounrey.dao.mapper.MidjourneyUserConversationMapper;
 import com.yhlxj.netty.websocket.midjounrey.dao.mapper.MidjourneyUserMapper;
 import com.yhlxj.netty.websocket.midjounrey.dao.model.entity.MidjourneyUser;
-import com.yhlxj.netty.websocket.midjounrey.dao.model.entity.MidjourneyUserConversation;
 import com.yhlxj.netty.websocket.midjounrey.redis.RedisService;
 import lombok.RequiredArgsConstructor;
 import lombok.extern.slf4j.Slf4j;
@@ -37,25 +36,16 @@ public class MidjourneyService{
 
    
     public MidjourneyUserConversation getConversationById(String id){
-
         //从缓存中查询 该任务是否存在
-        MidjourneyUserConversation conversation;
-        Object obj = redisService.get(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + id);
-        if(obj == null){
+        MidjourneyUserConversation conversation = (MidjourneyUserConversation)redisService.get(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + id);
+        log.info("conversation:{}",conversation);
+        if(conversation == null){
             conversation = conversationMapper.selectOne(QueryWrapperUtils.buildWrapper((wrapper)->{
                 wrapper.eq(MidjourneyUserConversation::getTaskId, id);
             }));
+            log.info("数据库查询: {}", conversation);
             redisService.set(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + id, conversation, RedisService.key.MIDJOURNEY_CONVERSATION.getTimeout());
-        }else {
-            try {
-                conversation = (MidjourneyUserConversation)obj;
-            }catch (Exception e){
-                String jsonStr = JSONUtil.toJsonStr(obj);
-                conversation = JSONUtil.parseObj(jsonStr).toBean(MidjourneyUserConversation.class);
-                redisService.set(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + id, conversation, RedisService.key.MIDJOURNEY_CONVERSATION.getTimeout());
-            }
         }
-
         return conversation;
     }
 

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

@@ -3,9 +3,8 @@ 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.cyksj.common.util.SpringCtxUtils;
+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.entity.MidjourneyUserConversation;
 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;
@@ -102,7 +101,7 @@ public class SocketServer {
 					session.sendText("pong");
 					return;
 				case MESSAGE:
-					body = JSONUtil.parseObj(message.substring(8));
+					body = JSONUtil.parseObj(message.substring(4));
 					break;
 				default:
 					log.error("[WS消息类型异常],{},{},{}", userToken, taskId, message);

+ 1 - 1
midjourney-netty-websocket-web/src/main/resources/application.yml

@@ -20,7 +20,7 @@ spring:
       - ${activeProfile}
   datasource:
     type: com.zaxxer.hikari.HikariDataSource
-    url: jdbc:mysql://sh-cdb-4tdnmrqa.sql.tencentcdb.com:63893/yhlxj?useAffectedRows=true
+    url: jdbc:mysql://mysql.yhlx.vip:3306/yhlxj?useAffectedRows=true
     username: yhlxj
     password: yhlxj123
     platform: mysql

+ 3 - 3
midjourney/src/main/java/com/yhlxj/service/midjourney/impl/MidjourneyServiceImpl.java

@@ -954,9 +954,6 @@ public class MidjourneyServiceImpl implements MidjourneyService {
         if (NumberUtil.isNumber(conversation.getAction())){
             conversation.setAction(action_list.get(jsons.getInt("action")));
         }
-        if (conversation.getPrompt() == null){
-            conversation.setPrompt(jsons.getStr("promptEn", jsons.getStr("promptFull")));
-        }
 
         String[] states = jsons.getStr("state").split(",");
         conversation.setUserId(Long.parseLong(states[0]));
@@ -965,6 +962,9 @@ public class MidjourneyServiceImpl implements MidjourneyService {
         conversation.setId(null);
         JSONObject jsonObject = new JSONObject(conversation.getProperties());
         String content = jsonObject.getStr("messageContent");
+        if (conversation.getPrompt() == null){
+            conversation.setPrompt(jsonObject.getStr("finalPrompt",jsons.getStr("promptEn",jsons.getStr("promptFull"))));
+        }
         if(StringUtils.isNotBlank(content)){
             jsonObject.set("messageContent","**" + conversation.getPrompt() + "** - <@mj>");
         }

+ 0 - 166
midjourney/src/main/java/com/yhlxj/web/wss/MidjourneyServerEndpoint.java

@@ -1,166 +0,0 @@
-package com.yhlxj.web.wss;
-
-import cn.hutool.json.JSONObject;
-import cn.hutool.json.JSONUtil;
-import com.cyksj.common.task.GlobalThreadPoolTaskExecutor;
-import com.yhlxj.dao.model.entity.MidjourneyUser;
-import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
-import com.yhlxj.dao.model.enums.WsMessageTypeEnum;
-import com.yhlxj.redis.RedisService;
-import com.yhlxj.service.midjourney.MidjourneyService;
-import com.yhlxj.service.midjourney.impl.MidjourneyServiceImpl;
-import com.yhlxj.util.SpringCtxUtils;
-import lombok.extern.slf4j.Slf4j;
-import org.apache.commons.lang3.StringUtils;
-import org.springframework.stereotype.Component;
-
-import javax.websocket.*;
-import javax.websocket.server.PathParam;
-import javax.websocket.server.ServerEndpoint;
-import java.io.IOException;
-import java.util.Map;
-import java.util.concurrent.ConcurrentHashMap;
-import java.util.concurrent.TimeUnit;
-
-/**
- * @author chan
- * @date 2024/6/28 13:59
- */
-@Component
-@ServerEndpoint("/ws/midjourney/{userToken}/{taskId}")
-@Slf4j
-public class MidjourneyServerEndpoint {
-    private static final GlobalThreadPoolTaskExecutor TASK_EXECUTOR = GlobalThreadPoolTaskExecutor.getInstance();
-
-    //webSocket
-    public static Map<Long, WssSession> WSS_SESSION_MAP = new ConcurrentHashMap<>();
-
-    /**
-     * 连接建立回调
-     *
-     * @param session
-     */
-    @OnOpen
-    public void onOpen(Session session, @PathParam("userToken") String userToken, @PathParam("taskId") String taskId) throws Exception {
-        log.info("[WS建立连接],token:{}, taskId:{}", userToken, taskId);
-
-        session.setMaxIdleTimeout(TimeUnit.MINUTES.toMillis(60));
-        //检查uniodId是否合法
-        if (StringUtils.isBlank(userToken) || StringUtils.isBlank(taskId)) {
-
-            log.error("[WS建立连接失败],[参数非法],{},{}", userToken, taskId);
-
-            session.getBasicRemote().sendText("入参错误");
-            session.close();
-        }
-        MidjourneyUser user = getUser(userToken);
-        if (user == null) {
-            log.error("[WS建立连接失败],[用户过期],{},{}", userToken, taskId);
-            session.getBasicRemote().sendText("用户过期");
-            session.close();
-        } else {
-            MidjourneyService midjourneyService = SpringCtxUtils.getBean(MidjourneyServiceImpl.class);
-            MidjourneyUserConversation conversation = midjourneyService.getConversationById(taskId);
-            WssSession wsSession = new WssSession(session);
-            if (conversation == null) {
-                wsSession.sendMessage(WsMessageTypeEnum.SERVER_TASK_NOT_FIND, "任务不存在.");
-            } else {
-                WSS_SESSION_MAP.put(user.getId(), wsSession);
-                wsSession.sendMessage(WsMessageTypeEnum.SERVER_TASK_INFO, new JSONObject(conversation));
-            }
-
-        }
-
-    }
-
-    /**
-     * 收到客户端消息回调
-     *
-     * @param message
-     * @param session
-     */
-    @OnMessage
-    public void OnMessage(String message, Session session, @PathParam("userToken") String userToken, @PathParam("taskId") String taskId) {
-        try {
-            log.info("[WS收到消息],{},{},{}", userToken, taskId, message);
-
-            WssSession wsSession = new WssSession(session);
-            JSONObject body = null;//消息具体内容
-
-            //处理消息类型,是轮询还是消息
-            switch (WsMessageTypeEnum.getMessageType(message.split(":")[0])) {
-                //ping pong
-                case PING:
-                    wsSession.pong();
-                    return;
-                case MESSAGE:
-                    body = JSONUtil.parseObj(message.substring(8));
-                    break;
-                default:
-                    log.error("[WS消息类型异常],{},{},{}", userToken, taskId, message);
-                    return;
-            }
-
-            //根据子类型进行处理
-            switch (WsMessageTypeEnum.getMessageType(body.getStr("type"))) {
-
-                case CLIENT_TASK_ID: //首次查看任务状态
-                    MidjourneyService midjourneyService = SpringCtxUtils.getBean(MidjourneyService.class);
-                    midjourneyService.getConversationById(body.getStr("taskId"));
-                    break;
-
-                case CLIENT_CLOSE:
-                    log.info("[WS连接关闭],{},{}", userToken, taskId);
-                    //dailyQuizService.close(uniodId);
-                    break;
-
-                default:
-                    log.error("[WS消息子类型异常],{},{},{}", userToken, taskId, message);
-
-            }
-
-        } catch (Exception e) {
-            log.error("[WS接收消息失败],{},{},{},{}", userToken, taskId, message, e);
-        }
-    }
-
-    /**
-     * 发生错误回调
-     *
-     * @param session
-     * @param error
-     */
-    @OnError
-    public void OnError(Session session, Throwable error, @PathParam("userToken") String userToken, @PathParam("taskId") String taskId) throws Exception {
-        log.error("[WS ONERROR],{}", error.getMessage());
-        if (session.isOpen()) {
-            session.close();
-        }
-        log.info("[WS ONERROR],移除连接池内连接");
-        MidjourneyUser user = getUser(userToken);
-        WSS_SESSION_MAP.remove(user.getId());
-    }
-
-    /**
-     * 关闭连接回调
-     *
-     * @param session
-     */
-    @OnClose
-    public void OnClose(Session session, @PathParam("userToken") String userToken, @PathParam("taskId") String taskId) throws Exception {
-
-        log.info("[WS连接关闭],{},{}", userToken, taskId);
-        if (session.isOpen()) {
-            session.close();
-        }
-        log.info("[WS连接关闭],移除连接池内连接");
-        MidjourneyUser user = getUser(userToken);
-        WSS_SESSION_MAP.remove(user.getId());
-    }
-
-    public MidjourneyUser getUser(String userToken) throws Exception{
-        MidjourneyService midjourneyService = SpringCtxUtils.getBean(MidjourneyServiceImpl.class);
-        return midjourneyService.getUser(userToken);
-    }
-
-}

+ 0 - 13
midjourney/src/main/java/com/yhlxj/web/wss/WssSendable.java

@@ -1,13 +0,0 @@
-package com.yhlxj.web.wss;
-
-import com.yhlxj.dao.model.enums.WsMessageTypeEnum;
-
-import java.io.IOException;
-
-public interface WssSendable {
-    void sendMessage(WsMessageTypeEnum type, Object object) throws IOException;
-
-    void close() throws IOException;
-
-    boolean isOpen() throws IOException;
-}

+ 0 - 143
midjourney/src/main/java/com/yhlxj/web/wss/WssSession.java

@@ -1,143 +0,0 @@
-package com.yhlxj.web.wss;
-
-import cn.hutool.json.JSONObject;
-import cn.hutool.json.JSONUtil;
-import com.yhlxj.dao.model.enums.WsMessageTypeEnum;
-import lombok.AllArgsConstructor;
-import lombok.Data;
-import org.slf4j.Logger;
-import org.slf4j.LoggerFactory;
-
-import javax.websocket.*;
-import java.io.IOException;
-import java.io.Serializable;
-import java.util.List;
-import java.util.Set;
-
-/**
- * @author chan
- * @date 2024/6/27 18:40
- */
-@Data
-@AllArgsConstructor
-public class WssSession implements WssSendable , Serializable {
-    private static Logger logger = LoggerFactory.getLogger(WssSession.class);
-
-    private Session session;
-
-    /**
-     * 发送自定义类型消息
-     *
-     * @param type
-     * @param object
-     */
-    @Override
-    public void sendMessage(WsMessageTypeEnum type, Object object) {
-        if (session.isOpen()) {
-            //组织成形如 message:{"type":"type", "data":{...}}想
-            JSONObject result = new JSONObject();
-            result.putOpt("type", type.type);
-            result.putOpt("content", JSONUtil.toJsonStr(object));
-
-            logger.info("[WS-SENDMESSAGE],{}", WsMessageTypeEnum.MESSAGE.type + ":" + result.toString());
-            try {
-                getSession().getBasicRemote().sendText(WsMessageTypeEnum.MESSAGE.type + ":" + result.toString());
-            } catch (Exception e) {
-                logger.info("[WS-SENDMESSAGE],[失败],{}", WsMessageTypeEnum.MESSAGE.type + ":" + result.toString());
-            }
-        }
-    }
-
-    /**
-     * 发送pong
-     *
-     * @throws IOException
-     */
-    public void pong() throws IOException {
-        getSession().getBasicRemote().sendText(WsMessageTypeEnum.PONG.type);
-    }
-
-
-    //-----------------DELEGATOR---------------------
-
-    public WebSocketContainer getContainer() {
-        return session.getContainer();
-    }
-
-    public void addMessageHandler(MessageHandler messageHandler) throws IllegalStateException {
-        session.addMessageHandler(messageHandler);
-    }
-
-    public Set<MessageHandler> getMessageHandlers() {
-        return session.getMessageHandlers();
-    }
-
-    public void removeMessageHandler(MessageHandler messageHandler) {
-        session.removeMessageHandler(messageHandler);
-    }
-
-    public String getProtocolVersion() {
-        return session.getProtocolVersion();
-    }
-
-    public String getNegotiatedSubprotocol() {
-        return session.getNegotiatedSubprotocol();
-    }
-
-    public List<Extension> getNegotiatedExtensions() {
-        return session.getNegotiatedExtensions();
-    }
-
-    public boolean isSecure() {
-        return session.isSecure();
-    }
-
-    public boolean isOpen() {
-        return session.isOpen();
-    }
-
-    public long getMaxIdleTimeout() {
-        return session.getMaxIdleTimeout();
-    }
-
-    public void setMaxIdleTimeout(long l) {
-        session.setMaxIdleTimeout(l);
-    }
-
-    public void setMaxBinaryMessageBufferSize(int i) {
-        session.setMaxBinaryMessageBufferSize(i);
-    }
-
-    public int getMaxBinaryMessageBufferSize() {
-        return session.getMaxBinaryMessageBufferSize();
-    }
-
-    public void setMaxTextMessageBufferSize(int i) {
-        session.setMaxTextMessageBufferSize(i);
-    }
-
-    public int getMaxTextMessageBufferSize() {
-        return session.getMaxTextMessageBufferSize();
-    }
-
-    public RemoteEndpoint.Async getAsyncRemote() {
-        return session.getAsyncRemote();
-    }
-
-    public RemoteEndpoint.Basic getBasicRemote() {
-        return session.getBasicRemote();
-    }
-
-    public String getId() {
-        return session.getId();
-    }
-
-    public void close() throws IOException {
-        session.close();
-    }
-
-    public void close(CloseReason closeReason) throws IOException {
-        session.close(closeReason);
-    }
-
-}