diff --git a/server/src/main/java/com/emotion/service/ScriptChatService.java b/server/src/main/java/com/emotion/service/ScriptChatService.java new file mode 100644 index 0000000..39ca05d --- /dev/null +++ b/server/src/main/java/com/emotion/service/ScriptChatService.java @@ -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); +} diff --git a/server/src/main/java/com/emotion/service/impl/ScriptChatServiceImpl.java b/server/src/main/java/com/emotion/service/impl/ScriptChatServiceImpl.java new file mode 100644 index 0000000..fb1fe24 --- /dev/null +++ b/server/src/main/java/com/emotion/service/impl/ScriptChatServiceImpl.java @@ -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 extraInputs = buildExtraInputs(operationType, conversationId, + request.getMessageId(), request.getUserMessageId(), scriptId); + + // 调用 AI(当前项目 AI 运行时不直接支持 SSE,先按非流式调用,再按字符分片发送) + String aiOutput = invokeAiRuntime(operationType, prompt, userId, extraInputs); + + // 保存 AI 消息 + Map 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 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 buildExtraInputs(String operationType, String conversationId, + String messageId, String userMessageId, String scriptId) { + Map 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 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 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 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 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); + } + } +}