chenbiao 2 лет назад
Родитель
Сommit
1435a958af

+ 5 - 0
midjourney/pom.xml

@@ -44,6 +44,11 @@
             <optional>true</optional>
             <optional>true</optional>
         </dependency>
         </dependency>
 
 
+        <dependency>
+            <groupId>org.springframework.boot</groupId>
+            <artifactId>spring-boot-starter-websocket</artifactId>
+        </dependency>
+
 
 
     </dependencies>
     </dependencies>
 
 

+ 27 - 0
midjourney/src/main/java/com/yhlxj/dao/model/dto/SubmitSeedDTO.java

@@ -0,0 +1,27 @@
+package com.yhlxj.dao.model.dto;
+
+import lombok.Data;
+
+import javax.validation.constraints.NotNull;
+
+/**
+ * @author chan
+ * @date 2024/6/28 21:48
+ */
+@Data
+public class SubmitSeedDTO {
+
+    /**
+     * 执行Imagine任务生成的ID
+     */
+    @NotNull(message = "任务ID不能为空")
+    private Long taskId;
+
+    /**
+     * 机器人类型
+     * bot类型,mj(默认)或niji,可用值:MID_JOURNEY,NIJI_JOURNEY,示例值(MID_JOURNEY)
+     * MID_JOURNEY
+     * NIJI_JOURNEY
+     */
+    private String botType = "MID_JOURNEY";
+}

+ 5 - 9
midjourney/src/main/java/com/yhlxj/dao/model/enums/WsMessageTypeEnum.java

@@ -12,21 +12,17 @@ import java.util.stream.Stream;
 public enum WsMessageTypeEnum {
 public enum WsMessageTypeEnum {
 
 
 
 
-    PING("ping"), PONG("pong"), MESSAGE("message"),
+    PING("ping"), PONG("pong"), MESSAGE("msg"),
 
 
     //客服端反馈子类型
     //客服端反馈子类型
-    CLIENT_Union_ID("task_id"),//开始连接
+    CLIENT_TASK_ID("task_id"),//开始连接
     CLIENT_CLOSE("close"),//关闭连接
     CLIENT_CLOSE("close"),//关闭连接
-    CLIENT_DAILY_RESULT("daily_result"),//完成
-    CLIENT_TUNNELID_CHANGED("tunnelId_changed"),//信道变更,暂未使用
+    CLIENT_ACK("ack"),//收到
     CLIENT_RESURGENCE("resurgence"),//客服端复活
     CLIENT_RESURGENCE("resurgence"),//客服端复活
 
 
     //服务端发送子类型
     //服务端发送子类型
-    SERVER_HAS_CHALLENGE("has_challenge"),
-    CONNECTION_TIMED_OUT("connection_timed_out"),
-    SERVER_QUESTION("question"),//服务端发送问题
-    SERVER_TUNNELID_CHANGED_DONE("tunnelId_change_done"),//信道切换完成
-    SERVER_GET_ANSWER("getAnswer"),
+    SERVER_TASK_INFO("task_info"),//服务端发送问题
+    SERVER_TASK_NOT_FIND("task_not_find"),
     SERVER_ERROR("error"),//错误信息
     SERVER_ERROR("error"),//错误信息
     SERVER_FAIL("fail"),//错误信息
     SERVER_FAIL("fail"),//错误信息
     ;
     ;

+ 1 - 0
midjourney/src/main/java/com/yhlxj/redis/RedisService.java

@@ -753,6 +753,7 @@ public class RedisService {
         MIDJOURNEY_CONVERSATION("midjourney:conversation:", "midjourney conversation", 60 * 60 * 2L),
         MIDJOURNEY_CONVERSATION("midjourney:conversation:", "midjourney conversation", 60 * 60 * 2L),
         MIDJOURNEY_ACCOUNT("midjourney:account:", "midjourney account", 60 * 60 * 48L),
         MIDJOURNEY_ACCOUNT("midjourney:account:", "midjourney account", 60 * 60 * 48L),
         MIDJOURNEY_QUERY("midjourney:query:", "midjourney query", 60 * 60 * 2L),
         MIDJOURNEY_QUERY("midjourney:query:", "midjourney query", 60 * 60 * 2L),
+        MIDJOURNEY_WS_SESSION("midjourney:ws:session", "midjourney ws session", 60 * 60 * 30L),
         ;
         ;
 
 
         private String name;
         private String name;

+ 4 - 0
midjourney/src/main/java/com/yhlxj/service/midjourney/MidjourneyService.java

@@ -23,8 +23,12 @@ public interface MidjourneyService {
 
 
     MidjourneyUserConversation submitShorten(MidjourneyUser user, String botType, String prompt) throws Exception;
     MidjourneyUserConversation submitShorten(MidjourneyUser user, String botType, String prompt) throws Exception;
 
 
+    MidjourneyUserConversation submitSeed(MidjourneyUser user, Long taskId, String botType) throws Exception;
+
     List<MidjourneyUserConversation> listConversationByIds(Integer mode, List<Long> ids) throws Exception;
     List<MidjourneyUserConversation> listConversationByIds(Integer mode, List<Long> ids) throws Exception;
 
 
+    MidjourneyUserConversation getConversationById(Long id);
+
     Object submitAction(MidjourneyUser user, Long taskId, String customId, Long num, String botType) throws Exception;
     Object submitAction(MidjourneyUser user, Long taskId, String customId, Long num, String botType) throws Exception;
 
 
     MidjourneyUserConversation cancelConversation(MidjourneyUser user, Long id);
     MidjourneyUserConversation cancelConversation(MidjourneyUser user, Long id);

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

@@ -3,6 +3,7 @@ package com.yhlxj.service.midjourney.impl;
 import cn.hutool.core.collection.CollectionUtil;
 import cn.hutool.core.collection.CollectionUtil;
 import cn.hutool.core.io.FileUtil;
 import cn.hutool.core.io.FileUtil;
 import cn.hutool.core.map.MapUtil;
 import cn.hutool.core.map.MapUtil;
+import cn.hutool.core.util.URLUtil;
 import cn.hutool.http.HttpRequest;
 import cn.hutool.http.HttpRequest;
 import cn.hutool.http.HttpUtil;
 import cn.hutool.http.HttpUtil;
 import cn.hutool.json.JSONArray;
 import cn.hutool.json.JSONArray;
@@ -21,11 +22,13 @@ import com.yhlxj.dao.model.dto.MessageButton;
 import com.yhlxj.dao.model.entity.MidjourneyAccount;
 import com.yhlxj.dao.model.entity.MidjourneyAccount;
 import com.yhlxj.dao.model.entity.MidjourneyUser;
 import com.yhlxj.dao.model.entity.MidjourneyUser;
 import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
 import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
+import com.yhlxj.dao.model.enums.WsMessageTypeEnum;
 import com.yhlxj.dao.model.response.SubmitResult;
 import com.yhlxj.dao.model.response.SubmitResult;
 import com.yhlxj.redis.RedisService;
 import com.yhlxj.redis.RedisService;
 import com.yhlxj.service.common.EnvCommonService;
 import com.yhlxj.service.common.EnvCommonService;
 import com.yhlxj.service.midjourney.MidjourneyAccountService;
 import com.yhlxj.service.midjourney.MidjourneyAccountService;
 import com.yhlxj.service.midjourney.MidjourneyService;
 import com.yhlxj.service.midjourney.MidjourneyService;
+import com.yhlxj.web.wss.WssSession;
 import lombok.RequiredArgsConstructor;
 import lombok.RequiredArgsConstructor;
 import lombok.extern.slf4j.Slf4j;
 import lombok.extern.slf4j.Slf4j;
 import org.apache.commons.lang3.StringUtils;
 import org.apache.commons.lang3.StringUtils;
@@ -34,6 +37,8 @@ import org.springframework.stereotype.Service;
 
 
 import javax.imageio.stream.FileImageOutputStream;
 import javax.imageio.stream.FileImageOutputStream;
 import java.io.*;
 import java.io.*;
+import java.net.MalformedURLException;
+import java.net.URL;
 import java.nio.file.Files;
 import java.nio.file.Files;
 import java.nio.file.Path;
 import java.nio.file.Path;
 import java.util.*;
 import java.util.*;
@@ -169,6 +174,52 @@ public class MidjourneyServiceImpl implements MidjourneyService {
         return saveConversation(user.getId(), user.getMode(), result,"SHORTEN", StringUtils.EMPTY, botType);
         return saveConversation(user.getId(), user.getMode(), result,"SHORTEN", StringUtils.EMPTY, botType);
     }
     }
 
 
+    @Override
+    public MidjourneyUserConversation submitSeed(MidjourneyUser user, Long taskId, String botType) throws Exception {
+        SubmitResult result = seed(user.getMode(),"seed", taskId);
+        return saveConversation(user.getId(), user.getMode(), result,"SEED", StringUtils.EMPTY, botType);
+    }
+
+    private SubmitResult seed(Integer mode, String seed, Long taskId) {
+        String url = "";
+        if (mode == 1) {
+            url = FAST_HOST;
+        } else if (mode == 2) {
+            url = RELAX_HOST;
+        }
+        url = url + String.format("/mj/task/%d/image-seed", taskId);
+        String body = HttpRequest.get(url).header("Authorization", FAST_TOKEN).execute().body();
+        log.info("seed body:{}", body);
+        SubmitResult submitResult = Jsons.parseObject(body, SubmitResult.class);
+        int code = submitResult.getCode();
+        if (code != 1 && code != 21 && code != 22) {
+            if (code == 3) {
+                if(mode == 2 && StringUtils.isNotBlank(accountWithMinUsage)){
+                    redisService.hdel(RedisService.key.MIDJOURNEY_ACCOUNT.getName(),accountWithMinUsage);
+                    String finalAccountWithMinUsage = accountWithMinUsage;
+                    TASK_EXECUTOR.execute(() -> {
+                        MidjourneyAccount account = midjourneyAccountService.getOne(Wrappers.lambdaQuery(MidjourneyAccount.class).eq(MidjourneyAccount::getInstanceId, finalAccountWithMinUsage).last("limit 1"));
+                        midjourneyAccountService.updateStatus(account.getId());
+                    });
+                }
+                throw BusinessRuntimeException.getInstance("账号不存在");
+            }
+            if (code == 4) {
+                throw BusinessRuntimeException.getInstance(submitResult.getDescription());
+            }
+            if (code == 24) {
+                JSONObject jsonObject = JSONUtil.parseObj(submitResult.getResult());
+                throw BusinessRuntimeException.getInstance("prompt包含敏感词:"+jsonObject.getStr("bannedWord"));
+            }
+            log.error("action:" + action + " error message:" + submitResult.getDescription());
+            throw BusinessRuntimeException.getInstance("队列已满,请稍后尝试");
+        }
+        if (StringUtils.isNotBlank(accountWithMinUsage) && mode == 2) {
+            submitResult.setInstanceId(Long.valueOf(accountWithMinUsage));
+        }
+        return submitResult;
+    }
+
     /**
     /**
      * 恢复次数
      * 恢复次数
      */
      */
@@ -387,6 +438,18 @@ public class MidjourneyServiceImpl implements MidjourneyService {
         return list;
         return list;
     }
     }
 
 
+    @Override
+    public MidjourneyUserConversation getConversationById(Long id){
+
+        //从缓存中查询 该任务是否存在
+        MidjourneyUserConversation conversation = (MidjourneyUserConversation)redisService.get(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + id);
+        if(conversation == null){
+            conversation = conversationMapper.selectById(id);
+            redisService.set(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + id, conversation, RedisService.key.MIDJOURNEY_CONVERSATION.getTimeout());
+        }
+
+        return conversation;
+    }
 
 
     public void sync(MidjourneyUserConversation conversation) throws IOException {
     public void sync(MidjourneyUserConversation conversation) throws IOException {
         if ((StringUtils.isNotBlank(conversation.getProgress()) && progress.contains(conversation.getProgress())) || status.contains(conversation.getStatus())) {
         if ((StringUtils.isNotBlank(conversation.getProgress()) && progress.contains(conversation.getProgress())) || status.contains(conversation.getStatus())) {
@@ -513,16 +576,30 @@ public class MidjourneyServiceImpl implements MidjourneyService {
     }
     }
 
 
     @Override
     @Override
-    public void notifyHook(String json) {
+    public void notifyHook(String json) throws MalformedURLException {
         log.info("notifyHook:{}", json);
         log.info("notifyHook:{}", json);
         JSONObject jsons = JSONUtil.parseObj(json);
         JSONObject jsons = JSONUtil.parseObj(json);
         MidjourneyUserConversation conversation = JSONUtil.toBean(jsons, MidjourneyUserConversation.class);
         MidjourneyUserConversation conversation = JSONUtil.toBean(jsons, MidjourneyUserConversation.class);
         conversation.setUserId(jsons.getLong("state"));
         conversation.setUserId(jsons.getLong("state"));
         conversation.setTaskId(jsons.getLong("id"));
         conversation.setTaskId(jsons.getLong("id"));
         if (StringUtil.isNotBlank(conversation.getImageUrl())) {
         if (StringUtil.isNotBlank(conversation.getImageUrl())) {
-            conversation.setImageUrl(conversation.getImageUrl().replace("cdn.discordapp.com", "mj.galaxydvd.com"));
+            URL url = new URL(conversation.getImageUrl());
+
+            // 使用 Hutool 提取各个部分并构建新 URL
+            String newUrl = URLUtil.toURI(new URL(url.getProtocol(), "cdn.mj.liuliangbang.vip", url.getPort(), url.getFile()) + "?" + url.getQuery()).toString();
+            //
+            conversation.setImageUrl(newUrl);
         }
         }
         redisService.set(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + conversation.getTaskId(), conversation,RedisService.key.MIDJOURNEY_CONVERSATION.getTimeout());
         redisService.set(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + conversation.getTaskId(), conversation,RedisService.key.MIDJOURNEY_CONVERSATION.getTimeout());
+
+        TASK_EXECUTOR.execute(()->{
+            //WebSocket 直接发送 任务状态
+            WssSession wssSession = (WssSession) redisService.hget(RedisService.key.MIDJOURNEY_WS_SESSION.getName(), conversation.getTaskId().toString());
+            if (wssSession != null) {
+                wssSession.sendMessage(WsMessageTypeEnum.SERVER_TASK_INFO, conversation);
+            }
+        });
+
         TASK_EXECUTOR.execute(() -> {
         TASK_EXECUTOR.execute(() -> {
             try {
             try {
                 sync(conversation);
                 sync(conversation);

+ 17 - 0
midjourney/src/main/java/com/yhlxj/web/config/WebSocketConfiguration.java

@@ -0,0 +1,17 @@
+package com.yhlxj.web.config;
+
+import org.springframework.context.annotation.Bean;
+import org.springframework.context.annotation.Configuration;
+import org.springframework.web.socket.server.standard.ServerEndpointExporter;
+
+/**
+ * @author chan
+ * @date 2024/6/28 17:33
+ */
+@Configuration
+public class WebSocketConfiguration {
+    @Bean
+    public ServerEndpointExporter serverEndpointExporter() {
+        return new ServerEndpointExporter();
+    }
+}

+ 23 - 0
midjourney/src/main/java/com/yhlxj/web/mirror/MidjourneyController.java

@@ -332,6 +332,28 @@ public class MidjourneyController {
         }
         }
         return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
         return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
     }
     }
+    /**
+     * 执行动作
+     */
+    @PostMapping("/submit/seed")
+    @NoSubmit
+    public Result<Object> seed(@RequestBody SubmitSeedDTO seedDTO) {
+        log.info("任务id:{},执行动作:{}",seedDTO.getTaskId(), "seed");
+        MidjourneyUser user = getUser();
+
+        Long num = checkUserLimit(user);
+        if (num != null && num < 0) {
+            throw BusinessRuntimeException.getInstance("次数已用完");
+        }
+        Object conversation;
+        try {
+            conversation = midjourneyService.submitSeed(user, seedDTO.getTaskId(),num, seedDTO.getBotType());
+        } catch (Exception e) {
+            recoverUserLimit(user.getId(), user.getMode(),num);
+            throw BusinessRuntimeException.getInstance(e.getMessage());
+        }
+        return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
+    }
 
 
 
 
     /**
     /**
@@ -344,6 +366,7 @@ public class MidjourneyController {
         MapBuilder builder = MapUtils.flatBuilder(request.getParameterMap());
         MapBuilder builder = MapUtils.flatBuilder(request.getParameterMap());
         return GatewayResponse.SUCCESS.newBuilder().toResult(beanSearcher.search(MidjourneyUserConversation.class,
         return GatewayResponse.SUCCESS.newBuilder().toResult(beanSearcher.search(MidjourneyUserConversation.class,
                 builder.field(MidjourneyUserConversation::getUserId, user.getId())
                 builder.field(MidjourneyUserConversation::getUserId, user.getId())
+                        .selectExclude(MidjourneyUserConversation::getChannelId,MidjourneyUserConversation::getInstanceId)
                         .orderBy(MidjourneyUserConversation::getId).desc().build()));
                         .orderBy(MidjourneyUserConversation::getId).desc().build()));
     }
     }
 
 

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

@@ -0,0 +1,184 @@
+package com.yhlxj.web.wss;
+
+import com.alibaba.fastjson.JSON;
+import com.alibaba.fastjson.JSONObject;
+import com.baomidou.mybatisplus.core.toolkit.Wrappers;
+import com.cyksj.common.exception.BusinessRuntimeException;
+import com.cyksj.common.task.GlobalThreadPoolTaskExecutor;
+import com.cyksj.common.util.QueryWrapperUtils;
+import com.cyksj.common.util.SpringCtxUtils;
+import com.yhlxj.dao.mapper.midjourney.MidjourneyUserConversationMapper;
+import com.yhlxj.dao.mapper.midjourney.MidjourneyUserMapper;
+import com.yhlxj.dao.model.entity.GroupsRelation;
+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 lombok.RequiredArgsConstructor;
+import org.apache.commons.lang3.StringUtils;
+import org.apache.tomcat.websocket.WsSession;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+import org.springframework.stereotype.Component;
+
+import javax.websocket.*;
+import javax.websocket.server.PathParam;
+import javax.websocket.server.ServerEndpoint;
+import java.io.IOException;
+
+/**
+ * @author chan
+ * @date 2024/6/28 13:59
+ */
+@Component
+@ServerEndpoint("/ws/midjourney/{userToken}/{taskId}")
+@RequiredArgsConstructor
+public class MidjourneyServerEndpoint {
+
+    private static Logger logger = LoggerFactory.getLogger(MidjourneyServerEndpoint.class);
+
+    private static final GlobalThreadPoolTaskExecutor TASK_EXECUTOR = GlobalThreadPoolTaskExecutor.getInstance();
+
+    private final MidjourneyService midjourneyService;
+
+    private final MidjourneyUserConversationMapper midjourneyUserConversationMapper;
+
+    private final MidjourneyUserMapper midjourneyUserMapper;
+
+    private final RedisService redisService;
+
+
+    /**
+     * 连接建立回调
+     *
+     * @param session
+     */
+    @OnOpen
+    public void onOpen(Session session, @PathParam("userToken") String userToken, @PathParam("taskId") String taskId) throws IOException {
+        logger.info("[WS建立连接],token:{}, taskId:{}", userToken, taskId);
+
+        //检查uniodId是否合法
+        if (StringUtils.isBlank(userToken) || StringUtils.isBlank(taskId)) {
+
+            logger.error("[WS建立连接失败],[参数非法],{},{}", userToken, taskId);
+
+            session.getBasicRemote().sendText("入参错误");
+            session.close();
+        }
+        MidjourneyUser user = getUser(userToken);
+        if (user == null) {
+            logger.error("[WS建立连接失败],[用户过期],{},{}", userToken, taskId);
+            session.getBasicRemote().sendText("用户过期");
+            session.close();
+        } else {
+            MidjourneyUserConversation conversation = midjourneyUserConversationMapper.selectOne(QueryWrapperUtils.buildWrapper((wrapper) -> {
+                wrapper.eq(MidjourneyUserConversation::getUserId, user.getId()).eq(MidjourneyUserConversation::getTaskId, taskId);
+            }));
+            WssSession wsSession = new WssSession(session);
+            if (conversation == null) {
+                wsSession.sendMessage(WsMessageTypeEnum.SERVER_TASK_NOT_FIND, "任务不存在.");
+                session.close();
+            } else {
+                if (!redisService.hHasKey(RedisService.key.MIDJOURNEY_WS_SESSION.getName(), taskId)) {
+                    redisService.hset(RedisService.key.MIDJOURNEY_WS_SESSION.getName(), taskId, wsSession, RedisService.key.MIDJOURNEY_WS_SESSION.getTimeout());
+                }
+                TASK_EXECUTOR.execute(() -> {
+                    redisService.set(RedisService.key.MIDJOURNEY_CONVERSATION.getName() + taskId, conversation, RedisService.key.MIDJOURNEY_CONVERSATION.getTimeout());
+                });
+            }
+
+        }
+
+    }
+
+    /**
+     * 收到客户端消息回调
+     *
+     * @param message
+     * @param session
+     */
+    @OnMessage
+    public void OnMessage(String message, Session session, @PathParam("userToken") String userToken, @PathParam("taskId") String taskId) {
+        try {
+            logger.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 = JSON.parseObject(message.substring(8));
+                    break;
+                default:
+                    logger.error("[WS消息类型异常],{},{},{}", userToken, taskId, message);
+                    return;
+            }
+
+            //根据子类型进行处理
+            switch (WsMessageTypeEnum.getMessageType(body.getString("type"))) {
+
+                case CLIENT_TASK_ID: //首次查看任务状态
+                    midjourneyService.getConversationById(body.getLong("taskId"));
+                    break;
+
+                case CLIENT_CLOSE:
+                    logger.info("[WS连接关闭],{},{}", userToken, taskId);
+                    //dailyQuizService.close(uniodId);
+                    break;
+
+                default:
+                    logger.error("[WS消息子类型异常],{},{},{}", userToken, taskId, message);
+
+            }
+
+        } catch (Exception e) {
+            logger.error("[WS接收消息失败],{},{},{},{}", userToken, taskId, message, e);
+        }
+    }
+
+    /**
+     * 发生错误回调
+     *
+     * @param session
+     * @param error
+     */
+    @OnError
+    public void OnError(Session session, Throwable error) throws IOException {
+        logger.error("[WS ONERROR],{}", error.getMessage());
+        if (session.isOpen()) {
+            session.close();
+        }
+    }
+
+    /**
+     * 关闭连接回调
+     *
+     * @param session
+     */
+    @OnClose
+    public void OnClose(Session session, @PathParam("userToken") String userToken, @PathParam("taskId") String taskId) throws IOException {
+
+        logger.info("[WS连接关闭],{},{}", userToken, taskId);
+        if (session.isOpen()) {
+            session.close();
+        }
+    }
+
+    public MidjourneyUser getUser(String userToken) {
+
+        MidjourneyUser midjourneyUser = midjourneyUserMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUser.class)
+                .eq(MidjourneyUser::getUserToken, userToken).last("limit 1"));
+        if (midjourneyUser == null) {
+            throw new BusinessRuntimeException("账号不存在,请重新登录");
+        }
+        return midjourneyUser;
+
+    }
+
+}

+ 23 - 0
netflix-common/src/main/java/com/cyksj/common/util/QueryWrapperUtils.java

@@ -0,0 +1,23 @@
+package com.cyksj.common.util;
+
+import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper;
+import com.baomidou.mybatisplus.core.conditions.update.LambdaUpdateWrapper;
+
+import java.util.function.Consumer;
+
+/**
+ * @author chan
+ * @date 2023/7/10 17:37
+ */
+public class QueryWrapperUtils {
+    public static <T> LambdaQueryWrapper<T> buildWrapper(Consumer<LambdaQueryWrapper<T>> consumer) {
+        LambdaQueryWrapper<T> queryWrapper = new LambdaQueryWrapper<>();
+        consumer.accept(queryWrapper);
+        return queryWrapper;
+    }
+    public static <T> LambdaUpdateWrapper<T> buildUpdateWrapper(Consumer<LambdaUpdateWrapper<T>> consumer) {
+        LambdaUpdateWrapper<T> queryWrapper = new LambdaUpdateWrapper<>();
+        consumer.accept(queryWrapper);
+        return queryWrapper;
+    }
+}