mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-24 16:42:24 +02:00
more
This commit is contained in:
@@ -49,6 +49,18 @@ func (g *AdvancedStateGraph) SetFinishNode(fn any) {
|
||||
g.finishPoint = NodeName(fn)
|
||||
}
|
||||
|
||||
func (g *AdvancedStateGraph) AddEntryNode(fn any) string {
|
||||
name := g.AddNode(fn)
|
||||
g.entryPoint = name
|
||||
return name
|
||||
}
|
||||
|
||||
func (g *AdvancedStateGraph) AddFinishNode(fn any) string {
|
||||
name := g.AddNode(fn)
|
||||
g.finishPoint = name
|
||||
return name
|
||||
}
|
||||
|
||||
func (g *AdvancedStateGraph) Compile() *CompiledGraph {
|
||||
return &CompiledGraph{
|
||||
nodes: g.nodes,
|
||||
@@ -96,7 +108,7 @@ func (h *Handler) WaitForResult() (map[string]any, error) {
|
||||
return res.state, res.err
|
||||
}
|
||||
|
||||
func (g *CompiledGraph) Start(initialState map[string]any) (*Handler, error) {
|
||||
func (g *CompiledGraph) Start(initialState map[string]any, initialInput any) (*Handler, error) {
|
||||
engine := NewRustEngine()
|
||||
for _, ch := range g.asyncChannels {
|
||||
if err := engine.AddAsyncChannel(ch); err != nil {
|
||||
@@ -114,6 +126,7 @@ func (g *CompiledGraph) Start(initialState map[string]any) (*Handler, error) {
|
||||
g.entryPoint,
|
||||
g.finishPoint,
|
||||
initialState,
|
||||
initialInput,
|
||||
func(node string, nodeInput any, fallbackState map[string]any) (Command, error) {
|
||||
fn, ok := g.nodes[node]
|
||||
if !ok {
|
||||
|
||||
@@ -139,18 +139,25 @@ func (e *RustEngine) RunGraph(
|
||||
entryPoint string,
|
||||
finishPoint string,
|
||||
initialState map[string]any,
|
||||
initialInput any,
|
||||
exec func(node string, nodeInput any, state map[string]any) (Command, error),
|
||||
) (map[string]any, error) {
|
||||
initialJSON, err := json.Marshal(initialState)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal initial state: %w", err)
|
||||
}
|
||||
initialInputJSON, err := json.Marshal(initialInput)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal initial input: %w", err)
|
||||
}
|
||||
centry := C.CString(entryPoint)
|
||||
cfinish := C.CString(finishPoint)
|
||||
cinitial := C.CString(string(initialJSON))
|
||||
cinitialInput := C.CString(string(initialInputJSON))
|
||||
defer C.free(unsafe.Pointer(centry))
|
||||
defer C.free(unsafe.Pointer(cfinish))
|
||||
defer C.free(unsafe.Pointer(cinitial))
|
||||
defer C.free(unsafe.Pointer(cinitialInput))
|
||||
|
||||
handle := cgo.NewHandle(&runGraphCallbackCtx{exec: exec})
|
||||
defer handle.Delete()
|
||||
@@ -160,6 +167,7 @@ func (e *RustEngine) RunGraph(
|
||||
centry,
|
||||
cfinish,
|
||||
cinitial,
|
||||
cinitialInput,
|
||||
C.ulong(handle),
|
||||
(C.rc_node_callback_t)(C.goNodeCallback),
|
||||
)
|
||||
|
||||
@@ -179,6 +179,7 @@ func TestSubAgentsEquivalentFlow(t *testing.T) {
|
||||
"output": []string{},
|
||||
"done": nil,
|
||||
},
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("start failed: %v", err)
|
||||
|
||||
@@ -67,16 +67,14 @@ func TestInputAndStatePrimitivesCompatible(t *testing.T) {
|
||||
workflow := &primitiveWorkflow{}
|
||||
graph := ag.NewAdvancedStateGraph()
|
||||
|
||||
graph.AddNode(workflow.startNode)
|
||||
graph.AddEntryNode(workflow.startNode)
|
||||
graph.AddNode(workflow.middleNode)
|
||||
graph.AddNode(workflow.finishNode)
|
||||
graph.SetEntryNode(workflow.startNode)
|
||||
graph.SetFinishNode(workflow.finishNode)
|
||||
graph.AddFinishNode(workflow.finishNode)
|
||||
|
||||
handler, err := graph.Compile().Start(map[string]any{
|
||||
"logs": []string{},
|
||||
"done": nil,
|
||||
})
|
||||
}, 100)
|
||||
if err != nil {
|
||||
t.Fatalf("start failed: %v", err)
|
||||
}
|
||||
@@ -89,7 +87,7 @@ func TestInputAndStatePrimitivesCompatible(t *testing.T) {
|
||||
t.Fatalf("unexpected done: %v", result["done"])
|
||||
}
|
||||
logs := logsSlice(result)
|
||||
if len(logs) != 3 || logs[0] != "start:0" || logs[1] != "middle:from_start" || logs[2] != "finish:from_middle" {
|
||||
if len(logs) != 3 || logs[0] != "start:100" || logs[1] != "middle:from_start" || logs[2] != "finish:from_middle" {
|
||||
t.Fatalf("unexpected logs: %#v", logs)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -24,6 +24,7 @@ char* rc_run_graph_json(
|
||||
const char* entry_point,
|
||||
const char* finish_point,
|
||||
const char* initial_state_json,
|
||||
const char* initial_input_json,
|
||||
unsigned long user_data,
|
||||
rc_node_callback_t callback
|
||||
);
|
||||
|
||||
+16
-1
@@ -136,6 +136,7 @@ fn run_graph_scheduler_json(
|
||||
entry_point: String,
|
||||
finish_point: String,
|
||||
initial_state: Value,
|
||||
initial_input: Value,
|
||||
user_data: CUserData,
|
||||
callback: CNodeCallback,
|
||||
) -> Result<Value, String> {
|
||||
@@ -145,7 +146,7 @@ fn run_graph_scheduler_json(
|
||||
let tx_for_spawn = tx.clone();
|
||||
let state_for_spawn = Arc::clone(&state);
|
||||
let state_for_merge = Arc::clone(&state);
|
||||
let initial_arg = Value::Null;
|
||||
let initial_arg = initial_input;
|
||||
run_scheduler_loop(
|
||||
entry_point,
|
||||
&finish_point,
|
||||
@@ -297,6 +298,7 @@ pub unsafe extern "C" fn rc_run_graph_json(
|
||||
entry_point: *const c_char,
|
||||
finish_point: *const c_char,
|
||||
initial_state_json: *const c_char,
|
||||
initial_input_json: *const c_char,
|
||||
user_data: libc::c_ulong,
|
||||
callback: Option<CNodeCallback>,
|
||||
) -> *mut c_char {
|
||||
@@ -326,6 +328,18 @@ pub unsafe extern "C" fn rc_run_graph_json(
|
||||
))
|
||||
}
|
||||
};
|
||||
let initial_input_json = match cstr_to_str(initial_input_json) {
|
||||
Ok(v) => v,
|
||||
Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
|
||||
};
|
||||
let initial_input: Value = match serde_json::from_str(initial_input_json) {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
return into_c_ptr(format!(
|
||||
"{{\"ok\":false,\"error\":\"invalid initial_input JSON: {e}\"}}"
|
||||
))
|
||||
}
|
||||
};
|
||||
|
||||
let (tx, rx) = mpsc::channel::<Result<Value, String>>();
|
||||
let user_data = CUserData(user_data);
|
||||
@@ -334,6 +348,7 @@ pub unsafe extern "C" fn rc_run_graph_json(
|
||||
entry_point,
|
||||
finish_point,
|
||||
initial_state,
|
||||
initial_input,
|
||||
user_data,
|
||||
callback,
|
||||
);
|
||||
|
||||
Reference in New Issue
Block a user