Ver Fonte

feat:使用graph实现具有循环的工作流
> 使用条件边将节点组成环从而实现循环的工作流
1. 生成笑话
2. 笑话打分
2.1 如果笑话不够优秀就重新生成再打分
3.结束

yangyi há 7 meses atrás
pai
commit
37eaf6b90c

+ 70 - 0
spring-ai-learn/src/main/java/space/anyi/springAiAlibabaLearn/config/GraphConfig.java

@@ -14,6 +14,7 @@ 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 space.anyi.springAiAlibabaLearn.nodeAction.LoopJudgeAsyncNodeAction;
 
 import java.util.List;
 import java.util.Map;
@@ -152,6 +153,18 @@ public class GraphConfig {
         //4.编译图,最终使用的graph
         return  stateGraph.compile();
     }
+
+    /**
+     * 有条件分支的工作流
+     * # 生成笑话的工作流
+     * 1.生成笑话
+     * 2.打分
+     * 2.1如果笑话不够优秀,则进行优化
+     * 3.结束
+     * @param chatModel
+     * @return
+     * @throws GraphStateException
+     */
     @Bean
     @Qualifier("conditionGraph")
     public CompiledGraph conditionGraph(ChatModel chatModel) throws GraphStateException {
@@ -197,4 +210,61 @@ public class GraphConfig {
         //4.编译图,最终使用的graph
         return  stateGraph.compile();
     }
+
+    /**
+     * 有循环的工作流
+     * # 生成笑话的工作流
+     * 1.生成笑话
+     * 2.打分
+     * 2.1笑话不够优秀则重新生成,然后再进行打分
+     * 3.结束
+     * @param chatModel
+     * @return
+     * @throws GraphStateException
+     */
+    @Bean
+    @Qualifier("loopGraph")
+    public CompiledGraph loopGraph(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,
+                        "loopCount",KeyStrategy.REPLACE
+                );
+            }
+        };
+        //1.创建状态图,用于存储工作流中的数据
+        StateGraph stateGraph = new StateGraph(keyStrategyFactory);
+
+        //2.向工作流添加节点
+        stateGraph.addNode("generation", new GenerationAsyncNodeAction(chatModel));
+        stateGraph.addNode("loopJudge", new LoopJudgeAsyncNodeAction(chatModel,5));
+        //3.定义图中的边,边的作用是指明节点间的流转关系
+        stateGraph.addEdge(StateGraph.START,"generation");
+        stateGraph.addEdge("generation","loopJudge");
+        /**
+         * 条件边
+         * 根据条件边构建一个环,这个环就是循环
+         */
+        stateGraph.addConditionalEdges("loopJudge",
+                //条件结果的处理
+                AsyncEdgeAction.edge_async(new EdgeAction() {
+                    @Override
+                    public String apply(OverAllState state) throws Exception {
+                        //取笑话的评分作为条件
+                        return state.value("result","break");
+                    }
+                }),
+                //根据条件值如何进行工作流节点流转
+                Map.of(
+                        //key为条件值,value为node的id
+                        "break",StateGraph.END,
+                        "loop","generation"
+                ));
+        //4.编译图,最终使用的graph
+        return  stateGraph.compile();
+    }
 }

+ 8 - 1
spring-ai-learn/src/main/java/space/anyi/springAiAlibabaLearn/controller/GraphController.java

@@ -22,13 +22,16 @@ public class GraphController {
     private final CompiledGraph compiledGraph;
     private final CompiledGraph createStanceGraph;
     private final CompiledGraph conditionGraph;
+    private final CompiledGraph loopGraph;
 
     public GraphController(@Autowired @Qualifier("compiledGraph")CompiledGraph compiledGraph,
                            @Autowired @Qualifier("createStanceGraph")CompiledGraph createStanceGraph,
-                           @Autowired@Qualifier("conditionGraph")CompiledGraph conditionGraph) {
+                           @Autowired@Qualifier("conditionGraph")CompiledGraph conditionGraph,
+                           @Autowired@Qualifier("loopGraph")CompiledGraph loopGraph) {
         this.compiledGraph = compiledGraph;
         this.createStanceGraph = createStanceGraph;
         this.conditionGraph = conditionGraph;
+        this.loopGraph = loopGraph;
     }
 
     @GetMapping("/quickStart")
@@ -51,4 +54,8 @@ public class GraphController {
     public Map<String, Object> condition(@RequestParam("topic")String topic){
         return conditionGraph.stream(Map.of("topic", topic)).log().last().block().state().data();
     }
+    @GetMapping("/loop")
+    public Map<String, Object> loop(@RequestParam("topic")String topic){
+        return loopGraph.stream(Map.of("topic", topic)).log().last().block().state().data();
+    }
 }

+ 49 - 0
spring-ai-learn/src/main/java/space/anyi/springAiAlibabaLearn/nodeAction/LoopJudgeAsyncNodeAction.java

@@ -0,0 +1,49 @@
+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 LoopJudgeAsyncNodeAction implements AsyncNodeAction {
+    private static final Logger log = LoggerFactory.getLogger(LoopJudgeAsyncNodeAction.class);
+    public final ChatClient chatClient;
+    //最大循环次数
+    public final int maxLoopCount;
+
+    public LoopJudgeAsyncNodeAction(ChatModel chatModel,int maxLoopCount) {
+        this.chatClient = ChatClient.builder(chatModel).build();
+        this.maxLoopCount = maxLoopCount;
+    }
+
+    @Override
+    public CompletableFuture<Map<String, Object>> apply(OverAllState state) {
+        return CompletableFuture.supplyAsync(() -> {
+            Integer loopCount = state.value("loopCount", 0);
+            loopCount++;
+            String joke = state.value("joke", "");
+            String content = chatClient
+                    .prompt()  // 开始构建请求
+                    .system("你是一个喜剧专家.专门给用户输入的笑话进行打分,打分范围:[0,10].仅输出打分的整数结果.")  // 设置系统消息
+                    .user(joke)  // 设置用户消息(主题)
+                    .call()  // 执行调用
+                    .content();  // 获取响应内容
+            Integer score = Integer.valueOf(content);
+            log.info("joke:{},source:{},loopCount:{}",joke,score,loopCount);
+            String result = "loop";
+            //跳出循环的控制变量
+            if (score>=8 || loopCount >= maxLoopCount){
+                result = "break";
+            }
+            return Map.of(
+                    "result", result,
+                    "loopCount",loopCount
+                    );
+        });
+    }
+}