Forráskód Böngészése

新增settings设置,自动修改常用命令的prompt ,以及同步mj-proxy的账号

chenbiao 2 éve
szülő
commit
b7cb805469

+ 13 - 0
midjourney/src/main/java/com/yhlxj/dao/mapper/AccountMapper.java

@@ -0,0 +1,13 @@
+package com.yhlxj.dao.mapper;
+
+
+import com.baomidou.mybatisplus.core.mapper.BaseMapper;
+import com.yhlxj.dao.model.entity.Account;
+
+/**
+ * @author chan
+ * @date 2022/9/27 6:22 PM
+ */
+public interface AccountMapper extends BaseMapper<Account> {
+
+}

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

@@ -0,0 +1,11 @@
+package com.yhlxj.dao.mapper.midjourney;
+
+import com.baomidou.mybatisplus.core.mapper.BaseMapper;
+import com.yhlxj.dao.model.entity.MidjourneyUserSettings;
+
+/**
+ * @author chan
+ * @date 2024/7/9 12:19
+ */
+public interface MidjourneyUserSettingsMapper extends BaseMapper<MidjourneyUserSettings> {
+}

+ 107 - 0
midjourney/src/main/java/com/yhlxj/dao/model/dto/MidjourneyProxyAccountInfo.java

@@ -0,0 +1,107 @@
+package com.yhlxj.dao.model.dto;
+
+import lombok.Data;
+
+import java.util.List;
+
+/**
+ * @author chan
+ * @date 2024/7/9 16:58
+ */
+
+@Data
+public class MidjourneyProxyAccountInfo {
+    private String nijiMode;
+    private String lifetimeUsage;
+    private boolean remixAutoSubmit;
+    private String remark;
+    private String mjBotChannelId;
+    private Displays displays;
+    private int timeoutMinutes;
+    private String variation;
+    private String mode;
+    private String billedWay;
+    private long dateCreated;
+
+    private boolean enable;
+
+    private List<String> nijiButtons;
+    private boolean nijiRemix;
+    private String relaxedUsage;
+    private String stylize;
+    private String id;
+    private String channelId;
+    private String email;
+    private String nijiBotChannelId;
+    private long renewDate;
+    private VersionSelector versionSelector;
+    private String subscribePlan;
+    private int queueSize;
+    private List<Button> buttons;
+    private int coreSize;
+    private boolean raw;
+    private String userAgent;
+    private String sessionId;
+    private String guildId;
+    private String userId;
+    private boolean publicMode;
+    private String version;
+    private String fastTimeRemaining;
+    private String userToken;
+    private String name;
+    private Properties properties;
+    private boolean remix;
+
+
+
+    @Data
+    public static class Displays {
+        private String mode;
+        private String nijiMode;
+        private String subscribePlan;
+        private String billedWay;
+        private String dateCreated;
+        private String stylize;
+        private String variation;
+        private String renewDate;
+
+    }
+
+    @Data
+    public static class VersionSelector {
+        private List<Option> options;
+        private String placeholder;
+        private int type;
+        private String customId;
+
+
+        @Data
+        public static class Option {
+            private String emoji;
+            private String description;
+            private String label;
+            private String value;
+            private boolean selected;
+
+        }
+    }
+
+    @Data
+    public static class Button {
+        private String emoji;
+        private int style;
+        private String label;
+        private int type;
+        private String customId;
+
+    }
+
+    @Data
+    public static class Properties {
+        private List<String> supportApplicationIds;
+        private int flags;
+        private int weight;
+        private String messageId;
+
+    }
+}

+ 106 - 0
midjourney/src/main/java/com/yhlxj/dao/model/entity/Account.java

@@ -0,0 +1,106 @@
+package com.yhlxj.dao.model.entity;
+
+import com.baomidou.mybatisplus.annotation.TableField;
+import com.yhlxj.dao.model.BaseEntity;
+import lombok.Getter;
+import lombok.Setter;
+
+import java.math.BigDecimal;
+import java.util.Date;
+
+/**
+ * @author chan
+ * @date 2022/9/27 6:17 PM
+ */
+@Getter
+@Setter
+public class Account extends BaseEntity {
+
+    private String title;
+
+    private String account;
+
+    private String password;
+
+    private Boolean isMonth;
+
+    private Long goodsId;
+
+    private String bankCard;
+
+    @TableField(exist = false)
+    private Long skuId;
+
+    private Date startTime;
+
+    private Date expiryTime;
+
+    private BigDecimal amount;
+
+    /**
+     * 国区id
+     */
+    private String cnAccount;
+
+    /**
+     * 美区id
+     */
+    private String usAccount;
+
+    /**
+     * 客服id
+     */
+    private Long customerServiceId;
+
+    /**
+     * chatGpt apiKey
+     */
+    private String apiKey;
+
+    /**
+     * chatGpt 邮箱密码
+     */
+    private String gptEmailPwd;
+
+    /**
+     * 辅助邮箱
+     */
+    private String email;
+
+    private String remark;
+
+    /**
+     * midJourney 登录安全码
+     */
+    private String verifyCodes;
+
+    /**
+     * 密保
+     */
+    private String secret;
+
+    /**
+     * gpt API 密钥
+     */
+    private String apiSecret;
+
+    /**
+     * 账号是否被弃用
+     */
+    private Boolean isAbandon;
+
+    /**
+     * gpt accesToken
+     */
+    private String gptRefreshToken;
+
+    /**
+     * 备用账号类型sku_ids
+     */
+    private String preSkuIds;
+
+    /**
+     * 类型 1通用 2备用
+     */
+    private Integer type;
+}

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

@@ -62,4 +62,8 @@ public class MidjourneyUser extends BaseEntity {
      * 拉黑
      */
     private Boolean isBlack;
+
+    @TableField(exist = false)
+    @DbIgnore
+    private MidjourneyUserSettings settings;
 }

+ 28 - 0
midjourney/src/main/java/com/yhlxj/dao/model/entity/MidjourneyUserSettings.java

@@ -0,0 +1,28 @@
+package com.yhlxj.dao.model.entity;
+
+import com.yhlxj.dao.model.BaseEntity;
+import lombok.Data;
+
+/**
+ * @author chan
+ * @date 2024/7/9 12:18
+ */
+@Data
+public class MidjourneyUserSettings extends BaseEntity {
+
+    /**
+     * 用户ID
+     */
+    private Long userId;
+
+    /**
+     * 设置
+     */
+    private String settings;
+
+    /**
+     * 提示词
+     */
+    private String prompt;
+
+}

+ 2 - 0
midjourney/src/main/java/com/yhlxj/service/midjourney/MidjourneyAccountService.java

@@ -21,4 +21,6 @@ public interface MidjourneyAccountService extends IService<MidjourneyAccount> {
 
     MidjourneyUser getMidjourneyUserToken(Long userId, Long relationId);
 
+    void syncAccountByMJPlus() throws Exception;
+
 }

+ 11 - 0
midjourney/src/main/java/com/yhlxj/service/midjourney/MidjourneyUserSettingsService.java

@@ -0,0 +1,11 @@
+package com.yhlxj.service.midjourney;
+
+import com.baomidou.mybatisplus.extension.service.IService;
+import com.yhlxj.dao.model.entity.MidjourneyUserSettings;
+
+/**
+ * @author chan
+ * @date 2024/7/9 12:22
+ */
+public interface MidjourneyUserSettingsService extends IService<MidjourneyUserSettings> {
+}

+ 53 - 4
midjourney/src/main/java/com/yhlxj/service/midjourney/impl/MidJourneyAccountServiceImpl.java

@@ -10,19 +10,19 @@ import com.baomidou.mybatisplus.core.toolkit.Wrappers;
 import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
 import com.cyksj.common.exception.BusinessRuntimeException;
 import com.cyksj.common.util.Jsons;
+import com.cyksj.common.util.QueryWrapperUtils;
 import com.cyksj.common.util.StringUtil;
-import com.yhlxj.dao.mapper.GoodsDonSkuMapper;
-import com.yhlxj.dao.mapper.GroupsMapper;
-import com.yhlxj.dao.mapper.GroupsRelationMapper;
-import com.yhlxj.dao.mapper.UserMapper;
+import com.yhlxj.dao.mapper.*;
 import com.yhlxj.dao.mapper.midjourney.MidjourneyAccountMapper;
 import com.yhlxj.dao.mapper.midjourney.MidjourneyUserMapper;
+import com.yhlxj.dao.model.dto.MidjourneyProxyAccountInfo;
 import com.yhlxj.dao.model.entity.*;
 import com.yhlxj.redis.RedisService;
 import com.yhlxj.service.midjourney.MidjourneyAccountService;
 import com.yhlxj.service.user.UserBindRelationService;
 import lombok.RequiredArgsConstructor;
 import lombok.extern.slf4j.Slf4j;
+import org.apache.commons.lang3.StringUtils;
 import org.springframework.stereotype.Service;
 import org.springframework.transaction.annotation.Transactional;
 
@@ -33,6 +33,7 @@ import java.nio.file.Path;
 import java.util.HashMap;
 import java.util.List;
 import java.util.Map;
+import java.util.stream.Collectors;
 
 /**
  * @author zwhui
@@ -57,6 +58,8 @@ public class MidJourneyAccountServiceImpl extends ServiceImpl<MidjourneyAccountM
 
     private final MidjourneyUserMapper midjourneyUserMapper;
 
+    private final AccountMapper accountMapper;
+
     private final RedisService redisService;
 
     @Override
@@ -285,4 +288,50 @@ public class MidJourneyAccountServiceImpl extends ServiceImpl<MidjourneyAccountM
         return result.getJSONObject("value").getJSONArray("saved").getJSONObject(0).getJSONObject("info").getStr("cdnUrl");
     }
 
+
+    @Transactional
+    @Override
+    public void syncAccountByMJPlus() throws Exception {
+        String body = HttpRequest.post(HOST + "/mj/account/query").body(Jsons.toJson(MapUtil.of("pageSize", 100))).execute().body();
+        log.info("getAllAccount body:{}", body);
+        JSONObject jsonObject = JSONUtil.parseObj(body);
+        List<MidjourneyProxyAccountInfo> content = JSONUtil.parseArray(jsonObject.getStr("content")).toList(MidjourneyProxyAccountInfo.class);
+        //如果 该账号不存在于 proxy 上,则在本地删除该账号
+
+        content.forEach((accountInfo)->{
+            if (StringUtils.isNotBlank(accountInfo.getRemark()) && StringUtils.isNumeric(accountInfo.getRemark())) {
+                MidjourneyAccount midjourneyAccount = getById(Long.parseLong(accountInfo.getRemark()));
+                if (midjourneyAccount == null) {
+                    midjourneyAccount = new MidjourneyAccount();
+                    midjourneyAccount.setId(Long.parseLong(accountInfo.getRemark()));
+                    midjourneyAccount.setGuildId(Long.parseLong(accountInfo.getGuildId()));
+                    midjourneyAccount.setChannelId(Long.parseLong(accountInfo.getChannelId()));
+                    midjourneyAccount.setUserToken(accountInfo.getUserToken());
+                }
+                midjourneyAccount.setInstanceId(Long.parseLong(accountInfo.getId()));
+                if (midjourneyAccount.getStatus() != accountInfo.isEnable()) {
+                    midjourneyAccount.setStatus(accountInfo.isEnable());
+                    if(!accountInfo.isEnable()){
+                        redisService.hdel(RedisService.key.MIDJOURNEY_ACCOUNT.getName(),midjourneyAccount.getInstanceId());
+                    }else {
+                        redisService.hincr(RedisService.key.MIDJOURNEY_ACCOUNT.getName(), midjourneyAccount.getInstanceId().toString(), 1.0);
+                    }
+
+                }
+
+                saveOrUpdate(midjourneyAccount);
+            }
+        });
+
+        //移除本地不存在的账号
+        List<String> ids = content.stream().map(MidjourneyProxyAccountInfo::getId).collect(Collectors.toList());
+        List<MidjourneyAccount> list = list();
+        list.forEach((account)->{
+            if (!ids.contains(account.getInstanceId().toString())) {
+                removeAccount(account.getId());
+                redisService.hdel(RedisService.key.MIDJOURNEY_ACCOUNT.getName(),account.getInstanceId());
+            }
+        });
+
+    }
 }

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

@@ -22,16 +22,14 @@ import com.yhlxj.dao.mapper.midjourney.MidjourneyUserMapper;
 import com.yhlxj.dao.model.dto.BlendDimensions;
 import com.yhlxj.dao.model.dto.MessageButton;
 import com.yhlxj.dao.model.dto.SubmitUploadDTO;
-import com.yhlxj.dao.model.entity.GroupsRelation;
-import com.yhlxj.dao.model.entity.MidjourneyAccount;
-import com.yhlxj.dao.model.entity.MidjourneyUser;
-import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
+import com.yhlxj.dao.model.entity.*;
 import com.yhlxj.dao.model.enums.WsMessageTypeEnum;
 import com.yhlxj.dao.model.response.SubmitResult;
 import com.yhlxj.redis.RedisService;
 import com.yhlxj.service.common.EnvCommonService;
 import com.yhlxj.service.midjourney.MidjourneyAccountService;
 import com.yhlxj.service.midjourney.MidjourneyService;
+import com.yhlxj.service.midjourney.MidjourneyUserSettingsService;
 import com.yhlxj.web.wss.WssSession;
 import lombok.Data;
 import lombok.RequiredArgsConstructor;
@@ -81,6 +79,9 @@ public class MidjourneyServiceImpl implements MidjourneyService {
 
     private static final List<String> status = List.of("MODAL","CANCEL","FAILURE");
 
+    private static final String PROMPT = "--v 6 ";
+    private static final String SETTINGS = "{\"version\":\"--v 6\",\"remix\":false,\"raw\":false,\"variation\":\"HIGH\",\"mode\":\"relax\",\"stylize\":\"MED\"}";
+
     private final RedisService redisService;
 
     private final MidjourneyUserMapper midjourneyUserMapper;
@@ -94,6 +95,8 @@ public class MidjourneyServiceImpl implements MidjourneyService {
 
     private final MidjourneyAccountService midjourneyAccountService;
 
+    private final MidjourneyUserSettingsService midjourneyUserSettingsService;
+
     private final GroupsRelationMapper groupsRelationMapper;
 
     @Override
@@ -113,6 +116,17 @@ public class MidjourneyServiceImpl implements MidjourneyService {
         Long relationId = midjourneyUser.getRelationId();
         GroupsRelation relation = groupsRelationMapper.selectById(relationId);
         midjourneyUser.setAqType(relation.getAqType());
+        MidjourneyUserSettings midjourneyUserSettings = midjourneyUserSettingsService.getOne(QueryWrapperUtils.buildWrapper((wrapper) -> {
+            wrapper.eq(MidjourneyUserSettings::getUserId, midjourneyUser.getId());
+        }));
+        if (midjourneyUserSettings == null) {
+            midjourneyUserSettings = new MidjourneyUserSettings();
+            midjourneyUserSettings.setSettings(SETTINGS);
+            midjourneyUserSettings.setPrompt(PROMPT);
+            midjourneyUserSettings.setUserId(midjourneyUser.getId());
+            midjourneyUserSettingsService.save(midjourneyUserSettings);
+        }
+        midjourneyUser.setSettings(midjourneyUserSettings);
         return midjourneyUser;
     }
 

+ 19 - 0
midjourney/src/main/java/com/yhlxj/service/midjourney/impl/MidjourneyUserSettingsServiceImpl.java

@@ -0,0 +1,19 @@
+package com.yhlxj.service.midjourney.impl;
+
+import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
+import com.yhlxj.dao.mapper.midjourney.MidjourneyUserSettingsMapper;
+import com.yhlxj.dao.model.entity.MidjourneyUserSettings;
+import com.yhlxj.service.midjourney.MidjourneyUserSettingsService;
+import lombok.RequiredArgsConstructor;
+import lombok.extern.slf4j.Slf4j;
+import org.springframework.stereotype.Service;
+
+/**
+ * @author chan
+ * @date 2024/7/9 12:23
+ */
+@Service
+@RequiredArgsConstructor
+@Slf4j
+public class MidjourneyUserSettingsServiceImpl extends ServiceImpl<MidjourneyUserSettingsMapper, MidjourneyUserSettings> implements MidjourneyUserSettingsService {
+}

+ 31 - 0
midjourney/src/main/java/com/yhlxj/service/task/Scheduler.java

@@ -0,0 +1,31 @@
+package com.yhlxj.service.task;//package com.cyksj.task;
+
+import com.yhlxj.service.midjourney.MidjourneyAccountService;
+import com.yhlxj.service.midjourney.MidjourneyService;
+import lombok.RequiredArgsConstructor;
+import lombok.extern.slf4j.Slf4j;
+import org.springframework.scheduling.annotation.Scheduled;
+import org.springframework.stereotype.Component;
+
+@Component
+@Slf4j
+@RequiredArgsConstructor
+public class Scheduler {
+
+    private final MidjourneyService midjourneyService;
+
+    private final MidjourneyAccountService midjourneyAccountService;
+
+    @Scheduled(cron = "0 0/2 * * * ?")
+    public void syscConversation() {
+        midjourneyService.syscConversation();
+    }
+
+    @Scheduled(cron = "0 0/15 * * * ?")
+    public void syncAccount() throws Exception {
+        midjourneyAccountService.syncAccountByMJPlus();
+    }
+
+
+}
+

+ 96 - 7
midjourney/src/main/java/com/yhlxj/web/mirror/MidjourneyController.java

@@ -16,16 +16,14 @@ import com.yhlxj.dao.mapper.MidjourneyPaintingPlazaMapper;
 import com.yhlxj.dao.mapper.midjourney.MidjourneyUserConversationMapper;
 import com.yhlxj.dao.mapper.midjourney.MidjourneyUserMapper;
 import com.yhlxj.dao.model.dto.*;
-import com.yhlxj.dao.model.entity.GroupsRelation;
-import com.yhlxj.dao.model.entity.MidjourneyPaintingPlaza;
-import com.yhlxj.dao.model.entity.MidjourneyUser;
-import com.yhlxj.dao.model.entity.MidjourneyUserConversation;
+import com.yhlxj.dao.model.entity.*;
 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.redis.RedisService;
 import com.yhlxj.service.midjourney.MidjourneyService;
+import com.yhlxj.service.midjourney.MidjourneyUserSettingsService;
 import lombok.RequiredArgsConstructor;
 import lombok.extern.slf4j.Slf4j;
 import org.springframework.beans.factory.annotation.Value;
@@ -39,7 +37,9 @@ import java.io.IOException;
 import java.io.InputStream;
 import java.nio.charset.StandardCharsets;
 import java.util.Date;
+import java.util.HashMap;
 import java.util.List;
+import java.util.Map;
 
 /** server/镜像服务/mj绘画
  * @author zwhui
@@ -72,6 +72,8 @@ public class MidjourneyController {
 
     private final MidjourneyUserConversationMapper conversationMapper;
 
+    private final MidjourneyUserSettingsService midjourneyUserSettingsService;
+
     public MidjourneyUser getUser(){
         String userToken = request.getHeader("user-token");
        return midjourneyService.getUser(userToken);
@@ -191,6 +193,11 @@ public class MidjourneyController {
         }
         SubmitResult conversation;
         try {
+            MidjourneyUserSettings settings = user.getSettings();
+            if(settings != null){
+                String prompt = settings.getPrompt();
+                submitImagineDTO.setPrompt(appendDefaults(submitImagineDTO.getPrompt(), prompt));
+            }
             conversation = midjourneyService.submitImagine(user, mode, submitImagineDTO.getPrompt(), submitImagineDTO.getBotType(), submitImagineDTO.getBase64Array());
         } catch (Exception e) {
             recoverUserLimit(user.getId(), mode,num);
@@ -468,9 +475,91 @@ public class MidjourneyController {
         return GatewayResponse.SUCCESS.newBuilder().toResult(search);
     }
 
-    @Scheduled(cron = "0 0/2 * * * ?")
-    public void syscConversation() {
-        midjourneyService.syscConversation();
+    /**
+     * 修改settings
+     */
+    @GetMapping("/get/settings")
+    public Result<MidjourneyUserSettings> getSettings(){
+        MidjourneyUser user = getUser();
+
+        return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyUserSettingsService.getById(user.getSettings().getId()));
+    }
+
+    /**
+     * 修改settings
+     */
+    @PostMapping("/put/settings")
+    public Result<MidjourneyUserSettings> upSettings(@RequestBody MidjourneyUserSettings settings){
+        MidjourneyUser user = getUser();
+        if (!user.getId().equals(settings.getUserId())) {
+            throw BusinessRuntimeException.getInstance("参数错误");
+        }
+        settings.setId(user.getSettings().getId());
+        midjourneyUserSettingsService.updateById(settings);
+        return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyUserSettingsService.getById(settings.getId()));
+    }
+
+
+    public static String appendDefaults(String userPrompt, String defaultOptions) {
+        // 解析默认选项
+        Map<String, String> defaultOptionsMap = parseOptions(defaultOptions);
+
+        // 修正用户输入中可能缺少空格的选项
+        userPrompt = fixMissingSpaces(userPrompt);
+
+        // 检查并追加缺失的选项
+        for (Map.Entry<String, String> entry : defaultOptionsMap.entrySet()) {
+            if (!containsOption(userPrompt, entry.getKey())) {
+                userPrompt += " " + entry.getValue();
+            }
+        }
+
+        // 单独处理 --s 和 --stylize 的情况
+        if (!containsOption(userPrompt, "--s") && !containsOption(userPrompt, "--stylize")) {
+            if (defaultOptionsMap.containsKey("--s")) {
+                userPrompt += " " + defaultOptionsMap.get("--s");
+            }
+        }
+
+        return userPrompt;
+    }
+
+    public static String fixMissingSpaces(String userPrompt) {
+        // 定义需要修正的正则表达式和替换格式
+        String[][] patterns = {
+                {"--v(\\d)", "--v $1"},
+                {"--niji(\\d)", "--niji $1"},
+                {"--s(\\d+)", "--s $1"},
+                {"--stylize(\\d+)", "--stylize $1"},
+                {"--style(\\w+)", "--style $1"},
+                {"--ar(\\d+)[::](\\d+)", "--ar $1:$2"},
+                {"--ar (\\d+):(\\d+)", "--ar $1:$2"}
+        };
+
+        // 修正用户输入
+        for (String[] pattern : patterns) {
+            userPrompt = userPrompt.replaceAll(pattern[0], pattern[1]);
+        }
+
+        return userPrompt;
+    }
+
+    public static boolean containsOption(String userPrompt, String option) {
+        return userPrompt.contains(option + " ") || userPrompt.matches(".*" + option + "\\d.*");
+    }
+
+    public static Map<String, String> parseOptions(String options) {
+        Map<String, String> optionsMap = new HashMap<>();
+        String[] parts = options.split("--");
+        for (String part : parts) {
+            if (!part.trim().isEmpty()) {
+                String[] keyValue = part.trim().split(" ", 2);
+                if (keyValue.length == 2) {
+                    optionsMap.put("--" + keyValue[0], "--" + keyValue[0] + " " + keyValue[1]);
+                }
+            }
+        }
+        return optionsMap;
     }