MidjourneyServiceImpl.java 13 KB

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