MidjourneyServiceImpl.java 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273
  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 = List.of("100%");
  47. private static final List<String> status = List.of("MODAL","CANCEL","FAILURE");
  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) throws Exception {
  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.setProperties(Jsons.toJson(properties));
  66. conversation.setChannelId(properties.get("discordChannelId") == null ? null : Long.parseLong(properties.get("discordChannelId").toString()))
  67. .setInstanceId(properties.get("discordInstanceId") == null ? null : Long.parseLong(properties.get("discordInstanceId").toString()));
  68. }
  69. conversationMapper.insert(conversation);
  70. return conversation;
  71. }
  72. @Override
  73. public MidjourneyUserConversation submitDescribe(MidjourneyUser user, String base64) throws Exception {
  74. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>()).put("state", user.getId()).put("base64", base64).build();
  75. SubmitResult result = submit(user.getMode(),"describe", param);
  76. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"DESCRIBE", StringUtils.EMPTY);
  77. }
  78. @Override
  79. public MidjourneyUserConversation submitBlend(MidjourneyUser user, BlendDimensions dimensions, List<String> base64Array) throws Exception {
  80. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  81. .put("base64Array", base64Array).put("state", user.getId()).build();
  82. if (dimensions != null) {
  83. param.put("dimensions", dimensions);
  84. }
  85. SubmitResult result = submit(user.getMode(),"blend", param);
  86. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"BLEND", StringUtils.EMPTY);
  87. }
  88. @Override
  89. public MidjourneyUserConversation submitModal(MidjourneyUser user, Long taskId, String prompt, String maskBase64) throws Exception {
  90. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  91. .put("taskId", taskId).put("state", user.getId()).build();
  92. if (StringUtils.isNotBlank(prompt)) {
  93. param.put("prompt", prompt);
  94. }
  95. if (StringUtils.isNotBlank(maskBase64)) {
  96. param.put("maskBase64", maskBase64);
  97. }
  98. SubmitResult result = submit(user.getMode(),"modal", param);
  99. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"MODAL", StringUtils.EMPTY);
  100. }
  101. @Override
  102. public MidjourneyUserConversation submitShorten(MidjourneyUser user, String prompt) throws Exception {
  103. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  104. .put("prompt", prompt).put("state", user.getId()).build();
  105. SubmitResult result = submit(user.getMode(),"shorten", param);
  106. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"SHORTEN", StringUtils.EMPTY);
  107. }
  108. @Override
  109. public MidjourneyUserConversation submitAction(MidjourneyUser user, Long taskId, String customId) throws Exception {
  110. MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, taskId).last("limit 1"));
  111. if (conversation == null) {
  112. throw BusinessRuntimeException.getInstance("关联任务不存在或已失效");
  113. }
  114. if (!user.getMode().equals(conversation.getMode())) {
  115. throw BusinessRuntimeException.getInstance("当前出图模式与关联任务出图模式不符");
  116. }
  117. if (StringUtils.isNotBlank(conversation.getButtons())){
  118. List<MessageButton> messageButtons = Jsons.parseList(conversation.getButtons(), MessageButton.class);
  119. messageButtons.forEach(button -> {
  120. if (button.getCustomId().equals(customId)) {
  121. log.info("action customId:{}", customId);
  122. button.setStyle(3);
  123. }
  124. });
  125. conversation.setButtons(Jsons.toJson(messageButtons));
  126. }
  127. conversation.setBookmark(customId.contains("BOOKMARK"));
  128. conversationMapper.updateById(conversation);
  129. if (conversation.getBookmark()) {
  130. return conversation;
  131. }
  132. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  133. .put("taskId", taskId).put("state", user.getId()).put("customId", customId).build();
  134. SubmitResult result = submit(user.getMode(),"action", param);
  135. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"ACTION", StringUtils.EMPTY);
  136. }
  137. public SubmitResult submit(Integer mode,String action, Map<String, Object> param) throws Exception {
  138. if (mode == 1) {
  139. param.put("mode", "FAST");
  140. }
  141. String url = (mode == 1 ? FAST_HOST : RELAX_HOST) + getActionUrl(action);
  142. String body = HttpRequest.post(url).body(Jsons.toJson(param)).header("Authorization", FAST_TOKEN).execute().body();
  143. log.info("action body:{}", body);
  144. SubmitResult submitResult = Jsons.parseObject(body, SubmitResult.class);
  145. int code = submitResult.getCode();
  146. if (code != 1 && code != 21 && code != 22) {
  147. throw BusinessRuntimeException.getInstance("action:" + action + " error message:" + submitResult.getDescription());
  148. }
  149. return submitResult;
  150. }
  151. private String getActionUrl(String action) {
  152. switch (action) {
  153. case "imagine":
  154. return "/mj/submit/imagine";
  155. case "action":
  156. return "/mj/submit/action";
  157. case "describe":
  158. return "/mj/submit/describe";
  159. case "blend":
  160. return"/mj/submit/blend";
  161. case "modal":
  162. return "/mj/submit/modal";
  163. case "shorten":
  164. return "/mj/submit/shorten";
  165. default: throw new IllegalArgumentException("Unknown action: " + action);
  166. }
  167. }
  168. @Override
  169. public List<MidjourneyUserConversation> listConversationByIds(Integer mode, List<Long> ids) throws Exception {
  170. List<MidjourneyUserConversation> list = new ArrayList<>();
  171. listByIds(mode,ids).forEach(json -> {
  172. JSONObject jsons = JSONUtil.parseObj(json);
  173. MidjourneyUserConversation conversation = JSONUtil.toBean(jsons, MidjourneyUserConversation.class);
  174. conversation.setUserId(jsons.getLong("state"));
  175. conversation.setTaskId(jsons.getLong("id"));
  176. list.add(conversation);
  177. TASK_EXECUTOR.execute(() -> {
  178. try {
  179. sync(conversation);
  180. } catch (IOException e) {
  181. throw new RuntimeException(e);
  182. }
  183. });
  184. });
  185. return list;
  186. }
  187. public void sync(MidjourneyUserConversation conversation) throws IOException {
  188. if ((StringUtils.isNotBlank(conversation.getProgress()) && progress.contains(conversation.getProgress())) || status.contains(conversation.getStatus())) {
  189. log.info("同步任务:{},进度:{}",conversation.getTaskId(),conversation.getProgress());
  190. MidjourneyUserConversation dbConversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, conversation.getTaskId())
  191. .eq(MidjourneyUserConversation::getUserId, conversation.getUserId()).last("limit 1"));
  192. if (dbConversation != null) {
  193. try {
  194. if (StringUtils.isNotBlank(conversation.getImageUrl())) {
  195. conversation.setImageUrl(uploadPic(conversation.getImageUrl(), "conversation"+dbConversation.getId()));
  196. }
  197. } catch (IOException e) {
  198. log.error("上传图片失败",e);
  199. }finally {
  200. conversation.setId(dbConversation.getId());
  201. conversationMapper.updateById(conversation);
  202. }
  203. }
  204. }
  205. }
  206. private static String uploadPic(String url, String prefix) throws IOException {
  207. byte[] body = HttpUtil.downloadBytes(url);
  208. Path tempFile = Files.createTempFile(prefix, ".png");
  209. try (FileImageOutputStream imageOutput = new FileImageOutputStream(tempFile.toFile())) {
  210. imageOutput.write(body, 0, body.length);
  211. }
  212. Map<String, Object> paramMap = new HashMap<>();
  213. paramMap.put("file", tempFile.toFile());
  214. JSONObject result = JSONUtil.parseObj(HttpUtil.post("https://files.liuliangbang.vip/pic/ups", paramMap));
  215. return result.getJSONObject("value").getJSONArray("saved").getJSONObject(0).getJSONObject("info").getStr("cdnUrl");
  216. }
  217. public JSONArray listByIds(Integer mode,List<Long> ids) throws Exception {
  218. Map<String, Object> param = MapUtil.builder(new HashMap<String, Object>()).put("ids", ids).build();
  219. 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();
  220. log.info("listByIds body:{}", body);
  221. return JSONUtil.parseArray(body);
  222. }
  223. @Override
  224. public MidjourneyUserConversation cancelConversation(MidjourneyUser user, Long id) {
  225. MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getId, id));
  226. if (conversation == null) {
  227. throw BusinessRuntimeException.getInstance("会话不存在");
  228. }
  229. if (!Objects.equals(conversation.getUserId(), user.getId())) {
  230. throw BusinessRuntimeException.getInstance("不是您的会话");
  231. }
  232. conversation.setStatus("CANCEL");
  233. conversationMapper.updateById(conversation);
  234. TASK_EXECUTOR.execute(() -> {
  235. cancel(user.getMode(),conversation.getTaskId());
  236. });
  237. return conversation;
  238. }
  239. public void cancel(Integer mode,Long taskId) {
  240. String body = HttpUtil.post(mode == 1 ? FAST_HOST : RELAX_HOST + "/mj/task/"+taskId+"/cancel",MapUtil.empty());
  241. log.info("cancel body:{}", body);
  242. JSONObject jsonObject = JSONUtil.parseObj(body);
  243. if (jsonObject.getInt("code") != 1) {
  244. throw BusinessRuntimeException.getInstance("cancel error message:" + jsonObject.getStr("description"));
  245. }
  246. }
  247. }