MidjourneyServiceImpl.java 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334
  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. import java.util.stream.Collectors;
  32. /**
  33. * @author zwhui
  34. * @date 2024/4/23 15:44
  35. */
  36. @Service
  37. @Slf4j
  38. @RequiredArgsConstructor
  39. public class MidjourneyServiceImpl implements MidjourneyService {
  40. private static final String RELAX_HOST = "http://43.154.230.104:8080";
  41. private static final String FAST_HOST = "https://aigc.api4midjourney.com/api";
  42. private static final String FAST_TOKEN = "14a6b469-7cb2-4b99-9238-676c7c5cffef";
  43. private final MidjourneyUserConversationMapper conversationMapper;
  44. private static final List<String> progress = List.of("100%");
  45. private static final List<String> status = List.of("MODAL","CANCEL","FAILURE");
  46. private final RedisService redisService;
  47. private static final GlobalThreadPoolTaskExecutor TASK_EXECUTOR = GlobalThreadPoolTaskExecutor.getInstance();
  48. @Override
  49. public MidjourneyUserConversation submitImagine(MidjourneyUser user, String prompt, List<String> base64Array) throws Exception {
  50. Map<String, Object> imagineParam = MapUtil.builder(new HashMap<String,Object>())
  51. .put("prompt", prompt).put("state", user.getId()).build();
  52. if (CollectionUtil.isNotEmpty(base64Array)) {
  53. imagineParam.put("base64Array", base64Array);
  54. }
  55. SubmitResult result = submit(user.getMode(),"imagine", imagineParam);
  56. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"IMAGINE", StringUtils.EMPTY);
  57. }
  58. public MidjourneyUserConversation saveConversation(Long userId,Integer mode,Long taskId,Map<String,Object> properties, String action,String prompt) throws Exception {
  59. MidjourneyUserConversation conversation = new MidjourneyUserConversation()
  60. .setUserId(userId).setTaskId(taskId).setAction(action).setPrompt(prompt)
  61. .setMode(mode).setStartTime(System.currentTimeMillis()).setProgress("0%")
  62. .setStatus("IN_PROGRESS").setTaskId(taskId);
  63. if (MapUtil.isNotEmpty(properties)) {
  64. conversation.setProperties(Jsons.toJson(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. Map<String, Object> properties = result.getProperties();
  86. if (user.getMode() == 1) {
  87. String finalPrompt = "%s --ar %s --style raw --s 250";
  88. List<String> picList = uploadBase64Pic(base64Array);
  89. String pics = picList.stream().map(pic -> "<" + pic + ">").collect(Collectors.joining(" "));
  90. properties.put("finalPrompt", String.format(finalPrompt, pics,dimensions.getValue()));
  91. }
  92. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),properties,"BLEND", StringUtils.EMPTY);
  93. }
  94. @Override
  95. public MidjourneyUserConversation submitModal(MidjourneyUser user, Long taskId, String prompt, String maskBase64) throws Exception {
  96. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  97. .put("taskId", taskId).put("state", user.getId()).build();
  98. if (StringUtils.isNotBlank(prompt)) {
  99. param.put("prompt", prompt);
  100. }
  101. if (StringUtils.isNotBlank(maskBase64)) {
  102. param.put("maskBase64", maskBase64);
  103. }
  104. SubmitResult result = submit(user.getMode(),"modal", param);
  105. return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"MODAL", StringUtils.EMPTY);
  106. }
  107. @Override
  108. public MidjourneyUserConversation submitShorten(MidjourneyUser user, String prompt) throws Exception {
  109. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  110. .put("prompt", prompt).put("state", user.getId()).build();
  111. SubmitResult result = submit(user.getMode(),"shorten", param);
  112. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"SHORTEN", StringUtils.EMPTY);
  113. }
  114. /**
  115. * 恢复次数
  116. */
  117. public void recoverUserLimit(Long id,Integer mode,Integer num){
  118. if (mode == 1){
  119. redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + id, 1L);
  120. }
  121. if (mode == 2){
  122. if (num != null) {
  123. redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + id, 1L);
  124. }
  125. }
  126. }
  127. @Override
  128. public MidjourneyUserConversation submitAction(MidjourneyUser user, Long taskId, String customId) throws Exception {
  129. MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, taskId).last("limit 1"));
  130. if (conversation == null) {
  131. throw BusinessRuntimeException.getInstance("关联任务不存在或已失效");
  132. }
  133. if (!user.getMode().equals(conversation.getMode())) {
  134. throw BusinessRuntimeException.getInstance("当前出图模式与关联任务出图模式不符");
  135. }
  136. AtomicBoolean flag = new AtomicBoolean(false);
  137. if (StringUtils.isNotBlank(conversation.getButtons())){
  138. List<MessageButton> messageButtons = Jsons.parseList(conversation.getButtons(), MessageButton.class);
  139. messageButtons.forEach(button -> {
  140. if (button.getCustomId().equals(customId)) {
  141. log.info("action customId:{}", customId);
  142. button.setStyle(3);
  143. if (button.getLabel().contains("Vary") || button.getCustomId().contains("::pan_")
  144. || button.getEmoji().equals("🔄") || button.getCustomId().contains("PromptAnalyzer:")
  145. || button.getCustomId().contains("PicReader::") || button.getCustomId().contains("::variation::")
  146. || button.getCustomId().contains("::CustomZoom::")) {
  147. flag.set(true);
  148. }
  149. }
  150. });
  151. conversation.setButtons(Jsons.toJson(messageButtons));
  152. }
  153. conversation.setBookmark(customId.contains("BOOKMARK"));
  154. conversationMapper.updateById(conversation);
  155. if (conversation.getBookmark()) {
  156. return conversation;
  157. }
  158. Map<String, Object> param = MapUtil.builder(new HashMap<String,Object>())
  159. .put("taskId", taskId).put("state", user.getId()).put("customId", customId).build();
  160. SubmitResult result = submit(user.getMode(),"action", param);
  161. if (flag.get()) {
  162. // 以上操作有弹窗确认,恢复次数
  163. recoverUserLimit(user.getId(), user.getMode(),user.getMjRelaxNum());
  164. }
  165. return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"ACTION", StringUtils.EMPTY);
  166. }
  167. public SubmitResult submit(Integer mode,String action, Map<String, Object> param) throws Exception {
  168. if (mode == 1) {
  169. param.put("mode", "FAST");
  170. }
  171. String url = (mode == 1 ? FAST_HOST : RELAX_HOST) + getActionUrl(action);
  172. String body = HttpRequest.post(url).body(Jsons.toJson(param)).header("Authorization", FAST_TOKEN).execute().body();
  173. log.info("action body:{}", body);
  174. SubmitResult submitResult = Jsons.parseObject(body, SubmitResult.class);
  175. int code = submitResult.getCode();
  176. if (code != 1 && code != 21 && code != 22) {
  177. if (code == 3) {
  178. throw BusinessRuntimeException.getInstance("账号不存在");
  179. }
  180. if (code == 4) {
  181. throw BusinessRuntimeException.getInstance("图片重复");
  182. }
  183. if (code == 24) {
  184. throw BusinessRuntimeException.getInstance("prompt包含敏感词");
  185. }
  186. log.error("action:" + action + " error message:" + submitResult.getDescription());
  187. throw BusinessRuntimeException.getInstance("队列已满,请稍后尝试");
  188. }
  189. return submitResult;
  190. }
  191. private String getActionUrl(String action) {
  192. switch (action) {
  193. case "imagine":
  194. return "/mj/submit/imagine";
  195. case "action":
  196. return "/mj/submit/action";
  197. case "describe":
  198. return "/mj/submit/describe";
  199. case "blend":
  200. return"/mj/submit/blend";
  201. case "modal":
  202. return "/mj/submit/modal";
  203. case "shorten":
  204. return "/mj/submit/shorten";
  205. default: throw new IllegalArgumentException("Unknown action: " + action);
  206. }
  207. }
  208. @Override
  209. public List<MidjourneyUserConversation> listConversationByIds(Integer mode, List<Long> ids) throws Exception {
  210. List<MidjourneyUserConversation> list = new ArrayList<>();
  211. listByIds(mode,ids).forEach(json -> {
  212. JSONObject jsons = JSONUtil.parseObj(json);
  213. MidjourneyUserConversation conversation = JSONUtil.toBean(jsons, MidjourneyUserConversation.class);
  214. conversation.setUserId(jsons.getLong("state"));
  215. conversation.setTaskId(jsons.getLong("id"));
  216. list.add(conversation);
  217. TASK_EXECUTOR.execute(() -> {
  218. try {
  219. sync(conversation);
  220. } catch (IOException e) {
  221. throw new RuntimeException(e);
  222. }
  223. });
  224. });
  225. return list;
  226. }
  227. public void sync(MidjourneyUserConversation conversation) throws IOException {
  228. if ((StringUtils.isNotBlank(conversation.getProgress()) && progress.contains(conversation.getProgress())) || status.contains(conversation.getStatus())) {
  229. log.info("同步任务:{},进度:{}",conversation.getTaskId(),conversation.getProgress());
  230. MidjourneyUserConversation dbConversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, conversation.getTaskId())
  231. .eq(MidjourneyUserConversation::getUserId, conversation.getUserId()).last("limit 1"));
  232. if (dbConversation != null) {
  233. try {
  234. if (StringUtils.isNotBlank(conversation.getImageUrl())) {
  235. conversation.setImageUrl(uploadPic(conversation.getImageUrl(), "conversation"+dbConversation.getId()));
  236. }
  237. } catch (IOException e) {
  238. log.error("上传图片失败",e);
  239. }finally {
  240. conversation.setId(dbConversation.getId());
  241. conversationMapper.updateById(conversation);
  242. }
  243. }
  244. }
  245. }
  246. private static String uploadPic(String url, String prefix) throws IOException {
  247. byte[] body = HttpUtil.downloadBytes(url);
  248. Path tempFile = Files.createTempFile(prefix, ".png");
  249. try (FileImageOutputStream imageOutput = new FileImageOutputStream(tempFile.toFile())) {
  250. imageOutput.write(body, 0, body.length);
  251. }
  252. Map<String, Object> paramMap = new HashMap<>();
  253. paramMap.put("file", tempFile.toFile());
  254. JSONObject result = JSONUtil.parseObj(HttpUtil.post("https://files.liuliangbang.vip/pic/ups", paramMap));
  255. return result.getJSONObject("value").getJSONArray("saved").getJSONObject(0).getJSONObject("info").getStr("cdnUrl");
  256. }
  257. private static List<String> uploadBase64Pic(List<String> base64Array) throws IOException {
  258. List<String> list = new ArrayList<>();
  259. for (String base64 : base64Array) {
  260. String fileExt = "jpeg";
  261. if (base64.contains(";")){
  262. fileExt = base64.split(";")[0].split("/")[1];
  263. base64 = base64.split(",")[1];
  264. }
  265. byte[] body = Base64.getDecoder().decode(base64);
  266. Path tempFile = Files.createTempFile(UUID.randomUUID().toString(), "." + fileExt);
  267. try (FileImageOutputStream imageOutput = new FileImageOutputStream(tempFile.toFile())) {
  268. imageOutput.write(body, 0, body.length);
  269. }
  270. Map<String, Object> paramMap = new HashMap<>();
  271. paramMap.put("file", tempFile.toFile());
  272. JSONObject result = JSONUtil.parseObj(HttpUtil.post("https://files.liuliangbang.vip/pic/ups", paramMap));
  273. list.add(result.getJSONObject("value").getJSONArray("saved").getJSONObject(0).getJSONObject("info").getStr("cdnUrl"));
  274. }
  275. return list;
  276. }
  277. public JSONArray listByIds(Integer mode,List<Long> ids) throws Exception {
  278. Map<String, Object> param = MapUtil.builder(new HashMap<String, Object>()).put("ids", ids).build();
  279. 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();
  280. log.info("listByIds body:{}", body);
  281. return JSONUtil.parseArray(body);
  282. }
  283. @Override
  284. public MidjourneyUserConversation cancelConversation(MidjourneyUser user, Long id) {
  285. MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getId, id));
  286. if (conversation == null) {
  287. throw BusinessRuntimeException.getInstance("会话不存在");
  288. }
  289. if (!Objects.equals(conversation.getUserId(), user.getId())) {
  290. throw BusinessRuntimeException.getInstance("不是您的会话");
  291. }
  292. conversation.setStatus("CANCEL");
  293. conversationMapper.updateById(conversation);
  294. TASK_EXECUTOR.execute(() -> {
  295. cancel(user.getMode(),conversation.getTaskId());
  296. });
  297. return conversation;
  298. }
  299. public void cancel(Integer mode,Long taskId) {
  300. String body = HttpUtil.post(mode == 1 ? FAST_HOST : RELAX_HOST + "/mj/task/"+taskId+"/cancel",MapUtil.empty());
  301. log.info("cancel body:{}", body);
  302. JSONObject jsonObject = JSONUtil.parseObj(body);
  303. if (jsonObject.getInt("code") != 1) {
  304. throw BusinessRuntimeException.getInstance("cancel error message:" + jsonObject.getStr("description"));
  305. }
  306. }
  307. }