QuestionAnsweringService.java 3.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596
  1. package com.seeckg.knowledgegraph_backend.service;
  2. import com.seeckg.knowledgegraph_backend.enums.ReturnCode;
  3. import com.seeckg.knowledgegraph_backend.enums.questionAnsweringType;
  4. import com.seeckg.knowledgegraph_backend.pojo.Knowledge;
  5. import com.seeckg.knowledgegraph_backend.pojo.Response;
  6. import com.seeckg.knowledgegraph_backend.remote.GPTRemoteService;
  7. import lombok.extern.slf4j.Slf4j;
  8. import org.springframework.stereotype.Service;
  9. import javax.annotation.Resource;
  10. import java.io.BufferedReader;
  11. import java.io.IOException;
  12. import java.io.InputStreamReader;
  13. import java.util.ArrayList;
  14. import java.util.List;
  15. @Service
  16. @Slf4j
  17. public class QuestionAnsweringService {
  18. @Resource
  19. GPTRemoteService gptRemoteService;
  20. KnowledgeService knowledgeService;
  21. /**
  22. * 选择本地问答或在线问答
  23. * 返回的数据格式是String
  24. * @return 包含上述数据的Response
  25. */
  26. public Response getAnswer(int type, String question) {
  27. if (type == questionAnsweringType.OFFLINE.ordinal()) {
  28. return getOfflineAnswer(question);
  29. }
  30. else if (type == questionAnsweringType.ONLINE.ordinal()) {
  31. return getOnlineAnswer(question);
  32. }
  33. return Response.buildFailed(ReturnCode.PARAM_ERROR.getCode(), "问答类型不存在!");
  34. }
  35. /**
  36. * 获取本地问答结果
  37. * 返回的数据格式是String
  38. * @return 包含上述数据的Response
  39. */
  40. public Response getOfflineAnswer(String question) {
  41. String answer = "";
  42. List<Knowledge> allKnowledge = knowledgeService.getAllKnowledge();
  43. List<String> allKnowledgeText = new ArrayList<>();
  44. for (Knowledge knowledge : allKnowledge) {
  45. String context = knowledge.getContext().replace('\n', ' ');
  46. allKnowledgeText.add(context);
  47. }
  48. String contexts = String.join(" ", allKnowledgeText);
  49. try {
  50. String[] args1 = new String[] { "python", "..\\util\\OfflineQuestionAnsweringServive.py", question, contexts};
  51. Process proc = Runtime.getRuntime().exec(args1);
  52. BufferedReader in = new BufferedReader(new InputStreamReader(proc.getInputStream()));
  53. String line = null;
  54. while ((line = in.readLine()) != null) {
  55. answer = line;
  56. // System.out.println(line);
  57. }
  58. in.close();
  59. proc.waitFor();
  60. } catch (IOException | InterruptedException e) {
  61. e.printStackTrace();
  62. return Response.buildFailed(ReturnCode.SERVER_ERROR.getCode(), "答案生成出错!");
  63. }
  64. return Response.buildSuccess(answer);
  65. }
  66. /**
  67. * 获取在线问答结果
  68. * 返回的数据格式是String
  69. * @return 包含上述数据的Response
  70. */
  71. public Response getOnlineAnswer(String question) {
  72. String answer = null;
  73. try {
  74. answer = gptRemoteService.getGPTAnswer(question);
  75. } catch (IOException e) {
  76. e.printStackTrace();
  77. return Response.buildFailed(ReturnCode.SERVER_ERROR.getCode(), "答案生成出错!");
  78. }
  79. return Response.buildSuccess(answer);
  80. }
  81. }