| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574 |
- package com.yhlxj.web.mirror;
- import cn.hutool.core.util.StrUtil;
- import com.baomidou.mybatisplus.core.toolkit.Wrappers;
- import com.cyksj.common.annotation.NoSubmit;
- import com.cyksj.common.exception.BusinessRuntimeException;
- import com.cyksj.common.util.IoKit;
- import com.cyksj.dto.Result;
- import com.cyksj.enums.GatewayResponse;
- import com.ejlchina.searcher.BeanSearcher;
- import com.ejlchina.searcher.SearchResult;
- import com.ejlchina.searcher.util.MapBuilder;
- import com.ejlchina.searcher.util.MapUtils;
- import com.yhlxj.dao.mapper.GroupsRelationMapper;
- import com.yhlxj.dao.mapper.MidjourneyPaintingPlazaMapper;
- import com.yhlxj.dao.mapper.midjourney.MidjourneyUserConversationMapper;
- import com.yhlxj.dao.mapper.midjourney.MidjourneyUserMapper;
- import com.yhlxj.dao.model.dto.*;
- import com.yhlxj.dao.model.entity.*;
- import com.yhlxj.dao.model.response.SubmitResult;
- import com.yhlxj.dao.model.views.MidjourneyPaintingUserView;
- import com.yhlxj.dao.model.views.MidjourneyUserPaintingDayRecordView;
- import com.yhlxj.dao.model.views.MidjourneyUserPaintingRecordView;
- import com.yhlxj.redis.RedisService;
- import com.yhlxj.service.midjourney.MidjourneyService;
- import com.yhlxj.service.midjourney.MidjourneyUserSettingsService;
- import lombok.RequiredArgsConstructor;
- import lombok.extern.slf4j.Slf4j;
- import org.springframework.beans.factory.annotation.Value;
- import org.springframework.scheduling.annotation.Scheduled;
- import org.springframework.web.bind.annotation.*;
- import javax.servlet.http.Cookie;
- import javax.servlet.http.HttpServletRequest;
- import javax.servlet.http.HttpServletResponse;
- import java.io.IOException;
- import java.io.InputStream;
- import java.nio.charset.StandardCharsets;
- import java.util.Date;
- import java.util.HashMap;
- import java.util.List;
- import java.util.Map;
- /** server/镜像服务/mj绘画
- * @author zwhui
- * @date 2024/4/23 10:31
- */
- @Slf4j
- @RequiredArgsConstructor
- @RestController
- @RequestMapping("/applets/midjourney")
- public class MidjourneyController {
- @Value("${midjourney.drawUrl}")
- private String midjourneyDrawUrl;
- private final HttpServletResponse response;
- private final MidjourneyUserMapper midjourneyUserMapper;
- private final HttpServletRequest request;
- private final MidjourneyService midjourneyService;
- private final BeanSearcher beanSearcher;
- private final RedisService redisService;
- private final GroupsRelationMapper groupsRelationMapper;
- private final MidjourneyPaintingPlazaMapper midjourneyPaintingPlazaMapper;
- private final MidjourneyUserConversationMapper conversationMapper;
- private final MidjourneyUserSettingsService midjourneyUserSettingsService;
- public MidjourneyUser getUser(){
- String userToken = request.getHeader("user-token");
- return midjourneyService.getUser(userToken);
- }
- @GetMapping("/midjourneyMirrorWithToken/{userToken}")
- public void midjourneyMirrorWithToken(@PathVariable String userToken) throws IOException {
- Cookie cookie = new Cookie("userToken", userToken);
- cookie.setMaxAge(31536000);
- cookie.setPath("/");
- response.addCookie(cookie);
- response.sendRedirect(midjourneyDrawUrl);
- }
- /**
- * 检查用户次数
- */
- public Long checkUserLimit(MidjourneyUser user,Integer mode){
- Long num = 0L;
- if (user.getExpireTime().before(new Date())){
- throw BusinessRuntimeException.getInstance("账号已过期");
- }
- if (mode == 1){
- num = redisService.decr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
- if (num < 0) {
- redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
- }else {
- int i = midjourneyUserMapper.decFastNum(user.getId(), 1, user.getMjFastNum());
- if (i == 0) {
- redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
- throw BusinessRuntimeException.getInstance("网络异常,请稍后再试");
- }
- }
- }
- if (mode == 2){
- Object relax = redisService.get(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId());
- if (relax != null){
- num = redisService.decr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId(), 1L);
- if (num < 0){
- redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId(), 1L);
- }else {
- int i = midjourneyUserMapper.decRelaxNum(user.getId(), 1, user.getMjRelaxNum());
- if (i == 0) {
- redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId(), 1L);
- throw BusinessRuntimeException.getInstance("网络异常,请稍后再试");
- }
- }
- }else {
- num = null;
- }
- }
- return num;
- }
- /**
- * 恢复次数
- */
- public void recoverUserLimit(Long id,Integer mode,Long num){
- if (mode == 1){
- redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + id, 1L);
- midjourneyUserMapper.incrFastNum(id, 1);
- }
- if (mode == 2){
- if (num != null){
- redisService.incr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + id, 1L);
- midjourneyUserMapper.incrRelaxNum(id, 1);
- }
- }
- }
- ///**
- // * 同步数据库
- // */
- //public void syncUser(Long id,Integer mode,Long num){
- // log.info("同步次数 id:{},mode:{},num:{}",id,mode,num);
- // if (num == null){
- // return;
- // }
- // LambdaUpdateWrapper<MidjourneyUser> wrapper = Wrappers.lambdaUpdate(MidjourneyUser.class)
- // .eq(MidjourneyUser::getId, id);
- // if (mode == 1){
- // wrapper.set(MidjourneyUser::getMjFastNum, num);
- // }
- // if (mode == 2){
- // wrapper.set(MidjourneyUser::getMjRelaxNum, num);
- // }
- // midjourneyUserMapper.update(null, wrapper);
- //}
- /**
- * 查询用户信息
- */
- @GetMapping("/whoami")
- public Result<MidjourneyUser> queryUser(){
- MidjourneyUser user = getUser();
- return GatewayResponse.SUCCESS.newBuilder().toResult(user);
- }
- /**
- * 提交Imagine任务
- */
- @PostMapping("/submit/imagine")
- @NoSubmit
- public Result<SubmitResult> submitImagine(@RequestBody SubmitImagineDTO submitImagineDTO) {
- log.info("提交Imagine任务,提示:{}",submitImagineDTO.getPrompt());
- MidjourneyUser user = getUser();
- Integer mode = submitImagineDTO.getMode();
- Long num = checkUserLimit(user, mode);
- if (num != null && num < 0){
- throw BusinessRuntimeException.getInstance("次数已用完");
- }
- SubmitResult conversation;
- try {
- MidjourneyUserSettings settings = user.getSettings();
- if(settings != null){
- String prompt = settings.getPrompt();
- submitImagineDTO.setPrompt(appendDefaults(submitImagineDTO.getPrompt(), prompt));
- }
- conversation = midjourneyService.submitImagine(user, mode, submitImagineDTO.getPrompt(), submitImagineDTO.getBotType(), submitImagineDTO.getBase64Array());
- } catch (Exception e) {
- recoverUserLimit(user.getId(), mode,num);
- throw BusinessRuntimeException.getInstance(e.getMessage());
- }
- //syncUser(user.getId(), user.getMode(), num);
- return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
- }
- /**
- * 提交Describe任务
- */
- @PostMapping("/submit/describe")
- @NoSubmit
- public Result<SubmitResult> submitDescribe(@RequestBody SubmitDescribeDTO submitDescribeDTO) {
- log.info("提交Describe任务");
- MidjourneyUser user = getUser();
- Integer mode = submitDescribeDTO.getMode();
- Long num = checkUserLimit(user,mode);
- if (num != null && num < 0){
- throw BusinessRuntimeException.getInstance("次数已用完");
- }
- SubmitResult conversation;
- try {
- conversation = midjourneyService.submitDescribe(user, mode,submitDescribeDTO.getBotType(),submitDescribeDTO.getBase64());
- } catch (Exception e) {
- recoverUserLimit(user.getId(), mode,num);
- throw BusinessRuntimeException.getInstance(e.getMessage());
- }
- //syncUser(user.getId(), user.getMode(), num);
- return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
- }
- /**
- * 提交Blend任务
- */
- @PostMapping("/submit/blend")
- @NoSubmit
- public Result<SubmitResult> submitBlend(@RequestBody SubmitBlendDTO submitBlendDTO) {
- log.info("提交Blend任务 dimensions:{},base64数组长度:{}", submitBlendDTO.getDimensions(), submitBlendDTO.getBase64Array().size());
- Integer mode = submitBlendDTO.getMode();
- MidjourneyUser user = getUser();
- Long num = checkUserLimit(user,mode);
- if (num != null && num < 0){
- throw BusinessRuntimeException.getInstance("次数已用完");
- }
- SubmitResult conversation;
- try {
- conversation = midjourneyService.submitBlend(user, mode,submitBlendDTO.getDimensions(), submitBlendDTO.getBotType(), submitBlendDTO.getBase64Array());
- } catch (Exception e) {
- recoverUserLimit(user.getId(), mode,num);
- throw BusinessRuntimeException.getInstance(e.getMessage());
- }
- //syncUser(user.getId(), user.getMode(), num);
- return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
- }
- /**
- * 提交Modal任务
- */
- @PostMapping("/submit/modal")
- @NoSubmit
- public Result<SubmitResult> submitModal(@RequestBody SubmitModalDTO submitModalDTO) {
- log.info("提交Modal任务,taskId:{},提示:{}",submitModalDTO.getTaskId(),submitModalDTO.getPrompt());
- Integer mode = submitModalDTO.getMode();
- MidjourneyUser user = getUser();
- Long num = checkUserLimit(user,mode);
- if (num != null && num < 0){
- throw BusinessRuntimeException.getInstance("次数已用完");
- }
- SubmitResult conversation;
- try {
- conversation = midjourneyService.submitModal(user, mode, submitModalDTO.getTaskId(), submitModalDTO.getPrompt(), submitModalDTO.getMaskBase64());
- } catch (Exception e) {
- recoverUserLimit(user.getId(), mode,num);
- throw BusinessRuntimeException.getInstance(e.getMessage());
- }
- //syncUser(user.getId(), user.getMode(), num);
- return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
- }
- /**
- * 提交Shorten任务
- */
- @PostMapping("/submit/shorten")
- @NoSubmit
- public Result<SubmitResult> submitShorten(@RequestBody SubmitShortenDTO submitShortenDTO) {
- log.info("提交Shorten任务 提示词:{}",submitShortenDTO.getPrompt());
- Integer mode = submitShortenDTO.getMode();
- MidjourneyUser user = getUser();
- Long num = checkUserLimit(user,mode);
- if (num != null && num < 0){
- throw BusinessRuntimeException.getInstance("次数已用完");
- }
- SubmitResult conversation;
- try {
- conversation = midjourneyService.submitShorten(user, mode, submitShortenDTO.getBotType(), submitShortenDTO.getPrompt());
- } catch (Exception e) {
- recoverUserLimit(user.getId(),mode,num);
- throw BusinessRuntimeException.getInstance(e.getMessage());
- }
- //syncUser(user.getId(), user.getMode(), num);
- return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
- }
- /**
- * 执行动作
- */
- @PostMapping("/submit/action")
- @NoSubmit
- public Result<SubmitResult> action(@RequestBody SubmitActionDTO actionDTO) {
- log.info("任务id:{},执行动作:{}",actionDTO.getTaskId(),actionDTO.getCustomId());
- MidjourneyUser user = getUser();
- MidjourneyUserConversation midjourneyUserConversation = conversationMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUserConversation.class).eq(MidjourneyUserConversation::getTaskId, actionDTO.getTaskId()).last("limit 1"));
- if (midjourneyUserConversation == null) {
- throw BusinessRuntimeException.getInstance("关联任务不存在或已失效");
- }
- Integer mode = midjourneyUserConversation.getMode();
- Long num = 0L;
- if (!actionDTO.getCustomId().contains("BOOKMARK")) {
- num = checkUserLimit(user,mode);
- if (num != null && num < 0) {
- throw BusinessRuntimeException.getInstance("次数已用完");
- }
- }
- SubmitResult conversation;
- try {
- conversation = midjourneyService.submitAction(user,midjourneyUserConversation,mode, actionDTO.getTaskId(), actionDTO.getCustomId(),num, actionDTO.getBotType());
- } catch (Exception e) {
- recoverUserLimit(user.getId(), mode,num);
- throw BusinessRuntimeException.getInstance(e.getMessage());
- }
- return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
- }
- /**
- * 执行动作
- */
- @PostMapping("/submit/seed")
- @NoSubmit
- public Result<Object> seed(@RequestBody SubmitSeedDTO seedDTO) {
- log.info("任务id:{},执行动作:{}",seedDTO.getTaskId(), "seed");
- MidjourneyUser user = getUser();
- Integer mode = seedDTO.getMode();
- Long num = checkUserLimit(user,mode);
- if (num != null && num < 0) {
- throw BusinessRuntimeException.getInstance("次数已用完");
- }
- Object conversation;
- try {
- conversation = midjourneyService.submitSeed(user, mode, seedDTO.getTaskId(), seedDTO.getBotType());
- } catch (Exception e) {
- recoverUserLimit(user.getId(), mode,num);
- throw BusinessRuntimeException.getInstance(e.getMessage());
- }
- return GatewayResponse.SUCCESS.newBuilder().toResult(conversation);
- }
- /**
- * 查询会话列表
- */
- @GetMapping("/conversation/list")
- public Result<SearchResult<MidjourneyUserConversation>> conversationList(){
- MidjourneyUser user = getUser();
- log.info("查询会话列表 userId:{}",user.getId());
- MapBuilder builder = MapUtils.flatBuilder(request.getParameterMap());
- return GatewayResponse.SUCCESS.newBuilder().toResult(beanSearcher.search(MidjourneyUserConversation.class,
- builder.field(MidjourneyUserConversation::getUserId, user.getId())
- .selectExclude(MidjourneyUserConversation::getChannelId,MidjourneyUserConversation::getInstanceId)
- .orderBy(MidjourneyUserConversation::getId).desc().build()));
- }
- /**
- * 根据ids查询会话
- */
- @GetMapping("/conversation/listByIds")
- public Result<List<MidjourneyUserConversation>> conversationListByIds(Integer mode,@RequestParam(value = "ids" ) List<Long> ids) throws Exception {
- MidjourneyUser user = getUser();
- log.info("根据ids查询会话 userId:{}",user.getId());
- return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.listConversationByIds(mode,ids));
- }
- /**
- * 取消任务
- */
- @PostMapping("/conversation/{id}/cancel")
- @NoSubmit
- public Result<Void> conversationCancel(@PathVariable("id") Long id){
- MidjourneyUser user = getUser();
- log.info("取消任务 id:{}",id);
- midjourneyService.cancelConversation(user,id);
- return GatewayResponse.SUCCESS.newBuilder().toResult();
- }
- /**
- * 修改midjourney账号模式
- */
- @PutMapping("/put/midjourney/user/{mode}")
- public Result<String> putMidjourneyUserMode(@PathVariable Integer mode, Long userId) {
- log.info("修改midjourney账号模式 userId:{},mode:{}", userId, mode);
- MidjourneyUser midjourneyUser = midjourneyUserMapper.selectById(userId);
- if (midjourneyUser == null) {
- throw BusinessRuntimeException.getInstance("参数错误");
- }
- //midjourneyUser.setMode(mode);
- midjourneyUserMapper.updateById(midjourneyUser);
- //redisService.del(RedisService.key.MIDJOURNEY_USER.getName() + midjourneyUser.getUserToken());
- return GatewayResponse.SUCCESS.newBuilder().toResult();
- }
- /**
- * midjourney 回调
- */
- @PostMapping("/notifyHook")
- public void notifyHook(HttpServletRequest request) throws Exception {
- InputStream inputStream = request.getInputStream();
- byte[] bytes = IoKit.toBytes(inputStream);
- String json = new String(bytes, StandardCharsets.UTF_8);
- midjourneyService.notifyHook(json);
- }
- /**
- * mj绘画广场
- */
- @GetMapping("/get/painting")
- public Result<SearchResult<MidjourneyPaintingUserView>> getPaintingCollect() {
- SearchResult<MidjourneyPaintingUserView> search = beanSearcher.search(MidjourneyPaintingUserView.class, MapUtils.flatBuilder(request.getParameterMap())
- .build());
- return GatewayResponse.SUCCESS.newBuilder().toResult(search);
- }
- /**
- * 复制mj绘画广场 提示词
- */
- @PostMapping("/painting/copy/{id}")
- public Result<String> resetPaintingCopyNum(@PathVariable Long id) {
- MidjourneyPaintingPlaza midjourneyPaintingPlaza = midjourneyPaintingPlazaMapper.selectById(id);
- if (midjourneyPaintingPlaza != null) {
- midjourneyPaintingPlazaMapper.updatePaintingCopyNum(id);
- }
- return GatewayResponse.SUCCESS.newBuilder().toResult();
- }
- /**
- * 用户绘画记录
- */
- @GetMapping("/get/painting/record")
- public Result<SearchResult<MidjourneyUserPaintingDayRecordView>> getPaintingRecord(String prompt, Boolean bookmark) {
- MidjourneyUser user = getUser();
- String condition = String.format(" and mc.user_id = %s", user.getId());
- if (StrUtil.isNotBlank(prompt)) {
- condition += String.format(" and mc.prompt like '%%%s%%'", prompt);
- }
- if (bookmark != null) {
- condition += String.format(" and mc.bookmark is %s", bookmark);
- }
- SearchResult<MidjourneyUserPaintingDayRecordView> search = beanSearcher.search(MidjourneyUserPaintingDayRecordView.class, MapUtils.flatBuilder(request.getParameterMap())
- .put("condition", condition)
- .orderBy(MidjourneyUserPaintingRecordView::getCreatedTime).desc()
- .build());
- String dateSql;
- for (MidjourneyUserPaintingDayRecordView e : search.getDataList()) {
- String date = e.getDate();
- dateSql = String.format(" and DATE_FORMAT(mc.created_time,'%%Y-%%m-%%d') = '%s'", date);
- List<MidjourneyUserPaintingRecordView> list = beanSearcher.searchAll(MidjourneyUserPaintingRecordView.class, MapUtils.builder()
- .put("condition", condition)
- .put("date", dateSql)
- .orderBy(MidjourneyUserPaintingRecordView::getCreatedTime).desc()
- .build());
- e.setList(list);
- }
- return GatewayResponse.SUCCESS.newBuilder().toResult(search);
- }
- /**
- * 获取settings
- */
- @GetMapping("/get/settings")
- public Result<MidjourneyUserSettings> getSettings(){
- MidjourneyUser user = getUser();
- return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyUserSettingsService.getById(user.getSettings().getId()));
- }
- /**
- * 修改settings
- */
- @PostMapping("/put/settings")
- public Result<MidjourneyUserSettings> upSettings(@RequestBody MidjourneyUserSettings settings){
- MidjourneyUser user = getUser();
- if (!user.getId().equals(settings.getUserId())) {
- throw BusinessRuntimeException.getInstance("参数错误");
- }
- settings.setId(user.getSettings().getId());
- midjourneyUserSettingsService.updateById(settings);
- return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyUserSettingsService.getById(settings.getId()));
- }
- public static String appendDefaults(String userPrompt, String defaultOptions) {
- // 解析默认选项
- Map<String, String> defaultOptionsMap = parseOptions(defaultOptions);
- // 修正用户输入中可能缺少空格的选项
- userPrompt = fixMissingSpaces(userPrompt);
- // 检查并追加缺失的选项
- for (Map.Entry<String, String> entry : defaultOptionsMap.entrySet()) {
- if (!containsOption(userPrompt, entry.getKey())) {
- userPrompt += " " + entry.getValue();
- }
- }
- // 单独处理 --s 和 --stylize 的情况
- if (!containsOption(userPrompt, "--s") && !containsOption(userPrompt, "--stylize")) {
- if (defaultOptionsMap.containsKey("--s")) {
- userPrompt += " " + defaultOptionsMap.get("--s");
- }
- }
- return userPrompt;
- }
- public static String fixMissingSpaces(String userPrompt) {
- // 定义需要修正的正则表达式和替换格式
- String[][] patterns = {
- {"--v(\\d)", "--v $1"},
- {"--niji(\\d)", "--niji $1"},
- {"--s(\\d+)", "--s $1"},
- {"--stylize(\\d+)", "--stylize $1"},
- {"--style(\\w+)", "--style $1"},
- {"--ar(\\d+)[::](\\d+)", "--ar $1:$2"},
- {"--ar (\\d+):(\\d+)", "--ar $1:$2"}
- };
- // 修正用户输入
- for (String[] pattern : patterns) {
- userPrompt = userPrompt.replaceAll(pattern[0], pattern[1]);
- }
- return userPrompt;
- }
- public static boolean containsOption(String userPrompt, String option) {
- return userPrompt.contains(option + " ") || userPrompt.matches(".*" + option + "\\d.*");
- }
- public static Map<String, String> parseOptions(String options) {
- Map<String, String> optionsMap = new HashMap<>();
- String[] parts = options.split("--");
- for (String part : parts) {
- if (!part.trim().isEmpty()) {
- String[] keyValue = part.trim().split(" ", 2);
- if (keyValue.length == 2) {
- optionsMap.put("--" + keyValue[0], "--" + keyValue[0] + " " + keyValue[1]);
- }
- }
- }
- return optionsMap;
- }
- /**
- * 上传文件到discord
- */
- @PostMapping("/upload")
- public Result<SubmitResult> upload(@RequestBody SubmitUploadDTO uploadDTO) throws Exception {
- log.info("上传文件到discord");
- return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.uploadFile(uploadDTO));
- }
- }
|