MidjourneyServiceImpl.java 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260
  1. package com.cyksj.service.midjourney.impl;
  2. import cn.hutool.core.bean.BeanUtil;
  3. import cn.hutool.core.collection.CollectionUtil;
  4. import cn.hutool.core.map.MapUtil;
  5. import cn.hutool.http.HttpRequest;
  6. import cn.hutool.http.HttpUtil;
  7. import cn.hutool.json.JSONArray;
  8. import cn.hutool.json.JSONObject;
  9. import cn.hutool.json.JSONUtil;
  10. import com.baomidou.mybatisplus.core.toolkit.Wrappers;
  11. import com.cyksj.common.exception.BusinessRuntimeException;
  12. import com.cyksj.common.task.GlobalThreadPoolTaskExecutor;
  13. import com.cyksj.common.util.Jsons;
  14. import com.cyksj.mapper.MidjourneyUserConversationMapper;
  15. import com.cyksj.model.dto.BlendDimensions;
  16. import com.cyksj.model.dto.MessageButton;
  17. import com.cyksj.model.entity.MidjourneyAccount;
  18. import com.cyksj.model.entity.MidjourneyUser;
  19. import com.cyksj.model.entity.MidjourneyUserConversation;
  20. import com.cyksj.model.response.SubmitResult;
  21. import com.cyksj.service.midjourney.MidjourneyService;
  22. import lombok.RequiredArgsConstructor;
  23. import lombok.extern.slf4j.Slf4j;
  24. import org.apache.commons.lang3.StringUtils;
  25. import org.springframework.stereotype.Service;
  26. import org.springframework.util.CollectionUtils;
  27. import javax.imageio.stream.FileImageOutputStream;
  28. import java.io.IOException;
  29. import java.nio.file.Files;
  30. import java.nio.file.Path;
  31. import java.util.*;
  32. import java.util.regex.Matcher;
  33. import java.util.regex.Pattern;
  34. /**
  35. * @author zwhui
  36. * @date 2024/4/23 15:44
  37. */
  38. @Service
  39. @Slf4j
  40. @RequiredArgsConstructor
  41. public class MidjourneyServiceImpl implements MidjourneyService {
  42. private static final String RELAX_HOST = "http://43.154.230.104:8080";
  43. private static final String FAST_HOST = "https://aigc.api4midjourney.com/api";
  44. private static final String FAST_TOKEN = "14a6b469-7cb2-4b99-9238-676c7c5cffef";
  45. private final MidjourneyUserConversationMapper conversationMapper;
  46. private static final List<String> progress = Arrays.asList("25%","35%", "45%","65%", "75%", "85%", "95%","100%");
  47. private static final List<String> status = Arrays.asList("MODAL","IN_PROGRESS", "SUCCESS", "FAILURE", "CANCEL");
  48. private static final GlobalThreadPoolTaskExecutor TASK_EXECUTOR = GlobalThreadPoolTaskExecutor.getInstance();
  49. @Override
  50. public MidjourneyUserConversation submitImagine(MidjourneyUser user, String prompt, List<String> base64Array) throws Exception {
  51. Map<String, Object> imagineParam = MapUtil.builder(new HashMap<String,Object>())
  52. .put("prompt", prompt).put("state", user.getId()).build();
  53. if (CollectionUtil.isNotEmpty(base64Array)) {
  54. imagineParam.put("base64Array", base64Array);
  55. }
  56. SubmitResult result = submit(user.getMode(),"imagine", imagineParam);
  57. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"IMAGINE", StringUtils.EMPTY);
  58. }
  59. public MidjourneyUserConversation saveConversation(Long userId,Integer mode,Long taskId,Map<String,Object> properties, String action,String prompt) {
  60. MidjourneyUserConversation conversation = new MidjourneyUserConversation()
  61. .setUserId(userId).setTaskId(taskId).setAction(action).setPrompt(prompt)
  62. .setMode(mode).setStartTime(System.currentTimeMillis()).setProgress("0%")
  63. .setStatus("IN_PROGRESS").setTaskId(taskId);
  64. if (MapUtil.isNotEmpty(properties)) {
  65. conversation.setChannelId(properties.get("discordChannelId") == null ? null : Long.parseLong(properties.get("discordChannelId").toString()))
  66. .setInstanceId(properties.get("discordInstanceId") == null ? null : Long.parseLong(properties.get("discordInstanceId").toString()));
  67. }
  68. conversationMapper.insert(conversation);
  69. return conversation;
  70. }
  71. @Override
  72. public MidjourneyUserConversation submitDescribe(MidjourneyUser user, String base64) throws Exception {
  73. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>()).put("state", user.getId()).put("base64", base64).build();
  74. SubmitResult result = submit(user.getMode(),"describe", param);
  75. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"DESCRIBE", StringUtils.EMPTY);
  76. }
  77. @Override
  78. public MidjourneyUserConversation submitBlend(MidjourneyUser user, BlendDimensions dimensions, List<String> base64Array) throws Exception {
  79. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  80. .put("base64Array", base64Array).put("state", user.getId()).build();
  81. if (dimensions != null) {
  82. param.put("dimensions", dimensions);
  83. }
  84. SubmitResult result = submit(user.getMode(),"blend", param);
  85. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"BLEND", StringUtils.EMPTY);
  86. }
  87. @Override
  88. public MidjourneyUserConversation submitModal(MidjourneyUser user, Long taskId, String prompt, String maskBase64) throws Exception {
  89. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  90. .put("taskId", taskId).put("state", user.getId()).build();
  91. if (StringUtils.isNotBlank(prompt)) {
  92. param.put("prompt", prompt);
  93. }
  94. if (StringUtils.isNotBlank(maskBase64)) {
  95. param.put("maskBase64", maskBase64);
  96. }
  97. SubmitResult result = submit(user.getMode(),"modal", param);
  98. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"MODAL", StringUtils.EMPTY);
  99. }
  100. @Override
  101. public MidjourneyUserConversation submitShorten(MidjourneyUser user, String prompt) throws Exception {
  102. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  103. .put("prompt", prompt).put("state", user.getId()).build();
  104. SubmitResult result = submit(user.getMode(),"shorten", param);
  105. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"SHORTEN", StringUtils.EMPTY);
  106. }
  107. @Override
  108. public MidjourneyUserConversation submitAction(MidjourneyUser user, Long taskId, String customId) throws Exception {
  109. MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, taskId).last("limit 1"));
  110. if (conversation == null) {
  111. throw BusinessRuntimeException.getInstance("关联任务不存在或已失效");
  112. }
  113. if (!user.getMode().equals(conversation.getMode())) {
  114. throw BusinessRuntimeException.getInstance("当前出图模式与关联任务出图模式不符");
  115. }
  116. if (StringUtils.isNotBlank(conversation.getButtons())){
  117. Jsons.parseList(conversation.getButtons(), MessageButton.class).forEach(button -> {
  118. if (button.getCustomId().equals(customId)) {
  119. button.setStyle(3);
  120. }
  121. });
  122. }
  123. conversationMapper.updateById(conversation);
  124. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  125. .put("taskId", taskId).put("state", user.getId()).put("customId", customId).build();
  126. SubmitResult result = submit(user.getMode(),"action", param);
  127. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"ACTION", StringUtils.EMPTY);
  128. }
  129. public SubmitResult submit(Integer mode,String action, Map<String, Object> param) throws Exception {
  130. if (mode == 1) {
  131. param.put("mode", "FAST");
  132. }
  133. String url = (mode == 1 ? FAST_HOST : RELAX_HOST) + getActionUrl(action);
  134. String body = HttpRequest.post(url).body(Jsons.toJson(param)).header("Authorization", FAST_TOKEN).execute().body();
  135. log.info("action body:{}", body);
  136. SubmitResult submitResult = Jsons.parseObject(body, SubmitResult.class);
  137. int code = submitResult.getCode();
  138. if (code != 1 && code != 21 && code != 22) {
  139. throw BusinessRuntimeException.getInstance("action:" + action + " error message:" + submitResult.getDescription());
  140. }
  141. return submitResult;
  142. }
  143. private String getActionUrl(String action) {
  144. switch (action) {
  145. case "imagine":
  146. return "/mj/submit/imagine";
  147. case "action":
  148. return "/mj/submit/action";
  149. case "describe":
  150. return "/mj/submit/describe";
  151. case "blend":
  152. return"/mj/submit/blend";
  153. case "modal":
  154. return "/mj/submit/modal";
  155. case "shorten":
  156. return "/mj/submit/shorten";
  157. default: throw new IllegalArgumentException("Unknown action: " + action);
  158. }
  159. }
  160. @Override
  161. public List<MidjourneyUserConversation> listConversationByIds(Integer mode, List<Long> ids) throws Exception {
  162. List<MidjourneyUserConversation> list = new ArrayList<>();
  163. listByIds(mode,ids).forEach(json -> {
  164. JSONObject jsons = JSONUtil.parseObj(json);
  165. MidjourneyUserConversation conversation = JSONUtil.toBean(jsons, MidjourneyUserConversation.class);
  166. conversation.setUserId(jsons.getLong("state"));
  167. conversation.setTaskId(jsons.getLong("id"));
  168. list.add(conversation);
  169. TASK_EXECUTOR.execute(() -> {
  170. try {
  171. sync(conversation);
  172. } catch (IOException e) {
  173. throw new RuntimeException(e);
  174. }
  175. });
  176. });
  177. return list;
  178. }
  179. public void sync(MidjourneyUserConversation conversation) throws IOException {
  180. if (progress.contains(conversation.getProgress()) || status.contains(conversation.getStatus())) {
  181. log.info("同步任务:{},进度:{}",conversation.getTaskId(),conversation.getProgress());
  182. MidjourneyUserConversation dbConversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, conversation.getTaskId())
  183. .eq(MidjourneyUserConversation::getUserId, conversation.getUserId()).last("limit 1"));
  184. if (dbConversation != null) {
  185. if (StringUtils.isNotBlank(conversation.getImageUrl())) {
  186. conversation.setImageUrl(uploadPic(conversation.getImageUrl(), "conversation"+dbConversation.getId()));
  187. }
  188. conversation.setId(dbConversation.getId());
  189. conversationMapper.updateById(conversation);
  190. }
  191. }
  192. }
  193. private static String uploadPic(String url, String prefix) throws IOException {
  194. byte[] body = HttpUtil.downloadBytes(url);
  195. Path tempFile = Files.createTempFile(prefix, ".png");
  196. try (FileImageOutputStream imageOutput = new FileImageOutputStream(tempFile.toFile())) {
  197. imageOutput.write(body, 0, body.length);
  198. }
  199. Map<String, Object> paramMap = new HashMap<>();
  200. paramMap.put("file", tempFile.toFile());
  201. JSONObject result = JSONUtil.parseObj(HttpUtil.post("https://files.liuliangbang.vip/pic/ups", paramMap));
  202. return result.getJSONObject("value").getJSONArray("saved").getJSONObject(0).getJSONObject("info").getStr("cdnUrl");
  203. }
  204. public JSONArray listByIds(Integer mode,List<Long> ids) throws Exception {
  205. Map<String, Object> param = MapUtil.builder(new HashMap<String, Object>()).put("ids", ids).build();
  206. String body = HttpRequest.post((mode == 1 ? FAST_HOST : RELAX_HOST) + "/mj/task/list-by-condition").header("Authorization", FAST_TOKEN).body(Jsons.toJson(param)).execute().body();
  207. log.info("listByIds body:{}", body);
  208. return JSONUtil.parseArray(body);
  209. }
  210. @Override
  211. public MidjourneyUserConversation cancelConversation(MidjourneyUser user, Long id) {
  212. MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getId, id));
  213. if (conversation == null) {
  214. throw BusinessRuntimeException.getInstance("会话不存在");
  215. }
  216. if (!Objects.equals(conversation.getUserId(), user.getId())) {
  217. throw BusinessRuntimeException.getInstance("不是您的会话");
  218. }
  219. conversation.setStatus("CANCEL");
  220. conversationMapper.updateById(conversation);
  221. TASK_EXECUTOR.execute(() -> {
  222. cancel(user.getMode(),conversation.getTaskId());
  223. });
  224. return conversation;
  225. }
  226. public void cancel(Integer mode,Long taskId) {
  227. String body = HttpUtil.post(mode == 1 ? FAST_HOST : RELAX_HOST + "/mj/task/"+taskId+"/cancel",MapUtil.empty());
  228. log.info("cancel body:{}", body);
  229. JSONObject jsonObject = JSONUtil.parseObj(body);
  230. if (jsonObject.getInt("code") != 1) {
  231. throw BusinessRuntimeException.getInstance("cancel error message:" + jsonObject.getStr("description"));
  232. }
  233. }
  234. }