Explorar o código

feat:探究stateGraph的存储,进行图编译的配置,进行图调用时的配置
### 手动进行图的编译配置
配置使用内存进行stateGraph存储的策略
### 调用graph时手动配置运行时的配置

yangyi hai 7 meses
pai
achega
8ca14f7fd3

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

@@ -2,6 +2,8 @@ package space.anyi.springAiAlibabaLearn.config;
 
 import com.alibaba.cloud.ai.graph.*;
 import com.alibaba.cloud.ai.graph.action.*;
+import com.alibaba.cloud.ai.graph.checkpoint.config.SaverConfig;
+import com.alibaba.cloud.ai.graph.checkpoint.savers.MemorySaver;
 import com.alibaba.cloud.ai.graph.exception.GraphStateException;
 import org.slf4j.Logger;
 import org.slf4j.LoggerFactory;
@@ -19,6 +21,8 @@ import space.anyi.springAiAlibabaLearn.nodeAction.LoopJudgeAsyncNodeAction;
 import java.util.List;
 import java.util.Map;
 import java.util.Optional;
+import java.util.concurrent.CompletableFuture;
+import java.util.function.Supplier;
 
 @Configuration
 public class GraphConfig {
@@ -267,4 +271,36 @@ public class GraphConfig {
         //4.编译图,最终使用的graph
         return  stateGraph.compile();
     }
+    @Bean
+    @Qualifier("saveGraph")
+    public CompiledGraph saveGraph() throws GraphStateException {
+        StateGraph stateGraph = new StateGraph(() -> Map.of(
+                "msg",KeyStrategy.REPLACE,
+                "history",KeyStrategy.APPEND
+        ));
+        stateGraph.addNode("save", state -> CompletableFuture.supplyAsync(() -> {
+            //取出msg
+            String msg = state.value("msg", "");
+            log.debug("msg:{}",msg);
+            //更新历史消息
+            return Map.of(
+                    "history",msg
+            );
+        }));
+        stateGraph.addEdge(StateGraph.START,"save");
+        stateGraph.addEdge("save",StateGraph.END);
+        //获取图的信息,mermaid格式,可以在markdown中渲染出来
+        GraphRepresentation graphRepresentation = stateGraph.getGraph(GraphRepresentation.Type.MERMAID);
+        log.info("graphRepresentation:{}",graphRepresentation.content());
+        //图的状态存储配置,使用内存的策略
+        SaverConfig saverConfig = SaverConfig.builder()
+                .register(new MemorySaver())
+                .build();
+        //图的编译配置
+        CompileConfig compileConfig = CompileConfig.builder()
+                .saverConfig(saverConfig)
+                .build();
+        //使用配置进行图的编译
+        return stateGraph.compile(compileConfig);
+    }
 }

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

@@ -2,6 +2,7 @@ package space.anyi.springAiAlibabaLearn.controller;
 
 import com.alibaba.cloud.ai.graph.CompiledGraph;
 import com.alibaba.cloud.ai.graph.NodeOutput;
+import com.alibaba.cloud.ai.graph.RunnableConfig;
 import org.slf4j.Logger;
 import org.slf4j.LoggerFactory;
 import org.springframework.beans.factory.annotation.Autowired;
@@ -23,15 +24,18 @@ public class GraphController {
     private final CompiledGraph createStanceGraph;
     private final CompiledGraph conditionGraph;
     private final CompiledGraph loopGraph;
+    private final CompiledGraph saveGraph;
 
     public GraphController(@Autowired @Qualifier("compiledGraph")CompiledGraph compiledGraph,
                            @Autowired @Qualifier("createStanceGraph")CompiledGraph createStanceGraph,
-                           @Autowired@Qualifier("conditionGraph")CompiledGraph conditionGraph,
-                           @Autowired@Qualifier("loopGraph")CompiledGraph loopGraph) {
+                           @Autowired @Qualifier("conditionGraph")CompiledGraph conditionGraph,
+                           @Autowired @Qualifier("loopGraph")CompiledGraph loopGraph,
+                           @Autowired @Qualifier("saveGraph")CompiledGraph saveGraph) {
         this.compiledGraph = compiledGraph;
         this.createStanceGraph = createStanceGraph;
         this.conditionGraph = conditionGraph;
         this.loopGraph = loopGraph;
+        this.saveGraph = saveGraph;
     }
 
     @GetMapping("/quickStart")
@@ -58,4 +62,17 @@ public class GraphController {
     public Map<String, Object> loop(@RequestParam("topic")String topic){
         return loopGraph.stream(Map.of("topic", topic)).log().last().block().state().data();
     }
+    @GetMapping("/saveTest")
+    public Map<String, Object> saveTest(@RequestParam("msg")String msg,@RequestParam("id") String id){
+        //graph工作流的运行配置
+        RunnableConfig runnableConfig = RunnableConfig.builder()
+                .threadId(id)
+                .build();
+        return saveGraph.stream(Map.of("msg", msg),runnableConfig)
+                .log()
+                .last()
+                .block()
+                .state()
+                .data();
+    }
 }