GraphController.java 3.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778
  1. package space.anyi.springAiAlibabaLearn.controller;
  2. import com.alibaba.cloud.ai.graph.CompiledGraph;
  3. import com.alibaba.cloud.ai.graph.NodeOutput;
  4. import com.alibaba.cloud.ai.graph.RunnableConfig;
  5. import org.slf4j.Logger;
  6. import org.slf4j.LoggerFactory;
  7. import org.springframework.beans.factory.annotation.Autowired;
  8. import org.springframework.beans.factory.annotation.Qualifier;
  9. import org.springframework.web.bind.annotation.GetMapping;
  10. import org.springframework.web.bind.annotation.RequestMapping;
  11. import org.springframework.web.bind.annotation.RequestParam;
  12. import org.springframework.web.bind.annotation.RestController;
  13. import reactor.core.publisher.Flux;
  14. import java.util.Map;
  15. import java.util.function.Consumer;
  16. @RestController
  17. @RequestMapping("/graph")
  18. public class GraphController {
  19. private static final Logger log = LoggerFactory.getLogger(GraphController.class);
  20. private final CompiledGraph compiledGraph;
  21. private final CompiledGraph createStanceGraph;
  22. private final CompiledGraph conditionGraph;
  23. private final CompiledGraph loopGraph;
  24. private final CompiledGraph saveGraph;
  25. public GraphController(@Autowired @Qualifier("compiledGraph")CompiledGraph compiledGraph,
  26. @Autowired @Qualifier("createStanceGraph")CompiledGraph createStanceGraph,
  27. @Autowired @Qualifier("conditionGraph")CompiledGraph conditionGraph,
  28. @Autowired @Qualifier("loopGraph")CompiledGraph loopGraph,
  29. @Autowired @Qualifier("saveGraph")CompiledGraph saveGraph) {
  30. this.compiledGraph = compiledGraph;
  31. this.createStanceGraph = createStanceGraph;
  32. this.conditionGraph = conditionGraph;
  33. this.loopGraph = loopGraph;
  34. this.saveGraph = saveGraph;
  35. }
  36. @GetMapping("/quickStart")
  37. public Flux<NodeOutput> quickStart(){
  38. return compiledGraph.stream().log().doOnNext(new Consumer<NodeOutput>() {
  39. @Override
  40. public void accept(NodeOutput nodeOutput) {
  41. log.debug("nodeOutput:{}",nodeOutput);
  42. //获取工作流中存储的数据
  43. Map<String, Object> data = nodeOutput.state().data();
  44. log.debug("data:{}",data);
  45. }
  46. });
  47. }
  48. @GetMapping("/createStance")
  49. public Map<String, Object> createStance(@RequestParam("word")String word){
  50. return createStanceGraph.stream(Map.of("word", word)).log().last().block().state().data();
  51. }
  52. @GetMapping("/condition")
  53. public Map<String, Object> condition(@RequestParam("topic")String topic){
  54. return conditionGraph.stream(Map.of("topic", topic)).log().last().block().state().data();
  55. }
  56. @GetMapping("/loop")
  57. public Map<String, Object> loop(@RequestParam("topic")String topic){
  58. return loopGraph.stream(Map.of("topic", topic)).log().last().block().state().data();
  59. }
  60. @GetMapping("/saveTest")
  61. public Map<String, Object> saveTest(@RequestParam("msg")String msg,@RequestParam("id") String id){
  62. //graph工作流的运行配置
  63. RunnableConfig runnableConfig = RunnableConfig.builder()
  64. .threadId(id)
  65. .build();
  66. return saveGraph.stream(Map.of("msg", msg),runnableConfig)
  67. .log()
  68. .last()
  69. .block()
  70. .state()
  71. .data();
  72. }
  73. }