mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-12 20:57:52 +02:00
more
This commit is contained in:
@@ -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}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -54,8 +54,8 @@ type WaitEvent struct {
|
||||
}
|
||||
|
||||
type Send struct {
|
||||
Node string
|
||||
Arg any
|
||||
Node string
|
||||
NodeInput any
|
||||
}
|
||||
|
||||
type Command struct {
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
);
|
||||
|
||||
|
||||
@@ -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() {
|
||||
|
||||
Reference in New Issue
Block a user