|
|
@@ -0,0 +1,90 @@
|
|
|
+package com.seeckg.knowledgegraph_backend.service;
|
|
|
+
|
|
|
+
|
|
|
+import com.seeckg.knowledgegraph_backend.enums.ReturnCode;
|
|
|
+import com.seeckg.knowledgegraph_backend.enums.questionAnsweringType;
|
|
|
+import com.seeckg.knowledgegraph_backend.pojo.Knowledge;
|
|
|
+import com.seeckg.knowledgegraph_backend.pojo.Response;
|
|
|
+import com.seeckg.knowledgegraph_backend.remote.GPTRemoteService;
|
|
|
+import lombok.extern.slf4j.Slf4j;
|
|
|
+import org.springframework.stereotype.Service;
|
|
|
+
|
|
|
+import javax.annotation.Resource;
|
|
|
+import java.io.BufferedReader;
|
|
|
+import java.io.IOException;
|
|
|
+import java.io.InputStreamReader;
|
|
|
+import java.util.ArrayList;
|
|
|
+import java.util.List;
|
|
|
+
|
|
|
+@Service
|
|
|
+@Slf4j
|
|
|
+public class QuestionAnsweringService {
|
|
|
+ @Resource
|
|
|
+ GPTRemoteService gptRemoteService;
|
|
|
+ KnowledgeService knowledgeService;
|
|
|
+
|
|
|
+
|
|
|
+ /**
|
|
|
+ * 选择离线问答或在线问答
|
|
|
+ * 返回的数据格式是String
|
|
|
+ * @return 包含上述数据的Response
|
|
|
+ */
|
|
|
+ public Response getAnswer(int type, String question) {
|
|
|
+ if (type == questionAnsweringType.OFFLINE.ordinal()) {
|
|
|
+ return getOfflineAnswer(question);
|
|
|
+ }
|
|
|
+ else if (type == questionAnsweringType.ONLINE.ordinal()) {
|
|
|
+ return getOnlineAnswer(question);
|
|
|
+ }
|
|
|
+ return Response.buildFailed(ReturnCode.PARAM_ERROR.getCode(), "问答类型不存在!");
|
|
|
+ }
|
|
|
+
|
|
|
+
|
|
|
+ /**
|
|
|
+ * 获取离线问答结果
|
|
|
+ * 返回的数据格式是String
|
|
|
+ * @return 包含上述数据的Response
|
|
|
+ */
|
|
|
+ public Response getOfflineAnswer(String question) {
|
|
|
+ String answer = "";
|
|
|
+
|
|
|
+ List<Knowledge> allKnowledge = knowledgeService.getAllKnowledge();
|
|
|
+ List<String> allKnowledgeText = new ArrayList<>();
|
|
|
+ for (Knowledge knowledge : allKnowledge) {
|
|
|
+ String context = knowledge.getContext().replace('\n', ' ');
|
|
|
+ allKnowledgeText.add(context);
|
|
|
+ }
|
|
|
+ String contexts = String.join(" ", allKnowledgeText);
|
|
|
+
|
|
|
+ try {
|
|
|
+ String[] args1 = new String[] { "python", "..\\util\\OfflineQuestionAnsweringServive.py", question, contexts};
|
|
|
+ Process proc = Runtime.getRuntime().exec(args1);
|
|
|
+
|
|
|
+ BufferedReader in = new BufferedReader(new InputStreamReader(proc.getInputStream()));
|
|
|
+ String line = null;
|
|
|
+ while ((line = in.readLine()) != null) {
|
|
|
+ answer = line;
|
|
|
+ // System.out.println(line);
|
|
|
+ }
|
|
|
+ in.close();
|
|
|
+ proc.waitFor();
|
|
|
+ } catch (IOException | InterruptedException e) {
|
|
|
+ e.printStackTrace();
|
|
|
+ return Response.buildFailed(ReturnCode.SERVER_ERROR.getCode(), "答案生成出错!");
|
|
|
+ }
|
|
|
+
|
|
|
+ return Response.buildSuccess(answer);
|
|
|
+ }
|
|
|
+
|
|
|
+
|
|
|
+ /**
|
|
|
+ * 获取在线问答结果
|
|
|
+ * 返回的数据格式是String
|
|
|
+ * @return 包含上述数据的Response
|
|
|
+ */
|
|
|
+ public Response getOnlineAnswer(String question) {
|
|
|
+ String answer = gptRemoteService.getGPTAnswer(question);
|
|
|
+ return Response.buildSuccess(answer);
|
|
|
+ }
|
|
|
+
|
|
|
+}
|