ChatCompletionService.java 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348
  1. package edu.nju.software.aipaasagent.service;
  2. import edu.nju.software.aipaasagent.agent.manager.AgentFactory;
  3. import edu.nju.software.aipaasagent.agent.core.base.BaseAgent;
  4. import edu.nju.software.aipaasagent.agent.config.AgentConfigRegistry;
  5. import edu.nju.software.aipaasagent.agent.config.AgentConfiguration;
  6. import edu.nju.software.aipaasagent.dto.ChatCompletionRequest;
  7. import edu.nju.software.aipaasagent.dto.ChatCompletionResponse;
  8. import edu.nju.software.aipaasagent.dto.StreamEvent;
  9. import lombok.extern.slf4j.Slf4j;
  10. import org.springframework.beans.factory.annotation.Autowired;
  11. import org.springframework.stereotype.Service;
  12. import reactor.core.publisher.Flux;
  13. import java.time.Instant;
  14. import java.util.List;
  15. import java.util.UUID;
  16. /**
  17. * Chat Completion 服务
  18. * 处理 OpenAI 兼容的聊天完成请求
  19. * 支持流式和非流式输出
  20. */
  21. @Slf4j
  22. @Service
  23. public class ChatCompletionService {
  24. @Autowired
  25. private AgentFactory agentFactory;
  26. @Autowired
  27. private AgentConfigRegistry agentConfigRegistry;
  28. @Autowired
  29. private MultimodalFileService multimodalFileService;
  30. /**
  31. * 非流式聊天完成
  32. *
  33. * @param request 聊天完成请求
  34. * @return 聊天完成响应
  35. */
  36. public ChatCompletionResponse chatCompletion(ChatCompletionRequest request) {
  37. String threadId = getOrCreateThreadId(request.getThreadId());
  38. String agentId = getOrCreateAgentId(request.getAgentId());
  39. log.info("非流式聊天完成 - threadId: {}, agentId: {}", threadId, agentId);
  40. // 获取或创建 Agent
  41. BaseAgent agent = getOrCreateAgent();
  42. // 转换消息格式并处理多模态文件
  43. String userMessage = extractLastUserMessage(request.getMessages());
  44. // 处理多模态文件(如果提供了 file_ids)
  45. if (request.getFileIds() != null && !request.getFileIds().isEmpty()) {
  46. userMessage = processMultimodalContent(userMessage, request.getFileIds());
  47. }
  48. String content = agent.doChat(userMessage, threadId, Boolean.TRUE.equals(request.getIsThink()));
  49. // 获取模型名称:接口参数 > ConfigMap 配置
  50. String model = resolveModelName(request.getModel(), agentId);
  51. // 构建响应
  52. return buildChatCompletionResponse(content, threadId, agentId, model);
  53. }
  54. /**
  55. * 流式聊天完成
  56. *
  57. * @param request 聊天完成请求
  58. * @return SSE 流
  59. */
  60. public Flux<ChatCompletionResponse> streamChatCompletion(ChatCompletionRequest request) {
  61. String threadId = getOrCreateThreadId(request.getThreadId());
  62. String agentId = getOrCreateAgentId(request.getAgentId());
  63. log.info("流式聊天完成 - threadId: {}, agentId: {}", threadId, agentId);
  64. // 获取或创建 Agent
  65. BaseAgent agent = getOrCreateAgent();
  66. // 转换消息格式并处理多模态文件
  67. String userMessage = extractLastUserMessage(request.getMessages());
  68. // 处理多模态文件(如果提供了 file_ids)
  69. if (request.getFileIds() != null && !request.getFileIds().isEmpty()) {
  70. userMessage = processMultimodalContent(userMessage, request.getFileIds());
  71. }
  72. // 生成响应 ID
  73. String responseId = "chatcmpl-" + UUID.randomUUID().toString().replace("-", "");
  74. // 获取模型名称:接口参数 > ConfigMap 配置
  75. String model = resolveModelName(request.getModel(), agentId);
  76. Flux<String> streamContent = agent.doStreamChat(userMessage, threadId, Boolean.TRUE.equals(request.getIsThink()));
  77. return streamContent
  78. .index()
  79. .map(tuple -> {
  80. long index = tuple.getT1();
  81. String content = tuple.getT2();
  82. return ChatCompletionResponse.builder()
  83. .id(responseId)
  84. .object("chat.completion.chunk")
  85. .created(Instant.now().getEpochSecond())
  86. .model(model)
  87. .choices(List.of(
  88. ChatCompletionResponse.Choice.builder()
  89. .index((int) index)
  90. .delta(ChatCompletionResponse.Delta.builder()
  91. .content(content)
  92. .build())
  93. .finishReason(null)
  94. .build()
  95. ))
  96. .threadId(threadId)
  97. .agentId(agentId)
  98. .build();
  99. })
  100. .concatWith(Flux.just(
  101. // 发送结束标记
  102. ChatCompletionResponse.builder()
  103. .id(responseId)
  104. .object("chat.completion.chunk")
  105. .created(Instant.now().getEpochSecond())
  106. .model(model)
  107. .choices(List.of(
  108. ChatCompletionResponse.Choice.builder()
  109. .index(0)
  110. .delta(new ChatCompletionResponse.Delta())
  111. .finishReason("stop")
  112. .build()
  113. ))
  114. .threadId(threadId)
  115. .agentId(agentId)
  116. .build()
  117. ));
  118. }
  119. /**
  120. * SSE 流式聊天完成(返回 StreamEvent 格式)
  121. * 用于 stream=true + isThink=true 的情况
  122. *
  123. * @param request 聊天完成请求
  124. * @return StreamEvent 流
  125. */
  126. public Flux<StreamEvent> streamChatCompletionSSE(ChatCompletionRequest request) {
  127. String threadId = getOrCreateThreadId(request.getThreadId());
  128. String agentId = getOrCreateAgentId(request.getAgentId());
  129. log.info("SSE流式聊天完成 - threadId: {}, agentId: {}", threadId, agentId);
  130. // 获取或创建 Agent
  131. BaseAgent agent = getOrCreateAgent();
  132. // 转换消息格式并处理多模态文件
  133. String userMessage = extractLastUserMessage(request.getMessages());
  134. // 处理多模态文件(如果提供了 file_ids)
  135. if (request.getFileIds() != null && !request.getFileIds().isEmpty()) {
  136. userMessage = processMultimodalContent(userMessage, request.getFileIds());
  137. }
  138. // 调用 SSE 流式方法
  139. return agent.doStreamChatSSE(userMessage, threadId, Boolean.TRUE.equals(request.getIsThink()));
  140. }
  141. /**
  142. * 获取或创建 Agent
  143. */
  144. private BaseAgent getOrCreateAgent() {
  145. // 先尝试获取已有实例
  146. BaseAgent agent = agentFactory.getCurrentAgent();
  147. if (agent == null) {
  148. // 未创建则创建新实例
  149. agent = agentFactory.createAgent();
  150. }
  151. return agent;
  152. }
  153. /**
  154. * 获取或创建 Thread ID
  155. */
  156. private String getOrCreateThreadId(String threadId) {
  157. if (threadId == null || threadId.isEmpty()) {
  158. return "thread-" + UUID.randomUUID().toString().replace("-", "");
  159. }
  160. return threadId;
  161. }
  162. /**
  163. * 获取或创建 Agent ID
  164. */
  165. private String getOrCreateAgentId(String agentId) {
  166. if (agentId == null || agentId.isEmpty()) {
  167. return "seeCoderManus";
  168. }
  169. return agentId;
  170. }
  171. /**
  172. * 解析模型名称
  173. * 优先级:接口参数 > ConfigMap 配置
  174. *
  175. * @param requestModel 请求中的模型名称(可选)
  176. * @param agentId Agent ID
  177. * @return 最终使用的模型名称
  178. */
  179. private String resolveModelName(String requestModel, String agentId) {
  180. // 如果接口传了 model,优先使用
  181. if (requestModel != null && !requestModel.isEmpty()) {
  182. log.debug("使用接口传入的模型名称: {}", requestModel);
  183. return requestModel;
  184. }
  185. // 否则从 ConfigMap 获取
  186. AgentConfiguration config = agentConfigRegistry.getConfig(agentId);
  187. if (config != null && config.getModel() != null && config.getModel().getName() != null) {
  188. String configModel = config.getModel().getName();
  189. log.debug("使用 ConfigMap 配置的模型名称: {}", configModel);
  190. return configModel;
  191. }
  192. // 默认模型
  193. log.warn("未找到模型配置,使用默认模型: gemma3:4b");
  194. return "gemma3:4b";
  195. }
  196. /**
  197. * 提取最后一条用户消息
  198. */
  199. private String extractLastUserMessage(List<ChatCompletionRequest.ChatMessage> messages) {
  200. if (messages == null || messages.isEmpty()) {
  201. return "";
  202. }
  203. // 找到最后一条用户消息
  204. for (int i = messages.size() - 1; i >= 0; i--) {
  205. ChatCompletionRequest.ChatMessage message = messages.get(i);
  206. if ("user".equals(message.getRole())) {
  207. return message.getContent();
  208. }
  209. }
  210. // 如果没有用户消息,返回最后一条消息的内容
  211. return messages.get(messages.size() - 1).getContent();
  212. }
  213. /**
  214. * 处理多模态内容
  215. * 将文件内容整合到用户消息中
  216. *
  217. * @param userMessage 用户原始消息
  218. * @param fileIds 文件 ID 列表
  219. * @return 整合后的消息内容
  220. */
  221. private String processMultimodalContent(String userMessage, List<String> fileIds) {
  222. log.info("处理多模态内容 - fileIds: {}", fileIds);
  223. // 读取文件内容
  224. List<MultimodalFileService.FileContent> fileContents = multimodalFileService.readFiles(fileIds);
  225. if (fileContents.isEmpty()) {
  226. log.warn("未能读取任何文件内容 - fileIds: {}", fileIds);
  227. return userMessage;
  228. }
  229. StringBuilder multimodalContent = new StringBuilder();
  230. // 添加用户原始消息
  231. if (userMessage != null && !userMessage.isEmpty()) {
  232. multimodalContent.append("用户消息:\n").append(userMessage).append("\n\n");
  233. }
  234. // 添加文件内容
  235. multimodalContent.append("文件内容:\n");
  236. for (int i = 0; i < fileContents.size(); i++) {
  237. MultimodalFileService.FileContent fileContent = fileContents.get(i);
  238. multimodalContent.append("--- 文件 ").append(i + 1).append(" ---\n");
  239. multimodalContent.append("文件名:").append(fileContent.getFileName()).append("\n");
  240. multimodalContent.append("文件类型:").append(fileContent.getContentType()).append("\n");
  241. if (fileContent.isText()) {
  242. // 文本文件,直接添加内容
  243. multimodalContent.append("文件内容:\n").append(fileContent.getTextContent()).append("\n");
  244. } else if (fileContent.isImage()) {
  245. // 图片文件,添加 Base64 编码的 Data URI
  246. // 注意:某些模型支持直接处理图片,这里将图片作为 Base64 提供
  247. multimodalContent.append("图片内容(Base64):\n");
  248. multimodalContent.append(fileContent.getDataUri()).append("\n");
  249. } else {
  250. // 其他类型文件,添加 Base64 内容
  251. multimodalContent.append("文件内容(Base64):\n");
  252. multimodalContent.append(fileContent.getBase64Content()).append("\n");
  253. }
  254. multimodalContent.append("\n");
  255. }
  256. String result = multimodalContent.toString();
  257. log.info("多模态内容处理完成 - 原始消息长度: {}, 处理后长度: {}",
  258. userMessage != null ? userMessage.length() : 0, result.length());
  259. return result;
  260. }
  261. /**
  262. * 构建聊天完成响应
  263. */
  264. private ChatCompletionResponse buildChatCompletionResponse(
  265. String content,
  266. String threadId,
  267. String agentId,
  268. String model) {
  269. String responseId = "chatcmpl-" + UUID.randomUUID().toString().replace("-", "");
  270. return ChatCompletionResponse.builder()
  271. .id(responseId)
  272. .object("chat.completion")
  273. .created(Instant.now().getEpochSecond())
  274. .model(model != null ? model : "default-model")
  275. .choices(List.of(
  276. ChatCompletionResponse.Choice.builder()
  277. .index(0)
  278. .message(ChatCompletionResponse.Message.builder()
  279. .role("assistant")
  280. .content(content)
  281. .build())
  282. .finishReason("stop")
  283. .build()
  284. ))
  285. .usage(ChatCompletionResponse.Usage.builder()
  286. .promptTokens(0) // TODO: 实现 token 计数
  287. .completionTokens(0)
  288. .totalTokens(0)
  289. .build())
  290. .threadId(threadId)
  291. .agentId(agentId)
  292. .build();
  293. }
  294. }