From d50801d87151de0c2465e742a5f5c568180ceb83 Mon Sep 17 00:00:00 2001 From: Quanzheng Long Date: Fri, 13 Mar 2026 13:16:11 -0700 Subject: [PATCH] more --- langgraph-go/advancedgraph/graph.go | 12 +-- langgraph-go/advancedgraph/rust_engine.go | 20 ++-- langgraph-go/advancedgraph/types.go | 4 +- langgraph-go/tests/sub_agents_test.go | 43 +++------ langgraph-go/tests/test_primitives_test.go | 95 +++++++++++++++++++ .../langgraph/advanced_graph/state.py | 87 +++++++++++++---- .../tests/advanced-graph/test_primitives.py | 46 +++++++++ rust-core/include/langgraph_rust_core.h | 4 +- rust-core/src/lib_c.rs | 16 ++-- 9 files changed, 246 insertions(+), 81 deletions(-) create mode 100644 langgraph-go/tests/test_primitives_test.go create mode 100644 libs/langgraph/tests/advanced-graph/test_primitives.py diff --git a/langgraph-go/advancedgraph/graph.go b/langgraph-go/advancedgraph/graph.go index 54a0113be..f98ead71c 100644 --- a/langgraph-go/advancedgraph/graph.go +++ b/langgraph-go/advancedgraph/graph.go @@ -7,7 +7,7 @@ import ( "strings" ) -type NodeFunc func(ctx *Context, state map[string]any) (Command, error) +type NodeFunc func(ctx *Context, input any, state map[string]any) (Command, error) type AdvancedStateGraph struct { nodes map[string]NodeFunc @@ -108,19 +108,15 @@ func (g *CompiledGraph) Start(initialState map[string]any) (*Handler, error) { g.entryPoint, g.finishPoint, initialState, - func(node string, arg any, fallbackState map[string]any) (Command, error) { + func(node string, nodeInput any, fallbackState map[string]any) (Command, error) { fn, ok := g.nodes[node] if !ok { return Command{}, fmt.Errorf("unknown node `%s`", node) } - stateArg, ok := arg.(map[string]any) - if !ok { - stateArg = fallbackState - } - if stateArg == nil { + if fallbackState == nil { return Command{}, fmt.Errorf("node `%s` expected map state argument", node) } - return fn(&Context{engine: engine}, stateArg) + return fn(&Context{engine: engine}, nodeInput, fallbackState) }, ) handler.done <- resultOrErr{state: state, err: err} diff --git a/langgraph-go/advancedgraph/rust_engine.go b/langgraph-go/advancedgraph/rust_engine.go index cccd67bf5..13cf73560 100644 --- a/langgraph-go/advancedgraph/rust_engine.go +++ b/langgraph-go/advancedgraph/rust_engine.go @@ -5,7 +5,7 @@ package advancedgraph #cgo LDFLAGS: -L${SRCDIR}/../../rust-core/target/debug -llanggraph_rust_core #include "langgraph_rust_core.h" #include -extern char* goNodeCallback(void* user_data, char* node, char* arg_json, char* state_json); +extern char* goNodeCallback(unsigned long user_data, char* node, char* arg_json, char* state_json); */ import "C" @@ -21,11 +21,11 @@ type RustEngine struct { } type runGraphCallbackCtx struct { - exec func(node string, arg any, state map[string]any) (Command, error) + exec func(node string, nodeInput any, state map[string]any) (Command, error) } //export goNodeCallback -func goNodeCallback(userData unsafe.Pointer, node *C.char, argJSON *C.char, stateJSON *C.char) *C.char { +func goNodeCallback(userData C.ulong, node *C.char, argJSON *C.char, stateJSON *C.char) *C.char { handle := cgo.Handle(uintptr(userData)) ctx, ok := handle.Value().(*runGraphCallbackCtx) if !ok { @@ -34,22 +34,22 @@ func goNodeCallback(userData unsafe.Pointer, node *C.char, argJSON *C.char, stat nodeName := C.GoString(node) - var arg any - if err := json.Unmarshal([]byte(C.GoString(argJSON)), &arg); err != nil { + var nodeInput any + if err := json.Unmarshal([]byte(C.GoString(argJSON)), &nodeInput); err != nil { return cCallbackEnvelopeError(fmt.Sprintf("decode arg failed for `%s`: %v", nodeName, err)) } var state map[string]any if err := json.Unmarshal([]byte(C.GoString(stateJSON)), &state); err != nil { return cCallbackEnvelopeError(fmt.Sprintf("decode state failed for `%s`: %v", nodeName, err)) } - arg = coerceJSONValue(arg) + nodeInput = coerceJSONValue(nodeInput) stateAny := coerceJSONValue(state) state, ok = stateAny.(map[string]any) if !ok { return cCallbackEnvelopeError(fmt.Sprintf("decoded state has unexpected type for `%s`", nodeName)) } - cmd, err := ctx.exec(nodeName, arg, state) + cmd, err := ctx.exec(nodeName, nodeInput, state) if err != nil { return cCallbackEnvelopeError(err.Error()) } @@ -58,7 +58,7 @@ func goNodeCallback(userData unsafe.Pointer, node *C.char, argJSON *C.char, stat for _, send := range cmd.Goto { sends = append(sends, map[string]any{ "node": send.Node, - "arg": send.Arg, + "arg": send.NodeInput, }) } payload := map[string]any{ @@ -139,7 +139,7 @@ func (e *RustEngine) RunGraph( entryPoint string, finishPoint string, initialState map[string]any, - exec func(node string, arg any, state map[string]any) (Command, error), + exec func(node string, nodeInput any, state map[string]any) (Command, error), ) (map[string]any, error) { initialJSON, err := json.Marshal(initialState) if err != nil { @@ -160,7 +160,7 @@ func (e *RustEngine) RunGraph( centry, cfinish, cinitial, - unsafe.Pointer(uintptr(handle)), + C.ulong(handle), (C.rc_node_callback_t)(C.goNodeCallback), ) defer C.rc_string_free(resp) diff --git a/langgraph-go/advancedgraph/types.go b/langgraph-go/advancedgraph/types.go index 8748b7be5..bcfc694bd 100644 --- a/langgraph-go/advancedgraph/types.go +++ b/langgraph-go/advancedgraph/types.go @@ -54,8 +54,8 @@ type WaitEvent struct { } type Send struct { - Node string - Arg any + Node string + NodeInput any } type Command struct { diff --git a/langgraph-go/tests/sub_agents_test.go b/langgraph-go/tests/sub_agents_test.go index d3727e637..6696a6622 100644 --- a/langgraph-go/tests/sub_agents_test.go +++ b/langgraph-go/tests/sub_agents_test.go @@ -34,15 +34,6 @@ type lunchWorkflow struct { names map[string]string } -func cloneState(state map[string]any) map[string]any { - out := make(map[string]any, len(state)+2) - for k, v := range state { - out[k] = v - } - out["output"] = append([]string(nil), outputSlice(state)...) - return out -} - func outputSlice(state map[string]any) []string { raw, ok := state["output"] if !ok || raw == nil { @@ -64,35 +55,29 @@ func outputSlice(state map[string]any) []string { } } -func (w *lunchWorkflow) llmNode(ctx *ag.Context, state map[string]any) (ag.Command, error) { +func (w *lunchWorkflow) llmNode(ctx *ag.Context, _ any, state map[string]any) (ag.Command, error) { decisions := w.planner.invoke() sends := make([]ag.Send, 0, 4) for _, d := range decisions { if d.Type == "end" { - next := cloneState(state) - next["complete"] = d.Complete return ag.Command{ Goto: []ag.Send{ - {Node: w.names["order"], Arg: next}, + {Node: w.names["order"], NodeInput: d.Complete}, }, }, nil } if d.Type == "sub_agent" { - next := cloneState(state) - next["sub_agent_input"] = d.SubAgent - sends = append(sends, ag.Send{Node: w.names["sub"], Arg: next}) + sends = append(sends, ag.Send{Node: w.names["sub"], NodeInput: d.SubAgent}) } if d.Type == "tool" { - next := cloneState(state) - next["tool_input"] = d.Tool - sends = append(sends, ag.Send{Node: w.names["tool"], Arg: next}) + sends = append(sends, ag.Send{Node: w.names["tool"], NodeInput: d.Tool}) } } - sends = append(sends, ag.Send{Node: w.names["wait"], Arg: cloneState(state)}) + sends = append(sends, ag.Send{Node: w.names["wait"]}) return ag.Command{Goto: sends}, nil } -func (w *lunchWorkflow) waitNode(ctx *ag.Context, state map[string]any) (ag.Command, error) { +func (w *lunchWorkflow) waitNode(ctx *ag.Context, _ any, state map[string]any) (ag.Command, error) { event, err := ctx.WaitFor( ag.AnyOf( ag.ChannelCondition{Channel: "tool_completion_channel", N: 1}, @@ -117,23 +102,23 @@ func (w *lunchWorkflow) waitNode(ctx *ag.Context, state map[string]any) (ag.Comm output = append(output, "user_input: "+payload) } state["output"] = output - return ag.Command{Goto: []ag.Send{{Node: w.names["llm"], Arg: cloneState(state)}}}, nil + return ag.Command{Goto: []ag.Send{{Node: w.names["llm"]}}, Update: state}, nil } output = append(output, "timer: no updates yet") state["output"] = output - return ag.Command{Goto: []ag.Send{{Node: w.names["wait"], Arg: cloneState(state)}}}, nil + return ag.Command{Goto: []ag.Send{{Node: w.names["wait"]}}, Update: state}, nil } -func (w *lunchWorkflow) toolNode(ctx *ag.Context, state map[string]any) (ag.Command, error) { - toolInput, _ := state["tool_input"].(string) +func (w *lunchWorkflow) toolNode(ctx *ag.Context, input any, _ map[string]any) (ag.Command, error) { + toolInput, _ := input.(string) time.Sleep(100 * time.Millisecond) err := ctx.PublishToChannel("tool_completion_channel", "tool completed for: "+toolInput) return ag.Command{}, err } -func (w *lunchWorkflow) subAgentNode(ctx *ag.Context, state map[string]any) (ag.Command, error) { - subInput, _ := state["sub_agent_input"].(string) +func (w *lunchWorkflow) subAgentNode(ctx *ag.Context, input any, _ map[string]any) (ag.Command, error) { + subInput, _ := input.(string) time.Sleep(5 * time.Second) err := ctx.PublishToChannel( "subagent_completion_channel", @@ -142,8 +127,8 @@ func (w *lunchWorkflow) subAgentNode(ctx *ag.Context, state map[string]any) (ag. return ag.Command{}, err } -func (w *lunchWorkflow) orderFoodNode(ctx *ag.Context, state map[string]any) (ag.Command, error) { - complete, _ := state["complete"].(string) +func (w *lunchWorkflow) orderFoodNode(ctx *ag.Context, input any, state map[string]any) (ag.Command, error) { + complete, _ := input.(string) output := outputSlice(state) output = append(output, "order_food: "+complete) state["output"] = output diff --git a/langgraph-go/tests/test_primitives_test.go b/langgraph-go/tests/test_primitives_test.go new file mode 100644 index 000000000..967e168a1 --- /dev/null +++ b/langgraph-go/tests/test_primitives_test.go @@ -0,0 +1,95 @@ +package tests + +import ( + "testing" + + ag "github.com/langchain-ai/langgraph/langgraph-go/advancedgraph" +) + +type primitiveWorkflow struct { + names map[string]string +} + +func logsSlice(state map[string]any) []string { + raw, ok := state["logs"] + if !ok || raw == nil { + return []string{} + } + switch v := raw.(type) { + case []string: + return v + case []any: + out := make([]string, 0, len(v)) + for _, item := range v { + if s, ok := item.(string); ok { + out = append(out, s) + } + } + return out + default: + return []string{} + } +} + +func (w *primitiveWorkflow) startNode(ctx *ag.Context, _ any, state map[string]any) (ag.Command, error) { + logs := logsSlice(state) + logs = append(logs, "start") + state["logs"] = logs + return ag.Command{ + Update: state, + Goto: []ag.Send{ + {Node: w.names["middle"], NodeInput: "from_start"}, + }, + }, nil +} + +func (w *primitiveWorkflow) middleNode(ctx *ag.Context, input any, state map[string]any) (ag.Command, error) { + logs := logsSlice(state) + logs = append(logs, "middle:"+input.(string)) + state["logs"] = logs + return ag.Command{ + Update: state, + Goto: []ag.Send{ + {Node: w.names["finish"], NodeInput: "from_middle"}, + }, + }, nil +} + +func (w *primitiveWorkflow) finishNode(ctx *ag.Context, input any, state map[string]any) (ag.Command, error) { + logs := logsSlice(state) + logs = append(logs, "finish:"+input.(string)) + state["logs"] = logs + state["done"] = input.(string) + return ag.Command{Update: state}, nil +} + +func TestInputAndStatePrimitivesCompatible(t *testing.T) { + workflow := &primitiveWorkflow{names: make(map[string]string)} + graph := ag.NewAdvancedStateGraph() + + workflow.names["start"] = graph.AddNode(workflow.startNode) + workflow.names["middle"] = graph.AddNode(workflow.middleNode) + workflow.names["finish"] = graph.AddNode(workflow.finishNode) + graph.SetEntryNode(workflow.startNode) + graph.SetFinishNode(workflow.finishNode) + + handler, err := graph.Compile().Start(map[string]any{ + "logs": []string{}, + "done": nil, + }) + if err != nil { + t.Fatalf("start failed: %v", err) + } + + result, err := handler.WaitForResult() + if err != nil { + t.Fatalf("result failed: %v", err) + } + if result["done"] != "from_middle" { + t.Fatalf("unexpected done: %v", result["done"]) + } + logs := logsSlice(result) + if len(logs) != 3 || logs[0] != "start" || logs[1] != "middle:from_start" || logs[2] != "finish:from_middle" { + t.Fatalf("unexpected logs: %#v", logs) + } +} diff --git a/libs/langgraph/langgraph/advanced_graph/state.py b/libs/langgraph/langgraph/advanced_graph/state.py index f192a6833..246ffafaa 100644 --- a/libs/langgraph/langgraph/advanced_graph/state.py +++ b/libs/langgraph/langgraph/advanced_graph/state.py @@ -249,25 +249,27 @@ class _GraphEngineRun: def _publish_sync(self, channel: str, value: Any) -> None: self._rust_engine.publish_obj(channel, value) - def _execute_node_for_rust(self, node_name: str, arg: Any, state: Any) -> dict[str, Any]: + def _execute_node_for_rust( + self, node_name: str, node_input: Any, state: Any + ) -> dict[str, Any]: if node_name not in self._nodes: raise ValueError(f"Unknown node `{node_name}`") node = self._nodes[node_name] - result = _invoke_node(node, self.context, arg) + result = _invoke_node(node, self.context, node_input, state) if inspect.isawaitable(result): result = asyncio.run(cast(Coroutine[Any, Any, Any], result)) if isinstance(result, Command): update = result.update - sends = _normalize_goto(result.goto, default_arg=state) + sends = _normalize_goto(result.goto, default_input=node_input) else: update = result - sends = _normalize_result_to_sends(result, default_arg=state) + sends = _normalize_result_to_sends(result, default_input=node_input) - if update is None and isinstance(arg, dict): + if update is None and isinstance(state, dict): # Preserve in-place state mutations for prototype nodes like wait_node. - update = arg + update = state return { "update": update, @@ -277,46 +279,46 @@ class _GraphEngineRun: ], } -def _normalize_result_to_sends(result: Any, *, default_arg: Any) -> list[Send]: +def _normalize_result_to_sends(result: Any, *, default_input: Any) -> list[Send]: if result is None: return [] if isinstance(result, Send): return [result] if callable(result): - return [Send(_infer_node_name(result), default_arg)] + return [Send(_infer_node_name(result), default_input)] if isinstance(result, str): - return [Send(result, default_arg)] + return [Send(result, default_input)] if isinstance(result, Sequence) and not isinstance(result, (str, bytes)): sends: list[Send] = [] for item in result: if isinstance(item, Send): sends.append(item) elif callable(item): - sends.append(Send(_infer_node_name(item), default_arg)) + sends.append(Send(_infer_node_name(item), default_input)) elif isinstance(item, str): - sends.append(Send(item, default_arg)) + sends.append(Send(item, default_input)) return sends return [] -def _normalize_goto(goto: Any, *, default_arg: Any) -> list[Send]: +def _normalize_goto(goto: Any, *, default_input: Any) -> list[Send]: if not goto: return [] if isinstance(goto, Send): return [goto] if callable(goto): - return [Send(_infer_node_name(goto), default_arg)] + return [Send(_infer_node_name(goto), default_input)] if isinstance(goto, str): - return [Send(goto, default_arg)] + return [Send(goto, default_input)] if isinstance(goto, Sequence): sends: list[Send] = [] for item in goto: if isinstance(item, Send): sends.append(item) elif callable(item): - sends.append(Send(_infer_node_name(item), default_arg)) + sends.append(Send(_infer_node_name(item), default_input)) elif isinstance(item, str): - sends.append(Send(item, default_arg)) + sends.append(Send(item, default_input)) return sends return [] @@ -383,14 +385,57 @@ def _resolve_target_name(target: Any) -> str: raise ValueError(f"Unsupported node target type: {type(target)!r}") -def _invoke_node(node: Callable[..., Any], ctx: Context, state: Any) -> Any: +def _invoke_node(node: Callable[..., Any], ctx: Context, node_input: Any, state: Any) -> Any: try: params = list(inspect.signature(node).parameters.values()) except (TypeError, ValueError): params = [] - if len(params) >= 2: - return node(ctx, state) + if not params: + return node() + + names = [param.name.lower() for param in params] + has_ctx = [("ctx" in name or "context" in name) for name in names] + has_state = [("state" in name) for name in names] + has_input = [("input" in name) for name in names] + + kwargs: dict[str, Any] = {} + unresolved = False + for idx, param in enumerate(params): + if has_ctx[idx]: + kwargs[param.name] = ctx + elif has_state[idx]: + kwargs[param.name] = state + elif has_input[idx]: + kwargs[param.name] = node_input + else: + unresolved = True + + if kwargs and not unresolved: + return node(**kwargs) + if len(params) == 1: - return node(state) - return node() + if has_ctx[0]: + return node(ctx) + if has_state[0]: + return node(state) + return node(node_input) + + if len(params) == 2: + if has_ctx[0] and has_state[1]: + return node(ctx, state) + if has_ctx[0] and has_input[1]: + return node(ctx, node_input) + if has_input[0] and has_state[1]: + return node(node_input, state) + if has_state[0] and has_input[1]: + return node(state, node_input) + if has_state[0]: + return node(state, node_input) + if has_state[1]: + return node(node_input, state) + if has_ctx[0]: + return node(ctx, node_input) + return node(node_input, state) + + return node(ctx, node_input, state) diff --git a/libs/langgraph/tests/advanced-graph/test_primitives.py b/libs/langgraph/tests/advanced-graph/test_primitives.py new file mode 100644 index 000000000..ddc41af59 --- /dev/null +++ b/libs/langgraph/tests/advanced-graph/test_primitives.py @@ -0,0 +1,46 @@ +import pytest +from typing_extensions import TypedDict + +from langgraph.advanced_graph import AdvancedStateGraph, Context +from langgraph.types import Command, Send + +pytestmark = pytest.mark.anyio + + +class PrimitiveState(TypedDict): + counter: int + logs: list[str] + done: str | None + + +async def test_input_and_state_primitives_are_compatible() -> None: + graph = AdvancedStateGraph(PrimitiveState) + + async def start_node(state: PrimitiveState) -> Command: + state["logs"].append(f"start:counter={state['counter']}") + return Command(goto=Send("middle_node", "from_start")) + + async def middle_node(ctx: Context, tool_input: str, state: PrimitiveState) -> Command: + state["logs"].append(f"middle:input={tool_input}") + return Command(update=state, goto=Send("finish_node", "from_middle")) + + async def finish_node(payload: str, state: PrimitiveState) -> dict[str, object]: + state["logs"].append(f"finish:input={payload}") + return { + "logs": state["logs"], + "counter": state["counter"], + "done": payload, + } + + graph.add_entry_node(start_node) + graph.add_node(middle_node) + graph.add_finish_node(finish_node) + + result = await graph.compile().ainvoke({"counter": 7, "logs": [], "done": None}) + assert result["counter"] == 7 + assert result["done"] == "from_middle" + assert result["logs"] == [ + "start:counter=7", + "middle:input=from_start", + "finish:input=from_middle", + ] diff --git a/rust-core/include/langgraph_rust_core.h b/rust-core/include/langgraph_rust_core.h index 81465a338..d7c9f6544 100644 --- a/rust-core/include/langgraph_rust_core.h +++ b/rust-core/include/langgraph_rust_core.h @@ -7,7 +7,7 @@ extern "C" { typedef struct Engine Engine; typedef char* (*rc_node_callback_t)( - void* user_data, + unsigned long user_data, char* node, char* arg_json, char* state_json @@ -24,7 +24,7 @@ char* rc_run_graph_json( const char* entry_point, const char* finish_point, const char* initial_state_json, - void* user_data, + 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 9677d0b15..b874fa545 100644 --- a/rust-core/src/lib_c.rs +++ b/rust-core/src/lib_c.rs @@ -5,7 +5,7 @@ use crate::engine::{ use serde::Deserialize; use serde_json::Value; use std::ffi::{CStr, CString}; -use std::os::raw::{c_char, c_void}; +use std::os::raw::c_char; use std::sync::mpsc; use std::sync::{Arc, Mutex}; @@ -33,16 +33,14 @@ struct CallbackEnvelopeIn { } type CNodeCallback = unsafe extern "C" fn( - user_data: *mut c_void, + user_data: libc::c_ulong, node: *mut c_char, arg_json: *mut c_char, state_json: *mut c_char, ) -> *mut c_char; #[derive(Clone, Copy)] -struct CUserData(*mut c_void); -unsafe impl Send for CUserData {} -unsafe impl Sync for CUserData {} +struct CUserData(libc::c_ulong); fn cstr_to_str<'a>(ptr: *const c_char) -> Result<&'a str, String> { if ptr.is_null() { @@ -95,7 +93,7 @@ fn spawn_json_node_task( arg: Value, state_snapshot: Value, tx: mpsc::Sender), String>>, - user_data_bits: usize, + user_data_bits: libc::c_ulong, callback: CNodeCallback, ) -> Result<(), String> { node_pool_execute(move || { @@ -112,7 +110,7 @@ fn spawn_json_node_task( CString::new(state_json).map_err(|e| format!("invalid state JSON bytes: {e}"))?; let out_ptr = unsafe { callback( - user_data_bits as *mut c_void, + user_data_bits, node_c.as_ptr() as *mut c_char, arg_c.as_ptr() as *mut c_char, state_c.as_ptr() as *mut c_char, @@ -143,7 +141,7 @@ fn run_graph_scheduler_json( ) -> Result { let (tx, rx) = mpsc::channel::), String>>(); let state = Arc::new(Mutex::new(initial_state)); - let user_data_bits = user_data.0 as usize; + let user_data_bits = user_data.0; let tx_for_spawn = tx.clone(); let state_for_spawn = Arc::clone(&state); let state_for_merge = Arc::clone(&state); @@ -299,7 +297,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, - user_data: *mut c_void, + user_data: libc::c_ulong, callback: Option, ) -> *mut c_char { if ptr.is_null() {