ChatClientController.java 3.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293
  1. package space.anyi.springAiAlibabaLearn.controller;
  2. import org.slf4j.Logger;
  3. import org.slf4j.LoggerFactory;
  4. import org.springframework.ai.chat.client.ChatClient;
  5. import org.springframework.ai.chat.messages.SystemMessage;
  6. import org.springframework.ai.chat.messages.UserMessage;
  7. import org.springframework.ai.chat.model.ChatModel;
  8. import org.springframework.ai.chat.model.ChatResponse;
  9. import org.springframework.ai.chat.prompt.ChatOptions;
  10. import org.springframework.ai.chat.prompt.Prompt;
  11. import org.springframework.web.bind.annotation.GetMapping;
  12. import org.springframework.web.bind.annotation.RequestMapping;
  13. import org.springframework.web.bind.annotation.RequestParam;
  14. import org.springframework.web.bind.annotation.RestController;
  15. import reactor.core.publisher.Flux;
  16. import space.anyi.springAiAlibabaLearn.model.Person;
  17. @RestController
  18. @RequestMapping("/chatClient")
  19. public class ChatClientController {
  20. private final Logger log = LoggerFactory.getLogger(ChatClientController.class);
  21. private final ChatClient chatClient;
  22. public ChatClientController(ChatModel chatModel) {
  23. this.chatClient = ChatClient.create(chatModel);
  24. }
  25. @GetMapping("/chat")
  26. public String call(@RequestParam("message") String message) {
  27. log.debug("message: {}", message);
  28. return chatClient
  29. .prompt(message)
  30. .call()
  31. .chatClientResponse()
  32. .chatResponse()
  33. .getResult()
  34. .getOutput()
  35. .getText();
  36. }
  37. @GetMapping("/prompt")
  38. public String prompt(@RequestParam("message") String message) {
  39. log.debug("message: {}", message);
  40. SystemMessage systemMessage = SystemMessage.builder().text("你是一个机器人!无论用户问什么内容,你只会回答'Java是世界上最好的语言!'.").build();
  41. log.debug("systemMessage:{}",systemMessage);
  42. UserMessage userMessage = new UserMessage(message);
  43. log.debug("userMessage:{}",userMessage);
  44. Prompt prompt = Prompt.builder()
  45. .messages(systemMessage,userMessage)
  46. .chatOptions(
  47. ChatOptions.builder()
  48. .maxTokens(4096)
  49. .temperature(0.2)
  50. .build()
  51. )
  52. .build();
  53. log.debug("prompt: {}", prompt);
  54. ChatResponse chatResponse = chatClient
  55. .prompt(prompt)
  56. .call()
  57. .chatClientResponse()
  58. .chatResponse();
  59. log.debug("chatResponse: {}", chatResponse);
  60. return chatResponse.getResult().getOutput().getText();
  61. }
  62. @GetMapping("/stream")
  63. public Flux<String> stream(@RequestParam("message") String message) {
  64. log.debug("message: {}", message);
  65. return chatClient
  66. .prompt(message)
  67. .stream()
  68. .chatClientResponse()
  69. .doOnNext(chatClientResponse -> {
  70. String text = chatClientResponse.chatResponse().getResult().getOutput().getText();
  71. log.debug("res:{}",text);
  72. })
  73. .map(chatClientResponse -> chatClientResponse.chatResponse().getResult().getOutput().getText());
  74. }
  75. @GetMapping("/entity")
  76. public Person entity(){
  77. //将LLM的输出转换为Java对象
  78. /**
  79. * 内部会在prompt中追加Person类的结构消息,并要求LLM输出符合格式的JSON字符串
  80. */
  81. return chatClient
  82. .prompt("生成一个人的消息.")
  83. .call()
  84. .entity(Person.class);
  85. }
  86. }