浏览代码

feat:实现graph多分支带条件边的demo;
### 实现了一个生成笑话的工作流
- node1根据主题生成笑话;
- node2给笑话打分;
- node3笑话优化;
> 如果笑话不够优秀就使用笑话优化节点进行优化,如何够优秀就直接结束然后输出.

yangyi 7 月之前
父节点
当前提交
4e7b1132bc

+ 49 - 2
spring-ai-learn/src/main/java/space/anyi/springAiAlibabaLearn/config/GraphConfig.java

@@ -1,8 +1,7 @@
 package space.anyi.springAiAlibabaLearn.config;
 
 import com.alibaba.cloud.ai.graph.*;
-import com.alibaba.cloud.ai.graph.action.AsyncNodeAction;
-import com.alibaba.cloud.ai.graph.action.NodeAction;
+import com.alibaba.cloud.ai.graph.action.*;
 import com.alibaba.cloud.ai.graph.exception.GraphStateException;
 import org.slf4j.Logger;
 import org.slf4j.LoggerFactory;
@@ -12,6 +11,9 @@ import org.springframework.ai.chat.model.ChatModel;
 import org.springframework.beans.factory.annotation.Qualifier;
 import org.springframework.context.annotation.Bean;
 import org.springframework.context.annotation.Configuration;
+import space.anyi.springAiAlibabaLearn.nodeAction.GenerationAsyncNodeAction;
+import space.anyi.springAiAlibabaLearn.nodeAction.ImproveJokeAsyncNodeAction;
+import space.anyi.springAiAlibabaLearn.nodeAction.JudgeJokeAsyncNodeAction;
 
 import java.util.List;
 import java.util.Map;
@@ -150,4 +152,49 @@ public class GraphConfig {
         //4.编译图,最终使用的graph
         return  stateGraph.compile();
     }
+    @Bean
+    @Qualifier("conditionGraph")
+    public CompiledGraph conditionGraph(ChatModel chatModel) throws GraphStateException {
+        KeyStrategyFactory keyStrategyFactory = new KeyStrategyFactory(){
+            @Override
+            public Map<String, KeyStrategy> apply() {
+                return Map.of(
+                        "topic",KeyStrategy.REPLACE,
+                        "joke",KeyStrategy.REPLACE,
+                        "result",KeyStrategy.REPLACE
+                );
+            }
+        };
+        //1.创建状态图,用于存储工作流中的数据
+        StateGraph stateGraph = new StateGraph(keyStrategyFactory);
+
+        //2.向工作流添加节点
+        stateGraph.addNode("generation", new GenerationAsyncNodeAction(chatModel));
+        stateGraph.addNode("judge", new JudgeJokeAsyncNodeAction(chatModel));
+        stateGraph.addNode("improve", new ImproveJokeAsyncNodeAction(chatModel));
+        //3.定义图中的边,边的作用是指明节点间的流转关系
+        stateGraph.addEdge(StateGraph.START,"generation");
+        stateGraph.addEdge("generation","judge");
+        /**
+         * 条件边,根据是上一个节点的数据条件连接下一个节点
+         */
+        stateGraph.addConditionalEdges("judge",
+                //条件结果的处理
+                AsyncEdgeAction.edge_async(new EdgeAction() {
+                    @Override
+                    public String apply(OverAllState state) throws Exception {
+                        //取笑话的评分作为条件
+                        return state.value("result","优秀");
+                    }
+                }),
+                //根据条件值如何进行工作流节点流转
+                Map.of(
+                        //key为条件值,value为node的id
+                        "优秀",StateGraph.END,
+                        "不优秀","improve"
+                ));
+        stateGraph.addEdge("improve",StateGraph.END);
+        //4.编译图,最终使用的graph
+        return  stateGraph.compile();
+    }
 }

+ 9 - 2
spring-ai-learn/src/main/java/space/anyi/springAiAlibabaLearn/controller/GraphController.java

@@ -2,7 +2,6 @@ package space.anyi.springAiAlibabaLearn.controller;
 
 import com.alibaba.cloud.ai.graph.CompiledGraph;
 import com.alibaba.cloud.ai.graph.NodeOutput;
-import jakarta.annotation.Resource;
 import org.slf4j.Logger;
 import org.slf4j.LoggerFactory;
 import org.springframework.beans.factory.annotation.Autowired;
@@ -22,10 +21,14 @@ public class GraphController {
     private static final Logger log = LoggerFactory.getLogger(GraphController.class);
     private final CompiledGraph compiledGraph;
     private final CompiledGraph createStanceGraph;
+    private final CompiledGraph conditionGraph;
 
-    public GraphController(@Autowired @Qualifier("compiledGraph")CompiledGraph compiledGraph, @Autowired @Qualifier("createStanceGraph")CompiledGraph createStanceGraph) {
+    public GraphController(@Autowired @Qualifier("compiledGraph")CompiledGraph compiledGraph,
+                           @Autowired @Qualifier("createStanceGraph")CompiledGraph createStanceGraph,
+                           @Autowired@Qualifier("conditionGraph")CompiledGraph conditionGraph) {
         this.compiledGraph = compiledGraph;
         this.createStanceGraph = createStanceGraph;
+        this.conditionGraph = conditionGraph;
     }
 
     @GetMapping("/quickStart")
@@ -44,4 +47,8 @@ public class GraphController {
     public Map<String, Object> createStance(@RequestParam("word")String word){
         return createStanceGraph.stream(Map.of("word", word)).log().last().block().state().data();
     }
+    @GetMapping("/condition")
+    public Map<String, Object> condition(@RequestParam("topic")String topic){
+        return conditionGraph.stream(Map.of("topic", topic)).log().last().block().state().data();
+    }
 }

+ 36 - 0
spring-ai-learn/src/main/java/space/anyi/springAiAlibabaLearn/nodeAction/GenerationAsyncNodeAction.java

@@ -0,0 +1,36 @@
+package space.anyi.springAiAlibabaLearn.nodeAction;
+
+import com.alibaba.cloud.ai.graph.OverAllState;
+import com.alibaba.cloud.ai.graph.action.AsyncNodeAction;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+import org.springframework.ai.chat.client.ChatClient;
+import org.springframework.ai.chat.model.ChatModel;
+
+import java.util.Map;
+import java.util.concurrent.CompletableFuture;
+
+/**
+ * 根据主题生成笑话的节点
+ */
+public class GenerationAsyncNodeAction implements AsyncNodeAction {
+    private static final Logger log = LoggerFactory.getLogger(GenerationAsyncNodeAction.class);
+    public final ChatClient chatClient;
+
+    public GenerationAsyncNodeAction(ChatModel chatModel) {
+        this.chatClient = ChatClient.builder(chatModel).build();
+    }
+    @Override
+    public CompletableFuture<Map<String, Object>> apply(OverAllState state) {
+        return CompletableFuture.supplyAsync(() -> {
+            String content = chatClient
+                    .prompt()  // 开始构建请求
+                    .system("你是一个喜剧专家.专门根据用户输入的主题来创作笑话.")  // 设置系统消息
+                    .user(state.value("topic", "生活"))  // 设置用户消息(主题)
+                    .call()  // 执行调用
+                    .content();  // 获取响应内容
+
+            return Map.of("joke", content != null ? content : "无法生成笑话");
+        });
+    }
+}

+ 36 - 0
spring-ai-learn/src/main/java/space/anyi/springAiAlibabaLearn/nodeAction/ImproveJokeAsyncNodeAction.java

@@ -0,0 +1,36 @@
+package space.anyi.springAiAlibabaLearn.nodeAction;
+
+import com.alibaba.cloud.ai.graph.OverAllState;
+import com.alibaba.cloud.ai.graph.action.AsyncNodeAction;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+import org.springframework.ai.chat.client.ChatClient;
+import org.springframework.ai.chat.model.ChatModel;
+
+import java.util.Map;
+import java.util.concurrent.CompletableFuture;
+
+/**
+ * 优化笑话的节点
+ */
+public class ImproveJokeAsyncNodeAction implements AsyncNodeAction {
+    private static final Logger log = LoggerFactory.getLogger(ImproveJokeAsyncNodeAction.class);
+    public final ChatClient chatClient;
+
+    public ImproveJokeAsyncNodeAction(ChatModel chatModel) {
+        this.chatClient = ChatClient.builder(chatModel).build();
+    }
+    @Override
+    public CompletableFuture<Map<String, Object>> apply(OverAllState state) {
+        return CompletableFuture.supplyAsync(() -> {
+            String content = chatClient
+                    .prompt()  // 开始构建请求
+                    .system("你是一个喜剧专家.专门负责优化用户创作的笑话.")  // 设置系统消息
+                    .user(state.value("joke", ""))  // 设置用户消息(主题)
+                    .call()  // 执行调用
+                    .content();  // 获取响应内容
+
+            return Map.of("joke", content != null ? content : "无法生成笑话");
+        });
+    }
+}

+ 37 - 0
spring-ai-learn/src/main/java/space/anyi/springAiAlibabaLearn/nodeAction/JudgeJokeAsyncNodeAction.java

@@ -0,0 +1,37 @@
+package space.anyi.springAiAlibabaLearn.nodeAction;
+
+import com.alibaba.cloud.ai.graph.OverAllState;
+import com.alibaba.cloud.ai.graph.action.AsyncNodeAction;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+import org.springframework.ai.chat.client.ChatClient;
+import org.springframework.ai.chat.model.ChatModel;
+
+import java.util.Map;
+import java.util.concurrent.CompletableFuture;
+
+/**
+ *给笑话打分的节点
+ */
+public class JudgeJokeAsyncNodeAction implements AsyncNodeAction {
+    private static final Logger log = LoggerFactory.getLogger(JudgeJokeAsyncNodeAction.class);
+    public final ChatClient chatClient;
+
+    public JudgeJokeAsyncNodeAction(ChatModel chatModel) {
+        this.chatClient = ChatClient.builder(chatModel).build();
+    }
+
+    @Override
+    public CompletableFuture<Map<String, Object>> apply(OverAllState state) {
+        return CompletableFuture.supplyAsync(() -> {
+            String content = chatClient
+                    .prompt()  // 开始构建请求
+                    .system("你是一个喜剧专家.专门给用户输入的笑话进行打分,打分结果:优秀,不优秀.仅输出打分结果.")  // 设置系统消息
+                    .user(state.value("joke", ""))  // 设置用户消息(主题)
+                    .call()  // 执行调用
+                    .content();  // 获取响应内容
+
+            return Map.of("result", content != null ? content : "优秀");
+        });
+    }
+}