package com.cyksj.service.midjourney.impl; import cn.hutool.core.collection.CollectionUtil; import cn.hutool.core.map.MapUtil; import cn.hutool.http.HttpRequest; import cn.hutool.http.HttpUtil; import cn.hutool.json.JSONArray; import cn.hutool.json.JSONObject; import cn.hutool.json.JSONUtil; import com.baomidou.mybatisplus.core.toolkit.Wrappers; import com.cyksj.common.exception.BusinessRuntimeException; import com.cyksj.common.task.GlobalThreadPoolTaskExecutor; import com.cyksj.common.util.Jsons; import com.cyksj.mapper.MidjourneyUserConversationMapper; import com.cyksj.model.dto.BlendDimensions; import com.cyksj.model.dto.MessageButton; import com.cyksj.model.entity.MidjourneyUser; import com.cyksj.model.entity.MidjourneyUserConversation; import com.cyksj.model.response.SubmitResult; import com.cyksj.redis.RedisService; import com.cyksj.service.midjourney.MidjourneyService; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; import org.springframework.stereotype.Service; import javax.imageio.stream.FileImageOutputStream; import java.io.IOException; import java.nio.file.Files; import java.nio.file.Path; import java.util.*; import java.util.concurrent.atomic.AtomicBoolean; import java.util.stream.Collectors; /** * @author zwhui * @date 2024/4/23 15:44 */ @Service @Slf4j @RequiredArgsConstructor public class MidjourneyServiceImpl implements MidjourneyService { private static final String RELAX_HOST = "http://43.154.230.104:8080"; private static final String FAST_HOST = "https://aigc.api4midjourney.com/api"; private static final String FAST_TOKEN = "14a6b469-7cb2-4b99-9238-676c7c5cffef"; private final MidjourneyUserConversationMapper conversationMapper; private static final List progress = List.of("100%"); private static final List status = List.of("MODAL","CANCEL","FAILURE"); private final RedisService redisService; private static final GlobalThreadPoolTaskExecutor TASK_EXECUTOR = GlobalThreadPoolTaskExecutor.getInstance(); @Override public MidjourneyUserConversation submitImagine(MidjourneyUser user, String prompt, List base64Array) throws Exception { Map imagineParam = MapUtil.builder(new HashMap()) .put("prompt", prompt).put("state", user.getId()).build(); if (CollectionUtil.isNotEmpty(base64Array)) { imagineParam.put("base64Array", base64Array); } SubmitResult result = submit(user.getMode(),"imagine", imagineParam); return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"IMAGINE", StringUtils.EMPTY); } public MidjourneyUserConversation saveConversation(Long userId,Integer mode,Long taskId,Map properties, String action,String prompt) throws Exception { MidjourneyUserConversation conversation = new MidjourneyUserConversation() .setUserId(userId).setTaskId(taskId).setAction(action).setPrompt(prompt) .setMode(mode).setStartTime(System.currentTimeMillis()).setProgress("0%") .setStatus("IN_PROGRESS").setTaskId(taskId); if (MapUtil.isNotEmpty(properties)) { conversation.setProperties(Jsons.toJson(properties)); conversation.setChannelId(properties.get("discordChannelId") == null ? null : Long.parseLong(properties.get("discordChannelId").toString())) .setInstanceId(properties.get("discordInstanceId") == null ? null : Long.parseLong(properties.get("discordInstanceId").toString())); } conversationMapper.insert(conversation); return conversation; } @Override public MidjourneyUserConversation submitDescribe(MidjourneyUser user, String base64) throws Exception { Map param = MapUtil.builder(new HashMap()).put("state", user.getId()).put("base64", base64).build(); SubmitResult result = submit(user.getMode(),"describe", param); return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"DESCRIBE", StringUtils.EMPTY); } @Override public MidjourneyUserConversation submitBlend(MidjourneyUser user, BlendDimensions dimensions, List base64Array) throws Exception { Map param = MapUtil.builder(new HashMap()) .put("base64Array", base64Array).put("state", user.getId()).build(); if (dimensions != null) { param.put("dimensions", dimensions); } SubmitResult result = submit(user.getMode(),"blend", param); Map properties = result.getProperties(); if (user.getMode() == 1) { String finalPrompt = "%s --ar %s --style raw --s 250"; List picList = uploadBase64Pic(base64Array); String pics = picList.stream().map(pic -> "<" + pic + ">").collect(Collectors.joining(" ")); properties.put("finalPrompt", String.format(finalPrompt, pics,dimensions.getValue())); } return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),properties,"BLEND", StringUtils.EMPTY); } @Override public MidjourneyUserConversation submitModal(MidjourneyUser user, Long taskId, String prompt, String maskBase64) throws Exception { Map param = MapUtil.builder(new HashMap()) .put("taskId", taskId).put("state", user.getId()).build(); if (StringUtils.isNotBlank(prompt)) { param.put("prompt", prompt); } if (StringUtils.isNotBlank(maskBase64)) { param.put("maskBase64", maskBase64); } SubmitResult result = submit(user.getMode(),"modal", param); return saveConversation(user.getId(), user.getMode(),Long.parseLong(result.getResult()),result.getProperties(),"MODAL", StringUtils.EMPTY); } @Override public MidjourneyUserConversation submitShorten(MidjourneyUser user, String prompt) throws Exception { Map param = MapUtil.builder(new HashMap()) .put("prompt", prompt).put("state", user.getId()).build(); SubmitResult result = submit(user.getMode(),"shorten", param); return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"SHORTEN", StringUtils.EMPTY); } /** * 恢复次数 */ public void recoverUserLimit(Long id,Integer mode,Integer num){ if (mode == 1){ redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + id, 1L); } if (mode == 2){ if (num != null) { redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + id, 1L); } } } @Override public MidjourneyUserConversation submitAction(MidjourneyUser user, Long taskId, String customId) throws Exception { MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, taskId).last("limit 1")); if (conversation == null) { throw BusinessRuntimeException.getInstance("关联任务不存在或已失效"); } if (!user.getMode().equals(conversation.getMode())) { throw BusinessRuntimeException.getInstance("当前出图模式与关联任务出图模式不符"); } AtomicBoolean flag = new AtomicBoolean(false); if (StringUtils.isNotBlank(conversation.getButtons())){ List messageButtons = Jsons.parseList(conversation.getButtons(), MessageButton.class); messageButtons.forEach(button -> { if (button.getCustomId().equals(customId)) { log.info("action customId:{}", customId); button.setStyle(3); if (button.getLabel().contains("Vary") || button.getCustomId().contains("::pan_") || button.getEmoji().equals("🔄") || button.getCustomId().contains("PromptAnalyzer:") || button.getCustomId().contains("PicReader::") || button.getCustomId().contains("::variation::") || button.getCustomId().contains("::CustomZoom::")) { flag.set(true); } } }); conversation.setButtons(Jsons.toJson(messageButtons)); } conversation.setBookmark(customId.contains("BOOKMARK")); conversationMapper.updateById(conversation); if (conversation.getBookmark()) { return conversation; } Map param = MapUtil.builder(new HashMap()) .put("taskId", taskId).put("state", user.getId()).put("customId", customId).build(); SubmitResult result = submit(user.getMode(),"action", param); if (flag.get()) { // 以上操作有弹窗确认,恢复次数 recoverUserLimit(user.getId(), user.getMode(),user.getMjRelaxNum()); } return saveConversation(user.getId(), user.getMode(), Long.parseLong(result.getResult()),result.getProperties(),"ACTION", StringUtils.EMPTY); } public SubmitResult submit(Integer mode,String action, Map param) throws Exception { if (mode == 1) { param.put("mode", "FAST"); } String url = (mode == 1 ? FAST_HOST : RELAX_HOST) + getActionUrl(action); String body = HttpRequest.post(url).body(Jsons.toJson(param)).header("Authorization", FAST_TOKEN).execute().body(); log.info("action body:{}", body); SubmitResult submitResult = Jsons.parseObject(body, SubmitResult.class); int code = submitResult.getCode(); if (code != 1 && code != 21 && code != 22) { if (code == 3) { throw BusinessRuntimeException.getInstance("账号不存在"); } if (code == 4) { throw BusinessRuntimeException.getInstance("图片重复"); } if (code == 24) { throw BusinessRuntimeException.getInstance("prompt包含敏感词"); } log.error("action:" + action + " error message:" + submitResult.getDescription()); throw BusinessRuntimeException.getInstance("队列已满,请稍后尝试"); } return submitResult; } private String getActionUrl(String action) { switch (action) { case "imagine": return "/mj/submit/imagine"; case "action": return "/mj/submit/action"; case "describe": return "/mj/submit/describe"; case "blend": return"/mj/submit/blend"; case "modal": return "/mj/submit/modal"; case "shorten": return "/mj/submit/shorten"; default: throw new IllegalArgumentException("Unknown action: " + action); } } @Override public List listConversationByIds(Integer mode, List ids) throws Exception { List list = new ArrayList<>(); listByIds(mode,ids).forEach(json -> { JSONObject jsons = JSONUtil.parseObj(json); MidjourneyUserConversation conversation = JSONUtil.toBean(jsons, MidjourneyUserConversation.class); conversation.setUserId(jsons.getLong("state")); conversation.setTaskId(jsons.getLong("id")); list.add(conversation); TASK_EXECUTOR.execute(() -> { try { sync(conversation); } catch (IOException e) { throw new RuntimeException(e); } }); }); return list; } public void sync(MidjourneyUserConversation conversation) throws IOException { if ((StringUtils.isNotBlank(conversation.getProgress()) && progress.contains(conversation.getProgress())) || status.contains(conversation.getStatus())) { log.info("同步任务:{},进度:{}",conversation.getTaskId(),conversation.getProgress()); MidjourneyUserConversation dbConversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, conversation.getTaskId()) .eq(MidjourneyUserConversation::getUserId, conversation.getUserId()).last("limit 1")); if (dbConversation != null) { try { if (StringUtils.isNotBlank(conversation.getImageUrl())) { conversation.setImageUrl(uploadPic(conversation.getImageUrl(), "conversation"+dbConversation.getId())); } } catch (IOException e) { log.error("上传图片失败",e); }finally { conversation.setId(dbConversation.getId()); conversationMapper.updateById(conversation); } } } } private static String uploadPic(String url, String prefix) throws IOException { byte[] body = HttpUtil.downloadBytes(url); Path tempFile = Files.createTempFile(prefix, ".png"); try (FileImageOutputStream imageOutput = new FileImageOutputStream(tempFile.toFile())) { imageOutput.write(body, 0, body.length); } Map paramMap = new HashMap<>(); paramMap.put("file", tempFile.toFile()); JSONObject result = JSONUtil.parseObj(HttpUtil.post("https://files.liuliangbang.vip/pic/ups", paramMap)); return result.getJSONObject("value").getJSONArray("saved").getJSONObject(0).getJSONObject("info").getStr("cdnUrl"); } private static List uploadBase64Pic(List base64Array) throws IOException { List list = new ArrayList<>(); for (String base64 : base64Array) { String fileExt = "jpeg"; if (base64.contains(";")){ fileExt = base64.split(";")[0].split("/")[1]; base64 = base64.split(",")[1]; } byte[] body = Base64.getDecoder().decode(base64); Path tempFile = Files.createTempFile(UUID.randomUUID().toString(), "." + fileExt); try (FileImageOutputStream imageOutput = new FileImageOutputStream(tempFile.toFile())) { imageOutput.write(body, 0, body.length); } Map paramMap = new HashMap<>(); paramMap.put("file", tempFile.toFile()); JSONObject result = JSONUtil.parseObj(HttpUtil.post("https://files.liuliangbang.vip/pic/ups", paramMap)); list.add(result.getJSONObject("value").getJSONArray("saved").getJSONObject(0).getJSONObject("info").getStr("cdnUrl")); } return list; } public JSONArray listByIds(Integer mode,List ids) throws Exception { Map param = MapUtil.builder(new HashMap()).put("ids", ids).build(); 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(); log.info("listByIds body:{}", body); return JSONUtil.parseArray(body); } @Override public MidjourneyUserConversation cancelConversation(MidjourneyUser user, Long id) { MidjourneyUserConversation conversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getId, id)); if (conversation == null) { throw BusinessRuntimeException.getInstance("会话不存在"); } if (!Objects.equals(conversation.getUserId(), user.getId())) { throw BusinessRuntimeException.getInstance("不是您的会话"); } conversation.setStatus("CANCEL"); conversationMapper.updateById(conversation); TASK_EXECUTOR.execute(() -> { cancel(user.getMode(),conversation.getTaskId()); }); return conversation; } public void cancel(Integer mode,Long taskId) { String body = HttpUtil.post(mode == 1 ? FAST_HOST : RELAX_HOST + "/mj/task/"+taskId+"/cancel",MapUtil.empty()); log.info("cancel body:{}", body); JSONObject jsonObject = JSONUtil.parseObj(body); if (jsonObject.getInt("code") != 1) { throw BusinessRuntimeException.getInstance("cancel error message:" + jsonObject.getStr("description")); } } }