Ver Fonte

fix MODAL延迟提交

zwhui há 2 anos atrás
pai
commit
27fc26f015

+ 2 - 1
midjourney/src/main/java/com/yhlxj/dao/model/dto/SubmitBlendDTO.java

@@ -3,6 +3,7 @@ package com.yhlxj.dao.model.dto;
 import lombok.Data;
 
 import javax.validation.constraints.NotBlank;
+import javax.validation.constraints.NotEmpty;
 import java.util.List;
 
 /**
@@ -15,7 +16,7 @@ public class SubmitBlendDTO {
 	/**
 	 * base64
 	 */
-	@NotBlank(message = "请上传至少一张图片")
+	@NotEmpty(message = "请上传至少两张图片")
 	private List<String> base64Array;
 
 	/**

+ 2 - 1
midjourney/src/main/java/com/yhlxj/dao/model/dto/SubmitUploadDTO.java

@@ -3,6 +3,7 @@ package com.yhlxj.dao.model.dto;
 import lombok.Data;
 
 import javax.validation.constraints.NotBlank;
+import javax.validation.constraints.NotEmpty;
 import java.util.List;
 
 /**
@@ -14,7 +15,7 @@ public class SubmitUploadDTO {
     /**
      * 垫图base64数组
      */
-    @NotBlank(message = "请上传至少一张图片")
+    @NotEmpty(message = "请上传至少一张图片")
     private List<String> base64Array;
 
 }

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

@@ -288,6 +288,18 @@ public class MidjourneyServiceImpl implements MidjourneyService {
             param.put("maskBase64", maskBase64);
         }
         Long apiChannelId = midjourneyUserConversation == null ? null : midjourneyUserConversation.getApiChannelId();
+        JSONArray jsonArray = listByIds(mode, List.of(taskId), apiChannelId);
+        if (!jsonArray.isEmpty()){
+            JSONObject jsonObject = JSONUtil.parseObj(Jsons.toJson(jsonArray.get(0)));
+            if (!"MODAL".equals(jsonObject.getStr("status"))){
+                log.error("task {} status is not modal", taskId);
+                if (System.currentTimeMillis() - jsonObject.getLong("start_time") > 60 * 1000L) {
+                    throw BusinessRuntimeException.getInstance("任务超时,请稍后再试");
+                }
+                Thread.sleep(3000);
+                return submitModal(user, mode, taskId, prompt, maskBase64);
+            }
+        }
         SubmitResult result = submit(user,mode,"modal", null, param, apiChannelId,false);
         saveConversation(user.getId(),  mode,result,"MODAL", StringUtils.EMPTY, StringUtils.EMPTY,apiChannelId);
         Map<String, Object> properties = result.getProperties();

+ 10 - 8
midjourney/src/main/java/com/yhlxj/web/mirror/MidjourneyController.java

@@ -31,11 +31,13 @@ import lombok.RequiredArgsConstructor;
 import lombok.extern.slf4j.Slf4j;
 import org.springframework.beans.factory.annotation.Value;
 import org.springframework.scheduling.annotation.Scheduled;
+import org.springframework.validation.annotation.Validated;
 import org.springframework.web.bind.annotation.*;
 
 import javax.servlet.http.Cookie;
 import javax.servlet.http.HttpServletRequest;
 import javax.servlet.http.HttpServletResponse;
+import javax.validation.Valid;
 import java.io.IOException;
 import java.io.InputStream;
 import java.nio.charset.StandardCharsets;
@@ -186,7 +188,7 @@ public class MidjourneyController {
      */
     @PostMapping("/submit/imagine")
     @NoSubmit
-    public Result<SubmitResult> submitImagine(@RequestBody SubmitImagineDTO submitImagineDTO) {
+    public Result<SubmitResult> submitImagine(@Validated @RequestBody SubmitImagineDTO submitImagineDTO) {
         log.info("提交Imagine任务,提示:{}",submitImagineDTO.getPrompt());
         MidjourneyUser user = getUser();
         Integer mode = submitImagineDTO.getMode();
@@ -215,7 +217,7 @@ public class MidjourneyController {
      */
     @PostMapping("/submit/describe")
     @NoSubmit
-    public Result<SubmitResult> submitDescribe(@RequestBody SubmitDescribeDTO submitDescribeDTO) {
+    public Result<SubmitResult> submitDescribe(@Validated @RequestBody SubmitDescribeDTO submitDescribeDTO) {
         log.info("提交Describe任务");
         MidjourneyUser user = getUser();
         Integer mode = submitDescribeDTO.getMode();
@@ -239,7 +241,7 @@ public class MidjourneyController {
      */
     @PostMapping("/submit/blend")
     @NoSubmit
-    public Result<SubmitResult> submitBlend(@RequestBody SubmitBlendDTO submitBlendDTO) {
+    public Result<SubmitResult> submitBlend(@Validated @RequestBody SubmitBlendDTO submitBlendDTO) {
         log.info("提交Blend任务 dimensions:{},base64数组长度:{}", submitBlendDTO.getDimensions(), submitBlendDTO.getBase64Array().size());
         Integer mode = submitBlendDTO.getMode();
         MidjourneyUser user = getUser();
@@ -263,7 +265,7 @@ public class MidjourneyController {
      */
     @PostMapping("/submit/modal")
     @NoSubmit
-    public Result<SubmitResult> submitModal(@RequestBody SubmitModalDTO submitModalDTO) {
+    public Result<SubmitResult> submitModal(@Validated @RequestBody SubmitModalDTO submitModalDTO) {
         log.info("提交Modal任务,taskId:{},提示:{}",submitModalDTO.getTaskId(),submitModalDTO.getPrompt());
         Integer mode = submitModalDTO.getMode();
         MidjourneyUser user = getUser();
@@ -291,7 +293,7 @@ public class MidjourneyController {
      */
     @PostMapping("/submit/shorten")
     @NoSubmit
-    public Result<SubmitResult> submitShorten(@RequestBody SubmitShortenDTO submitShortenDTO) {
+    public Result<SubmitResult> submitShorten(@Validated @RequestBody SubmitShortenDTO submitShortenDTO) {
         log.info("提交Shorten任务 提示词:{}",submitShortenDTO.getPrompt());
         Integer mode = submitShortenDTO.getMode();
         MidjourneyUser user = getUser();
@@ -315,7 +317,7 @@ public class MidjourneyController {
      */
     @PostMapping("/submit/action")
     @NoSubmit
-    public Result<SubmitResult> action(@RequestBody SubmitActionDTO actionDTO) {
+    public Result<SubmitResult> action(@Validated @RequestBody SubmitActionDTO actionDTO) {
         log.info("任务id:{},执行动作:{}",actionDTO.getTaskId(),actionDTO.getCustomId());
         MidjourneyUser user = getUser();
         MidjourneyUserConversation midjourneyUserConversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, actionDTO.getTaskId()).last("limit 1"));
@@ -345,7 +347,7 @@ public class MidjourneyController {
      */
     @PostMapping("/submit/seed")
     @NoSubmit
-    public Result<Object> seed(@RequestBody SubmitSeedDTO seedDTO) {
+    public Result<Object> seed(@Validated @RequestBody SubmitSeedDTO seedDTO) {
         log.info("任务id:{},执行动作:{}",seedDTO.getTaskId(), "seed");
         MidjourneyUser user = getUser();
         Integer mode = seedDTO.getMode();
@@ -583,7 +585,7 @@ public class MidjourneyController {
      * 上传文件到discord
      */
     @PostMapping("/upload")
-    public Result<MidjourneyUserConversation> upload(@RequestBody SubmitUploadDTO uploadDTO) throws Exception {
+    public Result<MidjourneyUserConversation> upload(@Validated @RequestBody SubmitUploadDTO uploadDTO) throws Exception {
         log.info("上传文件到discord");
         return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.uploadFile(uploadDTO,getUser()));
     }