feat: 新增 ScriptChatService 接口
This commit is contained in:
@@ -0,0 +1,15 @@
|
||||
package com.emotion.service;
|
||||
|
||||
import com.emotion.dto.request.ScriptChatStreamRequest;
|
||||
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
|
||||
|
||||
/**
|
||||
* 剧本对话流式服务
|
||||
*/
|
||||
public interface ScriptChatService {
|
||||
|
||||
/**
|
||||
* 流式处理改写/续写/聊天请求
|
||||
*/
|
||||
SseEmitter streamChat(ScriptChatStreamRequest request);
|
||||
}
|
||||
@@ -0,0 +1,286 @@
|
||||
package com.emotion.service.impl;
|
||||
|
||||
import com.alibaba.fastjson2.JSON;
|
||||
import com.alibaba.fastjson2.JSONObject;
|
||||
import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper;
|
||||
import com.emotion.dto.request.ScriptChatStreamRequest;
|
||||
import com.emotion.dto.request.ai.AiRuntimeRequest;
|
||||
import com.emotion.dto.response.ai.AiRuntimeTestResponse;
|
||||
import com.emotion.entity.Conversation;
|
||||
import com.emotion.entity.EpicScript;
|
||||
import com.emotion.entity.Message;
|
||||
import com.emotion.service.*;
|
||||
import com.emotion.util.SnowflakeIdGenerator;
|
||||
import com.emotion.util.UserContextHolder;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Service;
|
||||
import org.springframework.util.StringUtils;
|
||||
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.time.LocalDateTime;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Optional;
|
||||
import java.util.concurrent.CompletableFuture;
|
||||
import java.util.concurrent.Executor;
|
||||
|
||||
@Slf4j
|
||||
@Service
|
||||
public class ScriptChatServiceImpl implements ScriptChatService {
|
||||
|
||||
@Autowired
|
||||
private ConversationService conversationService;
|
||||
|
||||
@Autowired
|
||||
private EpicScriptService epicScriptService;
|
||||
|
||||
@Autowired
|
||||
private MessageService messageService;
|
||||
|
||||
@Autowired
|
||||
private AiRuntimeService aiRuntimeService;
|
||||
|
||||
@Autowired
|
||||
private SnowflakeIdGenerator snowflakeIdGenerator;
|
||||
|
||||
@Autowired
|
||||
private Executor taskExecutor;
|
||||
|
||||
@Override
|
||||
public SseEmitter streamChat(ScriptChatStreamRequest request) {
|
||||
String currentUserId = UserContextHolder.getCurrentUserId();
|
||||
if (currentUserId == null) {
|
||||
throw new IllegalStateException("用户未登录");
|
||||
}
|
||||
|
||||
Conversation conversation = conversationService.getById(request.getConversationId());
|
||||
if (conversation == null || !currentUserId.equals(conversation.getUserId())) {
|
||||
throw new IllegalStateException("对话不存在或无权限");
|
||||
}
|
||||
|
||||
EpicScript script = epicScriptService.getById(conversation.getScriptId());
|
||||
if (script == null) {
|
||||
throw new IllegalStateException("关联剧本不存在");
|
||||
}
|
||||
|
||||
SseEmitter emitter = new SseEmitter(300_000L);
|
||||
CompletableFuture.runAsync(() -> {
|
||||
try {
|
||||
processStream(request, currentUserId, conversation, script, emitter);
|
||||
} catch (Exception e) {
|
||||
log.error("流式处理失败", e);
|
||||
sendError(emitter, e.getMessage());
|
||||
}
|
||||
}, taskExecutor);
|
||||
return emitter;
|
||||
}
|
||||
|
||||
private void processStream(ScriptChatStreamRequest request, String userId,
|
||||
Conversation conversation, EpicScript script, SseEmitter emitter) {
|
||||
String operationType = request.getOperationType();
|
||||
String conversationId = conversation.getId();
|
||||
String scriptId = script.getId();
|
||||
|
||||
// 查询用户指令消息
|
||||
Message userMessage = messageService.getById(request.getUserMessageId());
|
||||
if (userMessage == null || !conversationId.equals(userMessage.getConversationId())) {
|
||||
throw new IllegalStateException("用户指令消息不存在");
|
||||
}
|
||||
|
||||
// 查询被改写的 AI 消息(rewrite/continue 需要)
|
||||
Message targetMessage = null;
|
||||
if ("rewrite".equals(operationType) || "continue".equals(operationType)) {
|
||||
targetMessage = messageService.getById(request.getMessageId());
|
||||
if (targetMessage == null || !"script".equals(targetMessage.getType())) {
|
||||
throw new IllegalStateException("被改写的剧本消息不存在");
|
||||
}
|
||||
}
|
||||
|
||||
// 组装 prompt 和上下文
|
||||
String prompt = buildPrompt(operationType, targetMessage, userMessage);
|
||||
Map<String, Object> extraInputs = buildExtraInputs(operationType, conversationId,
|
||||
request.getMessageId(), request.getUserMessageId(), scriptId);
|
||||
|
||||
// 调用 AI(当前项目 AI 运行时不直接支持 SSE,先按非流式调用,再按字符分片发送)
|
||||
String aiOutput = invokeAiRuntime(operationType, prompt, userId, extraInputs);
|
||||
|
||||
// 保存 AI 消息
|
||||
Map<String, Object> metadata = parseScriptOutput(aiOutput);
|
||||
String content = buildDisplayContent(metadata);
|
||||
int nextVersionNumber = calculateNextVersionNumber(targetMessage);
|
||||
|
||||
Message aiMessage = new Message();
|
||||
aiMessage.setId(snowflakeIdGenerator.nextIdAsString());
|
||||
aiMessage.setConversationId(conversationId);
|
||||
aiMessage.setScriptId(scriptId);
|
||||
aiMessage.setUserId(userId);
|
||||
aiMessage.setContent(content);
|
||||
aiMessage.setType("chat".equals(operationType) ? "chat" : "script");
|
||||
aiMessage.setSender("assistant");
|
||||
aiMessage.setTimestamp(LocalDateTime.now());
|
||||
aiMessage.setMessageOrder(calculateNextMessageOrder(conversationId));
|
||||
aiMessage.setParentMessageId(targetMessage != null ? targetMessage.getId() : null);
|
||||
aiMessage.setVersionNumber("chat".equals(operationType) ? 0 : nextVersionNumber);
|
||||
aiMessage.setStatus("sent");
|
||||
aiMessage.setMetadata("chat".equals(operationType) ? null : JSON.toJSONString(metadata));
|
||||
messageService.createMessage(aiMessage);
|
||||
|
||||
// 更新剧本快照(仅 script 类型)
|
||||
if (!"chat".equals(operationType)) {
|
||||
script.setCurrentVersionMessageId(aiMessage.getId());
|
||||
script.setTitle((String) metadata.getOrDefault("title", script.getTitle()));
|
||||
script.setPlotIntro((String) metadata.getOrDefault("plotIntro", ""));
|
||||
script.setPlotTurning((String) metadata.getOrDefault("plotTurning", ""));
|
||||
script.setPlotClimax((String) metadata.getOrDefault("plotClimax", ""));
|
||||
script.setPlotEnding((String) metadata.getOrDefault("plotEnding", ""));
|
||||
Map<String, Object> plotJson = script.getPlotJson();
|
||||
if (plotJson == null) {
|
||||
plotJson = new HashMap<>();
|
||||
}
|
||||
plotJson.putAll(metadata);
|
||||
script.setPlotJson(plotJson);
|
||||
epicScriptService.updateById(script);
|
||||
|
||||
conversation.setCurrentMessageId(aiMessage.getId());
|
||||
}
|
||||
conversation.setLastActiveTime(LocalDateTime.now());
|
||||
conversation.setMessageCount(conversation.getMessageCount() + 1);
|
||||
conversationService.updateById(conversation);
|
||||
|
||||
// 流式发送内容(简单分片模拟,后续可接入真正 SSE)
|
||||
try {
|
||||
emitter.send(SseEmitter.event().name("start").data("{}"));
|
||||
int chunkSize = 20;
|
||||
for (int i = 0; i < content.length(); i += chunkSize) {
|
||||
String chunk = content.substring(i, Math.min(i + chunkSize, content.length()));
|
||||
emitter.send(SseEmitter.event().name("delta").data(chunk));
|
||||
Thread.sleep(30);
|
||||
}
|
||||
emitter.send(SseEmitter.event().name("metadata").data(JSON.toJSONString(metadata)));
|
||||
emitter.send(SseEmitter.event().name("done").data(JSON.toJSONString(Map.of(
|
||||
"messageId", aiMessage.getId(),
|
||||
"versionNumber", aiMessage.getVersionNumber()
|
||||
))));
|
||||
emitter.complete();
|
||||
} catch (IOException | InterruptedException e) {
|
||||
log.error("发送 SSE 事件失败", e);
|
||||
emitter.completeWithError(e);
|
||||
}
|
||||
}
|
||||
|
||||
private String buildPrompt(String operationType, Message targetMessage, Message userMessage) {
|
||||
StringBuilder sb = new StringBuilder();
|
||||
sb.append("操作类型:").append(operationType).append("\n");
|
||||
if (targetMessage != null && StringUtils.hasText(targetMessage.getMetadata())) {
|
||||
sb.append("当前剧本内容:\n").append(targetMessage.getMetadata()).append("\n\n");
|
||||
}
|
||||
sb.append("用户指令:").append(userMessage.getContent());
|
||||
return sb.toString();
|
||||
}
|
||||
|
||||
private Map<String, Object> buildExtraInputs(String operationType, String conversationId,
|
||||
String messageId, String userMessageId, String scriptId) {
|
||||
Map<String, Object> extras = new HashMap<>();
|
||||
extras.put("operationType", operationType);
|
||||
extras.put("conversationId", conversationId);
|
||||
extras.put("messageId", messageId);
|
||||
extras.put("userMessageId", userMessageId);
|
||||
extras.put("scriptId", scriptId);
|
||||
return extras;
|
||||
}
|
||||
|
||||
private String invokeAiRuntime(String operationType, String input, String userId, Map<String, Object> extraInputs) {
|
||||
JSONObject inputs = new JSONObject();
|
||||
inputs.put("input", input);
|
||||
inputs.put("prompt", input);
|
||||
inputs.put("message", input);
|
||||
if (extraInputs != null) {
|
||||
inputs.putAll(extraInputs);
|
||||
}
|
||||
|
||||
AiRuntimeRequest runtimeRequest = new AiRuntimeRequest();
|
||||
runtimeRequest.setSceneCode("script_generate");
|
||||
runtimeRequest.setUserId(userId);
|
||||
runtimeRequest.setInputs(inputs);
|
||||
|
||||
AiRuntimeTestResponse response = aiRuntimeService.test(runtimeRequest);
|
||||
if (response == null || !"success".equals(response.getStatus()) || !StringUtils.hasText(response.getOutput())) {
|
||||
String message = response == null ? "AI_RUNTIME_EMPTY_RESPONSE" : response.getErrorMessage();
|
||||
throw new IllegalStateException(StringUtils.hasText(message) ? message : "AI_RUNTIME_FAILED");
|
||||
}
|
||||
return response.getOutput();
|
||||
}
|
||||
|
||||
private Map<String, Object> parseScriptOutput(String output) {
|
||||
if (!StringUtils.hasText(output)) {
|
||||
return new HashMap<>();
|
||||
}
|
||||
try {
|
||||
return JSON.parseObject(output, Map.class);
|
||||
} catch (Exception e) {
|
||||
log.warn("JSON解析失败,返回空Map: {}", output, e);
|
||||
return new HashMap<>();
|
||||
}
|
||||
}
|
||||
|
||||
private String buildDisplayContent(Map<String, Object> metadata) {
|
||||
StringBuilder sb = new StringBuilder();
|
||||
appendSection(sb, metadata.get("title"), null);
|
||||
appendSection(sb, metadata.get("plotIntro"), "序幕:低谷回响");
|
||||
appendSection(sb, metadata.get("plotTurning"), "转折:契机出现");
|
||||
appendSection(sb, metadata.get("plotClimax"), "高潮:命运抉择");
|
||||
appendSection(sb, metadata.get("plotEnding"), "结局:新的开始");
|
||||
return sb.toString();
|
||||
}
|
||||
|
||||
private void appendSection(StringBuilder sb, Object value, String label) {
|
||||
if (value == null) {
|
||||
return;
|
||||
}
|
||||
String text = String.valueOf(value);
|
||||
if (!StringUtils.hasText(text)) {
|
||||
return;
|
||||
}
|
||||
if (sb.length() > 0) {
|
||||
sb.append("\n\n");
|
||||
}
|
||||
if (StringUtils.hasText(label)) {
|
||||
sb.append("【").append(label).append("】\n");
|
||||
}
|
||||
sb.append(text);
|
||||
}
|
||||
|
||||
private int calculateNextVersionNumber(Message targetMessage) {
|
||||
if (targetMessage == null || !StringUtils.hasText(targetMessage.getParentMessageId())) {
|
||||
return 1;
|
||||
}
|
||||
// 查询同一 parentMessageId 下最大 version_number
|
||||
LambdaQueryWrapper<Message> wrapper = new LambdaQueryWrapper<>();
|
||||
wrapper.eq(Message::getParentMessageId, targetMessage.getParentMessageId())
|
||||
.orderByDesc(Message::getVersionNumber)
|
||||
.last("LIMIT 1");
|
||||
Message latest = messageService.getOne(wrapper);
|
||||
if (latest == null) {
|
||||
return 1;
|
||||
}
|
||||
return latest.getVersionNumber() + 1;
|
||||
}
|
||||
|
||||
private long calculateNextMessageOrder(String conversationId) {
|
||||
Message lastMessage = messageService.getLastMessageByConversationId(conversationId);
|
||||
return lastMessage == null ? 1L : lastMessage.getMessageOrder() + 1;
|
||||
}
|
||||
|
||||
private void sendError(SseEmitter emitter, String message) {
|
||||
try {
|
||||
emitter.send(SseEmitter.event().name("error").data(message));
|
||||
emitter.complete();
|
||||
} catch (IOException e) {
|
||||
log.error("发送错误事件失败", e);
|
||||
emitter.completeWithError(e);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user