diff --git a/server/pom.xml b/server/pom.xml index d9d646b..2cd8134 100644 --- a/server/pom.xml +++ b/server/pom.xml @@ -104,6 +104,13 @@ runtime + + + com.squareup.okhttp3 + okhttp + 4.12.0 + + com.alibaba.fastjson2 diff --git a/server/src/main/java/com/emotion/service/impl/ShortNovelServiceImpl.java b/server/src/main/java/com/emotion/service/impl/ShortNovelServiceImpl.java index c6c4b61..bf07488 100644 --- a/server/src/main/java/com/emotion/service/impl/ShortNovelServiceImpl.java +++ b/server/src/main/java/com/emotion/service/impl/ShortNovelServiceImpl.java @@ -8,30 +8,29 @@ import com.emotion.dto.request.ShortNovelStreamRequest; import com.emotion.exception.BusinessException; import com.emotion.service.ShortNovelService; import com.emotion.util.UserContextHolder; +import okhttp3.MediaType; +import okhttp3.OkHttpClient; +import okhttp3.Request; +import okhttp3.RequestBody; +import okhttp3.Response; +import okhttp3.ResponseBody; +import okio.BufferedSource; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.beans.factory.annotation.Autowired; -import org.springframework.http.HttpEntity; -import org.springframework.http.HttpHeaders; -import org.springframework.http.HttpMethod; -import org.springframework.http.MediaType; -import org.springframework.http.ResponseEntity; import org.springframework.stereotype.Service; -import org.springframework.web.client.RestTemplate; import org.springframework.web.servlet.mvc.method.annotation.SseEmitter; -import java.io.BufferedReader; -import java.io.InputStream; -import java.io.InputStreamReader; -import java.nio.charset.StandardCharsets; import java.util.HashMap; import java.util.Map; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; /** * 短篇小说外部服务代理实现 - * 使用 RestTemplate 读取外部 SSE 流并通过 SseEmitter 转发给小程序前端 + * 使用 OkHttp 逐行流式读取外部 SSE 响应,通过 SseEmitter 转发给小程序前端 + * OkHttp 的 ResponseBody.source() 是真正的网络流式读取,不会缓冲整个响应体 * * @author huazhongmin * @date 2026-07-19 @@ -48,7 +47,29 @@ public class ShortNovelServiceImpl implements ShortNovelService { @Autowired private EpicScriptDialogueServiceImpl epicScriptDialogueServiceImpl; - private final RestTemplate restTemplate = new RestTemplate(); + /** + * OkHttp 客户端:连接/读取超时与 Spring 配置对齐,支持 SSE 长连接流式读取 + * 延迟初始化,避免 @Autowired 注入前 config 还未填充 + */ + private volatile OkHttpClient okHttpClient; + + private OkHttpClient getOkHttpClient() { + if (okHttpClient == null) { + synchronized (this) { + if (okHttpClient == null) { + okHttpClient = new OkHttpClient.Builder() + .connectTimeout(config.getConnectTimeout(), TimeUnit.MILLISECONDS) + .readTimeout(config.getReadTimeout(), TimeUnit.MILLISECONDS) + .writeTimeout(config.getConnectTimeout(), TimeUnit.MILLISECONDS) + .build(); + } + } + } + return okHttpClient; + } + + private static final MediaType JSON_MEDIA_TYPE = MediaType.parse("application/json; charset=utf-8"); + private static final MediaType SSE_ACCEPT_TYPE = MediaType.parse("text/event-stream"); @Override public SseEmitter stream(ShortNovelStreamRequest request) { @@ -83,7 +104,7 @@ public class ShortNovelServiceImpl implements ShortNovelService { } /** - * 通用 SSE 转发逻辑 + * 通用 SSE 转发逻辑(OkHttp 流式读取) * * @param path 外部服务路径 * @param body 请求体 @@ -94,76 +115,84 @@ public class ShortNovelServiceImpl implements ShortNovelService { SseEmitter emitter = new SseEmitter(config.getReadTimeout().longValue()); EXECUTOR.execute(() -> { - HttpHeaders headers = new HttpHeaders(); - headers.setContentType(MediaType.APPLICATION_JSON); - headers.set("X-API-Token", config.getApiToken()); - headers.set("Accept", "text/event-stream"); - String url = config.getApiBaseUrl() + path; try { - ResponseEntity response = restTemplate.exchange( - url, HttpMethod.POST, - new HttpEntity<>(JSON.toJSONString(body), headers), - org.springframework.core.io.Resource.class); + // 构建 OkHttp 请求 + Request okhttpRequest = new Request.Builder() + .url(url) + .post(RequestBody.create(JSON.toJSONString(body), JSON_MEDIA_TYPE)) + .addHeader("X-API-Token", config.getApiToken()) + .addHeader("Accept", "text/event-stream") + .build(); - InputStream inputStream = response.getBody().getInputStream(); - BufferedReader reader = new BufferedReader(new InputStreamReader(inputStream, StandardCharsets.UTF_8)); + // OkHttp 执行请求,ResponseBody.source() 是真正的网络流 + try (Response response = getOkHttpClient().newCall(okhttpRequest).execute()) { + if (!response.isSuccessful()) { + throw new BusinessException("上游服务返回错误: HTTP " + response.code()); + } - String line; - StringBuilder dataBuffer = new StringBuilder(); + ResponseBody responseBody = response.body(); + if (responseBody == null) { + throw new BusinessException("上游服务返回空响应体"); + } - while ((line = reader.readLine()) != null) { - if (line.startsWith("data:")) { - dataBuffer.append(line.substring(5).trim()); - } else if (line.isEmpty() && dataBuffer.length() > 0) { - String dataStr = dataBuffer.toString(); - dataBuffer.setLength(0); + // source() 返回的 BufferedSource 逐行从网络 socket 读取,不会缓冲整个响应 + BufferedSource source = responseBody.source(); + StringBuilder dataBuffer = new StringBuilder(); - if ("[DONE]".equals(dataStr)) { - continue; - } + while (!source.exhausted()) { + String line = source.readUtf8Line(); + if (line == null) break; - try { - JSONObject event = JSON.parseObject(dataStr); - String type = event.getString("type"); - // 排查用:打印每个收到的事件类型,确认上游真实的"完成"事件名 - log.info("[ShortNovel SSE] 收到事件: type={}, session_id={}, keys={}", - type, event.getString("session_id"), event.keySet()); + if (line.startsWith("data:")) { + dataBuffer.append(line.substring(5).trim()); + } else if (line.isEmpty() && dataBuffer.length() > 0) { + String dataStr = dataBuffer.toString(); + dataBuffer.setLength(0); - // 拦截 novel_done 事件,保存到数据库 - if ("novel_done".equals(type)) { - JSONObject payload = event.getJSONObject("payload"); - if (payload != null && originalQuery != null) { - String fullText = payload.getString("full_text"); - if (fullText != null) { - Map metadata = new HashMap<>(); - if (payload.get("title") != null) { - metadata.put("title", payload.get("title")); - } - Map saveResult = epicScriptDialogueServiceImpl.saveNovelResult( - currentUserId, - originalQuery, - fullText, - metadata); - - // 注入 scriptId 到事件中 - payload.put("scriptId", saveResult.get("scriptId")); - payload.put("conversationId", saveResult.get("conversationId")); - payload.put("currentVersionMessageId", saveResult.get("currentVersionMessageId")); - } - } + if ("[DONE]".equals(dataStr)) { + continue; } - // 转发事件给前端 - emitter.send(SseEmitter.event().name(type).data(event.toJSONString())); - } catch (Exception parseEx) { - log.warn("SSE 事件解析失败: {}", parseEx.getMessage()); + try { + JSONObject event = JSON.parseObject(dataStr); + String type = event.getString("type"); + + // 拦截 novel_done 事件,保存到数据库 + if ("novel_done".equals(type)) { + JSONObject payload = event.getJSONObject("payload"); + if (payload != null && originalQuery != null) { + String fullText = payload.getString("full_text"); + if (fullText != null) { + Map metadata = new HashMap<>(); + if (payload.get("title") != null) { + metadata.put("title", payload.get("title")); + } + Map saveResult = epicScriptDialogueServiceImpl.saveNovelResult( + currentUserId, + originalQuery, + fullText, + metadata); + + // 注入 scriptId 到事件中 + payload.put("scriptId", saveResult.get("scriptId")); + payload.put("conversationId", saveResult.get("conversationId")); + payload.put("currentVersionMessageId", saveResult.get("currentVersionMessageId")); + } + } + } + + // 逐事件转发给前端(OkHttp 每读到一个完整 SSE 事件就立即转发) + emitter.send(SseEmitter.event().name(type).data(event.toJSONString())); + } catch (Exception parseEx) { + log.warn("SSE 事件解析失败: {}", parseEx.getMessage()); + } } } - } - emitter.complete(); + emitter.complete(); + } } catch (Exception e) { log.error("SSE 代理异常: {}", e.getMessage(), e); try {