This commit is contained in:
Quanzheng Long
2026-03-13 13:50:52 -07:00
parent d7614999b0
commit 85272db354
6 changed files with 44 additions and 8 deletions
+14 -1
View File
@@ -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),
)
+1
View File
@@ -179,6 +179,7 @@ func TestSubAgentsEquivalentFlow(t *testing.T) {
"output": []string{},
"done": nil,
},
nil,
)
if err != nil {
t.Fatalf("start failed: %v", err)
+4 -6
View File
@@ -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)
}
}
+1
View File
@@ -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
View File
@@ -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,
);