zwhui преди 1 година
родител
ревизия
aa2a662f92

+ 5 - 0
midjourney/src/main/java/com/yhlxj/dao/mapper/midjourney/MidjourneyUserMapper.java

@@ -18,6 +18,11 @@ public interface MidjourneyUserMapper extends BaseMapper<MidjourneyUser> {
     @Update("update midjourney_user set mj_fast_num = mj_fast_num - #{num} where id = #{id} and mj_fast_num = #{mjFastNum} and mj_fast_num > 0")
     int decFastNum(@Param("id") Long id, @Param("num") Integer num, @Param("mjFastNum") Integer mjFastNum);
 
+    @Update("update midjourney_user set mj_upgrade_num = mj_upgrade_num + #{num} where id = #{id}")
+    int incrUpgradeNum(@Param("id") Long id, @Param("num") Integer num);
+    @Update("update midjourney_user set mj_upgrade_num = mj_upgrade_num - #{num} where id = #{id} and mj_upgrade_num = #{mjUpgradeNum} and mj_upgrade_num > 0")
+    int decUpgradeNum(@Param("id") Long id, @Param("num") Integer num, @Param("mjUpgradeNum") Integer mjUpgradeNum);
+
     @Update("update midjourney_user set mj_relax_num = mj_relax_num + #{num} where id = #{id}")
     int incrRelaxNum(@Param("id") Long id, @Param("num") Integer num);
 

+ 11 - 0
midjourney/src/main/java/com/yhlxj/dao/mapper/midjourney/MidjourneyUserQueueMapper.java

@@ -0,0 +1,11 @@
+package com.yhlxj.dao.mapper.midjourney;
+
+import com.baomidou.mybatisplus.core.mapper.BaseMapper;
+import com.yhlxj.dao.model.entity.MidjourneyUserQueue;
+
+/**
+ * @author zwhui
+ * @date 2025/6/6 13:47
+ */
+public interface MidjourneyUserQueueMapper extends BaseMapper<MidjourneyUserQueue> {
+}

+ 5 - 0
midjourney/src/main/java/com/yhlxj/dao/model/entity/MidjourneyUser.java

@@ -47,6 +47,11 @@ public class MidjourneyUser extends BaseEntity {
      */
     private Integer mjFastNum;
 
+    /**
+     * mj 加量包次数
+     */
+    private Integer mjUpgradeNum;
+
     /**
      * mj relax次数
      */

+ 23 - 0
midjourney/src/main/java/com/yhlxj/dao/model/entity/MidjourneyUserQueue.java

@@ -0,0 +1,23 @@
+package com.yhlxj.dao.model.entity;
+
+import com.yhlxj.dao.model.BaseEntity;
+import lombok.Builder;
+import lombok.Data;
+
+/**
+ * @author zwhui
+ * @date 2025/6/6 13:42
+ */
+@Data
+@Builder
+public class MidjourneyUserQueue extends BaseEntity {
+    private Long userId;
+
+    private Long queueId;
+
+    private Long position;
+
+    private String prompt;
+
+    private String base64Array;
+}

+ 14 - 0
midjourney/src/main/java/com/yhlxj/dao/model/views/MidjourneyUserQueueView.java

@@ -0,0 +1,14 @@
+package com.yhlxj.dao.model.views;
+
+import lombok.Data;
+
+/**
+ * @author zwhui
+ * @date 2025/6/6 13:42
+ */
+@Data
+public class MidjourneyUserQueueView {
+    private Long queueId;
+
+    private Long position;
+}

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

@@ -701,6 +701,11 @@ public class RedisService {
         return redisTemplate.opsForZSet().rangeByScore(key, min, max, offset, count);
     }
 
+
+    public Set<Object> zRangeByScore(String key, double min, double max) {
+        return redisTemplate.opsForZSet().rangeByScore(key, min, max);
+    }
+
     /**
      * 从ZSet获取所有成员及其分数
      * @param key ZSet的键
@@ -737,6 +742,10 @@ public class RedisService {
         }
     }
 
+    public Long zRank(String key, Object value) {
+        return redisTemplate.opsForZSet().rank(key, value);
+    }
+
 
 
 
@@ -776,6 +785,7 @@ public class RedisService {
     @Getter
     public enum key {
         MIDJOURNEY_FAST_LIMIT("midjourney:fast:limit:", "midjourney fast次数", 60 * 60 * 48L),
+        MIDJOURNEY_UPGRADE_LIMIT("midjourney:upgrade:limit:", "midjourney upgrade次数", 60 * 60 * 48L),
         MIDJOURNEY_RELAX_LIMIT("midjourney:relax:limit:", "midjourney relax次数", 60 * 60 * 48L),
         MIDJOURNEY_EXPIRE_TIME("midjourney:expire:time:", "midjourney expire time", 60 * 60 * 48L),
         MIDJOURNEY_USER("midjourney:user:", "midjourney user", 60 * 60 * 2L),
@@ -789,6 +799,8 @@ public class RedisService {
         MIDJOURNEY_CAR_TASK("midjourney:car:task","midjourney car 任务数", 48 * 60 * 60L),
         MIDJOURNEY_USER_AGAIN("midjourney:user:again","midjourney用户重试次数", 30L),
         MIDJOURNEY_RELEASE_LIMIT("midjourney:release:limit:","midjourney解除限制次数", 30L),
+        MIDJOURNEY_QUEUE_LIMIT("midjourney:queue:limit:","midjourney慢速队列", 60 * 60 * 48L),
+        MIDJOURNEY_QUEUE_DETAIL_LIMIT("midjourney:queue:detail:limit:","midjourney慢速队列详情", 60 * 60 * 48L),
         ;
 
         private String name;

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

@@ -4,12 +4,16 @@ package com.yhlxj.service.midjourney;
 import com.yhlxj.dao.mapper.midjourney.MidjourneyUserConversationMapper;
 import com.yhlxj.dao.model.dto.AlphaSubmitRequestDTO;
 import com.yhlxj.dao.model.dto.BlendDimensions;
+import com.yhlxj.dao.model.dto.SubmitImagineDTO;
 import com.yhlxj.dao.model.dto.SubmitUploadDTO;
 import com.yhlxj.dao.model.entity.MidjourneyUser;
 import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
+import com.yhlxj.dao.model.entity.MidjourneyUserQueue;
 import com.yhlxj.dao.model.response.SubmitResult;
+import com.yhlxj.dao.model.views.MidjourneyUserQueueView;
 
 import java.util.List;
+import java.util.Map;
 import java.util.concurrent.ExecutionException;
 
 /**
@@ -56,4 +60,8 @@ public interface MidjourneyService {
     void userSubmitLimit(MidjourneyUser user,String action,Object customId,Boolean again);
 
     void releaseLimit(String userToken);
+
+    MidjourneyUserQueueView queueLimit(Long userId, SubmitImagineDTO submitImagineDTO);
+
+    List<MidjourneyUserQueue> getTaskPosition(Long userId);
 }

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

@@ -15,6 +15,7 @@ import com.baomidou.mybatisplus.core.conditions.update.LambdaUpdateWrapper;
 import com.baomidou.mybatisplus.core.toolkit.Wrappers;
 import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
 import com.cyksj.common.exception.BusinessRuntimeException;
+import com.cyksj.common.snowflake.Sequence;
 import com.cyksj.common.task.GlobalThreadPoolTaskExecutor;
 import com.cyksj.common.util.*;
 import com.yhlxj.dao.mapper.GroupsRelationMapper;
@@ -22,6 +23,7 @@ import com.yhlxj.dao.mapper.midjourney.*;
 import com.yhlxj.dao.model.dto.*;
 import com.yhlxj.dao.model.entity.*;
 import com.yhlxj.dao.model.response.SubmitResult;
+import com.yhlxj.dao.model.views.MidjourneyUserQueueView;
 import com.yhlxj.redis.RedisService;
 import com.yhlxj.service.common.EnvCommonService;
 import com.yhlxj.service.midjourney.MidjourneyAccountService;
@@ -34,6 +36,7 @@ import lombok.extern.slf4j.Slf4j;
 import net.coobird.thumbnailator.Thumbnails;
 import org.apache.commons.lang3.StringUtils;
 import org.springframework.beans.factory.annotation.Value;
+import org.springframework.data.redis.core.ZSetOperations;
 import org.springframework.stereotype.Service;
 
 import javax.imageio.ImageIO;
@@ -48,6 +51,7 @@ import java.io.InputStream;
 import java.net.URL;
 import java.nio.file.Files;
 import java.nio.file.Path;
+import java.time.LocalTime;
 import java.util.*;
 import java.util.concurrent.*;
 import java.util.regex.Matcher;
@@ -66,6 +70,8 @@ public class MidjourneyServiceImpl implements MidjourneyService {
 
     private final MidjourneyUserConversationMapper conversationMapper;
 
+    Sequence sequence = new Sequence(0);
+
     private static final List<String> progress = List.of("35%","65%","100%");
 
     private static final List<String> status = List.of("MODAL","CANCEL","FAILURE");
@@ -109,6 +115,12 @@ public class MidjourneyServiceImpl implements MidjourneyService {
     private final MidjourneySessionMapper midjourneySessionMapper;
 
 
+    private static final int MIN_PROCESSING_TIME = 10;
+    private static final int MAX_PROCESSING_TIME = 30;
+    private static final int AVG_PROCESSING_TIME = 20;
+
+    private final MidjourneyUserQueueMapper midjourneyUserQueueMapper;
+
 
     @Override
     public MidjourneyUser getUser(String userToken) throws ExecutionException, InterruptedException {
@@ -118,9 +130,11 @@ public class MidjourneyServiceImpl implements MidjourneyService {
         }
 
         if (!redisService.hasKey(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + midjourneyUser.getId())
+                || !redisService.hasKey(RedisService.key.MIDJOURNEY_UPGRADE_LIMIT.getName() + midjourneyUser.getId())
                 || !redisService.hasKey(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + midjourneyUser.getId())) {
                 redisService.multiSet(MapUtil.builder((new HashMap<String,Object>()))
                     .put(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + midjourneyUser.getId(), midjourneyUser.getMjFastNum())
+                    .put(RedisService.key.MIDJOURNEY_UPGRADE_LIMIT.getName() + midjourneyUser.getId(), midjourneyUser.getMjUpgradeNum())
                     .put(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + midjourneyUser.getId(), midjourneyUser.getMjRelaxNum())
                     .build());
         }
@@ -1358,4 +1372,135 @@ public class MidjourneyServiceImpl implements MidjourneyService {
             throw new BusinessRuntimeException("10003","账号异常,请联系客服");
         }
     }
+
+
+    @Override
+    public MidjourneyUserQueueView queueLimit(Long userId, SubmitImagineDTO submitImagineDTO) {
+        Long queueId = sequence.nextId();
+
+        RedisService.key key = RedisService.key.MIDJOURNEY_QUEUE_DETAIL_LIMIT;
+        String queueKey = key.getName() + userId;
+        MidjourneyUserQueue queue = MidjourneyUserQueue.builder()
+                .queueId(queueId).userId(userId).prompt(submitImagineDTO.getPrompt())
+                .base64Array(JSONUtil.toJsonStr(submitImagineDTO.getBase64Array())).build();
+
+        redisService.lSet(queueKey, JSONUtil.toJsonStr(queue),key.getTimeout());
+
+        // 计算当前时段需要的额外等待时间(分钟)
+        int additionalWaitMinutes = calculateAdditionalWaitTime();
+        // 清理过期的假任务,保持队列干净
+        cleanupExpiredTasks(userId);
+        // 获取队列中最后一个任务的过期时间
+        long taskExpireTime;
+        String TASK_QUEUE = RedisService.key.MIDJOURNEY_QUEUE_LIMIT.getName() + userId;
+        Set<ZSetOperations.TypedTuple<Object>> lastTask = redisService.zRangeWithScores(TASK_QUEUE, -1, -1);
+        if (lastTask != null && !lastTask.isEmpty()) {
+            // 如果队列不为空,使用最后一个任务的时间作为起点
+            taskExpireTime = lastTask.iterator().next().getScore().longValue();
+        } else {
+            // 队列为空,使用当前时间
+            taskExpireTime = System.currentTimeMillis() / 1000;
+        }
+        // 计算队列中应该存在的假任务总数
+        int requiredDummyTaskCount = (additionalWaitMinutes * 60) / AVG_PROCESSING_TIME;
+
+        // 获取当前队列中的假任务数量
+        long existingDummyTaskCount = countExistingDummyTasks(userId);
+
+        // 计算需要添加的假任务数量
+        int dummyTasksToAdd = requiredDummyTaskCount - (int)existingDummyTaskCount;
+
+        if (dummyTasksToAdd > 0) {
+            // 添加所需数量的假任务
+            for (int i = 0; i < dummyTasksToAdd; i++) {
+                String dummyTaskId = "dummy-" + UUID.randomUUID().toString();
+                // 均匀分布在当前时间到目标等待时间之间
+                int processingTime = MIN_PROCESSING_TIME +
+                        (int)(Math.random() * (MAX_PROCESSING_TIME - MIN_PROCESSING_TIME));
+
+                taskExpireTime += processingTime;
+                System.out.println("currentTime"+taskExpireTime);
+                redisService.zAdd(TASK_QUEUE, taskExpireTime, dummyTaskId);
+            }
+        }
+
+        // 真实任务的时间戳设置为"当前时间+额外等待时间"
+        int processingTime = MIN_PROCESSING_TIME +
+                (int)(Math.random() * (MAX_PROCESSING_TIME - MIN_PROCESSING_TIME));
+        taskExpireTime += processingTime;
+        System.out.println("taskExpireTime"+taskExpireTime);
+        redisService.zAdd(TASK_QUEUE, taskExpireTime, queueId);
+
+        // 获取任务在队列中的位置
+        Long rank = redisService.zRank(TASK_QUEUE, queueId);
+        long position = (rank != null) ? rank + 1 : 1;
+
+        MidjourneyUserQueueView view = new MidjourneyUserQueueView();
+        view.setQueueId(queueId);
+        view.setPosition(position);
+        return view;
+    }
+
+
+
+    /**
+     * 查询任务以及在队列中的位置
+     */
+    @Override
+    public List<MidjourneyUserQueue> getTaskPosition(Long userId) {
+        cleanupExpiredTasks(userId);
+        List<MidjourneyUserQueue> list = redisService.lGet(RedisService.key.MIDJOURNEY_QUEUE_DETAIL_LIMIT.getName() + userId, 0L, -1L)
+                .stream().map(item -> JSONUtil.toBean((String) item, MidjourneyUserQueue.class)).collect(Collectors.toList());
+        list.forEach(item -> {
+            Long rank = redisService.zRank(RedisService.key.MIDJOURNEY_QUEUE_LIMIT.getName() + userId, item.getQueueId());
+            item.setPosition((rank != null) ? rank + 1 : 1);
+        });
+        return list;
+    }
+
+
+
+    // 计算当前队列中假任务的数量
+    private long countExistingDummyTasks(Long userId) {
+        Set<Object> dummyTasks = redisService.zRangeByScore(RedisService.key.MIDJOURNEY_QUEUE_LIMIT.getName() + userId, 0, Double.POSITIVE_INFINITY);
+        if (dummyTasks == null) return 0;
+
+        return dummyTasks.stream()
+                .filter(id -> id instanceof String && ((String)id).startsWith("dummy-"))
+                .count();
+    }
+
+    /**
+     * 根据当前时间计算额外等待时间(分钟)
+     */
+    private int calculateAdditionalWaitTime() {
+        LocalTime now = LocalTime.now();
+        int hour = now.getHour();
+
+        if (hour >= 0 && hour < 9) {
+            // 00:00-08:59 多等4分钟
+            return 4;
+        } else if (hour >= 9 && hour < 18) {
+            // 09:00-17:59 多等12分钟
+            return 12;
+        } else {
+            // 18:00-23:59 多等8分钟
+            return 8;
+        }
+    }
+
+
+    /**
+     * 清理过期任务
+     * 定期调用此方法清理队列中的过期任务
+     */
+    public void cleanupExpiredTasks(Long userId) {
+        // 获取当前时间戳(秒)
+        long currentTimeSeconds = System.currentTimeMillis() / 1000;
+
+        // 移除过期的任务(分数小于当前时间的所有元素)
+        redisService.zRemoveRangeByScore(RedisService.key.MIDJOURNEY_QUEUE_LIMIT.getName() + userId, 0, currentTimeSeconds);
+    }
+
+
 }

+ 83 - 17
midjourney/src/main/java/com/yhlxj/web/mirror/MidjourneyController.java

@@ -30,6 +30,7 @@ import com.yhlxj.dao.model.response.SubmitResult;
 import com.yhlxj.dao.model.views.MidjourneyPaintingUserView;
 import com.yhlxj.dao.model.views.MidjourneyUserPaintingDayRecordView;
 import com.yhlxj.dao.model.views.MidjourneyUserPaintingRecordView;
+import com.yhlxj.dao.model.views.MidjourneyUserQueueView;
 import com.yhlxj.redis.RedisService;
 import com.yhlxj.service.midjourney.MidjourneyService;
 import com.yhlxj.service.midjourney.MidjourneyUserSettingsService;
@@ -110,8 +111,9 @@ public class MidjourneyController {
     /**
      * 检查用户次数
      */
-    public Long checkUserLimit(MidjourneyUser user,Integer mode){
+    public Map<String,Object> checkUserLimit(MidjourneyUser user,Integer mode){
         Long num = 0L;
+        boolean upgrade = false;
         if (user.getExpireTime().before(new Date())){
             throw BusinessRuntimeException.getInstance("账号已过期");
         }
@@ -121,6 +123,18 @@ public class MidjourneyController {
 
             if (num < 0) {
                 redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
+                //快速没有次数,扣除加量包次数
+                num = redisService.decr(RedisService.key.MIDJOURNEY_UPGRADE_LIMIT.getName() + user.getId(), 1L);
+                if (num < 0) {
+                    redisService.incr(RedisService.key.MIDJOURNEY_UPGRADE_LIMIT.getName() + user.getId(), 1L);
+                }else {
+                    upgrade = true;
+                    int i = midjourneyUserMapper.decUpgradeNum(user.getId(), 1, user.getMjUpgradeNum());
+                    if (i == 0) {
+                        redisService.incr(RedisService.key.MIDJOURNEY_UPGRADE_LIMIT.getName() + user.getId(), 1L);
+                        throw BusinessRuntimeException.getInstance("网络异常,请稍后再试");
+                    }
+                }
             }else {
                 int i = midjourneyUserMapper.decFastNum(user.getId(), 1, user.getMjFastNum());
                 if (i == 0) {
@@ -147,16 +161,24 @@ public class MidjourneyController {
                 num = null;
             }
         }
-        return num;
+        Map<String, Object> result = new HashMap<>();
+        result.put("num", num);
+        result.put("upgrade", upgrade);
+        return result;
     }
 
     /**
      * 恢复次数
      */
-    public void recoverUserLimit(Long id,Integer mode,Long num){
+    public void recoverUserLimit(Long id,Integer mode,Long num,Boolean upgrade){
         if (mode == 1){
-            redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + id, 1L);
-            midjourneyUserMapper.incrFastNum(id, 1);
+            if (upgrade){
+                redisService.incr(RedisService.key.MIDJOURNEY_UPGRADE_LIMIT.getName() + id, 1L);
+                midjourneyUserMapper.incrUpgradeNum(id, 1);
+            }else {
+                redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + id, 1L);
+                midjourneyUserMapper.incrFastNum(id, 1);
+            }
         }
         if (mode == 2){
             if (num != null){
@@ -206,7 +228,9 @@ public class MidjourneyController {
         log.info("提交Imagine任务,提示:{}",submitImagineDTO.getPrompt());
         MidjourneyUser user = getUser();
         Integer mode = submitImagineDTO.getMode();
-        Long num = checkUserLimit(user, mode);
+        Map<String, Object> checkResult = checkUserLimit(user, mode);
+        Long num = (Long) checkResult.get("num");
+        Boolean upgrade = (Boolean) checkResult.get("upgrade");
         if (num != null && num < 0){
             throw BusinessRuntimeException.getInstance("次数已用完");
         }
@@ -219,7 +243,7 @@ public class MidjourneyController {
             }
             conversation = midjourneyService.submitImagine(user, mode, submitImagineDTO.getPrompt(), submitImagineDTO.getBotType(), submitImagineDTO.getBase64Array(), num);
         } catch (Exception e) {
-            recoverUserLimit(user.getId(), mode,num);
+            recoverUserLimit(user.getId(), mode,num,upgrade);
             log.error("提交Imagine任务异常,{}", StringUtil.getErrorText(e));
             throw BusinessRuntimeException.getInstance(e.getMessage());
         }
@@ -236,7 +260,9 @@ public class MidjourneyController {
         log.info("提交Describe任务");
         MidjourneyUser user = getUser();
         Integer mode = submitDescribeDTO.getMode();
-        Long num = checkUserLimit(user,mode);
+        Map<String, Object> checkResult = checkUserLimit(user, mode);
+        Long num = (Long) checkResult.get("num");
+        Boolean upgrade = (Boolean) checkResult.get("upgrade");
         if (num != null && num < 0){
             throw BusinessRuntimeException.getInstance("次数已用完");
         }
@@ -244,7 +270,7 @@ public class MidjourneyController {
         try {
             conversation = midjourneyService.submitDescribe(user, mode,submitDescribeDTO.getBotType(),submitDescribeDTO.getBase64(), num);
         } catch (Exception e) {
-            recoverUserLimit(user.getId(), mode,num);
+            recoverUserLimit(user.getId(), mode,num,upgrade);
             log.error("提交Describe任务异常,{}", StringUtil.getErrorText(e));
             throw BusinessRuntimeException.getInstance(e.getMessage());
         }
@@ -261,7 +287,9 @@ public class MidjourneyController {
         log.info("提交Blend任务 dimensions:{},base64数组长度:{}", submitBlendDTO.getDimensions(), submitBlendDTO.getBase64Array().size());
         Integer mode = submitBlendDTO.getMode();
         MidjourneyUser user = getUser();
-        Long num = checkUserLimit(user,mode);
+        Map<String, Object> checkResult = checkUserLimit(user, mode);
+        Long num = (Long) checkResult.get("num");
+        Boolean upgrade = (Boolean) checkResult.get("upgrade");
         if (num != null && num < 0){
             throw BusinessRuntimeException.getInstance("次数已用完");
         }
@@ -269,7 +297,7 @@ public class MidjourneyController {
         try {
             conversation = midjourneyService.submitBlend(user, mode,submitBlendDTO.getDimensions(), submitBlendDTO.getBotType(), submitBlendDTO.getBase64Array(), num);
         } catch (Exception e) {
-            recoverUserLimit(user.getId(), mode,num);
+            recoverUserLimit(user.getId(), mode,num,upgrade);
             log.error("提交Blend任务异常,{}", StringUtil.getErrorText(e));
             throw BusinessRuntimeException.getInstance(e.getMessage());
         }
@@ -290,7 +318,9 @@ public class MidjourneyController {
             MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, submitModalDTO.getBeforeTaskId()).last("limit 1"));
             mode = conversation.getMode();
         }
-        Long num = checkUserLimit(user,mode);
+        Map<String, Object> checkResult = checkUserLimit(user, mode);
+        Long num = (Long) checkResult.get("num");
+        Boolean upgrade = (Boolean) checkResult.get("upgrade");
         if (num != null && num < 0){
             throw BusinessRuntimeException.getInstance("次数已用完");
         }
@@ -298,7 +328,7 @@ public class MidjourneyController {
         try {
             conversation = midjourneyService.submitModal(user, mode, submitModalDTO.getTaskId(), submitModalDTO.getPrompt(), submitModalDTO.getMaskBase64(), num);
         } catch (Exception e) {
-            recoverUserLimit(user.getId(), mode,num);
+            recoverUserLimit(user.getId(), mode,num,upgrade);
             log.error("提交Modal任务异常,{}", StringUtil.getErrorText(e));
             throw BusinessRuntimeException.getInstance(e.getMessage());
         }
@@ -315,7 +345,9 @@ public class MidjourneyController {
         log.info("提交Shorten任务 提示词:{}",submitShortenDTO.getPrompt());
         Integer mode = submitShortenDTO.getMode();
         MidjourneyUser user = getUser();
-        Long num = checkUserLimit(user,mode);
+        Map<String, Object> checkResult = checkUserLimit(user, mode);
+        Long num = (Long) checkResult.get("num");
+        Boolean upgrade = (Boolean) checkResult.get("upgrade");
         if (num != null && num < 0){
             throw BusinessRuntimeException.getInstance("次数已用完");
         }
@@ -323,7 +355,7 @@ public class MidjourneyController {
         try {
             conversation = midjourneyService.submitShorten(user, mode, submitShortenDTO.getBotType(), submitShortenDTO.getPrompt(), num);
         } catch (Exception e) {
-            recoverUserLimit(user.getId(),mode,num);
+            recoverUserLimit(user.getId(),mode,num,upgrade);
             throw BusinessRuntimeException.getInstance(e.getMessage());
         }
         //syncUser(user.getId(), user.getMode(), num);
@@ -344,8 +376,11 @@ public class MidjourneyController {
         }
         Integer mode = midjourneyUserConversation.getMode();
         Long num = 0L;
+        Boolean upgrade = false;
         if (!actionDTO.getCustomId().contains("BOOKMARK")) {
-            num = checkUserLimit(user,mode);
+            Map<String, Object> checkResult = checkUserLimit(user, mode);
+            num = (Long) checkResult.get("num");
+            upgrade = (Boolean) checkResult.get("upgrade");
             if (num != null && num < 0) {
                 throw BusinessRuntimeException.getInstance("次数已用完");
             }
@@ -354,7 +389,7 @@ public class MidjourneyController {
         try {
             conversation = midjourneyService.submitAction(user,midjourneyUserConversation,mode, actionDTO.getTaskId(), actionDTO.getCustomId(),num, actionDTO.getBotType());
         } catch (Exception e) {
-            recoverUserLimit(user.getId(), mode,num);
+            recoverUserLimit(user.getId(), mode,num,upgrade);
             log.error("执行动作异常,{}", StringUtil.getErrorText(e));
             throw BusinessRuntimeException.getInstance(e.getMessage());
         }
@@ -631,4 +666,35 @@ public class MidjourneyController {
         midjourneyService.releaseLimit(userToken);
         return GatewayResponse.SUCCESS.newBuilder().toResult();
     }
+
+
+    /**
+     * 保存排队任务
+     */
+    @PostMapping("/queue")
+    public Result<MidjourneyUserQueueView> saveQueue(@Validated @RequestBody SubmitImagineDTO submitImagineDTO) throws Exception {
+        log.info("提交Imagine任务,提示:{}",submitImagineDTO.getPrompt());
+        MidjourneyUser user = getUser();
+        Integer mode = submitImagineDTO.getMode();
+        if (mode != 2){
+            throw BusinessRuntimeException.getInstance("网络异常!");
+        }
+        Long relax = (Long) redisService.get(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId());
+        if (relax != null && relax <= 0){
+            throw BusinessRuntimeException.getInstance("次数已用完");
+        }
+
+
+        return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.queueLimit(user.getId(),submitImagineDTO));
+    }
+
+
+    /**
+     * 获取任务以及位置
+     */
+    @GetMapping("/getTaskPosition")
+    public Result<List<MidjourneyUserQueue>> getTaskPosition() throws Exception{
+        MidjourneyUser user = getUser();
+        return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.getTaskPosition(user.getId()));
+    }
 }