diff --git a/langgraph-go/advancedgraph/graph.go b/langgraph-go/advancedgraph/graph.go index f592f9578..1ab681fb8 100644 --- a/langgraph-go/advancedgraph/graph.go +++ b/langgraph-go/advancedgraph/graph.go @@ -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 { diff --git a/langgraph-go/advancedgraph/rust_engine.go b/langgraph-go/advancedgraph/rust_engine.go index 01607235e..24ebdad92 100644 --- a/langgraph-go/advancedgraph/rust_engine.go +++ b/langgraph-go/advancedgraph/rust_engine.go @@ -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), ) diff --git a/langgraph-go/tests/sub_agents_test.go b/langgraph-go/tests/sub_agents_test.go index 38f6467d5..a03c16523 100644 --- a/langgraph-go/tests/sub_agents_test.go +++ b/langgraph-go/tests/sub_agents_test.go @@ -179,6 +179,7 @@ func TestSubAgentsEquivalentFlow(t *testing.T) { "output": []string{}, "done": nil, }, + nil, ) if err != nil { t.Fatalf("start failed: %v", err) diff --git a/langgraph-go/tests/test_primitives_test.go b/langgraph-go/tests/test_primitives_test.go index e56bd6a1f..b58b75fc8 100644 --- a/langgraph-go/tests/test_primitives_test.go +++ b/langgraph-go/tests/test_primitives_test.go @@ -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) } } diff --git a/rust-core/include/langgraph_rust_core.h b/rust-core/include/langgraph_rust_core.h index d7c9f6544..c00880bde 100644 --- a/rust-core/include/langgraph_rust_core.h +++ b/rust-core/include/langgraph_rust_core.h @@ -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 ); diff --git a/rust-core/src/lib_c.rs b/rust-core/src/lib_c.rs index 8d79240ec..c4da9cfdc 100644 --- a/rust-core/src/lib_c.rs +++ b/rust-core/src/lib_c.rs @@ -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 { @@ -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, ) -> *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::>(); 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, );