This commit is contained in:
Quanzheng Long
2026-03-13 13:16:11 -07:00
parent 0327a86b80
commit d50801d871
9 changed files with 246 additions and 81 deletions
+4 -8
View File
@@ -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}
+10 -10
View File
@@ -5,7 +5,7 @@ package advancedgraph
#cgo LDFLAGS: -L${SRCDIR}/../../rust-core/target/debug -llanggraph_rust_core
#include "langgraph_rust_core.h"
#include <stdlib.h>
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)
+2 -2
View File
@@ -54,8 +54,8 @@ type WaitEvent struct {
}
type Send struct {
Node string
Arg any
Node string
NodeInput any
}
type Command struct {
+14 -29
View File
@@ -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
@@ -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)
}
}
@@ -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)
@@ -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",
]
+2 -2
View File
@@ -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
);
+7 -9
View File
@@ -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<Result<(String, NodeExecResult<Value, Value>), 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<Value, String> {
let (tx, rx) = mpsc::channel::<Result<(String, NodeExecResult<Value, Value>), 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<CNodeCallback>,
) -> *mut c_char {
if ptr.is_null() {