MidjourneyServiceImpl.java 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302
  1. package com.cyksj.service.midjourney.impl;
  2. import cn.hutool.core.collection.CollectionUtil;
  3. import cn.hutool.core.map.MapUtil;
  4. import cn.hutool.http.HttpRequest;
  5. import cn.hutool.http.HttpUtil;
  6. import cn.hutool.json.JSONArray;
  7. import cn.hutool.json.JSONObject;
  8. import cn.hutool.json.JSONUtil;
  9. import com.baomidou.mybatisplus.core.toolkit.Wrappers;
  10. import com.cyksj.common.exception.BusinessRuntimeException;
  11. import com.cyksj.common.task.GlobalThreadPoolTaskExecutor;
  12. import com.cyksj.common.util.Jsons;
  13. import com.cyksj.mapper.MidjourneyUserConversationMapper;
  14. import com.cyksj.model.dto.BlendDimensions;
  15. import com.cyksj.model.dto.MessageButton;
  16. import com.cyksj.model.entity.MidjourneyUser;
  17. import com.cyksj.model.entity.MidjourneyUserConversation;
  18. import com.cyksj.model.response.SubmitResult;
  19. import com.cyksj.redis.RedisService;
  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 javax.imageio.stream.FileImageOutputStream;
  26. import java.io.IOException;
  27. import java.nio.file.Files;
  28. import java.nio.file.Path;
  29. import java.util.*;
  30. import java.util.concurrent.atomic.AtomicBoolean;
  31. /**
  32. * @author zwhui
  33. * @date 2024/4/23 15:44
  34. */
  35. @Service
  36. @Slf4j
  37. @RequiredArgsConstructor
  38. public class MidjourneyServiceImpl implements MidjourneyService {
  39. private static final String RELAX_HOST = "http://43.154.230.104:8080";
  40. private static final String FAST_HOST = "https://aigc.api4midjourney.com/api";
  41. private static final String FAST_TOKEN = "14a6b469-7cb2-4b99-9238-676c7c5cffef";
  42. private final MidjourneyUserConversationMapper conversationMapper;
  43. private static final List<String> progress = List.of("100%");
  44. private static final List<String> status = List.of("MODAL","CANCEL","FAILURE");
  45. private final RedisService redisService;
  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) throws Exception {
  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.setProperties(Jsons.toJson(properties));
  64. conversation.setChannelId(properties.get("discordChannelId") == null ? null : Long.parseLong(properties.get("discordChannelId").toString()))
  65. .setInstanceId(properties.get("discordInstanceId") == null ? null : Long.parseLong(properties.get("discordInstanceId").toString()));
  66. }
  67. conversationMapper.insert(conversation);
  68. return conversation;
  69. }
  70. @Override
  71. public MidjourneyUserConversation submitDescribe(MidjourneyUser user, String base64) throws Exception {
  72. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>()).put("state", user.getId()).put("base64", base64).build();
  73. SubmitResult result = submit(user.getMode(),"describe", param);
  74. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"DESCRIBE", StringUtils.EMPTY);
  75. }
  76. @Override
  77. public MidjourneyUserConversation submitBlend(MidjourneyUser user, BlendDimensions dimensions, List<String> base64Array) throws Exception {
  78. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  79. .put("base64Array", base64Array).put("state", user.getId()).build();
  80. if (dimensions != null) {
  81. param.put("dimensions", dimensions);
  82. }
  83. SubmitResult result = submit(user.getMode(),"blend", param);
  84. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"BLEND", StringUtils.EMPTY);
  85. }
  86. @Override
  87. public MidjourneyUserConversation submitModal(MidjourneyUser user, Long taskId, String prompt, String maskBase64) throws Exception {
  88. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  89. .put("taskId", taskId).put("state", user.getId()).build();
  90. if (StringUtils.isNotBlank(prompt)) {
  91. param.put("prompt", prompt);
  92. }
  93. if (StringUtils.isNotBlank(maskBase64)) {
  94. param.put("maskBase64", maskBase64);
  95. }
  96. SubmitResult result = submit(user.getMode(),"modal", param);
  97. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"MODAL", StringUtils.EMPTY);
  98. }
  99. @Override
  100. public MidjourneyUserConversation submitShorten(MidjourneyUser user, String prompt) throws Exception {
  101. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  102. .put("prompt", prompt).put("state", user.getId()).build();
  103. SubmitResult result = submit(user.getMode(),"shorten", param);
  104. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"SHORTEN", StringUtils.EMPTY);
  105. }
  106. /**
  107. * 恢复次数
  108. */
  109. public void recoverUserLimit(Long id,Integer mode,Integer num){
  110. if (mode == 1){
  111. redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + id, 1L);
  112. }
  113. if (mode == 2){
  114. if (num != null) {
  115. redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + id, 1L);
  116. }
  117. }
  118. }
  119. @Override
  120. public MidjourneyUserConversation submitAction(MidjourneyUser user, Long taskId, String customId) throws Exception {
  121. MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, taskId).last("limit 1"));
  122. if (conversation == null) {
  123. throw BusinessRuntimeException.getInstance("关联任务不存在或已失效");
  124. }
  125. if (!user.getMode().equals(conversation.getMode())) {
  126. throw BusinessRuntimeException.getInstance("当前出图模式与关联任务出图模式不符");
  127. }
  128. AtomicBoolean flag = new AtomicBoolean(false);
  129. if (StringUtils.isNotBlank(conversation.getButtons())){
  130. List<MessageButton> messageButtons = Jsons.parseList(conversation.getButtons(), MessageButton.class);
  131. messageButtons.forEach(button -> {
  132. if (button.getCustomId().equals(customId)) {
  133. log.info("action customId:{}", customId);
  134. button.setStyle(3);
  135. if (button.getLabel().contains("Vary") || button.getCustomId().contains("::pan_")
  136. || button.getEmoji().equals("🔄") || button.getCustomId().contains("PromptAnalyzer:")
  137. || button.getCustomId().contains("PicReader::") || button.getCustomId().contains("::variation::")
  138. || button.getCustomId().contains("::CustomZoom::")) {
  139. flag.set(true);
  140. }
  141. }
  142. });
  143. conversation.setButtons(Jsons.toJson(messageButtons));
  144. }
  145. conversation.setBookmark(customId.contains("BOOKMARK"));
  146. conversationMapper.updateById(conversation);
  147. if (conversation.getBookmark()) {
  148. return conversation;
  149. }
  150. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  151. .put("taskId", taskId).put("state", user.getId()).put("customId", customId).build();
  152. SubmitResult result = submit(user.getMode(),"action", param);
  153. if (flag.get()) {
  154. // 以上操作有弹窗确认,恢复次数
  155. recoverUserLimit(user.getId(), user.getMode(),user.getMjRelaxNum());
  156. }
  157. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"ACTION", StringUtils.EMPTY);
  158. }
  159. public SubmitResult submit(Integer mode,String action, Map<String, Object> param) throws Exception {
  160. if (mode == 1) {
  161. param.put("mode", "FAST");
  162. }
  163. String url = (mode == 1 ? FAST_HOST : RELAX_HOST) + getActionUrl(action);
  164. String body = HttpRequest.post(url).body(Jsons.toJson(param)).header("Authorization", FAST_TOKEN).execute().body();
  165. log.info("action body:{}", body);
  166. SubmitResult submitResult = Jsons.parseObject(body, SubmitResult.class);
  167. int code = submitResult.getCode();
  168. if (code != 1 && code != 21 && code != 22) {
  169. if (code == 3) {
  170. throw BusinessRuntimeException.getInstance("账号不存在");
  171. }
  172. if (code == 24) {
  173. throw BusinessRuntimeException.getInstance("prompt包含敏感词");
  174. }
  175. log.error("action:" + action + " error message:" + submitResult.getDescription());
  176. throw BusinessRuntimeException.getInstance("队列已满,请稍后尝试");
  177. }
  178. return submitResult;
  179. }
  180. private String getActionUrl(String action) {
  181. switch (action) {
  182. case "imagine":
  183. return "/mj/submit/imagine";
  184. case "action":
  185. return "/mj/submit/action";
  186. case "describe":
  187. return "/mj/submit/describe";
  188. case "blend":
  189. return"/mj/submit/blend";
  190. case "modal":
  191. return "/mj/submit/modal";
  192. case "shorten":
  193. return "/mj/submit/shorten";
  194. default: throw new IllegalArgumentException("Unknown action: " + action);
  195. }
  196. }
  197. @Override
  198. public List<MidjourneyUserConversation> listConversationByIds(Integer mode, List<Long> ids) throws Exception {
  199. List<MidjourneyUserConversation> list = new ArrayList<>();
  200. listByIds(mode,ids).forEach(json -> {
  201. JSONObject jsons = JSONUtil.parseObj(json);
  202. MidjourneyUserConversation conversation = JSONUtil.toBean(jsons, MidjourneyUserConversation.class);
  203. conversation.setUserId(jsons.getLong("state"));
  204. conversation.setTaskId(jsons.getLong("id"));
  205. list.add(conversation);
  206. TASK_EXECUTOR.execute(() -> {
  207. try {
  208. sync(conversation);
  209. } catch (IOException e) {
  210. throw new RuntimeException(e);
  211. }
  212. });
  213. });
  214. return list;
  215. }
  216. public void sync(MidjourneyUserConversation conversation) throws IOException {
  217. if ((StringUtils.isNotBlank(conversation.getProgress()) && progress.contains(conversation.getProgress())) || status.contains(conversation.getStatus())) {
  218. log.info("同步任务:{},进度:{}",conversation.getTaskId(),conversation.getProgress());
  219. MidjourneyUserConversation dbConversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, conversation.getTaskId())
  220. .eq(MidjourneyUserConversation::getUserId, conversation.getUserId()).last("limit 1"));
  221. if (dbConversation != null) {
  222. try {
  223. if (StringUtils.isNotBlank(conversation.getImageUrl())) {
  224. conversation.setImageUrl(uploadPic(conversation.getImageUrl(), "conversation"+dbConversation.getId()));
  225. }
  226. } catch (IOException e) {
  227. log.error("上传图片失败",e);
  228. }finally {
  229. conversation.setId(dbConversation.getId());
  230. conversationMapper.updateById(conversation);
  231. }
  232. }
  233. }
  234. }
  235. private static String uploadPic(String url, String prefix) throws IOException {
  236. byte[] body = HttpUtil.downloadBytes(url);
  237. Path tempFile = Files.createTempFile(prefix, ".png");
  238. try (FileImageOutputStream imageOutput = new FileImageOutputStream(tempFile.toFile())) {
  239. imageOutput.write(body, 0, body.length);
  240. }
  241. Map<String, Object> paramMap = new HashMap<>();
  242. paramMap.put("file", tempFile.toFile());
  243. JSONObject result = JSONUtil.parseObj(HttpUtil.post("https://files.liuliangbang.vip/pic/ups", paramMap));
  244. return result.getJSONObject("value").getJSONArray("saved").getJSONObject(0).getJSONObject("info").getStr("cdnUrl");
  245. }
  246. public JSONArray listByIds(Integer mode,List<Long> ids) throws Exception {
  247. Map<String, Object> param = MapUtil.builder(new HashMap<String, Object>()).put("ids", ids).build();
  248. 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();
  249. log.info("listByIds body:{}", body);
  250. return JSONUtil.parseArray(body);
  251. }
  252. @Override
  253. public MidjourneyUserConversation cancelConversation(MidjourneyUser user, Long id) {
  254. MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getId, id));
  255. if (conversation == null) {
  256. throw BusinessRuntimeException.getInstance("会话不存在");
  257. }
  258. if (!Objects.equals(conversation.getUserId(), user.getId())) {
  259. throw BusinessRuntimeException.getInstance("不是您的会话");
  260. }
  261. conversation.setStatus("CANCEL");
  262. conversationMapper.updateById(conversation);
  263. TASK_EXECUTOR.execute(() -> {
  264. cancel(user.getMode(),conversation.getTaskId());
  265. });
  266. return conversation;
  267. }
  268. public void cancel(Integer mode,Long taskId) {
  269. String body = HttpUtil.post(mode == 1 ? FAST_HOST : RELAX_HOST + "/mj/task/"+taskId+"/cancel",MapUtil.empty());
  270. log.info("cancel body:{}", body);
  271. JSONObject jsonObject = JSONUtil.parseObj(body);
  272. if (jsonObject.getInt("code") != 1) {
  273. throw BusinessRuntimeException.getInstance("cancel error message:" + jsonObject.getStr("description"));
  274. }
  275. }
  276. }