|
|
@@ -8,6 +8,7 @@ import com.cyksj.enums.GatewayResponse;
|
|
|
import com.cyksj.mapper.MidjourneyUserMapper;
|
|
|
import com.cyksj.model.entity.MidjourneyUser;
|
|
|
import com.cyksj.model.entity.MidjourneyUserConversation;
|
|
|
+import com.cyksj.redis.RedisService;
|
|
|
import com.cyksj.service.midjourney.MidjourneyAccountService;
|
|
|
import com.cyksj.service.midjourney.MidjourneyService;
|
|
|
import com.ejlchina.searcher.BeanSearcher;
|
|
|
@@ -39,33 +40,52 @@ public class MidjourneyController {
|
|
|
private final MidjourneyService midjourneyService;
|
|
|
|
|
|
private final BeanSearcher beanSearcher;
|
|
|
+
|
|
|
+ private final RedisService redisService;
|
|
|
|
|
|
public MidjourneyUser getUser(){
|
|
|
- return midjourneyUserMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUser.class)
|
|
|
- .eq(MidjourneyUser::getUserToken, request.getHeader("user-token")).last("limit 1"));
|
|
|
+ String userToken = request.getHeader("user-token");
|
|
|
+ MidjourneyUser midjourneyUser = (MidjourneyUser) redisService.get(RedisService.key.MIDJOURNEY_USER.getName() + userToken);
|
|
|
+ if (midjourneyUser == null){
|
|
|
+ midjourneyUser = midjourneyUserMapper.selectOne(Wrappers.lambdaQuery(MidjourneyUser.class)
|
|
|
+ .eq(MidjourneyUser::getUserToken, userToken).last("limit 1"));
|
|
|
+ }
|
|
|
+ Object fastNum = redisService.get(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + midjourneyUser.getId());
|
|
|
+ if (fastNum == null){
|
|
|
+ redisService.set(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + midjourneyUser.getId(), midjourneyUser.getMjFastNum(), RedisService.key.MIDJOURNEY_FAST_LIMIT.getTimeout());
|
|
|
+ }
|
|
|
+ Object relaxNum = redisService.get(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + midjourneyUser.getId());
|
|
|
+ if (relaxNum == null){
|
|
|
+ if (midjourneyUser.getMjRelaxNum() != null){
|
|
|
+ redisService.set(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + midjourneyUser.getId(), midjourneyUser.getMjRelaxNum(), RedisService.key.MIDJOURNEY_RELAX_LIMIT.getTimeout());
|
|
|
+ }
|
|
|
+ }
|
|
|
+ return midjourneyUser;
|
|
|
}
|
|
|
|
|
|
|
|
|
/**
|
|
|
- * 检查用户fast次数
|
|
|
+ * 检查用户次数
|
|
|
*/
|
|
|
- public void checkUserFast(MidjourneyUser user){
|
|
|
+ public void checkUserLimit(MidjourneyUser user){
|
|
|
boolean flag = true;
|
|
|
+ Long fastNum = null;
|
|
|
+ Long relaxNum = null;
|
|
|
if (user.getMode() == 1){
|
|
|
if (user.getExpireTime().before(new Date())){
|
|
|
flag = false;
|
|
|
}
|
|
|
- if (user.getMjFastNum() < 0) {
|
|
|
+ fastNum = redisService.decr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
|
|
|
+ if (fastNum < 0) {
|
|
|
user.setMode(2);
|
|
|
- midjourneyUserMapper.updateById(user);
|
|
|
- }else {
|
|
|
- midjourneyUserMapper.update(null,Wrappers.lambdaUpdate(MidjourneyUser.class).set(MidjourneyUser::getMjFastNum, user.getMjFastNum() - 1).eq(MidjourneyUser::getId,user.getId()));
|
|
|
+ fastNum = redisService.incr(RedisService.key.MIDJOURNEY_FAST_LIMIT.getName() + user.getId(), 1L);
|
|
|
}
|
|
|
}
|
|
|
if (user.getMode() == 2){
|
|
|
- if (user.getMjRelaxNum() != null){
|
|
|
- if (user.getMjRelaxNum() > 0) {
|
|
|
- midjourneyUserMapper.update(null, Wrappers.lambdaUpdate(MidjourneyUser.class).set(MidjourneyUser::getMjRelaxNum, user.getMjRelaxNum() - 1).eq(MidjourneyUser::getId, user.getId()));
|
|
|
+ relaxNum = (Long) redisService.get(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId());
|
|
|
+ if (relaxNum != null){
|
|
|
+ if (relaxNum > 0){
|
|
|
+ relaxNum = redisService.decr(RedisService.key.MIDJOURNEY_RELAX_LIMIT.getName() + user.getId(), 1L);
|
|
|
}else {
|
|
|
flag = false;
|
|
|
}
|
|
|
@@ -73,6 +93,11 @@ public class MidjourneyController {
|
|
|
}
|
|
|
if (!flag){
|
|
|
throw BusinessRuntimeException.getInstance("可用次数不足");
|
|
|
+ }else {
|
|
|
+ midjourneyUserMapper.update(null, Wrappers.lambdaUpdate(MidjourneyUser.class)
|
|
|
+ .eq(MidjourneyUser::getId, user.getId()).set(MidjourneyUser::getMode, user.getMode())
|
|
|
+ .set(fastNum != null,MidjourneyUser::getMjFastNum, user.getMjFastNum())
|
|
|
+ .set(relaxNum != null,MidjourneyUser::getMjRelaxNum, user.getMjRelaxNum()));
|
|
|
}
|
|
|
}
|
|
|
|
|
|
@@ -85,7 +110,7 @@ public class MidjourneyController {
|
|
|
public Result<MidjourneyUserConversation> submitImagine(@RequestBody SubmitImagineDTO submitImagineDTO)throws Exception{
|
|
|
log.info("提交Imagine任务,提示:{},base64数组长度:{}",submitImagineDTO.getPrompt(),submitImagineDTO.getBase64Array());
|
|
|
MidjourneyUser user = getUser();
|
|
|
- checkUserFast(user);
|
|
|
+ checkUserLimit(user);
|
|
|
return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.submitImagine(user,submitImagineDTO.getPrompt(),submitImagineDTO.getBase64Array()));
|
|
|
}
|
|
|
|
|
|
@@ -96,7 +121,7 @@ public class MidjourneyController {
|
|
|
public Result<MidjourneyUserConversation> submitDescribe(@RequestBody SubmitDescribeDTO submitDescribeDTO)throws Exception{
|
|
|
log.info("提交Describe任务");
|
|
|
MidjourneyUser user = getUser();
|
|
|
- checkUserFast(user);
|
|
|
+ checkUserLimit(user);
|
|
|
return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.submitDescribe(user,submitDescribeDTO.getBase64()));
|
|
|
}
|
|
|
|
|
|
@@ -107,7 +132,7 @@ public class MidjourneyController {
|
|
|
public Result<MidjourneyUserConversation> submitBlend(@RequestBody SubmitBlendDTO submitBlendDTO) throws Exception{
|
|
|
log.info("提交Blend任务 dimensions:{},base64数组长度:{}", submitBlendDTO.getDimensions(), submitBlendDTO.getBase64Array().size());
|
|
|
MidjourneyUser user = getUser();
|
|
|
- checkUserFast(user);
|
|
|
+ checkUserLimit(user);
|
|
|
return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.submitBlend(user, submitBlendDTO.getDimensions(), submitBlendDTO.getBase64Array()));
|
|
|
}
|
|
|
|
|
|
@@ -118,7 +143,7 @@ public class MidjourneyController {
|
|
|
public Result<MidjourneyUserConversation> submitModal(@RequestBody SubmitModalDTO submitModalDTO) throws Exception{
|
|
|
log.info("提交Modal任务,taskId:{},提示:{},base64数组长度:{}",submitModalDTO.getTaskId(),submitModalDTO.getPrompt(),submitModalDTO.getMaskBase64().length());
|
|
|
MidjourneyUser user = getUser();
|
|
|
- checkUserFast(user);
|
|
|
+ checkUserLimit(user);
|
|
|
return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.submitModal(user,submitModalDTO.getTaskId(),submitModalDTO.getPrompt(),submitModalDTO.getMaskBase64()));
|
|
|
}
|
|
|
|
|
|
@@ -129,7 +154,7 @@ public class MidjourneyController {
|
|
|
public Result<MidjourneyUserConversation> submitShorten(@RequestBody SubmitShortenDTO submitShortenDTO) throws Exception {
|
|
|
log.info("提交Shorten任务 提示词:{}",submitShortenDTO.getPrompt());
|
|
|
MidjourneyUser user = getUser();
|
|
|
- checkUserFast(user);
|
|
|
+ checkUserLimit(user);
|
|
|
return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.submitShorten(user,submitShortenDTO.getPrompt()));
|
|
|
}
|
|
|
|
|
|
@@ -140,7 +165,7 @@ public class MidjourneyController {
|
|
|
public Result<MidjourneyUserConversation> action(@RequestBody SubmitActionDTO actionDTO) throws Exception {
|
|
|
log.info("任务id:{},执行动作:{}",actionDTO.getTaskId(),actionDTO.getCustomId());
|
|
|
MidjourneyUser user = getUser();
|
|
|
- checkUserFast(user);
|
|
|
+ checkUserLimit(user);
|
|
|
return GatewayResponse.SUCCESS.newBuilder().toResult(midjourneyService.submitAction(user,actionDTO.getTaskId(), actionDTO.getCustomId()));
|
|
|
}
|
|
|
|