MidjourneyController.java 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574
  1. package com.yhlxj.web.mirror;
  2. import cn.hutool.core.util.StrUtil;
  3. import com.baomidou.mybatisplus.core.toolkit.Wrappers;
  4. import com.cyksj.common.annotation.NoSubmit;
  5. import com.cyksj.common.exception.BusinessRuntimeException;
  6. import com.cyksj.common.util.IoKit;
  7. import com.cyksj.dto.Result;
  8. import com.cyksj.enums.GatewayResponse;
  9. import com.ejlchina.searcher.BeanSearcher;
  10. import com.ejlchina.searcher.SearchResult;
  11. import com.ejlchina.searcher.util.MapBuilder;
  12. import com.ejlchina.searcher.util.MapUtils;
  13. import com.yhlxj.dao.mapper.GroupsRelationMapper;
  14. import com.yhlxj.dao.mapper.MidjourneyPaintingPlazaMapper;
  15. import com.yhlxj.dao.mapper.midjourney.MidjourneyUserConversationMapper;
  16. import com.yhlxj.dao.mapper.midjourney.MidjourneyUserMapper;
  17. import com.yhlxj.dao.model.dto.*;
  18. import com.yhlxj.dao.model.entity.*;
  19. import com.yhlxj.dao.model.response.SubmitResult;
  20. import com.yhlxj.dao.model.views.MidjourneyPaintingUserView;
  21. import com.yhlxj.dao.model.views.MidjourneyUserPaintingDayRecordView;
  22. import com.yhlxj.dao.model.views.MidjourneyUserPaintingRecordView;
  23. import com.yhlxj.redis.RedisService;
  24. import com.yhlxj.service.midjourney.MidjourneyService;
  25. import com.yhlxj.service.midjourney.MidjourneyUserSettingsService;
  26. import lombok.RequiredArgsConstructor;
  27. import lombok.extern.slf4j.Slf4j;
  28. import org.springframework.beans.factory.annotation.Value;
  29. import org.springframework.scheduling.annotation.Scheduled;
  30. import org.springframework.web.bind.annotation.*;
  31. import javax.servlet.http.Cookie;
  32. import javax.servlet.http.HttpServletRequest;
  33. import javax.servlet.http.HttpServletResponse;
  34. import java.io.IOException;
  35. import java.io.InputStream;
  36. import java.nio.charset.StandardCharsets;
  37. import java.util.Date;
  38. import java.util.HashMap;
  39. import java.util.List;
  40. import java.util.Map;
  41. /** server/镜像服务/mj绘画
  42. * @author zwhui
  43. * @date 2024/4/23 10:31
  44. */
  45. @Slf4j
  46. @RequiredArgsConstructor
  47. @RestController
  48. @RequestMapping("/applets/midjourney")
  49. public class MidjourneyController {
  50. @Value("${midjourney.drawUrl}")
  51. private String midjourneyDrawUrl;
  52. private final HttpServletResponse response;
  53. private final MidjourneyUserMapper midjourneyUserMapper;
  54. private final HttpServletRequest request;
  55. private final MidjourneyService midjourneyService;
  56. private final BeanSearcher beanSearcher;
  57. private final RedisService redisService;
  58. private final GroupsRelationMapper groupsRelationMapper;
  59. private final MidjourneyPaintingPlazaMapper midjourneyPaintingPlazaMapper;
  60. private final MidjourneyUserConversationMapper conversationMapper;
  61. private final MidjourneyUserSettingsService midjourneyUserSettingsService;
  62. public MidjourneyUser getUser(){
  63. String userToken = request.getHeader("user-token");
  64. return midjourneyService.getUser(userToken);
  65. }
  66. @GetMapping("/midjourneyMirrorWithToken/{userToken}")
  67. public void midjourneyMirrorWithToken(@PathVariable String userToken) throws IOException {
  68. Cookie cookie = new Cookie("userToken", userToken);
  69. cookie.setMaxAge(31536000);
  70. cookie.setPath("/");
  71. response.addCookie(cookie);
  72. response.sendRedirect(midjourneyDrawUrl);
  73. }
  74. /**
  75. * 检查用户次数
  76. */
  77. public Long checkUserLimit(MidjourneyUser user,Integer mode){
  78. Long num = 0L;
  79. if (user.getExpireTime().before(new Date())){
  80. throw BusinessRuntimeException.getInstance("账号已过期");
  81. }
  82. if (mode == 1){
  83. num = redisService.decr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
  84. if (num < 0) {
  85. redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
  86. }else {
  87. int i = midjourneyUserMapper.decFastNum(user.getId(), 1, user.getMjFastNum());
  88. if (i == 0) {
  89. redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
  90. throw BusinessRuntimeException.getInstance("网络异常,请稍后再试");
  91. }
  92. }
  93. }
  94. if (mode == 2){
  95. Object relax = redisService.get(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId());
  96. if (relax != null){
  97. num = redisService.decr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId(), 1L);
  98. if (num < 0){
  99. redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId(), 1L);
  100. }else {
  101. int i = midjourneyUserMapper.decRelaxNum(user.getId(), 1, user.getMjRelaxNum());
  102. if (i == 0) {
  103. redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId(), 1L);
  104. throw BusinessRuntimeException.getInstance("网络异常,请稍后再试");
  105. }
  106. }
  107. }else {
  108. num = null;
  109. }
  110. }
  111. return num;
  112. }
  113. /**
  114. * 恢复次数
  115. */
  116. public void recoverUserLimit(Long id,Integer mode,Long num){
  117. if (mode == 1){
  118. redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + id, 1L);
  119. midjourneyUserMapper.incrFastNum(id, 1);
  120. }
  121. if (mode == 2){
  122. if (num != null){
  123. redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + id, 1L);
  124. midjourneyUserMapper.incrRelaxNum(id, 1);
  125. }
  126. }
  127. }
  128. ///**
  129. // * 同步数据库
  130. // */
  131. //public void syncUser(Long id,Integer mode,Long num){
  132. // log.info("同步次数 id:{},mode:{},num:{}",id,mode,num);
  133. // if (num == null){
  134. // return;
  135. // }
  136. // LambdaUpdateWrapper<MidjourneyUser> wrapper = Wrappers.lambdaUpdate(MidjourneyUser.class)
  137. // .eq(MidjourneyUser::getId, id);
  138. // if (mode == 1){
  139. // wrapper.set(MidjourneyUser::getMjFastNum, num);
  140. // }
  141. // if (mode == 2){
  142. // wrapper.set(MidjourneyUser::getMjRelaxNum, num);
  143. // }
  144. // midjourneyUserMapper.update(null, wrapper);
  145. //}
  146. /**
  147. * 查询用户信息
  148. */
  149. @GetMapping("/whoami")
  150. public Result<MidjourneyUser> queryUser(){
  151. MidjourneyUser user = getUser();
  152. return GatewayResponse.SUCCESS.newBuilder().toResult(user);
  153. }
  154. /**
  155. * 提交Imagine任务
  156. */
  157. @PostMapping("/submit/imagine")
  158. @NoSubmit
  159. public Result<SubmitResult> submitImagine(@RequestBody SubmitImagineDTO submitImagineDTO) {
  160. log.info("提交Imagine任务,提示:{}",submitImagineDTO.getPrompt());
  161. MidjourneyUser user = getUser();
  162. Integer mode = submitImagineDTO.getMode();
  163. Long num = checkUserLimit(user, mode);
  164. if (num != null && num < 0){
  165. throw BusinessRuntimeException.getInstance("次数已用完");
  166. }
  167. SubmitResult conversation;
  168. try {
  169. MidjourneyUserSettings settings = user.getSettings();
  170. if(settings != null){
  171. String prompt = settings.getPrompt();
  172. submitImagineDTO.setPrompt(appendDefaults(submitImagineDTO.getPrompt(), prompt));
  173. }
  174. conversation = midjourneyService.submitImagine(user, mode, submitImagineDTO.getPrompt(), submitImagineDTO.getBotType(), submitImagineDTO.getBase64Array());
  175. } catch (Exception e) {
  176. recoverUserLimit(user.getId(), mode,num);
  177. throw BusinessRuntimeException.getInstance(e.getMessage());
  178. }
  179. //syncUser(user.getId(), user.getMode(), num);
  180. return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
  181. }
  182. /**
  183. * 提交Describe任务
  184. */
  185. @PostMapping("/submit/describe")
  186. @NoSubmit
  187. public Result<SubmitResult> submitDescribe(@RequestBody SubmitDescribeDTO submitDescribeDTO) {
  188. log.info("提交Describe任务");
  189. MidjourneyUser user = getUser();
  190. Integer mode = submitDescribeDTO.getMode();
  191. Long num = checkUserLimit(user,mode);
  192. if (num != null && num < 0){
  193. throw BusinessRuntimeException.getInstance("次数已用完");
  194. }
  195. SubmitResult conversation;
  196. try {
  197. conversation = midjourneyService.submitDescribe(user, mode,submitDescribeDTO.getBotType(),submitDescribeDTO.getBase64());
  198. } catch (Exception e) {
  199. recoverUserLimit(user.getId(), mode,num);
  200. throw BusinessRuntimeException.getInstance(e.getMessage());
  201. }
  202. //syncUser(user.getId(), user.getMode(), num);
  203. return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
  204. }
  205. /**
  206. * 提交Blend任务
  207. */
  208. @PostMapping("/submit/blend")
  209. @NoSubmit
  210. public Result<SubmitResult> submitBlend(@RequestBody SubmitBlendDTO submitBlendDTO) {
  211. log.info("提交Blend任务 dimensions:{},base64数组长度:{}", submitBlendDTO.getDimensions(), submitBlendDTO.getBase64Array().size());
  212. Integer mode = submitBlendDTO.getMode();
  213. MidjourneyUser user = getUser();
  214. Long num = checkUserLimit(user,mode);
  215. if (num != null && num < 0){
  216. throw BusinessRuntimeException.getInstance("次数已用完");
  217. }
  218. SubmitResult conversation;
  219. try {
  220. conversation = midjourneyService.submitBlend(user, mode,submitBlendDTO.getDimensions(), submitBlendDTO.getBotType(), submitBlendDTO.getBase64Array());
  221. } catch (Exception e) {
  222. recoverUserLimit(user.getId(), mode,num);
  223. throw BusinessRuntimeException.getInstance(e.getMessage());
  224. }
  225. //syncUser(user.getId(), user.getMode(), num);
  226. return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
  227. }
  228. /**
  229. * 提交Modal任务
  230. */
  231. @PostMapping("/submit/modal")
  232. @NoSubmit
  233. public Result<SubmitResult> submitModal(@RequestBody SubmitModalDTO submitModalDTO) {
  234. log.info("提交Modal任务,taskId:{},提示:{}",submitModalDTO.getTaskId(),submitModalDTO.getPrompt());
  235. Integer mode = submitModalDTO.getMode();
  236. MidjourneyUser user = getUser();
  237. Long num = checkUserLimit(user,mode);
  238. if (num != null && num < 0){
  239. throw BusinessRuntimeException.getInstance("次数已用完");
  240. }
  241. SubmitResult conversation;
  242. try {
  243. conversation = midjourneyService.submitModal(user, mode, submitModalDTO.getTaskId(), submitModalDTO.getPrompt(), submitModalDTO.getMaskBase64());
  244. } catch (Exception e) {
  245. recoverUserLimit(user.getId(), mode,num);
  246. throw BusinessRuntimeException.getInstance(e.getMessage());
  247. }
  248. //syncUser(user.getId(), user.getMode(), num);
  249. return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
  250. }
  251. /**
  252. * 提交Shorten任务
  253. */
  254. @PostMapping("/submit/shorten")
  255. @NoSubmit
  256. public Result<SubmitResult> submitShorten(@RequestBody SubmitShortenDTO submitShortenDTO) {
  257. log.info("提交Shorten任务 提示词:{}",submitShortenDTO.getPrompt());
  258. Integer mode = submitShortenDTO.getMode();
  259. MidjourneyUser user = getUser();
  260. Long num = checkUserLimit(user,mode);
  261. if (num != null && num < 0){
  262. throw BusinessRuntimeException.getInstance("次数已用完");
  263. }
  264. SubmitResult conversation;
  265. try {
  266. conversation = midjourneyService.submitShorten(user, mode, submitShortenDTO.getBotType(), submitShortenDTO.getPrompt());
  267. } catch (Exception e) {
  268. recoverUserLimit(user.getId(),mode,num);
  269. throw BusinessRuntimeException.getInstance(e.getMessage());
  270. }
  271. //syncUser(user.getId(), user.getMode(), num);
  272. return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
  273. }
  274. /**
  275. * 执行动作
  276. */
  277. @PostMapping("/submit/action")
  278. @NoSubmit
  279. public Result<SubmitResult> action(@RequestBody SubmitActionDTO actionDTO) {
  280. log.info("任务id:{},执行动作:{}",actionDTO.getTaskId(),actionDTO.getCustomId());
  281. MidjourneyUser user = getUser();
  282. MidjourneyUserConversation midjourneyUserConversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, actionDTO.getTaskId()).last("limit 1"));
  283. if (midjourneyUserConversation == null) {
  284. throw BusinessRuntimeException.getInstance("关联任务不存在或已失效");
  285. }
  286. Integer mode = midjourneyUserConversation.getMode();
  287. Long num = 0L;
  288. if (!actionDTO.getCustomId().contains("BOOKMARK")) {
  289. num = checkUserLimit(user,mode);
  290. if (num != null && num < 0) {
  291. throw BusinessRuntimeException.getInstance("次数已用完");
  292. }
  293. }
  294. SubmitResult conversation;
  295. try {
  296. conversation = midjourneyService.submitAction(user,midjourneyUserConversation,mode, actionDTO.getTaskId(), actionDTO.getCustomId(),num, actionDTO.getBotType());
  297. } catch (Exception e) {
  298. recoverUserLimit(user.getId(), mode,num);
  299. throw BusinessRuntimeException.getInstance(e.getMessage());
  300. }
  301. return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
  302. }
  303. /**
  304. * 执行动作
  305. */
  306. @PostMapping("/submit/seed")
  307. @NoSubmit
  308. public Result<Object> seed(@RequestBody SubmitSeedDTO seedDTO) {
  309. log.info("任务id:{},执行动作:{}",seedDTO.getTaskId(), "seed");
  310. MidjourneyUser user = getUser();
  311. Integer mode = seedDTO.getMode();
  312. Long num = checkUserLimit(user,mode);
  313. if (num != null && num < 0) {
  314. throw BusinessRuntimeException.getInstance("次数已用完");
  315. }
  316. Object conversation;
  317. try {
  318. conversation = midjourneyService.submitSeed(user, mode, seedDTO.getTaskId(), seedDTO.getBotType());
  319. } catch (Exception e) {
  320. recoverUserLimit(user.getId(), mode,num);
  321. throw BusinessRuntimeException.getInstance(e.getMessage());
  322. }
  323. return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
  324. }
  325. /**
  326. * 查询会话列表
  327. */
  328. @GetMapping("/conversation/list")
  329. public Result<SearchResult<MidjourneyUserConversation>> conversationList(){
  330. MidjourneyUser user = getUser();
  331. log.info("查询会话列表 userId:{}",user.getId());
  332. MapBuilder builder = MapUtils.flatBuilder(request.getParameterMap());
  333. return GatewayResponse.SUCCESS.newBuilder().toResult(beanSearcher.search(MidjourneyUserConversation.class,
  334. builder.field(MidjourneyUserConversation::getUserId, user.getId())
  335. .selectExclude(MidjourneyUserConversation::getChannelId,MidjourneyUserConversation::getInstanceId)
  336. .orderBy(MidjourneyUserConversation::getId).desc().build()));
  337. }
  338. /**
  339. * 根据ids查询会话
  340. */
  341. @GetMapping("/conversation/listByIds")
  342. public Result<List<MidjourneyUserConversation>> conversationListByIds(Integer mode,@RequestParam(value = "ids" ) List<Long> ids) throws Exception {
  343. MidjourneyUser user = getUser();
  344. log.info("根据ids查询会话 userId:{}",user.getId());
  345. return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.listConversationByIds(mode,ids));
  346. }
  347. /**
  348. * 取消任务
  349. */
  350. @PostMapping("/conversation/{id}/cancel")
  351. @NoSubmit
  352. public Result<Void> conversationCancel(@PathVariable("id") Long id){
  353. MidjourneyUser user = getUser();
  354. log.info("取消任务 id:{}",id);
  355. midjourneyService.cancelConversation(user,id);
  356. return GatewayResponse.SUCCESS.newBuilder().toResult();
  357. }
  358. /**
  359. * 修改midjourney账号模式
  360. */
  361. @PutMapping("/put/midjourney/user/{mode}")
  362. public Result<String> putMidjourneyUserMode(@PathVariable Integer mode, Long userId) {
  363. log.info("修改midjourney账号模式 userId:{},mode:{}", userId, mode);
  364. MidjourneyUser midjourneyUser = midjourneyUserMapper.selectById(userId);
  365. if (midjourneyUser == null) {
  366. throw BusinessRuntimeException.getInstance("参数错误");
  367. }
  368. //midjourneyUser.setMode(mode);
  369. midjourneyUserMapper.updateById(midjourneyUser);
  370. //redisService.del(RedisService.key.MIDJOURNEY_USER.getName() + midjourneyUser.getUserToken());
  371. return GatewayResponse.SUCCESS.newBuilder().toResult();
  372. }
  373. /**
  374. * midjourney 回调
  375. */
  376. @PostMapping("/notifyHook")
  377. public void notifyHook(HttpServletRequest request) throws Exception {
  378. InputStream inputStream = request.getInputStream();
  379. byte[] bytes = IoKit.toBytes(inputStream);
  380. String json = new String(bytes, StandardCharsets.UTF_8);
  381. midjourneyService.notifyHook(json);
  382. }
  383. /**
  384. * mj绘画广场
  385. */
  386. @GetMapping("/get/painting")
  387. public Result<SearchResult<MidjourneyPaintingUserView>> getPaintingCollect() {
  388. SearchResult<MidjourneyPaintingUserView> search = beanSearcher.search(MidjourneyPaintingUserView.class, MapUtils.flatBuilder(request.getParameterMap())
  389. .build());
  390. return GatewayResponse.SUCCESS.newBuilder().toResult(search);
  391. }
  392. /**
  393. * 复制mj绘画广场 提示词
  394. */
  395. @PostMapping("/painting/copy/{id}")
  396. public Result<String> resetPaintingCopyNum(@PathVariable Long id) {
  397. MidjourneyPaintingPlaza midjourneyPaintingPlaza = midjourneyPaintingPlazaMapper.selectById(id);
  398. if (midjourneyPaintingPlaza != null) {
  399. midjourneyPaintingPlazaMapper.updatePaintingCopyNum(id);
  400. }
  401. return GatewayResponse.SUCCESS.newBuilder().toResult();
  402. }
  403. /**
  404. * 用户绘画记录
  405. */
  406. @GetMapping("/get/painting/record")
  407. public Result<SearchResult<MidjourneyUserPaintingDayRecordView>> getPaintingRecord(String prompt, Boolean bookmark) {
  408. MidjourneyUser user = getUser();
  409. String condition = String.format(" and mc.user_id = %s", user.getId());
  410. if (StrUtil.isNotBlank(prompt)) {
  411. condition += String.format(" and mc.prompt like '%%%s%%'", prompt);
  412. }
  413. if (bookmark != null) {
  414. condition += String.format(" and mc.bookmark is %s", bookmark);
  415. }
  416. SearchResult<MidjourneyUserPaintingDayRecordView> search = beanSearcher.search(MidjourneyUserPaintingDayRecordView.class, MapUtils.flatBuilder(request.getParameterMap())
  417. .put("condition", condition)
  418. .orderBy(MidjourneyUserPaintingRecordView::getCreatedTime).desc()
  419. .build());
  420. String dateSql;
  421. for (MidjourneyUserPaintingDayRecordView e : search.getDataList()) {
  422. String date = e.getDate();
  423. dateSql = String.format(" and DATE_FORMAT(mc.created_time,'%%Y-%%m-%%d') = '%s'", date);
  424. List<MidjourneyUserPaintingRecordView> list = beanSearcher.searchAll(MidjourneyUserPaintingRecordView.class, MapUtils.builder()
  425. .put("condition", condition)
  426. .put("date", dateSql)
  427. .orderBy(MidjourneyUserPaintingRecordView::getCreatedTime).desc()
  428. .build());
  429. e.setList(list);
  430. }
  431. return GatewayResponse.SUCCESS.newBuilder().toResult(search);
  432. }
  433. /**
  434. * 获取settings
  435. */
  436. @GetMapping("/get/settings")
  437. public Result<MidjourneyUserSettings> getSettings(){
  438. MidjourneyUser user = getUser();
  439. return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyUserSettingsService.getById(user.getSettings().getId()));
  440. }
  441. /**
  442. * 修改settings
  443. */
  444. @PostMapping("/put/settings")
  445. public Result<MidjourneyUserSettings> upSettings(@RequestBody MidjourneyUserSettings settings){
  446. MidjourneyUser user = getUser();
  447. if (!user.getId().equals(settings.getUserId())) {
  448. throw BusinessRuntimeException.getInstance("参数错误");
  449. }
  450. settings.setId(user.getSettings().getId());
  451. midjourneyUserSettingsService.updateById(settings);
  452. return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyUserSettingsService.getById(settings.getId()));
  453. }
  454. public static String appendDefaults(String userPrompt, String defaultOptions) {
  455. // 解析默认选项
  456. Map<String, String> defaultOptionsMap = parseOptions(defaultOptions);
  457. // 修正用户输入中可能缺少空格的选项
  458. userPrompt = fixMissingSpaces(userPrompt);
  459. // 检查并追加缺失的选项
  460. for (Map.Entry<String, String> entry : defaultOptionsMap.entrySet()) {
  461. if (!containsOption(userPrompt, entry.getKey())) {
  462. userPrompt += " " + entry.getValue();
  463. }
  464. }
  465. // 单独处理 --s 和 --stylize 的情况
  466. if (!containsOption(userPrompt, "--s") && !containsOption(userPrompt, "--stylize")) {
  467. if (defaultOptionsMap.containsKey("--s")) {
  468. userPrompt += " " + defaultOptionsMap.get("--s");
  469. }
  470. }
  471. return userPrompt;
  472. }
  473. public static String fixMissingSpaces(String userPrompt) {
  474. // 定义需要修正的正则表达式和替换格式
  475. String[][] patterns = {
  476. {"--v(\\d)", "--v $1"},
  477. {"--niji(\\d)", "--niji $1"},
  478. {"--s(\\d+)", "--s $1"},
  479. {"--stylize(\\d+)", "--stylize $1"},
  480. {"--style(\\w+)", "--style $1"},
  481. {"--ar(\\d+)[::](\\d+)", "--ar $1:$2"},
  482. {"--ar (\\d+):(\\d+)", "--ar $1:$2"}
  483. };
  484. // 修正用户输入
  485. for (String[] pattern : patterns) {
  486. userPrompt = userPrompt.replaceAll(pattern[0], pattern[1]);
  487. }
  488. return userPrompt;
  489. }
  490. public static boolean containsOption(String userPrompt, String option) {
  491. return userPrompt.contains(option + " ") || userPrompt.matches(".*" + option + "\\d.*");
  492. }
  493. public static Map<String, String> parseOptions(String options) {
  494. Map<String, String> optionsMap = new HashMap<>();
  495. String[] parts = options.split("--");
  496. for (String part : parts) {
  497. if (!part.trim().isEmpty()) {
  498. String[] keyValue = part.trim().split(" ", 2);
  499. if (keyValue.length == 2) {
  500. optionsMap.put("--" + keyValue[0], "--" + keyValue[0] + " " + keyValue[1]);
  501. }
  502. }
  503. }
  504. return optionsMap;
  505. }
  506. /**
  507. * 上传文件到discord
  508. */
  509. @PostMapping("/upload")
  510. public Result<SubmitResult> upload(@RequestBody SubmitUploadDTO uploadDTO) throws Exception {
  511. log.info("上传文件到discord");
  512. return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.uploadFile(uploadDTO));
  513. }
  514. }