Compare commits

...
Author SHA1 Message Date
Quanzheng Long 4f5b775819 no-finish-node 2026-03-16 15:55:35 -07:00
Quanzheng Long 76bee17ec4 fix 2026-03-14 17:58:28 -07:00
Quanzheng Long d4335683e6 rewrite-state-graph 2026-03-14 17:56:30 -07:00
Quanzheng Long b43107ec99 rewrite-state-graph 2026-03-14 17:50:55 -07:00
Quanzheng Long 3987f9cd63 rewrite-state-graph 2026-03-14 17:47:33 -07:00
Quanzheng Long 2bbc3bb1da more 2026-03-14 11:30:25 -07:00
Quanzheng Long b3474d2db1 more 2026-03-14 10:29:42 -07:00
Quanzheng Long 56e9fe1b10 more 2026-03-13 23:31:52 -07:00
Quanzheng Long 1b3d075dbb tokio-rewrite 2026-03-13 23:16:40 -07:00
Quanzheng Long c2a9661d2a more 2026-03-13 21:54:02 -07:00
Quanzheng Long 9e723642c0 more 2026-03-13 15:46:07 -07:00
Quanzheng Long 7c0a9275f9 more 2026-03-13 15:44:41 -07:00
Quanzheng Long 783b3d3435 more 2026-03-13 15:40:49 -07:00
Quanzheng Long 597b3402e6 optimizeUpdateState 2026-03-13 15:07:41 -07:00
Quanzheng Long 54b723775c more 2026-03-13 14:23:54 -07:00
Quanzheng Long 93c4a0a2d5 more 2026-03-13 14:14:58 -07:00
Quanzheng Long a5b90fbb95 more 2026-03-13 13:52:28 -07:00
Quanzheng Long 85272db354 more 2026-03-13 13:50:52 -07:00
Quanzheng Long d7614999b0 more 2026-03-13 13:44:44 -07:00
Quanzheng Long af5f74f2ad more 2026-03-13 13:27:32 -07:00
Quanzheng Long 708c0dff7f more 2026-03-13 13:21:11 -07:00
Quanzheng Long d50801d871 more 2026-03-13 13:16:11 -07:00
Quanzheng Long 0327a86b80 rename 2026-03-12 22:42:47 -07:00
Quanzheng Long 13847a32c0 refactor-go 2026-03-12 22:41:32 -07:00
Quanzheng Long 03ad7a011c imp 2026-03-12 21:48:32 -07:00
Quanzheng Long 896c1a8054 draft 2026-03-12 18:44:18 -07:00
23 changed files with 4336 additions and 147 deletions
+350
View File
@@ -0,0 +1,350 @@
package advancedgraph
import (
"encoding/json"
"fmt"
"reflect"
"runtime"
"strings"
)
type nodeExecutor func(ctx *Context, input any, state map[string]any) (Command, error)
type AdvancedStateGraph[StateT any] struct {
nodes map[string]nodeExecutor
asyncChannels []string
entryPoint string
finishPoint string
stateType reflect.Type
}
func NewAdvancedStateGraph[StateT any]() *AdvancedStateGraph[StateT] {
stateType := mustTypeOf[StateT]()
if stateType.Kind() != reflect.Struct {
panic(fmt.Sprintf("StateT must be a struct, got %s", stateType.String()))
}
return &AdvancedStateGraph[StateT]{
nodes: make(map[string]nodeExecutor),
stateType: stateType,
}
}
// AddNode keeps `fn` as `any` because advanced graph nodes can have different
// input argument types per node, while only `StateT` is globally constrained.
// We validate and adapt node signatures at runtime in compileNodeExecutor.
func (g *AdvancedStateGraph[StateT]) AddNode(fn any) string {
name := NodeName(fn)
return g.AddNodeAs(name, fn)
}
func (g *AdvancedStateGraph[StateT]) AddNodeAs(name string, fn any) string {
if _, exists := g.nodes[name]; exists {
panic(fmt.Sprintf("node `%s` already exists", name))
}
exec, err := compileNodeExecutor(fn, g.stateType)
if err != nil {
panic(err)
}
g.nodes[name] = exec
return name
}
func (g *AdvancedStateGraph[StateT]) AddAsyncChannel(name string) {
g.asyncChannels = append(g.asyncChannels, name)
}
func (g *AdvancedStateGraph[StateT]) AddEntryNode(fn any) string {
name := NodeName(fn)
return g.AddEntryNodeAs(name, fn)
}
func (g *AdvancedStateGraph[StateT]) AddEntryNodeAs(name string, fn any) string {
name = g.AddNodeAs(name, fn)
g.entryPoint = name
return name
}
func (g *AdvancedStateGraph[StateT]) AddFinishNode(fn any) string {
name := NodeName(fn)
return g.AddFinishNodeAs(name, fn)
}
func (g *AdvancedStateGraph[StateT]) AddFinishNodeAs(name string, fn any) string {
name = g.AddNodeAs(name, fn)
g.finishPoint = name
return name
}
func (g *AdvancedStateGraph[StateT]) Compile() *CompiledGraph[StateT] {
return &CompiledGraph[StateT]{
nodes: g.nodes,
asyncChannels: g.asyncChannels,
entryPoint: g.entryPoint,
finishPoint: g.finishPoint,
stateType: g.stateType,
}
}
type CompiledGraph[StateT any] struct {
nodes map[string]nodeExecutor
asyncChannels []string
entryPoint string
finishPoint string
stateType reflect.Type
}
type Context struct {
engine *RustEngine
resumeEvent *WaitEvent
}
func (c *Context) WaitFor(cond AnyOfCondition) (WaitEvent, error) {
if c.resumeEvent != nil {
event := *c.resumeEvent
c.resumeEvent = nil
return event, nil
}
return WaitEvent{}, ErrWaitRequested{Condition: cond}
}
func (c *Context) PublishToChannel(channel string, value any) error {
return c.engine.Publish(channel, value)
}
type Handler[StateT any] struct {
engine *RustEngine
done chan resultOrErr[StateT]
}
type resultOrErr[StateT any] struct {
state StateT
err error
}
func (h *Handler[StateT]) PublishToChannel(channel string, value any) error {
return h.engine.Publish(channel, value)
}
func (h *Handler[StateT]) WaitForResult() (StateT, error) {
res := <-h.done
return res.state, res.err
}
func (g *CompiledGraph[StateT]) Start(initialInput any, initialState StateT) (*Handler[StateT], error) {
engine := NewRustEngine()
for _, ch := range g.asyncChannels {
if err := engine.AddAsyncChannel(ch); err != nil {
return nil, err
}
}
handler := &Handler[StateT]{
engine: engine,
done: make(chan resultOrErr[StateT], 1),
}
go func() {
defer engine.Close()
rawState, err := engine.RunGraph(
g.entryPoint,
g.finishPoint,
initialState,
initialInput,
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)
}
if fallbackState == nil {
return Command{}, fmt.Errorf("node `%s` expected map state argument", node)
}
resolvedInput, resumeEvent := unwrapResumeInput(nodeInput)
return fn(&Context{engine: engine, resumeEvent: resumeEvent}, resolvedInput, fallbackState)
},
)
if err != nil {
handler.done <- resultOrErr[StateT]{err: err}
close(handler.done)
return
}
state, err := mapToState[StateT](rawState)
handler.done <- resultOrErr[StateT]{state: state, err: err}
close(handler.done)
}()
return handler, nil
}
func NodeName(fn any) string {
rv := reflect.ValueOf(fn)
if !rv.IsValid() || rv.Kind() != reflect.Func {
panic("cannot infer node name from non-function value")
}
pc := rv.Pointer()
f := runtime.FuncForPC(pc)
if f == nil {
panic("cannot infer node name from nil function")
}
full := f.Name()
if strings.Contains(full, ".func") {
panic("anonymous functions are not allowed as nodes")
}
short := full
if i := strings.LastIndex(short, "/"); i >= 0 {
short = short[i+1:]
}
if i := strings.LastIndex(short, "."); i >= 0 {
short = short[i+1:]
}
short = strings.TrimSuffix(short, "-fm")
if short == "" || strings.Contains(short, "func") {
panic(fmt.Sprintf("cannot infer stable node name from `%s`", full))
}
return short
}
func compileNodeExecutor(fn any, expectedStateType reflect.Type) (nodeExecutor, error) {
rv := reflect.ValueOf(fn)
if !rv.IsValid() || rv.Kind() != reflect.Func {
return nil, fmt.Errorf("node must be a function")
}
rt := rv.Type()
if rt.NumIn() != 3 {
return nil, fmt.Errorf("node `%s` must accept exactly 3 args: (*Context, input, state)", NodeName(fn))
}
ctxType := reflect.TypeOf((*Context)(nil))
if rt.In(0) != ctxType {
return nil, fmt.Errorf("node `%s` first arg must be *Context", NodeName(fn))
}
if rt.NumOut() != 2 {
return nil, fmt.Errorf("node `%s` must return (Command, error)", NodeName(fn))
}
cmdType := reflect.TypeOf(Command{})
if rt.Out(0) != cmdType {
return nil, fmt.Errorf("node `%s` first return must be Command", NodeName(fn))
}
errType := reflect.TypeOf((*error)(nil)).Elem()
if !rt.Out(1).Implements(errType) {
return nil, fmt.Errorf("node `%s` second return must be error", NodeName(fn))
}
inputType := rt.In(1)
stateType := rt.In(2)
if stateType != expectedStateType {
return nil, fmt.Errorf(
"node `%s` state type mismatch: got %s, graph expects %s",
NodeName(fn),
stateType.String(),
expectedStateType.String(),
)
}
return func(ctx *Context, input any, state map[string]any) (Command, error) {
stateArg, err := convertStateArg(state, stateType)
if err != nil {
return Command{}, fmt.Errorf("node `%s` state decode failed: %w", NodeName(fn), err)
}
args := []reflect.Value{
reflect.ValueOf(ctx),
reflect.Zero(inputType),
stateArg,
}
if input != nil {
inVal := reflect.ValueOf(input)
if inVal.Type().AssignableTo(inputType) {
args[1] = inVal
} else if inVal.Type().ConvertibleTo(inputType) {
args[1] = inVal.Convert(inputType)
} else {
return Command{}, fmt.Errorf(
"node `%s` input type mismatch: got %T, want %s",
NodeName(fn),
input,
inputType.String(),
)
}
}
out := rv.Call(args)
cmd := out[0].Interface().(Command)
if cmd.Update != nil {
updateType := reflect.TypeOf(cmd.Update)
if updateType != stateType {
return Command{}, fmt.Errorf(
"node `%s` update type mismatch: got %s, graph expects %s",
NodeName(fn),
updateType.String(),
stateType.String(),
)
}
}
if out[1].IsNil() {
return cmd, nil
}
return cmd, out[1].Interface().(error)
}, nil
}
func convertStateArg(state map[string]any, stateType reflect.Type) (reflect.Value, error) {
if stateType == reflect.TypeOf(map[string]any{}) {
return reflect.ValueOf(state), nil
}
raw, err := json.Marshal(state)
if err != nil {
return reflect.Value{}, fmt.Errorf("marshal state: %w", err)
}
if stateType.Kind() == reflect.Ptr {
target := reflect.New(stateType.Elem())
if err := json.Unmarshal(raw, target.Interface()); err != nil {
return reflect.Value{}, fmt.Errorf("unmarshal state into %s: %w", stateType.String(), err)
}
return target, nil
}
target := reflect.New(stateType)
if err := json.Unmarshal(raw, target.Interface()); err != nil {
return reflect.Value{}, fmt.Errorf("unmarshal state into %s: %w", stateType.String(), err)
}
return target.Elem(), nil
}
func mapToState[StateT any](raw map[string]any) (StateT, error) {
var out StateT
if anyVal, ok := any(raw).(StateT); ok {
return anyVal, nil
}
payload, err := json.Marshal(raw)
if err != nil {
return out, fmt.Errorf("marshal state: %w", err)
}
if err := json.Unmarshal(payload, &out); err != nil {
return out, fmt.Errorf("unmarshal state: %w", err)
}
return out, nil
}
func mustTypeOf[T any]() reflect.Type {
var zero T
t := reflect.TypeOf(zero)
if t != nil {
return t
}
// Handles nil-able types where zero value has no dynamic type.
return reflect.TypeOf((*T)(nil)).Elem()
}
func unwrapResumeInput(input any) (any, *WaitEvent) {
wrapper, ok := input.(map[string]any)
if !ok {
return input, nil
}
rawArg, hasArg := wrapper["__lg_resume_arg__"]
rawEvent, hasEvent := wrapper["__lg_resume_event__"]
if !hasArg || !hasEvent {
return input, nil
}
eventPayload, err := json.Marshal(rawEvent)
if err != nil {
return rawArg, nil
}
var event WaitEvent
if err := json.Unmarshal(eventPayload, &event); err != nil {
return rawArg, nil
}
return rawArg, &event
}
+309
View File
@@ -0,0 +1,309 @@
package advancedgraph
/*
#cgo CFLAGS: -I${SRCDIR}/../../rust-core/include
#cgo LDFLAGS: -L${SRCDIR}/../../rust-core/target/debug -llanggraph_rust_core
#include "langgraph_rust_core.h"
#include <stdlib.h>
extern char* goNodeCallback(unsigned long user_data, char* node, char* arg_json, char* state_json);
*/
import "C"
import (
"encoding/json"
"fmt"
"reflect"
"sync"
"sync/atomic"
"unsafe"
)
type RustEngine struct {
ptr *C.Engine
}
type runGraphCallbackCtx struct {
exec func(node string, nodeInput any, state map[string]any) (Command, error)
}
var (
callbackRegistryMu sync.RWMutex
callbackRegistry = map[uint64]*runGraphCallbackCtx{}
callbackNextID uint64
)
func registerRunGraphCallbackCtx(ctx *runGraphCallbackCtx) uint64 {
id := atomic.AddUint64(&callbackNextID, 1)
callbackRegistryMu.Lock()
callbackRegistry[id] = ctx
callbackRegistryMu.Unlock()
return id
}
func unregisterRunGraphCallbackCtx(id uint64) {
callbackRegistryMu.Lock()
delete(callbackRegistry, id)
callbackRegistryMu.Unlock()
}
func getRunGraphCallbackCtx(id uint64) (*runGraphCallbackCtx, bool) {
callbackRegistryMu.RLock()
ctx, ok := callbackRegistry[id]
callbackRegistryMu.RUnlock()
return ctx, ok
}
//export goNodeCallback
func goNodeCallback(userData C.ulong, node *C.char, argJSON *C.char, stateJSON *C.char) *C.char {
ctx, ok := getRunGraphCallbackCtx(uint64(userData))
if !ok {
return cCallbackEnvelopeError("invalid callback context (possibly stale callback)")
}
nodeName := C.GoString(node)
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))
}
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, nodeInput, state)
if err != nil {
if waitReq, ok := AsErrWaitRequested(err); ok {
return cCallbackEnvelopeSuspend(waitReq.Condition)
}
return cCallbackEnvelopeError(err.Error())
}
sends := make([]map[string]any, 0, len(cmd.Goto))
for _, send := range cmd.Goto {
targetNode, err := resolveSendTarget(send.Node)
if err != nil {
return cCallbackEnvelopeError(err.Error())
}
sends = append(sends, map[string]any{
"node": targetNode,
"arg": send.NodeInput,
})
}
payload := map[string]any{
"update": cmd.Update,
"sends": sends,
}
raw, err := json.Marshal(map[string]any{
"ok": true,
"payload": payload,
})
if err != nil {
return cCallbackEnvelopeError(fmt.Sprintf("encode callback payload failed: %v", err))
}
return C.CString(string(raw))
}
func NewRustEngine() *RustEngine {
return &RustEngine{ptr: C.rc_engine_new()}
}
func (e *RustEngine) Close() {
if e.ptr != nil {
C.rc_engine_free(e.ptr)
e.ptr = nil
}
}
func (e *RustEngine) AddAsyncChannel(channel string) error {
cch := C.CString(channel)
defer C.free(unsafe.Pointer(cch))
resp := C.rc_add_async_channel(e.ptr, cch)
return parseRustStatus(resp)
}
func (e *RustEngine) Publish(channel string, value any) error {
payload, err := json.Marshal(value)
if err != nil {
return fmt.Errorf("marshal publish value: %w", err)
}
cch := C.CString(channel)
cval := C.CString(string(payload))
defer C.free(unsafe.Pointer(cch))
defer C.free(unsafe.Pointer(cval))
resp := C.rc_publish_json(e.ptr, cch, cval)
return parseRustStatus(resp)
}
func (e *RustEngine) WaitAnyOf(cond AnyOfCondition) (WaitEvent, error) {
payload, err := json.Marshal(cond)
if err != nil {
return WaitEvent{}, fmt.Errorf("marshal any_of: %w", err)
}
cpayload := C.CString(string(payload))
defer C.free(unsafe.Pointer(cpayload))
resp := C.rc_wait_any_of_json(e.ptr, cpayload)
defer C.rc_string_free(resp)
raw := C.GoString(resp)
var status struct {
OK bool `json:"ok"`
Error string `json:"error"`
Event json.RawMessage `json:"event"`
}
if err := json.Unmarshal([]byte(raw), &status); err != nil {
return WaitEvent{}, fmt.Errorf("decode rust wait response: %w", err)
}
if !status.OK {
return WaitEvent{}, fmt.Errorf("rust wait failed: %s", status.Error)
}
var event WaitEvent
if err := json.Unmarshal(status.Event, &event); err != nil {
return WaitEvent{}, fmt.Errorf("decode wait event: %w", err)
}
return event, nil
}
func (e *RustEngine) RunGraph(
entryPoint string,
finishPoint string,
initialState 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))
callbackID := registerRunGraphCallbackCtx(&runGraphCallbackCtx{exec: exec})
defer unregisterRunGraphCallbackCtx(callbackID)
resp := C.rc_run_graph_json(
e.ptr,
centry,
cfinish,
cinitial,
cinitialInput,
C.ulong(callbackID),
(C.rc_node_callback_t)(C.goNodeCallback),
)
defer C.rc_string_free(resp)
raw := C.GoString(resp)
var status struct {
OK bool `json:"ok"`
Error string `json:"error"`
State map[string]any `json:"state"`
}
if err := json.Unmarshal([]byte(raw), &status); err != nil {
return nil, fmt.Errorf("decode rust run response: %w", err)
}
if !status.OK {
return nil, fmt.Errorf("rust run failed: %s", status.Error)
}
coerced := coerceJSONValue(status.State)
typed, ok := coerced.(map[string]any)
if !ok {
return nil, fmt.Errorf("unexpected state type from rust run")
}
return typed, nil
}
func cCallbackEnvelopeError(message string) *C.char {
raw, _ := json.Marshal(map[string]any{
"ok": false,
"error": message,
})
return C.CString(string(raw))
}
func cCallbackEnvelopeSuspend(cond AnyOfCondition) *C.char {
raw, _ := json.Marshal(map[string]any{
"ok": true,
"suspend": map[string]any{
"kind": "any_of",
"any_of": cond,
},
})
return C.CString(string(raw))
}
func resolveSendTarget(target any) (string, error) {
if name, ok := target.(string); ok {
if name == "" {
return "", fmt.Errorf("send target cannot be empty string")
}
return name, nil
}
rv := reflect.ValueOf(target)
if rv.IsValid() && rv.Kind() == reflect.Func {
return NodeName(target), nil
}
return "", fmt.Errorf("unsupported send target type %T", target)
}
func coerceJSONValue(v any) any {
switch t := v.(type) {
case map[string]any:
out := make(map[string]any, len(t))
for k, val := range t {
out[k] = coerceJSONValue(val)
}
return out
case []any:
coerced := make([]any, len(t))
allStrings := true
for i, val := range t {
cv := coerceJSONValue(val)
coerced[i] = cv
if _, ok := cv.(string); !ok {
allStrings = false
}
}
if allStrings {
out := make([]string, len(coerced))
for i, item := range coerced {
out[i] = item.(string)
}
return out
}
return coerced
default:
return v
}
}
func parseRustStatus(resp *C.char) error {
defer C.rc_string_free(resp)
raw := C.GoString(resp)
var status struct {
OK bool `json:"ok"`
Error string `json:"error"`
}
if err := json.Unmarshal([]byte(raw), &status); err != nil {
return fmt.Errorf("decode rust response: %w", err)
}
if !status.OK {
return fmt.Errorf("rust error: %s", status.Error)
}
return nil
}
+89
View File
@@ -0,0 +1,89 @@
package advancedgraph
import (
"encoding/json"
"errors"
)
type WaitCondition interface {
toAny() map[string]any
}
type ChannelCondition struct {
Channel string
N int
}
func (c ChannelCondition) toAny() map[string]any {
n := c.N
if n <= 0 {
n = 1
}
return map[string]any{
"kind": "channel",
"channel": c.Channel,
"n": n,
}
}
type TimerCondition struct {
Seconds float64
}
func (t TimerCondition) toAny() map[string]any {
return map[string]any{
"kind": "timer",
"seconds": t.Seconds,
}
}
type AnyOfCondition struct {
Conditions []map[string]any `json:"conditions"`
}
func AnyOf(conditions ...WaitCondition) AnyOfCondition {
result := AnyOfCondition{Conditions: make([]map[string]any, 0, len(conditions))}
for _, cond := range conditions {
result.Conditions = append(result.Conditions, cond.toAny())
}
return result
}
type WaitEvent struct {
Condition string `json:"condition"`
Channel string `json:"channel,omitempty"`
Value json.RawMessage `json:"value,omitempty"`
Seconds float64 `json:"seconds,omitempty"`
}
type Send struct {
Node any
NodeInput any
}
type Command struct {
Update any
Goto []Send
}
type ErrWaitRequested struct {
Condition AnyOfCondition
}
func (e ErrWaitRequested) Error() string {
return "wait requested"
}
func AsErrWaitRequested(err error) (ErrWaitRequested, bool) {
var target ErrWaitRequested
if !errors.As(err, &target) {
return target, false
}
return target, true
}
func DecodeString(raw json.RawMessage) string {
var s string
_ = json.Unmarshal(raw, &s)
return s
}
+3
View File
@@ -0,0 +1,3 @@
module github.com/langchain-ai/langgraph/langgraph-go
go 1.25
+320
View File
@@ -0,0 +1,320 @@
package stategraph
import (
"encoding/json"
"errors"
"fmt"
"slices"
ag "github.com/langchain-ai/langgraph/langgraph-go/advancedgraph"
)
type StateNodeFunc[StateT any] func(ctx *Context, state StateT) (StateT, error)
const (
internalBarrierChannel = "__stategraph_barrier"
internalInterruptChannel = "__stategraph_interrupt"
)
type Context struct {
inner *ag.Context
}
func (c *Context) Interrupt(name string) (any, error) {
if name == "" {
return nil, fmt.Errorf("interrupt name cannot be empty")
}
event, err := c.inner.WaitFor(ag.AnyOf(ag.ChannelCondition{
Channel: internalInterruptChannel,
N: 1,
}))
if err != nil {
if waitReq, ok := ag.AsErrWaitRequested(err); ok {
return nil, errInterruptRequested{
Name: name,
Condition: waitReq.Condition,
}
}
return nil, err
}
if len(event.Value) == 0 {
return nil, nil
}
var payload interruptPayload
if err := json.Unmarshal(event.Value, &payload); err != nil {
var value any
if err := json.Unmarshal(event.Value, &value); err != nil {
return nil, fmt.Errorf("decode interrupt `%s` value: %w", name, err)
}
return value, nil
}
if payload.Name != "" && payload.Name != name {
return nil, fmt.Errorf("interrupt name mismatch: expected `%s`, got `%s`", name, payload.Name)
}
if len(payload.Value) == 0 {
return nil, nil
}
var value any
if err := json.Unmarshal(payload.Value, &value); err != nil {
return nil, fmt.Errorf("decode interrupt `%s` payload: %w", name, err)
}
return value, nil
}
type errInterruptRequested struct {
Name string
Condition ag.AnyOfCondition
}
func (e errInterruptRequested) Error() string {
if e.Name == "" {
return "interrupt requested"
}
return fmt.Sprintf("interrupt requested: %s", e.Name)
}
func asErrInterruptRequested(err error) (errInterruptRequested, bool) {
var target errInterruptRequested
if !errors.As(err, &target) {
return target, false
}
return target, true
}
type BasicStateGraph[StateT any] struct {
nodes map[string]StateNodeFunc[StateT]
edges map[string][]string
}
type interruptPayload struct {
Name string `json:"name"`
Value json.RawMessage `json:"value"`
}
func NewBasicStateGraph[StateT any]() *BasicStateGraph[StateT] {
return &BasicStateGraph[StateT]{
nodes: make(map[string]StateNodeFunc[StateT]),
edges: make(map[string][]string),
}
}
func (g *BasicStateGraph[StateT]) AddNode(fn StateNodeFunc[StateT]) string {
name := ag.NodeName(fn)
if _, exists := g.nodes[name]; exists {
panic(fmt.Sprintf("node `%s` already exists", name))
}
g.nodes[name] = fn
return name
}
func (g *BasicStateGraph[StateT]) AddEdge(from StateNodeFunc[StateT], to StateNodeFunc[StateT]) {
fromName := ag.NodeName(from)
toName := ag.NodeName(to)
if _, ok := g.nodes[fromName]; !ok {
panic(fmt.Sprintf("source node `%s` does not exist", fromName))
}
if _, ok := g.nodes[toName]; !ok {
panic(fmt.Sprintf("target node `%s` does not exist", toName))
}
g.edges[fromName] = append(g.edges[fromName], toName)
}
type CompiledBasicStateGraph[StateT any] struct {
inner *ag.CompiledGraph[StateT]
}
type Handler[StateT any] struct {
inner *ag.Handler[StateT]
}
func (h *Handler[StateT]) WaitForResult() (StateT, error) {
return h.inner.WaitForResult()
}
func (h *Handler[StateT]) Resume(name string, value any) error {
if name == "" {
return fmt.Errorf("interrupt name cannot be empty")
}
return h.inner.PublishToChannel(internalInterruptChannel, map[string]any{
"name": name,
"value": value,
})
}
func (g *BasicStateGraph[StateT]) Compile() *CompiledBasicStateGraph[StateT] {
if len(g.nodes) == 0 {
panic("graph has no nodes")
}
levels, err := g.computeSupersteps()
if err != nil {
panic(err)
}
adv := ag.NewAdvancedStateGraph[StateT]()
adv.AddAsyncChannel(internalBarrierChannel)
adv.AddAsyncChannel(internalInterruptChannel)
const finalNodeName = "__stategraph_finish"
finalNode := func(_ *ag.Context, _ any, state StateT) (ag.Command, error) {
return ag.Command{Update: state}, nil
}
adv.AddFinishNodeAs(finalNodeName, finalNode)
for stepIdx, stepNodes := range levels {
for _, nodeName := range stepNodes {
userFn := g.nodes[nodeName]
nextBarrier := fmt.Sprintf("__stategraph_barrier_%d", stepIdx+1)
wrapper := func(ctx *ag.Context, _ any, state StateT) (ag.Command, error) {
updated, err := userFn(&Context{inner: ctx}, state)
if err != nil {
if interruptReq, ok := asErrInterruptRequested(err); ok {
cond := interruptReq.Condition
if len(cond.Conditions) == 0 {
cond = ag.AnyOf(ag.ChannelCondition{
Channel: internalInterruptChannel,
N: 1,
})
}
return ag.Command{}, ag.ErrWaitRequested{Condition: cond}
}
return ag.Command{}, err
}
if err := ctx.PublishToChannel(internalBarrierChannel, map[string]any{
"step": stepIdx,
}); err != nil {
return ag.Command{}, err
}
return ag.Command{
Update: updated,
Goto: []ag.Send{{Node: nextBarrier}},
}, nil
}
adv.AddNodeAs(fmt.Sprintf("__stategraph_node_%s", nodeName), wrapper)
}
}
lastBarrier := len(levels)
for barrierStep := 0; barrierStep <= lastBarrier; barrierStep++ {
barrierName := fmt.Sprintf("__stategraph_barrier_%d", barrierStep)
nextStep := barrierStep
barrier := func(ctx *ag.Context, _ any, state StateT) (ag.Command, error) {
if nextStep > 0 {
needed := len(levels[nextStep-1])
_, err := ctx.WaitFor(ag.AnyOf(ag.ChannelCondition{
Channel: internalBarrierChannel,
N: needed,
}))
if err != nil {
return ag.Command{}, err
}
}
if nextStep >= len(levels) {
return ag.Command{
Update: state,
Goto: []ag.Send{{Node: finalNodeName}},
}, nil
}
sends := make([]ag.Send, 0, len(levels[nextStep]))
for _, nodeName := range levels[nextStep] {
sends = append(sends, ag.Send{
Node: fmt.Sprintf("__stategraph_node_%s", nodeName),
})
}
return ag.Command{
Update: state,
Goto: sends,
}, nil
}
if barrierStep == 0 {
adv.AddEntryNodeAs(barrierName, barrier)
} else {
adv.AddNodeAs(barrierName, barrier)
}
}
return &CompiledBasicStateGraph[StateT]{
inner: adv.Compile(),
}
}
func (g *CompiledBasicStateGraph[StateT]) Start(initialState StateT) (*Handler[StateT], error) {
raw, err := g.inner.Start(nil, initialState)
if err != nil {
return nil, err
}
return &Handler[StateT]{inner: raw}, nil
}
func (g *CompiledBasicStateGraph[StateT]) Invoke(initialState StateT) (StateT, error) {
handler, err := g.Start(initialState)
if err != nil {
var zero StateT
return zero, err
}
return handler.WaitForResult()
}
func (g *BasicStateGraph[StateT]) computeSupersteps() ([][]string, error) {
indegree := make(map[string]int, len(g.nodes))
for name := range g.nodes {
indegree[name] = 0
}
for from, tos := range g.edges {
if _, ok := g.nodes[from]; !ok {
return nil, fmt.Errorf("edge source `%s` does not exist", from)
}
for _, to := range tos {
if _, ok := g.nodes[to]; !ok {
return nil, fmt.Errorf("edge target `%s` does not exist", to)
}
indegree[to]++
}
}
queue := make([]string, 0, len(g.nodes))
level := make(map[string]int, len(g.nodes))
for name, deg := range indegree {
if deg == 0 {
queue = append(queue, name)
}
}
if len(queue) == 0 {
return nil, fmt.Errorf("graph has no entry nodes (cycle suspected)")
}
processed := 0
for len(queue) > 0 {
curr := queue[0]
queue = queue[1:]
processed++
currLevel := level[curr]
for _, to := range g.edges[curr] {
if level[to] < currLevel+1 {
level[to] = currLevel + 1
}
indegree[to]--
if indegree[to] == 0 {
queue = append(queue, to)
}
}
}
if processed != len(g.nodes) {
return nil, fmt.Errorf("graph contains a cycle")
}
maxLevel := 0
for _, lv := range level {
if lv > maxLevel {
maxLevel = lv
}
}
levels := make([][]string, maxLevel+1)
for nodeName := range g.nodes {
lv := level[nodeName]
levels[lv] = append(levels[lv], nodeName)
}
for i := range levels {
slices.Sort(levels[i])
}
return levels, nil
}
+218
View File
@@ -0,0 +1,218 @@
package stategraph_test
import (
"fmt"
"sync"
"sync/atomic"
"testing"
"time"
sg "github.com/langchain-ai/langgraph/langgraph-go/stategraph"
)
type stateGraphState struct {
Noop bool `json:"noop"`
}
type orderRecorder struct {
mu sync.Mutex
orders map[string]int32
seq int32
}
func newOrderRecorder() *orderRecorder {
return &orderRecorder{orders: make(map[string]int32)}
}
func (f *orderRecorder) record(name string) {
idx := atomic.AddInt32(&f.seq, 1)
f.mu.Lock()
f.orders[name] = idx
f.mu.Unlock()
}
type orderGraph struct {
recorder *orderRecorder
}
func newOrderGraph() *orderGraph {
return &orderGraph{
recorder: newOrderRecorder(),
}
}
func (g *orderGraph) A(_ *sg.Context, state stateGraphState) (stateGraphState, error) {
g.recorder.record("A")
return state, nil
}
func (g *orderGraph) B1(_ *sg.Context, state stateGraphState) (stateGraphState, error) {
g.recorder.record("B1")
return state, nil
}
func (g *orderGraph) B2(_ *sg.Context, state stateGraphState) (stateGraphState, error) {
g.recorder.record("B2")
return state, nil
}
func (g *orderGraph) C1(_ *sg.Context, state stateGraphState) (stateGraphState, error) {
g.recorder.record("C1")
return state, nil
}
func (g *orderGraph) C2(_ *sg.Context, state stateGraphState) (stateGraphState, error) {
g.recorder.record("C2")
return state, nil
}
func (g *orderGraph) C3(_ *sg.Context, state stateGraphState) (stateGraphState, error) {
g.recorder.record("C3")
return state, nil
}
func (g *orderGraph) D(_ *sg.Context, state stateGraphState) (stateGraphState, error) {
g.recorder.record("D")
return state, nil
}
func TestBasicStateGraphWithoutInterrupt(t *testing.T) {
fixture := newOrderGraph()
graph := sg.NewBasicStateGraph[stateGraphState]()
graph.AddNode(fixture.A)
graph.AddNode(fixture.B1)
graph.AddNode(fixture.B2)
graph.AddNode(fixture.C1)
graph.AddNode(fixture.C2)
graph.AddNode(fixture.C3)
graph.AddNode(fixture.D)
graph.AddEdge(fixture.A, fixture.B1)
graph.AddEdge(fixture.A, fixture.B2)
graph.AddEdge(fixture.B1, fixture.C1)
graph.AddEdge(fixture.B1, fixture.C2)
graph.AddEdge(fixture.B2, fixture.C3)
graph.AddEdge(fixture.C1, fixture.D)
graph.AddEdge(fixture.C2, fixture.D)
graph.AddEdge(fixture.C3, fixture.D)
_, err := graph.Compile().Invoke(stateGraphState{})
if err != nil {
t.Fatalf("invoke failed: %v", err)
}
fixture.recorder.mu.Lock()
orders := make(map[string]int32, len(fixture.recorder.orders))
for k, v := range fixture.recorder.orders {
orders[k] = v
}
fixture.recorder.mu.Unlock()
for _, name := range []string{"A", "B1", "B2", "C1", "C2", "C3", "D"} {
if _, ok := orders[name]; !ok {
t.Fatalf("node %s did not execute; orders=%v", name, orders)
}
}
maxB := maxInt32(orders["B1"], orders["B2"])
minC := minInt32(orders["C1"], minInt32(orders["C2"], orders["C3"]))
maxC := maxInt32(orders["C1"], maxInt32(orders["C2"], orders["C3"]))
if !(orders["A"] < orders["B1"] && orders["A"] < orders["B2"]) {
t.Fatalf("A should run before B-step, orders=%v", orders)
}
if !(maxB < minC) {
t.Fatalf("B-step should finish before C-step, orders=%v", orders)
}
if !(maxC < orders["D"]) {
t.Fatalf("C-step should finish before D, orders=%v", orders)
}
}
type interruptState struct {
A bool `json:"a"`
B bool `json:"b"`
}
type interruptFixture struct{}
func (f *interruptFixture) A(_ *sg.Context, state interruptState) (interruptState, error) {
state.A = true
return state, nil
}
func (f *interruptFixture) B(ctx *sg.Context, state interruptState) (interruptState, error) {
if !state.A {
return state, fmt.Errorf("B should observe A=true")
}
value, err := ctx.Interrupt("resume_channel")
if err != nil {
return state, err
}
s, ok := value.(string)
if !ok || s != "go" {
return state, fmt.Errorf("unexpected interrupt payload: %#v", value)
}
state.B = true
return state, nil
}
func TestBasicStateGraphWithInterrupt(t *testing.T) {
fixture := &interruptFixture{}
graph := sg.NewBasicStateGraph[interruptState]()
graph.AddNode(fixture.A)
graph.AddNode(fixture.B)
graph.AddEdge(fixture.A, fixture.B)
handler, err := graph.Compile().Start(interruptState{})
if err != nil {
t.Fatalf("start failed: %v", err)
}
doneCh := make(chan interruptState, 1)
errCh := make(chan error, 1)
go func() {
result, runErr := handler.WaitForResult()
if runErr != nil {
errCh <- runErr
return
}
doneCh <- result
}()
select {
case <-doneCh:
t.Fatalf("run should pause for interrupt, but completed early")
case err := <-errCh:
t.Fatalf("run should pause for interrupt, but failed early: %v", err)
case <-time.After(120 * time.Millisecond):
// expected: paused
}
if err := handler.Resume("resume_channel", "go"); err != nil {
t.Fatalf("resume interrupt failed: %v", err)
}
select {
case err := <-errCh:
t.Fatalf("run failed after interrupt: %v", err)
case result := <-doneCh:
if !(result.A && result.B) {
t.Fatalf("unexpected final state: %#v", result)
}
case <-time.After(2 * time.Second):
t.Fatalf("timeout waiting for resumed run completion")
}
}
func minInt32(a int32, b int32) int32 {
if a < b {
return a
}
return b
}
func maxInt32(a int32, b int32) int32 {
if a > b {
return a
}
return b
}
+204
View File
@@ -0,0 +1,204 @@
package tests
import (
"slices"
"testing"
"time"
ag "github.com/langchain-ai/langgraph/langgraph-go/advancedgraph"
)
type decision struct {
Type string
SubAgent string
Tool string
Complete string
}
type mockLLM struct {
responses [][]decision
i int
}
func (m *mockLLM) invoke() []decision {
if m.i >= len(m.responses) {
return []decision{}
}
resp := m.responses[m.i]
m.i++
return resp
}
type lunchWorkflow struct {
planner *mockLLM
}
type lunchState struct {
Input string `json:"input"`
Output []string `json:"output"`
Done string `json:"done"`
}
func (w *lunchWorkflow) llmNode(ctx *ag.Context, _ any, _ lunchState) (ag.Command, error) {
decisions := w.planner.invoke()
sends := make([]ag.Send, 0, 4)
for _, d := range decisions {
if d.Type == "end" {
return ag.Command{
Goto: []ag.Send{
{Node: w.orderFoodNode, NodeInput: d.Complete},
},
}, nil
}
if d.Type == "sub_agent" {
sends = append(sends, ag.Send{Node: w.subAgentNode, NodeInput: d.SubAgent})
}
if d.Type == "tool" {
sends = append(sends, ag.Send{Node: w.toolNode, NodeInput: d.Tool})
}
}
sends = append(sends, ag.Send{Node: w.waitNode})
return ag.Command{Goto: sends}, nil
}
func (w *lunchWorkflow) waitNode(ctx *ag.Context, _ any, state lunchState) (ag.Command, error) {
event, err := ctx.WaitFor(
ag.AnyOf(
ag.ChannelCondition{Channel: "tool_completion_channel", N: 1},
ag.ChannelCondition{Channel: "subagent_completion_channel", N: 1},
ag.ChannelCondition{Channel: "user_input_channel", N: 1},
ag.TimerCondition{Seconds: 1},
),
)
if err != nil {
return ag.Command{}, err
}
output := append([]string(nil), state.Output...)
if event.Condition == "channel" {
payload := ag.DecodeString(event.Value)
switch event.Channel {
case "tool_completion_channel":
output = append(output, "tool: "+payload)
case "subagent_completion_channel":
output = append(output, "sub_agent: "+payload)
case "user_input_channel":
output = append(output, "user_input: "+payload)
}
state.Output = output
return ag.Command{Goto: []ag.Send{{Node: w.llmNode}}, Update: state}, nil
}
output = append(output, "timer: no updates yet")
state.Output = output
return ag.Command{Goto: []ag.Send{{Node: w.waitNode}}, Update: state}, nil
}
func (w *lunchWorkflow) toolNode(ctx *ag.Context, input any, _ lunchState) (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, input any, _ lunchState) (ag.Command, error) {
subInput, _ := input.(string)
time.Sleep(5 * time.Second)
err := ctx.PublishToChannel(
"subagent_completion_channel",
"research sub agent completed for: "+subInput,
)
return ag.Command{}, err
}
func (w *lunchWorkflow) orderFoodNode(ctx *ag.Context, input any, state lunchState) (ag.Command, error) {
complete, _ := input.(string)
output := append([]string(nil), state.Output...)
output = append(output, "order_food: "+complete)
state.Output = output
state.Done = complete
return ag.Command{Update: state}, nil
}
func TestSubAgentsEquivalentFlow(t *testing.T) {
planner := &mockLLM{
responses: [][]decision{
{
{Type: "sub_agent", SubAgent: "research lunch options"},
{Type: "tool", Tool: "slack_tool"},
},
{},
{},
{{Type: "sub_agent", SubAgent: "find vegetarian fallback"}},
{{Type: "end", Complete: "order submitted"}},
},
}
workflow := &lunchWorkflow{
planner: planner,
}
graph := ag.NewAdvancedStateGraph[lunchState]()
graph.AddAsyncChannel("tool_completion_channel")
graph.AddAsyncChannel("subagent_completion_channel")
graph.AddAsyncChannel("user_input_channel")
graph.AddEntryNode(workflow.llmNode)
graph.AddNode(workflow.waitNode)
graph.AddNode(workflow.toolNode)
graph.AddNode(workflow.subAgentNode)
graph.AddFinishNode(workflow.orderFoodNode)
handler, err := graph.Compile().Start(
nil,
lunchState{
Input: "help me get something for lunch",
Output: []string{},
Done: "",
},
)
if err != nil {
t.Fatalf("start failed: %v", err)
}
time.Sleep(10 * time.Millisecond)
if err := handler.PublishToChannel("user_input_channel", "No spicy food please"); err != nil {
t.Fatalf("publish failed: %v", err)
}
result, err := handler.WaitForResult()
if err != nil {
t.Fatalf("result failed: %v", err)
}
output := result.Output
if len(output) == 0 {
t.Fatalf("output is empty, full result=%#v", result)
}
if result.Done != "order submitted" {
t.Fatalf("unexpected done: %v", result.Done)
}
if !slices.Contains(output, "user_input: No spicy food please") {
t.Fatalf("missing user input output: %#v", output)
}
if !slices.Contains(output, "tool: tool completed for: slack_tool") {
t.Fatalf("missing tool output: %#v", output)
}
if !slices.Contains(output, "sub_agent: research sub agent completed for: research lunch options") {
t.Fatalf("missing first sub-agent output: %#v", output)
}
if !slices.Contains(output, "sub_agent: research sub agent completed for: find vegetarian fallback") {
t.Fatalf("missing second sub-agent output: %#v", output)
}
timerCount := 0
for _, line := range output {
if line == "timer: no updates yet" {
timerCount++
}
}
if timerCount < 3 {
t.Fatalf("expected >=3 timer outputs, got %d, output=%#v", timerCount, output)
}
if output[len(output)-1] != "order_food: order submitted" {
t.Fatalf("unexpected last output: %#v", output[len(output)-1])
}
}
+124
View File
@@ -0,0 +1,124 @@
package tests
import (
"fmt"
"testing"
ag "github.com/langchain-ai/langgraph/langgraph-go/advancedgraph"
)
type primitiveWorkflow struct {
}
type primitiveState struct {
Count int `json:"count"`
Logs []string `json:"logs"`
Done string `json:"done"`
}
func (w *primitiveWorkflow) startNode(ctx *ag.Context, input int, state primitiveState) (ag.Command, error) {
state.Logs = append(state.Logs, fmt.Sprintf("start:%d", input))
return ag.Command{
Update: state,
Goto: []ag.Send{
{Node: w.middleNode, NodeInput: "from_start"},
},
}, nil
}
func (w *primitiveWorkflow) middleNode(ctx *ag.Context, input string, state primitiveState) (ag.Command, error) {
state.Logs = append(state.Logs, "middle:"+input)
return ag.Command{
Update: state,
Goto: []ag.Send{
{Node: w.finishNode, NodeInput: "from_middle"},
},
}, nil
}
func (w *primitiveWorkflow) finishNode(ctx *ag.Context, input string, state primitiveState) (ag.Command, error) {
state.Logs = append(state.Logs, "finish:"+input)
state.Done = input
return ag.Command{Update: state}, nil
}
func TestInputAndStatePrimitivesCompatible(t *testing.T) {
workflow := &primitiveWorkflow{}
graph := ag.NewAdvancedStateGraph[primitiveState]()
graph.AddEntryNode(workflow.startNode)
graph.AddNode(workflow.middleNode)
graph.AddFinishNode(workflow.finishNode)
handler, err := graph.Compile().Start(100, primitiveState{
Count: 1,
Logs: []string{},
Done: "",
})
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)
}
if result.Count != 1 {
t.Fatalf("unexpected count: %v", result.Count)
}
if len(result.Logs) != 3 || result.Logs[0] != "start:100" || result.Logs[1] != "middle:from_start" || result.Logs[2] != "finish:from_middle" {
t.Fatalf("unexpected logs: %#v", result.Logs)
}
}
func (w *primitiveWorkflow) startNoFinishNode(ctx *ag.Context, _ any, state primitiveState) (ag.Command, error) {
state.Logs = append(state.Logs, "start")
return ag.Command{
Update: state,
Goto: []ag.Send{
{Node: w.middleNoFinishNode, NodeInput: "from_start"},
},
}, nil
}
func (w *primitiveWorkflow) middleNoFinishNode(ctx *ag.Context, input string, state primitiveState) (ag.Command, error) {
state.Logs = append(state.Logs, "middle:"+input)
state.Count += 1
state.Done = "stopped"
// No goto and no finish node configured: run should end automatically.
return ag.Command{Update: state}, nil
}
func TestRunEndsWithoutFinishNode(t *testing.T) {
workflow := &primitiveWorkflow{}
graph := ag.NewAdvancedStateGraph[primitiveState]()
graph.AddEntryNode(workflow.startNoFinishNode)
graph.AddNode(workflow.middleNoFinishNode)
handler, err := graph.Compile().Start(nil, primitiveState{
Count: 7,
Logs: []string{},
Done: "",
})
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 != "stopped" {
t.Fatalf("unexpected done: %v", result.Done)
}
if result.Count != 8 {
t.Fatalf("unexpected count: %v", result.Count)
}
if len(result.Logs) != 2 || result.Logs[0] != "start" || result.Logs[1] != "middle:from_start" {
t.Fatalf("unexpected logs: %#v", result.Logs)
}
}
@@ -0,0 +1,54 @@
package tests
import (
"strings"
"testing"
ag "github.com/langchain-ai/langgraph/langgraph-go/advancedgraph"
)
func TestNewAdvancedStateGraphRejectsNonStructState(t *testing.T) {
defer func() {
if r := recover(); r == nil {
t.Fatalf("expected panic for non-struct StateT")
}
}()
_ = ag.NewAdvancedStateGraph[map[string]any]()
}
type stateTypeA struct {
X int `json:"x"`
}
type stateTypeB struct {
X int `json:"x"`
}
type wrongUpdateWorkflow struct{}
func (w *wrongUpdateWorkflow) startNode(ctx *ag.Context, _ any, _ stateTypeA) (ag.Command, error) {
return ag.Command{Goto: []ag.Send{{Node: w.badNode}}}, nil
}
func (w *wrongUpdateWorkflow) badNode(ctx *ag.Context, _ any, _ stateTypeA) (ag.Command, error) {
return ag.Command{Update: stateTypeB{X: 1}}, nil
}
func TestNodeUpdateTypeMustMatchGraphStateType(t *testing.T) {
workflow := &wrongUpdateWorkflow{}
graph := ag.NewAdvancedStateGraph[stateTypeA]()
graph.AddEntryNode(workflow.startNode)
graph.AddFinishNode(workflow.badNode)
handler, err := graph.Compile().Start(nil, stateTypeA{X: 0})
if err != nil {
t.Fatalf("start failed: %v", err)
}
_, err = handler.WaitForResult()
if err == nil {
t.Fatalf("expected runtime error for wrong update type")
}
if !strings.Contains(err.Error(), "update type mismatch") {
t.Fatalf("unexpected error: %v", err)
}
}
@@ -0,0 +1,135 @@
package tests
import (
"reflect"
"testing"
"time"
ag "github.com/langchain-ai/langgraph/langgraph-go/advancedgraph"
)
type updateElisionState struct {
X int `json:"x"`
S updateStruct `json:"s"`
M map[string]int `json:"m"`
L []int `json:"l"`
PS *updateStruct `json:"ps"`
PM *map[string]int `json:"pm"`
PL *[]int `json:"pl"`
}
type updateStruct struct {
V int `json:"v"`
}
type updateElisionWorkflow struct{}
func makeState(v int) updateElisionState {
m := map[string]int{"n": v}
l := []int{v}
return updateElisionState{
X: v,
S: updateStruct{V: v},
M: map[string]int{"n": v},
L: []int{v},
PS: &updateStruct{V: v},
PM: &m,
PL: &l,
}
}
func (w *updateElisionWorkflow) startNoopNode(ctx *ag.Context, _ any, _ updateElisionState) (ag.Command, error) {
return ag.Command{
Goto: []ag.Send{
{Node: w.fastNode},
{Node: w.slowNoopNode},
},
}, nil
}
func (w *updateElisionWorkflow) startChangedNode(ctx *ag.Context, _ any, _ updateElisionState) (ag.Command, error) {
return ag.Command{
Goto: []ag.Send{
{Node: w.fastNode},
{Node: w.slowChangedNode},
},
}, nil
}
func (w *updateElisionWorkflow) fastNode(ctx *ag.Context, _ any, state updateElisionState) (ag.Command, error) {
_ = state
return ag.Command{Update: makeState(1)}, nil
}
func (w *updateElisionWorkflow) slowNoopNode(ctx *ag.Context, _ any, state updateElisionState) (ag.Command, error) {
time.Sleep(100 * time.Millisecond)
// Returns same state as initial snapshot; without runtime elision this can overwrite newer updates.
return ag.Command{Update: state}, nil
}
func (w *updateElisionWorkflow) slowChangedNode(ctx *ag.Context, _ any, _ updateElisionState) (ag.Command, error) {
time.Sleep(100 * time.Millisecond)
// Real change should not be elided.
return ag.Command{Update: makeState(2)}, nil
}
func assertStateEquals(t *testing.T, got updateElisionState, expected updateElisionState) {
t.Helper()
if got.X != expected.X {
t.Fatalf("unexpected X: got=%d want=%d", got.X, expected.X)
}
if got.S != expected.S {
t.Fatalf("unexpected S: got=%#v want=%#v", got.S, expected.S)
}
if !reflect.DeepEqual(got.M, expected.M) {
t.Fatalf("unexpected M: got=%#v want=%#v", got.M, expected.M)
}
if !reflect.DeepEqual(got.L, expected.L) {
t.Fatalf("unexpected L: got=%#v want=%#v", got.L, expected.L)
}
if got.PS == nil || expected.PS == nil || *got.PS != *expected.PS {
t.Fatalf("unexpected PS: got=%#v want=%#v", got.PS, expected.PS)
}
if got.PM == nil || expected.PM == nil || !reflect.DeepEqual(*got.PM, *expected.PM) {
t.Fatalf("unexpected PM: got=%#v want=%#v", got.PM, expected.PM)
}
if got.PL == nil || expected.PL == nil || !reflect.DeepEqual(*got.PL, *expected.PL) {
t.Fatalf("unexpected PL: got=%#v want=%#v", got.PL, expected.PL)
}
}
func TestNoopSlowUpdateCanOverrideFastUpdate(t *testing.T) {
workflow := &updateElisionWorkflow{}
graph := ag.NewAdvancedStateGraph[updateElisionState]()
graph.AddEntryNode(workflow.startNoopNode)
graph.AddNode(workflow.fastNode)
graph.AddFinishNode(workflow.slowNoopNode)
handler, err := graph.Compile().Start(nil, makeState(0))
if err != nil {
t.Fatalf("start failed: %v", err)
}
result, err := handler.WaitForResult()
if err != nil {
t.Fatalf("result failed: %v", err)
}
assertStateEquals(t, result, makeState(0))
}
func TestChangedSlowUpdateOverridesFastUpdate(t *testing.T) {
workflow := &updateElisionWorkflow{}
graph := ag.NewAdvancedStateGraph[updateElisionState]()
graph.AddEntryNode(workflow.startChangedNode)
graph.AddNode(workflow.fastNode)
graph.AddFinishNode(workflow.slowChangedNode)
handler, err := graph.Compile().Start(nil, makeState(0))
if err != nil {
t.Fatalf("start failed: %v", err)
}
result, err := handler.WaitForResult()
if err != nil {
t.Fatalf("result failed: %v", err)
}
assertStateEquals(t, result, makeState(2))
}
+229 -130
View File
@@ -1,12 +1,18 @@
from __future__ import annotations
import atexit
import asyncio
import os
import inspect
from collections.abc import Callable, Mapping, Sequence
import threading
from concurrent.futures import ThreadPoolExecutor
from collections.abc import Callable, Coroutine, Sequence
from dataclasses import dataclass
from datetime import timedelta
from typing import Any, Generic, TypeVar, cast
from langgraph_rust_core import PyRustEngine # type: ignore[import-untyped]
from langgraph.types import Command, Send
StateT = TypeVar("StateT")
@@ -35,6 +41,37 @@ class AnyOfCondition:
WaitCondition = ChannelCondition | TimerCondition
_EXECUTOR_LOCK = threading.Lock()
_EXECUTOR: ThreadPoolExecutor | None = None
def _advanced_graph_executor() -> ThreadPoolExecutor:
global _EXECUTOR
with _EXECUTOR_LOCK:
if _EXECUTOR is None:
worker_count = int(os.getenv("LANGGRAPH_ADVANCED_GRAPH_PY_THREADS", "256"))
worker_count = max(worker_count, 1)
_EXECUTOR = ThreadPoolExecutor(
max_workers=worker_count,
thread_name_prefix="langgraph-advanced-py",
)
atexit.register(_shutdown_advanced_graph_executor)
return _EXECUTOR
def _shutdown_advanced_graph_executor() -> None:
global _EXECUTOR
with _EXECUTOR_LOCK:
if _EXECUTOR is not None:
_EXECUTOR.shutdown(wait=False, cancel_futures=False)
_EXECUTOR = None
class WaitRequested(Exception):
def __init__(self, payload: dict[str, Any]) -> None:
super().__init__("wait requested")
self.payload = payload
class AdvancedStateGraph(Generic[StateT]):
"""Experimental in-memory graph engine with async channels."""
@@ -72,20 +109,14 @@ class AdvancedStateGraph(Generic[StateT]):
raise ValueError(f"Channel `{name}` already exists")
self._async_channels[name] = _ChannelSpec(typ=typ)
def set_entry_point(self, name_or_node: str | Callable[..., Any]) -> None:
self._entry_point = self._resolve_node_name(name_or_node)
def set_finish_point(self, name_or_node: str | Callable[..., Any]) -> None:
self._finish_point = self._resolve_node_name(name_or_node)
def add_entry_node(self, node: Callable[..., Any]) -> str:
node_name = self.add_node(node)
self.set_entry_point(node_name)
self._entry_point = self._resolve_node_name(node_name)
return node_name
def add_finish_node(self, node: Callable[..., Any]) -> str:
node_name = self.add_node(node)
self.set_finish_point(node_name)
self._finish_point = self._resolve_node_name(node_name)
return node_name
def _resolve_node_name(self, name_or_node: str | Callable[..., Any]) -> str:
@@ -99,11 +130,9 @@ class AdvancedStateGraph(Generic[StateT]):
def compile(self) -> CompiledGraphEngine[StateT]:
if self._entry_point is None:
raise ValueError("Entry point is not set")
if self._finish_point is None:
raise ValueError("Finish point is not set")
if self._entry_point not in self._nodes:
raise ValueError(f"Entry point node `{self._entry_point}` does not exist")
if self._finish_point not in self._nodes:
if self._finish_point is not None and self._finish_point not in self._nodes:
raise ValueError(f"Finish point node `{self._finish_point}` does not exist")
return CompiledGraphEngine(
nodes=dict(self._nodes),
@@ -122,7 +151,7 @@ class CompiledGraphEngine(Generic[StateT]):
nodes: dict[str, Callable[..., Any]],
async_channels: dict[str, _ChannelSpec],
entry_point: str,
finish_point: str,
finish_point: str | None,
) -> None:
self._nodes = nodes
self._async_channels = async_channels
@@ -151,7 +180,10 @@ class Context:
self._run = run
async def wait_for(self, target: WaitCondition | AnyOfCondition) -> Any:
return await self._run.wait_for(target)
resumed = self._run._consume_resume_event(target)
if resumed is not None:
return resumed
raise WaitRequested(_target_to_suspend_payload(target))
def publish_to_channel(self, channel: str, value: Any) -> None:
self._run.publish_nowait(channel, value)
@@ -186,49 +218,45 @@ class _GraphEngineRun:
nodes: dict[str, Callable[..., Any]],
async_channel_specs: dict[str, _ChannelSpec],
entry_point: str,
finish_point: str,
finish_point: str | None,
) -> None:
self._nodes = nodes
self._entry_point = entry_point
self._finish_point = finish_point
self._async_channels: dict[str, asyncio.Queue[Any]] = {
name: asyncio.Queue() for name, _spec in async_channel_specs.items()
}
self._rust_engine = PyRustEngine()
for name in async_channel_specs:
self._rust_engine.add_async_channel(name)
self._tasks: set[asyncio.Task[list[Send]]] = set()
self._finished = False
self._state: Any = None
self._local = threading.local()
self.context = Context(self)
async def run(self, initial_state: StateT) -> StateT:
self._state = initial_state
self._schedule(Send(self._entry_point, initial_state))
try:
while self._tasks and not self._finished:
done, _ = await asyncio.wait(
self._tasks, return_when=asyncio.FIRST_COMPLETED
)
for task in done:
self._tasks.remove(task)
exc = task.exception()
if exc is not None:
await self._cancel_all_tasks()
raise exc
sends = task.result()
for send in sends:
self._schedule(send)
if self._finished:
await self._cancel_all_tasks()
return cast(StateT, self._state)
finally:
await self._cancel_all_tasks()
finish_point = self._finish_point or ""
loop = asyncio.get_running_loop()
result_obj = await loop.run_in_executor(
_advanced_graph_executor(),
self._rust_engine.run_graph_py,
self._entry_point,
finish_point,
initial_state,
self._execute_node_for_rust,
)
self._state = result_obj
return cast(StateT, self._state)
async def publish(self, channel: str, value: Any) -> None:
queue = self._get_async_channel(channel)
await queue.put(value)
loop = asyncio.get_running_loop()
await loop.run_in_executor(
_advanced_graph_executor(),
self._publish_sync,
channel,
value,
)
def publish_nowait(self, channel: str, value: Any) -> None:
queue = self._get_async_channel(channel)
queue.put_nowait(value)
self._publish_sync(channel, value)
async def wait_for(self, target: WaitCondition | AnyOfCondition) -> Any:
if isinstance(target, ChannelCondition):
@@ -239,8 +267,12 @@ class _GraphEngineRun:
"value": value,
}
if isinstance(target, TimerCondition):
await asyncio.sleep(target.seconds)
return {"condition": "timer", "seconds": target.seconds}
loop = asyncio.get_running_loop()
return await loop.run_in_executor(
_advanced_graph_executor(),
self._rust_engine.wait_timer,
target.seconds,
)
if isinstance(target, AnyOfCondition):
return await self._wait_for_any_of(target)
raise ValueError(f"Unsupported wait condition type: {type(target)!r}")
@@ -248,132 +280,125 @@ class _GraphEngineRun:
async def _wait_for_channel_values(self, channel: str, n: int) -> Any:
if n < 1:
raise ValueError("wait_for count `n` must be >= 1")
queue = self._get_async_channel(channel)
if n == 1:
return await queue.get()
values: list[Any] = []
for _ in range(n):
values.append(await queue.get())
return values
loop = asyncio.get_running_loop()
event = await loop.run_in_executor(
_advanced_graph_executor(),
self._rust_engine.wait_channel,
channel,
n,
)
return event["value"]
async def _wait_for_any_of(self, condition: AnyOfCondition) -> Any:
if not condition.conditions:
raise ValueError("any_of() requires at least one condition")
payload = {
"conditions": [_condition_to_rust(cond) for cond in condition.conditions]
}
loop = asyncio.get_running_loop()
return await loop.run_in_executor(
_advanced_graph_executor(),
self._rust_engine.wait_any_of_obj,
payload,
)
tasks = [
asyncio.create_task(self.wait_for(inner_condition))
for inner_condition in condition.conditions
]
done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
for task in pending:
task.cancel()
await asyncio.gather(*pending, return_exceptions=True)
first = done.pop()
return first.result()
def _publish_sync(self, channel: str, value: Any) -> None:
self._rust_engine.publish_obj(channel, value)
def _get_async_channel(self, channel: str) -> asyncio.Queue[Any]:
if channel not in self._async_channels:
raise ValueError(f"Unknown channel `{channel}`")
return self._async_channels[channel]
def _schedule(self, send: Send) -> None:
if self._finished:
return
task: asyncio.Task[list[Send]] = asyncio.create_task(self._execute_send(send))
self._tasks.add(task)
async def _cancel_all_tasks(self) -> None:
if not self._tasks:
return
to_cancel = list(self._tasks)
for task in to_cancel:
task.cancel()
await asyncio.gather(*to_cancel, return_exceptions=True)
self._tasks.clear()
async def _execute_send(self, send: Send) -> list[Send]:
node_name = _resolve_target_name(send.node)
def _execute_node_for_rust(
self, node_name: str, node_input: Any, state: Any
) -> dict[str, Any]:
node_input, resume_event = _unwrap_resume_input(node_input)
self._set_resume_event(resume_event)
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, send.arg)
if inspect.isawaitable(result):
result = await result
try:
result = _invoke_node(node, self.context, node_input, state)
if inspect.isawaitable(result):
result = self._run_awaitable_in_worker(
cast(Coroutine[Any, Any, Any], result)
)
except WaitRequested as suspend:
return {"suspend": suspend.payload}
finally:
self._set_resume_event(None)
if isinstance(result, Command):
self._apply_update(result.update)
next_sends = _normalize_goto(result.goto, default_arg=self._state)
update = result.update
sends = _normalize_goto(result.goto, default_input=node_input)
else:
self._apply_update(result)
next_sends = _normalize_result_to_sends(result, default_arg=self._state)
update = result
sends = _normalize_result_to_sends(result, default_input=node_input)
if node_name == self._finish_point:
self._finished = True
return []
return next_sends
return {
"update": update,
"sends": [
{"node": _resolve_target_name(send.node), "arg": send.arg}
for send in sends
],
}
def _apply_update(self, update: Any) -> None:
if update is None:
return
if isinstance(update, Mapping):
if isinstance(self._state, Mapping):
# Keep semantics simple: in-place update for mapping-like state.
cast(dict[str, Any], self._state).update(update)
return
if isinstance(update, Sequence) and not isinstance(update, (str, bytes)):
pairs = list(update)
if all(
isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str)
for item in pairs
):
if isinstance(self._state, Mapping):
cast(dict[str, Any], self._state).update(
cast(dict[str, Any], pairs)
)
return
def _set_resume_event(self, event: dict[str, Any] | None) -> None:
self._local.resume_event = event
def _consume_resume_event(self, target: WaitCondition | AnyOfCondition) -> Any | None:
event = cast(dict[str, Any] | None, getattr(self._local, "resume_event", None))
if event is None:
return None
self._local.resume_event = None
return event
def _normalize_result_to_sends(result: Any, *, default_arg: Any) -> list[Send]:
def _run_awaitable_in_worker(self, awaitable: Coroutine[Any, Any, Any]) -> Any:
loop = cast(
asyncio.AbstractEventLoop | None,
getattr(self._local, "worker_loop", None),
)
if loop is None or loop.is_closed():
loop = asyncio.new_event_loop()
self._local.worker_loop = loop
return loop.run_until_complete(awaitable)
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 []
@@ -417,6 +442,37 @@ def any_of(*conditions: WaitCondition) -> AnyOfCondition:
return AnyOfCondition(conditions=tuple(conditions))
def _condition_to_rust(condition: WaitCondition) -> dict[str, Any]:
if isinstance(condition, ChannelCondition):
return {"kind": "channel", "channel": condition.channel, "n": condition.n}
if isinstance(condition, TimerCondition):
return {"kind": "timer", "seconds": condition.seconds}
raise TypeError(f"Unsupported condition type: {type(condition)!r}")
def _target_to_suspend_payload(target: WaitCondition | AnyOfCondition) -> dict[str, Any]:
if isinstance(target, AnyOfCondition):
return {
"kind": "any_of",
"any_of": {
"conditions": [_condition_to_rust(cond) for cond in target.conditions]
},
}
return {"kind": "condition", "condition": _condition_to_rust(target)}
def _unwrap_resume_input(node_input: Any) -> tuple[Any, dict[str, Any] | None]:
if not isinstance(node_input, dict):
return node_input, None
if "__lg_resume_arg__" not in node_input or "__lg_resume_event__" not in node_input:
return node_input, None
resume_arg = node_input["__lg_resume_arg__"]
resume_event = node_input["__lg_resume_event__"]
if isinstance(resume_event, dict):
return resume_arg, resume_event
return resume_arg, None
def _infer_node_name(node: Callable[..., Any]) -> str:
node_name = getattr(node, "__name__", "")
if not node_name or node_name == "<lambda>":
@@ -432,14 +488,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,306 @@
from __future__ import annotations
import asyncio
import os
import time
from typing import Any
# Configure advanced-graph runtime pools for this benchmark run.
os.environ["LANGGRAPH_RUN_POOL_SIZE"] = "10"
os.environ["LANGGRAPH_NODE_POOL_SIZE"] = "1000"
from langgraph.advanced_graph import (
AdvancedStateGraph,
channel_condition,
timer_condition,
)
from langgraph.graph import END, START, StateGraph
from langgraph.types import Command, Send
RUNS = 100
MIDDLE_COUNT = 10
SLEEP_SECONDS = 2.0
BLOCKING_SECONDS = 0.1
STATE_BYTES = 10 * 1024
def make_initial_state() -> dict[str, Any]:
return {"payload": "x" * STATE_BYTES, "done": False}
def build_advanced_parallel() -> Any:
graph: AdvancedStateGraph[dict[str, Any]] = AdvancedStateGraph(dict)
done_channel = "__bench_done_channel"
graph.add_async_channel(done_channel, str)
async def start_node(state: dict[str, Any]) -> Command:
_ = state
sends = [Send(f"middle_{i}", None) for i in range(MIDDLE_COUNT)]
sends.append(Send("end_node", None))
return Command(goto=sends)
async def end_node(ctx: Any, state: dict[str, Any]) -> dict[str, Any]:
await ctx.wait_for(channel_condition(done_channel, n=MIDDLE_COUNT))
out = dict(state)
out["done"] = True
return out
graph.add_entry_node(start_node)
for i in range(MIDDLE_COUNT):
async def middle_node(ctx: Any, state: dict[str, Any], idx: int = i) -> None:
_ = idx
_ = state
await ctx.wait_for(timer_condition(seconds=SLEEP_SECONDS))
ctx.publish_to_channel(done_channel, "done")
graph.add_node(f"middle_{i}", middle_node)
graph.add_finish_node(end_node)
return graph.compile()
def build_advanced_sequential() -> Any:
graph: AdvancedStateGraph[dict[str, Any]] = AdvancedStateGraph(dict)
async def start_node(state: dict[str, Any]) -> Command:
_ = state
return Command(goto=Send("middle_0", None))
async def end_node(state: dict[str, Any]) -> dict[str, Any]:
out = dict(state)
out["done"] = True
return out
graph.add_entry_node(start_node)
def make_middle(target: str):
async def middle_node(ctx: Any, state: dict[str, Any]) -> Command:
_ = state
await ctx.wait_for(timer_condition(seconds=SLEEP_SECONDS))
return Command(goto=Send(target, None))
return middle_node
for i in range(MIDDLE_COUNT):
next_name = "end_node" if i == MIDDLE_COUNT - 1 else f"middle_{i+1}"
graph.add_node(f"middle_{i}", make_middle(next_name))
graph.add_finish_node(end_node)
return graph.compile()
def build_stategraph_parallel() -> Any:
graph = StateGraph(dict)
async def start_node(state: dict[str, Any]) -> None:
_ = state
async def end_node(state: dict[str, Any]) -> dict[str, Any]:
out = dict(state)
out["done"] = True
return out
graph.add_node("start_node", start_node)
for i in range(MIDDLE_COUNT):
async def middle_node(state: dict[str, Any], idx: int = i) -> None:
_ = idx
_ = state
await asyncio.sleep(SLEEP_SECONDS)
graph.add_node(f"middle_{i}", middle_node)
graph.add_node("end_node", end_node)
graph.add_edge(START, "start_node")
for i in range(MIDDLE_COUNT):
graph.add_edge("start_node", f"middle_{i}")
graph.add_edge(f"middle_{i}", "end_node")
graph.add_edge("end_node", END)
return graph.compile()
def build_stategraph_sequential() -> Any:
graph = StateGraph(dict)
async def start_node(state: dict[str, Any]) -> None:
_ = state
async def end_node(state: dict[str, Any]) -> dict[str, Any]:
out = dict(state)
out["done"] = True
return out
graph.add_node("start_node", start_node)
for i in range(MIDDLE_COUNT):
async def middle_node(state: dict[str, Any], idx: int = i) -> None:
_ = idx
_ = state
await asyncio.sleep(SLEEP_SECONDS)
graph.add_node(f"middle_{i}", middle_node)
graph.add_node("end_node", end_node)
graph.add_edge(START, "start_node")
graph.add_edge("start_node", "middle_0")
for i in range(MIDDLE_COUNT - 1):
graph.add_edge(f"middle_{i}", f"middle_{i+1}")
graph.add_edge(f"middle_{MIDDLE_COUNT - 1}", "end_node")
graph.add_edge("end_node", END)
return graph.compile()
def build_advanced_parallel_blocking() -> Any:
graph: AdvancedStateGraph[dict[str, Any]] = AdvancedStateGraph(dict)
done_channel = "__bench_done_channel_blocking"
graph.add_async_channel(done_channel, str)
async def start_node(state: dict[str, Any]) -> Command:
_ = state
sends = [Send(f"middle_blocking_{i}", None) for i in range(MIDDLE_COUNT)]
sends.append(Send("end_node_blocking", None))
return Command(goto=sends)
async def end_node_blocking(ctx: Any, state: dict[str, Any]) -> dict[str, Any]:
await ctx.wait_for(channel_condition(done_channel, n=MIDDLE_COUNT))
out = dict(state)
out["done"] = True
return out
graph.add_entry_node(start_node)
for i in range(MIDDLE_COUNT):
async def middle_blocking(
ctx: Any, state: dict[str, Any], idx: int = i
) -> None:
_ = idx
_ = state
time.sleep(BLOCKING_SECONDS)
ctx.publish_to_channel(done_channel, "done")
graph.add_node(f"middle_blocking_{i}", middle_blocking)
graph.add_finish_node(end_node_blocking)
return graph.compile()
def build_advanced_sequential_blocking() -> Any:
graph: AdvancedStateGraph[dict[str, Any]] = AdvancedStateGraph(dict)
async def start_node(state: dict[str, Any]) -> Command:
_ = state
return Command(goto=Send("middle_blocking_seq_0", None))
async def end_node_blocking_seq(state: dict[str, Any]) -> dict[str, Any]:
out = dict(state)
out["done"] = True
return out
graph.add_entry_node(start_node)
def make_middle(target: str):
async def middle_blocking_seq(state: dict[str, Any]) -> Command:
_ = state
time.sleep(BLOCKING_SECONDS)
return Command(goto=Send(target, None))
return middle_blocking_seq
for i in range(MIDDLE_COUNT):
next_name = (
"end_node_blocking_seq"
if i == MIDDLE_COUNT - 1
else f"middle_blocking_seq_{i+1}"
)
graph.add_node(f"middle_blocking_seq_{i}", make_middle(next_name))
graph.add_finish_node(end_node_blocking_seq)
return graph.compile()
def build_stategraph_parallel_blocking() -> Any:
graph = StateGraph(dict)
async def start_node(state: dict[str, Any]) -> None:
_ = state
async def end_node(state: dict[str, Any]) -> dict[str, Any]:
out = dict(state)
out["done"] = True
return out
graph.add_node("start_node", start_node)
for i in range(MIDDLE_COUNT):
async def middle_blocking(state: dict[str, Any], idx: int = i) -> None:
_ = idx
_ = state
time.sleep(BLOCKING_SECONDS)
graph.add_node(f"middle_blocking_{i}", middle_blocking)
graph.add_node("end_node", end_node)
graph.add_edge(START, "start_node")
for i in range(MIDDLE_COUNT):
graph.add_edge("start_node", f"middle_blocking_{i}")
graph.add_edge(f"middle_blocking_{i}", "end_node")
graph.add_edge("end_node", END)
return graph.compile()
def build_stategraph_sequential_blocking() -> Any:
graph = StateGraph(dict)
async def start_node(state: dict[str, Any]) -> None:
_ = state
async def end_node(state: dict[str, Any]) -> dict[str, Any]:
out = dict(state)
out["done"] = True
return out
graph.add_node("start_node", start_node)
for i in range(MIDDLE_COUNT):
async def middle_blocking_seq(state: dict[str, Any], idx: int = i) -> None:
_ = idx
_ = state
time.sleep(BLOCKING_SECONDS)
graph.add_node(f"middle_blocking_seq_{i}", middle_blocking_seq)
graph.add_node("end_node", end_node)
graph.add_edge(START, "start_node")
graph.add_edge("start_node", "middle_blocking_seq_0")
for i in range(MIDDLE_COUNT - 1):
graph.add_edge(
f"middle_blocking_seq_{i}",
f"middle_blocking_seq_{i+1}",
)
graph.add_edge(f"middle_blocking_seq_{MIDDLE_COUNT - 1}", "end_node")
graph.add_edge("end_node", END)
return graph.compile()
async def run_benchmark(name: str, compiled: Any) -> float:
started = time.perf_counter()
tasks = [asyncio.create_task(compiled.ainvoke(make_initial_state())) for _ in range(RUNS)]
results = await asyncio.gather(*tasks)
elapsed = time.perf_counter() - started
if not all(item.get("done") is True for item in results):
raise RuntimeError(f"{name} produced unfinished runs")
return elapsed
async def main() -> None:
suites = [
("advanced-graph-parallel", build_advanced_parallel()),
("advanced-graph-sequential", build_advanced_sequential()),
("state-graph-parallel", build_stategraph_parallel()),
("state-graph-sequential", build_stategraph_sequential()),
("advanced-graph-parallel-blocking", build_advanced_parallel_blocking()),
("advanced-graph-sequential-blocking", build_advanced_sequential_blocking()),
("state-graph-parallel-blocking", build_stategraph_parallel_blocking()),
("state-graph-sequential-blocking", build_stategraph_sequential_blocking()),
]
print(
f"runs={RUNS}, middle_nodes={MIDDLE_COUNT}, sleep={SLEEP_SECONDS}s, "
f"blocking_sleep={BLOCKING_SECONDS}s, state_bytes={STATE_BYTES}"
)
for name, compiled in suites:
elapsed = await run_benchmark(name, compiled)
print(f"{name}: {elapsed:.3f}s")
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,66 @@
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",
]
async def test_run_ends_without_finish_node() -> None:
graph = AdvancedStateGraph(PrimitiveState)
async def start_node(state: PrimitiveState) -> Command:
state["logs"].append("start")
return Command(update=state, goto=Send("middle_node", "from_start"))
async def middle_node(input: str, state: PrimitiveState) -> dict[str, object]:
state["logs"].append(f"middle:{input}")
return {"counter": state["counter"] + 1, "logs": state["logs"], "done": "stopped"}
graph.add_entry_node(start_node)
graph.add_node(middle_node)
result = await graph.compile().ainvoke({"counter": 7, "logs": [], "done": None})
assert result["counter"] == 8
assert result["done"] == "stopped"
assert result["logs"] == ["start", "middle:from_start"]
@@ -0,0 +1,55 @@
import os
import subprocess
import sys
def test_run_pool_size_one_still_allows_parallel_runs() -> None:
script = r"""
import asyncio
import time
from typing_extensions import TypedDict
from langgraph.advanced_graph import AdvancedStateGraph, Context, timer_condition
from langgraph.types import Command, Send
class RunState(TypedDict):
done: bool
async def wait_node(ctx: Context, _: object, state: RunState) -> Command:
await ctx.wait_for(timer_condition(seconds=0.2))
return Command(goto=Send("finish_node", None), update=state)
async def finish_node(_: object, state: RunState) -> dict[str, bool]:
return {"done": True}
async def main() -> None:
graph = AdvancedStateGraph(RunState)
graph.add_entry_node(wait_node)
graph.add_finish_node(finish_node)
compiled = graph.compile()
started = time.perf_counter()
await asyncio.gather(
compiled.ainvoke({"done": False}),
compiled.ainvoke({"done": False}),
)
elapsed = time.perf_counter() - started
print(f"{elapsed:.6f}")
asyncio.run(main())
"""
env = os.environ.copy()
env["LANGGRAPH_RUN_POOL_SIZE"] = "1"
env.setdefault("LANGGRAPH_NODE_POOL_SIZE", "2")
completed = subprocess.run(
[sys.executable, "-c", script],
env=env,
capture_output=True,
text=True,
check=True,
)
elapsed = float(completed.stdout.strip().splitlines()[-1])
assert elapsed < 0.35, completed.stdout
@@ -38,7 +38,7 @@ class Decision:
complete: str | None = None
class MockPlanner:
class MockLLM:
def __init__(self) -> None:
self.responses: list[list[Decision]] = []
self._idx = 0
@@ -66,7 +66,7 @@ def build_sub_agent() -> Any:
return sub_agent.compile()
def build_main_agent(planner: MockPlanner, sub_agent: Any) -> Any:
def build_main_agent(planner: MockLLM, sub_agent: Any) -> Any:
async def llm_node(state: MainAgentState) -> Command:
# Planner decides whether to call a tool, spawn a sub-agent, or finish.
decisions = await planner.ainvoke(state)
@@ -78,10 +78,7 @@ def build_main_agent(planner: MockPlanner, sub_agent: Any) -> Any:
return Command(
goto=Send(
order_food_node,
{
"state": state,
"complete": decision.complete or "order flow completed",
},
decision.complete or "order flow completed",
)
)
if decision.type == "sub_agent" and decision.sub_agent:
@@ -89,7 +86,7 @@ def build_main_agent(planner: MockPlanner, sub_agent: Any) -> Any:
if decision.type == "tool" and decision.tool:
sends.append(Send("tool_node", decision.tool))
# Keep the main loop responsive: wait for one inbound message and continue.
sends.append(Send("wait_node", state))
sends.append(Send("wait_node", None))
return Command(goto=sends)
async def wait_node(ctx: Context, state: MainAgentState) -> Command:
@@ -112,11 +109,11 @@ def build_main_agent(planner: MockPlanner, sub_agent: Any) -> Any:
elif channel == "user_input_channel":
state["output"].append(f"user_input: {payload}")
# State changed -> ask planner what to do next.
return Command(goto=Send("llm_node", state))
return Command(update=state, goto=Send("llm_node", None))
else:
state["output"].append("timer: no updates yet")
# No meaningful state change -> keep waiting without calling planner.
return Command(goto=Send("wait_node", state))
return Command(update=state, goto=Send("wait_node", None))
async def tool_node(ctx: Context, tool_input: str) -> None:
await asyncio.sleep(0.1)
@@ -138,9 +135,8 @@ def build_main_agent(planner: MockPlanner, sub_agent: Any) -> Any:
sub_agent_output["output"],
)
async def order_food_node(payload: dict[str, Any]) -> dict[str, Any]:
state = payload["state"]
complete_message = payload["complete"]
async def order_food_node(input: str, state: MainAgentState) -> dict[str, Any]:
complete_message = input
return {
"done": complete_message,
"output": [*state["output"], f"order_food: {complete_message}"],
@@ -162,11 +158,11 @@ def build_main_agent(planner: MockPlanner, sub_agent: Any) -> Any:
async def test_async_sub_graph() -> None:
planner = MockPlanner()
llm = MockLLM()
sub_agent = build_sub_agent()
main_agent = build_main_agent(planner, sub_agent)
main_agent = build_main_agent(llm, sub_agent)
planner.responses = [
llm.responses = [
[
# First planner pass triggers one slow sub-agent.
Decision(type="sub_agent", sub_agent="research lunch options"),
@@ -202,7 +198,8 @@ async def test_async_sub_graph() -> None:
"sub_agent: research sub agent completed for: research lunch options" in output
)
assert (
"sub_agent: research sub agent completed for: find vegetarian fallback" in output
"sub_agent: research sub agent completed for: find vegetarian fallback"
in output
)
assert output[-1] == "order_food: order submitted"
@@ -214,6 +211,7 @@ async def test_async_sub_graph() -> None:
)
order_food_idx = output.index("order_food: order submitted")
assert first_sub_idx < second_sub_idx < order_food_idx
assert planner._idx == len(planner.responses)
assert llm._idx == len(llm.responses)
import json
print(json.dumps(result, ensure_ascii=False, indent=2))
@@ -0,0 +1,126 @@
import asyncio
from dataclasses import dataclass
from pydantic import BaseModel
from typing_extensions import TypedDict
import pytest
from langgraph.advanced_graph import AdvancedStateGraph, CompiledGraphEngine
from langgraph.types import Command, Send
pytestmark = pytest.mark.anyio
@dataclass
class DataClassPayload:
value: int
class PydanticPayload(BaseModel):
value: int
class InnerTypedDict(TypedDict):
flag: bool
n: int
@dataclass
class UpdateElisionState:
x: int
dc: DataClassPayload
model: PydanticPayload
td: InnerTypedDict
obj: dict[str, int]
items: list[int]
def _initial_state() -> UpdateElisionState:
return UpdateElisionState(
x=0,
dc=DataClassPayload(0),
model=PydanticPayload(value=0),
td={"flag": False, "n": 0},
obj={"n": 0},
items=[0],
)
async def test_noop_slow_update_does_not_override_fast_update() -> None:
graph: AdvancedStateGraph[UpdateElisionState] = AdvancedStateGraph(UpdateElisionState)
async def start_node(state: UpdateElisionState) -> Command:
return Command(goto=[Send("fast_node", None), Send("slow_node", None)])
async def fast_node(state: UpdateElisionState) -> UpdateElisionState:
state.x = 1
state.dc.value = 1
state.model.value = 1
state.td["flag"] = True
state.td["n"] = 1
state.obj["n"] = 1
state.items.append(1)
return state
async def slow_node(state: UpdateElisionState) -> UpdateElisionState:
await asyncio.sleep(0.1)
# Returns the same values as the initial snapshot.
return state
graph.add_entry_node(start_node)
graph.add_node(fast_node)
graph.add_finish_node(slow_node)
compiled: CompiledGraphEngine[UpdateElisionState] = graph.compile()
initial_state: UpdateElisionState = _initial_state()
result: UpdateElisionState = await compiled.ainvoke(initial_state)
assert result.x == 1
assert result.dc == DataClassPayload(1)
assert result.model.value == 1
assert result.td == {"flag": True, "n": 1}
assert result.obj == {"n": 1}
assert result.items == [0, 1]
async def test_changed_slow_update_overrides_fast_update() -> None:
graph: AdvancedStateGraph[UpdateElisionState] = AdvancedStateGraph(UpdateElisionState)
async def start_node(state: UpdateElisionState) -> Command:
return Command(goto=[Send("fast_node", None), Send("slow_node", None)])
async def fast_node(state: UpdateElisionState) -> UpdateElisionState:
state.x = 1
state.dc.value = 1
state.model.value = 1
state.td["flag"] = True
state.td["n"] = 1
state.obj["n"] = 1
state.items.append(1)
return state
async def slow_node(state: UpdateElisionState) -> UpdateElisionState:
await asyncio.sleep(0.1)
# Slow node makes real changes for all field types.
state.x = 2
state.dc.value = 2
state.model.value = 2
state.td["flag"] = False
state.td["n"] = 2
state.obj["n"] = 2
state.items.append(2)
return state
graph.add_entry_node(start_node)
graph.add_node(fast_node)
graph.add_finish_node(slow_node)
compiled: CompiledGraphEngine[UpdateElisionState] = graph.compile()
initial_state: UpdateElisionState = _initial_state()
result: UpdateElisionState = await compiled.ainvoke(initial_state)
assert result.x == 2
assert result.dc == DataClassPayload(2)
assert result.model.value == 2
assert result.td == {"flag": False, "n": 2}
assert result.obj == {"n": 2}
assert result.items == [0, 1, 2]
+338
View File
@@ -0,0 +1,338 @@
# This file is automatically @generated by Cargo.
# It is not intended for manual editing.
version = 4
[[package]]
name = "autocfg"
version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8"
[[package]]
name = "bitflags"
version = "2.11.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af"
[[package]]
name = "cfg-if"
version = "1.0.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801"
[[package]]
name = "heck"
version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea"
[[package]]
name = "indoc"
version = "2.0.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "79cf5c93f93228cf8efb3ba362535fb11199ac548a09ce117c9b1adc3030d706"
dependencies = [
"rustversion",
]
[[package]]
name = "itoa"
version = "1.0.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "92ecc6618181def0457392ccd0ee51198e065e016d1d527a7ac1b6dc7c1f09d2"
[[package]]
name = "langgraph_rust_core"
version = "0.1.0"
dependencies = [
"libc",
"parking_lot",
"pyo3",
"serde",
"serde_json",
"tokio",
]
[[package]]
name = "libc"
version = "0.2.183"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b5b646652bf6661599e1da8901b3b9522896f01e736bad5f723fe7a3a27f899d"
[[package]]
name = "lock_api"
version = "0.4.14"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965"
dependencies = [
"scopeguard",
]
[[package]]
name = "memchr"
version = "2.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79"
[[package]]
name = "memoffset"
version = "0.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "488016bfae457b036d996092f6cb448677611ce4449e970ceaf42695203f218a"
dependencies = [
"autocfg",
]
[[package]]
name = "once_cell"
version = "1.21.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50"
[[package]]
name = "parking_lot"
version = "0.12.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "93857453250e3077bd71ff98b6a65ea6621a19bb0f559a85248955ac12c45a1a"
dependencies = [
"lock_api",
"parking_lot_core",
]
[[package]]
name = "parking_lot_core"
version = "0.9.12"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1"
dependencies = [
"cfg-if",
"libc",
"redox_syscall",
"smallvec",
"windows-link",
]
[[package]]
name = "pin-project-lite"
version = "0.2.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd"
[[package]]
name = "portable-atomic"
version = "1.13.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49"
[[package]]
name = "proc-macro2"
version = "1.0.106"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934"
dependencies = [
"unicode-ident",
]
[[package]]
name = "pyo3"
version = "0.23.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7778bffd85cf38175ac1f545509665d0b9b92a198ca7941f131f85f7a4f9a872"
dependencies = [
"cfg-if",
"indoc",
"libc",
"memoffset",
"once_cell",
"portable-atomic",
"pyo3-build-config",
"pyo3-ffi",
"pyo3-macros",
"unindent",
]
[[package]]
name = "pyo3-build-config"
version = "0.23.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "94f6cbe86ef3bf18998d9df6e0f3fc1050a8c5efa409bf712e661a4366e010fb"
dependencies = [
"once_cell",
"target-lexicon",
]
[[package]]
name = "pyo3-ffi"
version = "0.23.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e9f1b4c431c0bb1c8fb0a338709859eed0d030ff6daa34368d3b152a63dfdd8d"
dependencies = [
"libc",
"pyo3-build-config",
]
[[package]]
name = "pyo3-macros"
version = "0.23.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fbc2201328f63c4710f68abdf653c89d8dbc2858b88c5d88b0ff38a75288a9da"
dependencies = [
"proc-macro2",
"pyo3-macros-backend",
"quote",
"syn",
]
[[package]]
name = "pyo3-macros-backend"
version = "0.23.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fca6726ad0f3da9c9de093d6f116a93c1a38e417ed73bf138472cf4064f72028"
dependencies = [
"heck",
"proc-macro2",
"pyo3-build-config",
"quote",
"syn",
]
[[package]]
name = "quote"
version = "1.0.45"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "41f2619966050689382d2b44f664f4bc593e129785a36d6ee376ddf37259b924"
dependencies = [
"proc-macro2",
]
[[package]]
name = "redox_syscall"
version = "0.5.18"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d"
dependencies = [
"bitflags",
]
[[package]]
name = "rustversion"
version = "1.0.22"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d"
[[package]]
name = "scopeguard"
version = "1.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49"
[[package]]
name = "serde"
version = "1.0.228"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e"
dependencies = [
"serde_core",
"serde_derive",
]
[[package]]
name = "serde_core"
version = "1.0.228"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad"
dependencies = [
"serde_derive",
]
[[package]]
name = "serde_derive"
version = "1.0.228"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79"
dependencies = [
"proc-macro2",
"quote",
"syn",
]
[[package]]
name = "serde_json"
version = "1.0.149"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86"
dependencies = [
"itoa",
"memchr",
"serde",
"serde_core",
"zmij",
]
[[package]]
name = "smallvec"
version = "1.15.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03"
[[package]]
name = "syn"
version = "2.0.117"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e665b8803e7b1d2a727f4023456bbbbe74da67099c585258af0ad9c5013b9b99"
dependencies = [
"proc-macro2",
"quote",
"unicode-ident",
]
[[package]]
name = "target-lexicon"
version = "0.12.16"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "61c41af27dd6d1e27b1b16b489db798443478cef1f06a660c96db617ba5de3b1"
[[package]]
name = "tokio"
version = "1.50.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "27ad5e34374e03cfffefc301becb44e9dc3c17584f414349ebe29ed26661822d"
dependencies = [
"pin-project-lite",
"tokio-macros",
]
[[package]]
name = "tokio-macros"
version = "2.6.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5c55a2eff8b69ce66c84f85e1da1c233edc36ceb85a2058d11b0d6a3c7e7569c"
dependencies = [
"proc-macro2",
"quote",
"syn",
]
[[package]]
name = "unicode-ident"
version = "1.0.24"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75"
[[package]]
name = "unindent"
version = "0.2.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7264e107f553ccae879d21fbea1d6724ac785e8c3bfc762137959b5802826ef3"
[[package]]
name = "windows-link"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
[[package]]
name = "zmij"
version = "1.0.21"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa"
+21
View File
@@ -0,0 +1,21 @@
[package]
name = "langgraph_rust_core"
version = "0.1.0"
edition = "2021"
[lib]
name = "langgraph_rust_core"
crate-type = ["cdylib", "rlib"]
[features]
default = ["python-bindings"]
python-bindings = ["dep:pyo3"]
[dependencies]
pyo3 = { version = "0.23.5", features = ["extension-module"], optional = true }
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
parking_lot = "0.12"
libc = "0.2"
tokio = { version = "1", features = ["macros", "rt-multi-thread", "sync", "time"] }
+39
View File
@@ -0,0 +1,39 @@
#ifndef LANGGRAPH_RUST_CORE_H
#define LANGGRAPH_RUST_CORE_H
#ifdef __cplusplus
extern "C" {
#endif
typedef struct Engine Engine;
typedef char* (*rc_node_callback_t)(
unsigned long user_data,
char* node,
char* arg_json,
char* state_json
);
Engine* rc_engine_new(void);
void rc_engine_free(Engine* ptr);
char* rc_add_async_channel(Engine* ptr, const char* channel);
char* rc_publish_json(Engine* ptr, const char* channel, const char* value_json);
char* rc_wait_any_of_json(Engine* ptr, const char* any_of_json);
char* rc_run_graph_json(
Engine* ptr,
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
);
void rc_string_free(char* ptr);
#ifdef __cplusplus
}
#endif
#endif
+433
View File
@@ -0,0 +1,433 @@
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::{HashMap, VecDeque};
use std::env;
use std::future::Future;
use std::sync::mpsc;
use std::sync::Arc;
use std::sync::Mutex as StdMutex;
use std::sync::OnceLock;
use std::thread;
use std::time::{Duration, Instant};
use tokio::runtime::Runtime;
use tokio::sync::Notify;
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(tag = "kind")]
pub enum WaitCondition {
#[serde(rename = "channel")]
Channel { channel: String, n: usize },
#[serde(rename = "timer")]
Timer { seconds: f64 },
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct AnyOfCondition {
pub conditions: Vec<WaitCondition>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(tag = "condition")]
pub enum WaitEvent {
#[serde(rename = "channel")]
Channel {
channel: String,
value: serde_json::Value,
},
#[serde(rename = "timer")]
Timer { seconds: f64 },
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(tag = "kind")]
pub enum WaitRequest {
#[serde(rename = "condition")]
Condition { condition: WaitCondition },
#[serde(rename = "any_of")]
AnyOf { any_of: AnyOfCondition },
}
pub struct SendPayload<A> {
pub node: String,
pub arg: A,
}
pub struct NodeExecResult<U, A> {
pub update: Option<U>,
pub sends: Vec<SendPayload<A>>,
}
pub enum NodeOutcome<U, A> {
Completed(NodeExecResult<U, A>),
Suspended { wait: WaitRequest },
}
type Task = Box<dyn FnOnce() + Send + 'static>;
fn debug_enabled() -> bool {
static DEBUG: OnceLock<bool> = OnceLock::new();
*DEBUG.get_or_init(|| {
matches!(
env::var("DEBUG")
.unwrap_or_default()
.trim()
.to_ascii_lowercase()
.as_str(),
"1" | "true" | "yes" | "on"
)
})
}
fn debug_log(message: &str) {
if debug_enabled() {
let current = thread::current();
let thread_name = current.name().unwrap_or("unnamed");
println!("[advanced-graph][{thread_name}] {message}");
}
}
fn pool_size_from_env(var_name: &str, default: usize, min: usize) -> usize {
let parsed = env::var(var_name)
.ok()
.and_then(|raw| raw.trim().parse::<usize>().ok());
parsed.unwrap_or(default).max(min)
}
struct ThreadPool {
tx: mpsc::Sender<Task>,
_workers: Vec<thread::JoinHandle<()>>,
}
impl ThreadPool {
fn new(size: usize, label: &str) -> Self {
let (tx, rx) = mpsc::channel::<Task>();
let rx = Arc::new(StdMutex::new(rx));
let mut workers = Vec::with_capacity(size);
for idx in 0..size {
let thread_name = format!("{label}-{idx}");
let rx = Arc::clone(&rx);
let handle = thread::Builder::new()
.name(thread_name)
.spawn(move || loop {
let task = {
let guard = rx.lock().expect("thread-pool receiver mutex poisoned");
guard.recv()
};
match task {
Ok(task) => task(),
Err(_) => break,
}
})
.expect("failed to spawn thread-pool worker");
workers.push(handle);
}
Self {
tx,
_workers: workers,
}
}
fn execute<F>(&self, task: F) -> Result<(), String>
where
F: FnOnce() + Send + 'static,
{
debug_log("thread-pool execute() called");
self.tx
.send(Box::new(task))
.map_err(|e| format!("thread-pool send failed: {e}"))
}
}
pub fn node_pool_execute<F>(task: F) -> Result<(), String>
where
F: FnOnce() + Send + 'static,
{
debug_log("node_pool_execute() called");
static NODE_POOL: OnceLock<ThreadPool> = OnceLock::new();
let pool = NODE_POOL.get_or_init(|| {
let default_size = thread::available_parallelism()
.map(|n| n.get().max(2))
.unwrap_or(4);
let size = pool_size_from_env("LANGGRAPH_NODE_POOL_SIZE", default_size, 1);
ThreadPool::new(size, "langgraph-node")
});
pool.execute(task)
}
pub fn run_loop_pool_execute<F>(task: F) -> Result<(), String>
where
F: FnOnce() + Send + 'static,
{
debug_log("run_loop_pool_execute() called");
run_runtime().spawn_blocking(task);
Ok(())
}
pub fn run_loop_spawn<F>(future: F) -> Result<(), String>
where
F: Future<Output = ()> + Send + 'static,
{
run_runtime().spawn(future);
Ok(())
}
pub fn run_loop_block_on<F>(future: F) -> F::Output
where
F: Future,
{
run_runtime().block_on(future)
}
fn run_runtime() -> &'static Runtime {
static RUNTIME: OnceLock<Runtime> = OnceLock::new();
RUNTIME.get_or_init(|| {
let default_size = thread::available_parallelism()
.map(|n| n.get().max(2))
.unwrap_or(2);
let worker_threads = pool_size_from_env("LANGGRAPH_RUN_POOL_SIZE", default_size, 1);
tokio::runtime::Builder::new_multi_thread()
.worker_threads(worker_threads)
.thread_name("langgraph-runloop")
.enable_all()
.build()
.expect("failed to build tokio runtime")
})
}
pub fn run_scheduler_loop<U: Send + 'static, A: Send + 'static, FSpawn, FMerge>(
entry_point: String,
finish_point: &str,
initial_arg: A,
mut spawn: FSpawn,
mut merge: FMerge,
rx: mpsc::Receiver<Result<(String, NodeExecResult<U, A>), String>>,
) -> Result<(), String>
where
FSpawn: FnMut(String, A) -> Result<(), String>,
FMerge: FnMut(&str, Option<U>) -> Result<(), String>,
{
debug_log("run_scheduler_loop() started");
let mut active: usize = 1;
debug_log("scheduling initial entry node");
spawn(entry_point, initial_arg)?;
while active > 0 {
debug_log(&format!(
"scheduler waiting for node result (active={active})"
));
let item = rx
.recv()
.map_err(|e| format!("scheduler recv failed: {e}"))?;
active = active.saturating_sub(1);
let (node_name, node_result) = item.map_err(|e| format!("node execution failed: {e}"))?;
debug_log(&format!("scheduler received result from node={node_name}"));
merge(&node_name, node_result.update)?;
debug_log(&format!("merged update from node={node_name}"));
if node_name == finish_point {
debug_log("finish node reached, stopping scheduler loop");
break;
}
for send in node_result.sends {
active += 1;
debug_log(&format!(
"scheduling next node={} (active={active})",
send.node
));
spawn(send.node, send.arg)?;
}
}
debug_log("run_scheduler_loop() finished");
Ok(())
}
pub fn merge_json_update(state: &mut Value, update: Option<Value>) {
debug_log("merge_json_update() called");
let Some(update_value) = update else {
debug_log("merge_json_update(): no update payload");
return;
};
match (&mut *state, update_value) {
(Value::Object(state_obj), Value::Object(update_obj)) => {
debug_log(&format!(
"merge_json_update(): object merge with {} keys",
update_obj.len()
));
for (k, v) in update_obj {
state_obj.insert(k, v);
}
}
(Value::Object(state_obj), Value::Array(entries)) => {
debug_log(&format!(
"merge_json_update(): tuple-list merge with {} entries",
entries.len()
));
for entry in entries {
if let Value::Array(pair) = entry {
if pair.len() == 2 {
if let Value::String(key) = &pair[0] {
state_obj.insert(key.clone(), pair[1].clone());
}
}
}
}
}
_ => debug_log("merge_json_update(): unsupported update shape, ignored"),
}
}
#[derive(Clone, Default)]
pub struct Engine {
channels: Arc<StdMutex<HashMap<String, VecDeque<serde_json::Value>>>>,
channel_notify: Arc<Notify>,
}
impl Engine {
pub fn new() -> Self {
debug_log("Engine::new()");
Self::default()
}
pub fn add_async_channel(&self, name: &str) {
debug_log(&format!("Engine::add_async_channel(name={name})"));
let mut channels = self.channels.lock().expect("channels mutex poisoned");
channels.entry(name.to_owned()).or_default();
}
pub fn publish_json(&self, channel: &str, value: serde_json::Value) -> Result<(), String> {
debug_log(&format!("Engine::publish_json(channel={channel})"));
let mut channels = self.channels.lock().expect("channels mutex poisoned");
let queue = channels
.get_mut(channel)
.ok_or_else(|| format!("Unknown channel `{channel}`"))?;
queue.push_back(value);
self.channel_notify.notify_waiters();
Ok(())
}
pub async fn wait_request_async(&self, wait: &WaitRequest) -> Result<WaitEvent, String> {
match wait {
WaitRequest::Condition { condition } => self.wait_for_async(condition).await,
WaitRequest::AnyOf { any_of } => self.wait_for_any_of_async(any_of).await,
}
}
pub async fn wait_for_async(&self, cond: &WaitCondition) -> Result<WaitEvent, String> {
debug_log(&format!("Engine::wait_for_async(cond={cond:?})"));
match cond {
WaitCondition::Channel { channel, n } => {
if *n < 1 {
return Err("channel condition n must be >= 1".to_string());
}
loop {
if let Some(event) = self.try_take_channel_event(channel, *n)? {
return Ok(event);
}
self.channel_notify.notified().await;
}
}
WaitCondition::Timer { seconds } => {
if *seconds <= 0.0 {
return Err("timer condition must be > 0".to_string());
}
tokio::time::sleep(Duration::from_secs_f64(*seconds)).await;
Ok(WaitEvent::Timer { seconds: *seconds })
}
}
}
pub async fn wait_for_any_of_async(&self, any_of: &AnyOfCondition) -> Result<WaitEvent, String> {
debug_log(&format!(
"Engine::wait_for_any_of_async(conditions={})",
any_of.conditions.len()
));
if any_of.conditions.is_empty() {
return Err("any_of requires at least one condition".to_string());
}
let started = Instant::now();
let mut min_timer: Option<f64> = None;
for cond in &any_of.conditions {
if let WaitCondition::Timer { seconds } = cond {
if *seconds <= 0.0 {
return Err("timer condition must be > 0".to_string());
}
min_timer = Some(min_timer.map_or(*seconds, |x| x.min(*seconds)));
}
}
loop {
for cond in &any_of.conditions {
if let WaitCondition::Channel { channel, n } = cond {
if *n < 1 {
return Err("channel condition n must be >= 1".to_string());
}
if let Some(event) = self.try_take_channel_event(channel, *n)? {
return Ok(event);
}
}
}
if let Some(seconds) = min_timer {
let timeout = Duration::from_secs_f64(seconds);
let elapsed = started.elapsed();
if elapsed >= timeout {
return Ok(WaitEvent::Timer { seconds });
}
let remaining = timeout.saturating_sub(elapsed);
tokio::select! {
_ = self.channel_notify.notified() => {}
_ = tokio::time::sleep(remaining) => {
return Ok(WaitEvent::Timer { seconds });
}
}
} else {
self.channel_notify.notified().await;
}
}
}
pub fn wait_for(&self, cond: &WaitCondition) -> Result<WaitEvent, String> {
run_loop_block_on(self.wait_for_async(cond))
}
pub fn wait_for_any_of(&self, any_of: &AnyOfCondition) -> Result<WaitEvent, String> {
run_loop_block_on(self.wait_for_any_of_async(any_of))
}
fn try_take_channel_event(
&self,
channel: &str,
n: usize,
) -> Result<Option<WaitEvent>, String> {
let mut channels = self.channels.lock().expect("channels mutex poisoned");
let queue = channels
.get_mut(channel)
.ok_or_else(|| format!("Unknown channel `{channel}`"))?;
if queue.len() < n {
return Ok(None);
}
if n == 1 {
if let Some(value) = queue.pop_front() {
return Ok(Some(WaitEvent::Channel {
channel: channel.to_string(),
value,
}));
}
return Ok(None);
}
let mut values = Vec::with_capacity(n);
for _ in 0..n {
if let Some(v) = queue.pop_front() {
values.push(v);
}
}
Ok(Some(WaitEvent::Channel {
channel: channel.to_string(),
value: serde_json::Value::Array(values),
}))
}
}
+4
View File
@@ -0,0 +1,4 @@
mod engine;
mod lib_c;
#[cfg(feature = "python-bindings")]
mod lib_py;
+475
View File
@@ -0,0 +1,475 @@
use crate::engine::{
merge_json_update, node_pool_execute, run_loop_block_on, run_loop_spawn, AnyOfCondition,
Engine, NodeExecResult, NodeOutcome, SendPayload, WaitEvent, WaitRequest,
};
use serde::Deserialize;
use serde_json::Value;
use std::ffi::{CStr, CString};
use std::os::raw::c_char;
use std::sync::mpsc;
use std::sync::{Arc, Mutex};
use tokio::sync::mpsc as tokio_mpsc;
#[derive(Debug, Deserialize)]
struct SendPayloadJson {
node: String,
#[serde(default)]
arg: Value,
}
#[derive(Debug, Deserialize)]
struct NodeExecResultJsonWire {
update: Option<Value>,
#[serde(default)]
sends: Vec<SendPayloadJson>,
}
#[derive(Debug, Deserialize)]
struct CallbackEnvelopeIn {
ok: bool,
#[serde(default)]
payload: Option<NodeExecResultJsonWire>,
#[serde(default)]
suspend: Option<WaitRequest>,
#[serde(default)]
error: Option<String>,
}
type CNodeCallback = unsafe extern "C" fn(
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(libc::c_ulong);
fn cstr_to_str<'a>(ptr: *const c_char) -> Result<&'a str, String> {
if ptr.is_null() {
return Err("Received null pointer".to_string());
}
let cstr = unsafe { CStr::from_ptr(ptr) };
cstr.to_str()
.map_err(|e| format!("Invalid UTF-8 input string: {e}"))
}
fn into_c_ptr(s: String) -> *mut c_char {
match CString::new(s) {
Ok(c) => c.into_raw(),
Err(_) => CString::new("{\"error\":\"NUL byte in output\"}")
.expect("static string is valid")
.into_raw(),
}
}
fn parse_c_callback_result(
raw: String,
node_name: &str,
) -> Result<NodeOutcome<Value, Value>, String> {
let parsed: CallbackEnvelopeIn = serde_json::from_str(&raw)
.map_err(|e| format!("decode callback envelope for `{node_name}` failed: {e}"))?;
if !parsed.ok {
return Err(parsed
.error
.unwrap_or_else(|| format!("callback reported error for `{node_name}`")));
}
if let Some(wait) = parsed.suspend {
return Ok(NodeOutcome::Suspended { wait });
}
let payload = parsed
.payload
.ok_or_else(|| format!("callback payload missing for `{node_name}`"))?;
let sends = payload
.sends
.into_iter()
.map(|s| SendPayload {
node: s.node,
arg: s.arg,
})
.collect();
Ok(NodeOutcome::Completed(NodeExecResult {
update: payload.update,
sends,
}))
}
enum SchedulerEventJson {
Node(Result<NodeExecutionJson, String>),
Resume {
node: String,
arg: Value,
event: WaitEvent,
},
WaitError(String),
}
struct NodeExecutionJson {
node: String,
arg: Value,
outcome: NodeOutcome<Value, Value>,
}
fn spawn_json_node_task(
node: String,
arg: Value,
state_snapshot: Value,
tx: tokio_mpsc::UnboundedSender<SchedulerEventJson>,
user_data_bits: libc::c_ulong,
callback: CNodeCallback,
) -> Result<(), String> {
node_pool_execute(move || {
let node_for_result = node.clone();
let arg_for_result = arg.clone();
let result = (|| -> Result<NodeExecutionJson, String> {
let node_c =
CString::new(node.clone()).map_err(|e| format!("invalid node name: {e}"))?;
let arg_json = serde_json::to_string(&arg)
.map_err(|e| format!("serialize arg for `{node}` failed: {e}"))?;
let state_json = serde_json::to_string(&state_snapshot)
.map_err(|e| format!("serialize state for `{node}` failed: {e}"))?;
let arg_c =
CString::new(arg_json).map_err(|e| format!("invalid arg JSON bytes: {e}"))?;
let state_c =
CString::new(state_json).map_err(|e| format!("invalid state JSON bytes: {e}"))?;
let out_ptr = unsafe {
callback(
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,
)
};
if out_ptr.is_null() {
return Err(format!("callback returned null for `{node}`"));
}
let out_raw = unsafe { CStr::from_ptr(out_ptr) }
.to_string_lossy()
.into_owned();
unsafe {
libc::free(out_ptr.cast());
}
let payload = parse_c_callback_result(out_raw, &node)?;
Ok(NodeExecutionJson {
node: node_for_result,
arg: arg_for_result,
outcome: payload,
})
})();
let _ = tx.send(SchedulerEventJson::Node(result));
})
}
async fn run_graph_scheduler_json(
entry_point: String,
finish_point: String,
initial_state: Value,
initial_input: Value,
engine: Engine,
user_data: CUserData,
callback: CNodeCallback,
) -> Result<Value, String> {
let (tx, mut rx) = tokio_mpsc::unbounded_channel::<SchedulerEventJson>();
let state = Arc::new(Mutex::new(initial_state));
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);
let mut active: usize = 1;
let mut waiting: usize = 0;
spawn_json_node_task(
entry_point,
initial_input,
state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone(),
tx_for_spawn.clone(),
user_data_bits,
callback,
)?;
while active > 0 || waiting > 0 {
let evt = rx
.recv()
.await
.ok_or_else(|| "scheduler event channel closed".to_string())?;
match evt {
SchedulerEventJson::Node(result) => {
active = active.saturating_sub(1);
let exec = result?;
match exec.outcome {
NodeOutcome::Completed(node_result) => {
let mut guard = state_for_merge.lock().expect("state mutex poisoned");
merge_json_update(&mut guard, node_result.update);
drop(guard);
if exec.node == finish_point {
break;
}
for send in node_result.sends {
active += 1;
let snapshot = state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone();
spawn_json_node_task(
send.node,
send.arg,
snapshot,
tx_for_spawn.clone(),
user_data_bits,
callback,
)?;
}
}
NodeOutcome::Suspended { wait } => {
waiting += 1;
let tx_wait = tx_for_spawn.clone();
let node = exec.node;
let arg = exec.arg;
let engine_for_wait = engine.clone();
tokio::spawn(async move {
match engine_for_wait.wait_request_async(&wait).await {
Ok(event) => {
let _ = tx_wait.send(SchedulerEventJson::Resume { node, arg, event });
}
Err(e) => {
let _ = tx_wait.send(SchedulerEventJson::WaitError(e));
}
}
});
}
}
}
SchedulerEventJson::Resume { node, arg, event } => {
waiting = waiting.saturating_sub(1);
active += 1;
let snapshot = state_for_spawn
.lock()
.expect("state mutex poisoned")
.clone();
spawn_json_node_task(
node,
wrap_resume_arg(arg, event),
snapshot,
tx_for_spawn.clone(),
user_data_bits,
callback,
)?;
}
SchedulerEventJson::WaitError(e) => return Err(e),
}
}
let final_state = state.lock().expect("state mutex poisoned").clone();
Ok(final_state)
}
fn wrap_resume_arg(arg: Value, event: WaitEvent) -> Value {
serde_json::json!({
"__lg_resume_arg__": arg,
"__lg_resume_event__": event,
})
}
#[no_mangle]
pub extern "C" fn rc_engine_new() -> *mut Engine {
Box::into_raw(Box::new(Engine::new()))
}
#[no_mangle]
/// # Safety
/// `ptr` must be either null or a valid pointer returned by `rc_engine_new`.
pub unsafe extern "C" fn rc_engine_free(ptr: *mut Engine) {
if ptr.is_null() {
return;
}
drop(Box::from_raw(ptr));
}
#[no_mangle]
/// # Safety
/// `ptr` must be either null or a valid pointer returned by this library.
pub unsafe extern "C" fn rc_string_free(ptr: *mut c_char) {
if ptr.is_null() {
return;
}
drop(CString::from_raw(ptr));
}
#[no_mangle]
/// # Safety
/// `ptr` must be a valid engine pointer from `rc_engine_new`.
/// `channel` must be a valid null-terminated UTF-8 string pointer.
pub unsafe extern "C" fn rc_add_async_channel(
ptr: *mut Engine,
channel: *const c_char,
) -> *mut c_char {
if ptr.is_null() {
return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string());
}
let channel = match cstr_to_str(channel) {
Ok(v) => v,
Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
};
(*ptr).add_async_channel(channel);
into_c_ptr("{\"ok\":true}".to_string())
}
#[no_mangle]
/// # Safety
/// `ptr` must be a valid engine pointer from `rc_engine_new`.
/// `channel` and `value_json` must be valid null-terminated UTF-8 string pointers.
pub unsafe extern "C" fn rc_publish_json(
ptr: *mut Engine,
channel: *const c_char,
value_json: *const c_char,
) -> *mut c_char {
if ptr.is_null() {
return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string());
}
let channel = match cstr_to_str(channel) {
Ok(v) => v,
Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
};
let value_json = match cstr_to_str(value_json) {
Ok(v) => v,
Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
};
let value: Value = match serde_json::from_str(value_json) {
Ok(v) => v,
Err(e) => {
return into_c_ptr(format!(
"{{\"ok\":false,\"error\":\"invalid JSON value: {e}\"}}"
))
}
};
let result = (*ptr).publish_json(channel, value);
match result {
Ok(()) => into_c_ptr("{\"ok\":true}".to_string()),
Err(e) => into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
}
}
#[no_mangle]
/// # Safety
/// `ptr` must be a valid engine pointer from `rc_engine_new`.
/// `any_of_json` must be a valid null-terminated UTF-8 string pointer.
pub unsafe extern "C" fn rc_wait_any_of_json(
ptr: *mut Engine,
any_of_json: *const c_char,
) -> *mut c_char {
if ptr.is_null() {
return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string());
}
let any_of_json = match cstr_to_str(any_of_json) {
Ok(v) => v,
Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
};
let any_of: AnyOfCondition = match serde_json::from_str(any_of_json) {
Ok(v) => v,
Err(e) => {
return into_c_ptr(format!(
"{{\"ok\":false,\"error\":\"invalid any_of JSON: {e}\"}}"
))
}
};
let result = run_loop_block_on((*ptr).wait_for_any_of_async(&any_of));
match result {
Ok(event) => match serde_json::to_string(&event) {
Ok(s) => into_c_ptr(format!("{{\"ok\":true,\"event\":{s}}}")),
Err(e) => into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
},
Err(e) => into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
}
}
#[no_mangle]
/// # Safety
/// `ptr` must be a valid engine pointer from `rc_engine_new`.
/// `entry_point`, `finish_point`, and `initial_state_json` must be valid null-terminated UTF-8 pointers.
/// `callback` must be a valid function pointer that returns a malloc-allocated C string.
pub unsafe extern "C" fn rc_run_graph_json(
ptr: *mut Engine,
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 {
if ptr.is_null() {
return into_c_ptr("{\"ok\":false,\"error\":\"null engine pointer\"}".to_string());
}
let Some(callback) = callback else {
return into_c_ptr("{\"ok\":false,\"error\":\"null callback pointer\"}".to_string());
};
let entry_point = match cstr_to_str(entry_point) {
Ok(v) => v.to_string(),
Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
};
let finish_point = match cstr_to_str(finish_point) {
Ok(v) => v.to_string(),
Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
};
let initial_state_json = match cstr_to_str(initial_state_json) {
Ok(v) => v,
Err(e) => return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
};
let initial_state: Value = match serde_json::from_str(initial_state_json) {
Ok(v) => v,
Err(e) => {
return into_c_ptr(format!(
"{{\"ok\":false,\"error\":\"invalid initial_state JSON: {e}\"}}"
))
}
};
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);
let run_engine = (*ptr).clone();
let submit = run_loop_spawn(async move {
let out = run_graph_scheduler_json(
entry_point,
finish_point,
initial_state,
initial_input,
run_engine,
user_data,
callback,
)
.await;
let _ = tx.send(out);
});
if let Err(e) = submit {
return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}"));
}
let run_result = match rx.recv() {
Ok(v) => v,
Err(e) => {
return into_c_ptr(format!(
"{{\"ok\":false,\"error\":\"run-loop recv failed: {e}\"}}"
))
}
};
match run_result {
Ok(state) => match serde_json::to_string(&state) {
Ok(s) => into_c_ptr(format!("{{\"ok\":true,\"state\":{s}}}")),
Err(e) => into_c_ptr(format!(
"{{\"ok\":false,\"error\":\"serialize state failed: {e}\"}}"
)),
},
Err(e) => into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}")),
}
}
+423
View File
@@ -0,0 +1,423 @@
#[cfg(feature = "python-bindings")]
use crate::engine::{
node_pool_execute, run_loop_block_on, run_loop_spawn, AnyOfCondition, Engine, NodeExecResult,
NodeOutcome, SendPayload, WaitCondition, WaitEvent, WaitRequest,
};
#[cfg(feature = "python-bindings")]
use pyo3::exceptions::PyValueError;
#[cfg(feature = "python-bindings")]
use pyo3::prelude::*;
#[cfg(feature = "python-bindings")]
use pyo3::types::PyAny;
#[cfg(feature = "python-bindings")]
use pyo3::types::{PyDict, PyList, PyTuple};
#[cfg(feature = "python-bindings")]
use serde_json::Value;
#[cfg(feature = "python-bindings")]
use std::sync::mpsc;
#[cfg(feature = "python-bindings")]
use std::sync::Arc;
#[cfg(feature = "python-bindings")]
use tokio::sync::mpsc as tokio_mpsc;
#[cfg(feature = "python-bindings")]
#[pyclass]
struct PyRustEngine {
inner: Engine,
}
#[cfg(feature = "python-bindings")]
#[pymethods]
impl PyRustEngine {
#[new]
fn new() -> Self {
Self {
inner: Engine::new(),
}
}
fn add_async_channel(&self, name: &str) {
self.inner.add_async_channel(name);
}
fn publish_json(&self, channel: &str, value_json: &str) -> PyResult<()> {
let value: Value = serde_json::from_str(value_json)
.map_err(|e| PyValueError::new_err(format!("Invalid JSON value: {e}")))?;
self.inner
.publish_json(channel, value)
.map_err(PyValueError::new_err)
}
fn publish_obj(&self, py: Python<'_>, channel: &str, value: Py<PyAny>) -> PyResult<()> {
let value_json = py_obj_to_json_string(py, &value.bind(py))?;
let parsed: Value = serde_json::from_str(&value_json)
.map_err(|e| PyValueError::new_err(format!("Invalid Python JSON value: {e}")))?;
self.inner
.publish_json(channel, parsed)
.map_err(PyValueError::new_err)
}
fn wait_any_of_json(&self, any_of_json: &str) -> PyResult<String> {
let any_of: AnyOfCondition = serde_json::from_str(any_of_json)
.map_err(|e| PyValueError::new_err(format!("Invalid any_of JSON: {e}")))?;
let event = run_loop_block_on(self.inner.wait_for_any_of_async(&any_of))
.map_err(PyValueError::new_err)?;
serde_json::to_string(&event)
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))
}
fn wait_channel(&self, py: Python<'_>, channel: &str, n: usize) -> PyResult<Py<PyAny>> {
let cond = WaitCondition::Channel {
channel: channel.to_string(),
n,
};
let event =
run_loop_block_on(self.inner.wait_for_async(&cond)).map_err(PyValueError::new_err)?;
let event_json = serde_json::to_string(&event)
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))?;
json_string_to_py_obj(py, &event_json)
}
fn wait_timer(&self, py: Python<'_>, seconds: f64) -> PyResult<Py<PyAny>> {
let cond = WaitCondition::Timer { seconds };
let event =
run_loop_block_on(self.inner.wait_for_async(&cond)).map_err(PyValueError::new_err)?;
let event_json = serde_json::to_string(&event)
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))?;
json_string_to_py_obj(py, &event_json)
}
fn wait_any_of_obj(&self, py: Python<'_>, any_of_payload: Py<PyAny>) -> PyResult<Py<PyAny>> {
let payload_json = py_obj_to_json_string(py, &any_of_payload.bind(py))?;
let any_of: AnyOfCondition = serde_json::from_str(&payload_json)
.map_err(|e| PyValueError::new_err(format!("Invalid any_of payload: {e}")))?;
let event = run_loop_block_on(self.inner.wait_for_any_of_async(&any_of))
.map_err(PyValueError::new_err)?;
let event_json = serde_json::to_string(&event)
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))?;
json_string_to_py_obj(py, &event_json)
}
fn wait_condition_json(&self, cond_json: &str) -> PyResult<String> {
let cond: WaitCondition = serde_json::from_str(cond_json)
.map_err(|e| PyValueError::new_err(format!("Invalid condition JSON: {e}")))?;
let event =
run_loop_block_on(self.inner.wait_for_async(&cond)).map_err(PyValueError::new_err)?;
serde_json::to_string(&event)
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))
}
fn run_graph_py(
&self,
py: Python<'_>,
entry_point: &str,
finish_point: &str,
initial_state: Py<PyAny>,
callback: Py<PyAny>,
) -> PyResult<Py<PyAny>> {
let state = Arc::new(initial_state);
let callback = Arc::new(callback);
let entry_point = entry_point.to_string();
let finish_point = finish_point.to_string();
let (done_tx, done_rx) = mpsc::channel::<Result<(), String>>();
let state_for_run = Arc::clone(&state);
let engine = self.inner.clone();
run_loop_spawn(async move {
let run_result =
run_graph_scheduler(entry_point, finish_point, callback, state_for_run, engine)
.await;
let _ = done_tx.send(run_result);
})
.map_err(PyValueError::new_err)?;
let run_result = py
.allow_threads(move || done_rx.recv())
.map_err(|e| PyValueError::new_err(format!("run-loop recv failed: {e}")))?;
run_result.map_err(PyValueError::new_err)?;
Ok((*state).clone_ref(py))
}
}
#[cfg(feature = "python-bindings")]
enum SchedulerEventPy {
Node(Result<NodeExecutionPy, String>),
Resume {
node: String,
arg: Py<PyAny>,
event: WaitEvent,
},
WaitError(String),
}
struct NodeExecutionPy {
node: String,
arg: Py<PyAny>,
outcome: NodeOutcome<Py<PyAny>, Py<PyAny>>,
}
#[cfg(feature = "python-bindings")]
async fn run_graph_scheduler(
entry_point: String,
finish_point: String,
callback: Arc<Py<PyAny>>,
state: Arc<Py<PyAny>>,
engine: Engine,
) -> Result<(), String> {
let (tx, mut rx) = tokio_mpsc::unbounded_channel::<SchedulerEventPy>();
let initial_arg = Python::with_gil(|py| (*state).clone_ref(py));
let callback_for_spawn = Arc::clone(&callback);
let state_for_spawn = Arc::clone(&state);
let tx_for_spawn = tx.clone();
let state_for_merge = Arc::clone(&state);
let mut active: usize = 1;
let mut waiting: usize = 0;
spawn_node_task(
entry_point,
initial_arg,
tx_for_spawn.clone(),
Arc::clone(&callback_for_spawn),
Arc::clone(&state_for_spawn),
)?;
while active > 0 || waiting > 0 {
let event = rx
.recv()
.await
.ok_or_else(|| "scheduler event channel closed".to_string())?;
match event {
SchedulerEventPy::Node(result) => {
active = active.saturating_sub(1);
let exec = result?;
match exec.outcome {
NodeOutcome::Completed(node_result) => {
Python::with_gil(|py| -> Result<(), String> {
if let Some(update) = node_result.update {
apply_update_to_state(py, state_for_merge.as_ref(), &update)
.map_err(|e| {
format!("state merge failed for `{}`: {e}", exec.node)
})?;
}
Ok(())
})?;
if exec.node == finish_point {
break;
}
for send in node_result.sends {
active += 1;
spawn_node_task(
send.node,
send.arg,
tx_for_spawn.clone(),
Arc::clone(&callback_for_spawn),
Arc::clone(&state_for_spawn),
)?;
}
}
NodeOutcome::Suspended { wait } => {
waiting += 1;
let tx_wait = tx_for_spawn.clone();
let node = exec.node;
let arg = exec.arg;
let engine_for_wait = engine.clone();
tokio::spawn(async move {
let outcome = engine_for_wait.wait_request_async(&wait).await;
match outcome {
Ok(event) => {
let _ = tx_wait.send(SchedulerEventPy::Resume { node, arg, event });
}
Err(e) => {
let _ = tx_wait.send(SchedulerEventPy::WaitError(e));
}
}
});
}
}
}
SchedulerEventPy::Resume { node, arg, event } => {
waiting = waiting.saturating_sub(1);
active += 1;
let resume_arg = wrap_resume_arg(&arg, &event)?;
spawn_node_task(
node,
resume_arg,
tx_for_spawn.clone(),
Arc::clone(&callback_for_spawn),
Arc::clone(&state_for_spawn),
)?;
}
SchedulerEventPy::WaitError(e) => return Err(e),
}
}
Ok(())
}
#[cfg(feature = "python-bindings")]
fn spawn_node_task(
node: String,
arg: Py<PyAny>,
tx: tokio_mpsc::UnboundedSender<SchedulerEventPy>,
callback: Arc<Py<PyAny>>,
state_for_task: Arc<Py<PyAny>>,
) -> Result<(), String> {
node_pool_execute(move || {
let node_for_result = node.clone();
let arg_for_result = Python::with_gil(|py| arg.clone_ref(py));
let outcome = Python::with_gil(
|py| -> Result<NodeExecutionPy, String> {
let callback_bound = callback.as_ref().bind(py);
let payload_obj = callback_bound
.call1((node.as_str(), arg, (*state_for_task).clone_ref(py)))
.map_err(|e| format!("callback failed for node `{node}`: {e}"))?;
let payload = parse_node_outcome(py, &payload_obj)
.map_err(|e| format!("invalid callback payload for `{node}`: {e}"))?;
Ok(NodeExecutionPy {
node: node_for_result,
arg: arg_for_result,
outcome: payload,
})
},
);
let _ = tx.send(SchedulerEventPy::Node(outcome));
})
}
#[cfg(feature = "python-bindings")]
fn parse_node_outcome(
py: Python<'_>,
payload_obj: &Bound<'_, PyAny>,
) -> Result<NodeOutcome<Py<PyAny>, Py<PyAny>>, String> {
let payload_dict = payload_obj
.downcast::<PyDict>()
.map_err(|_| "payload must be a dict".to_string())?;
let suspended_item = payload_dict
.get_item("suspend")
.map_err(|e| format!("failed to read suspend: {e}"))?;
if let Some(wait_obj) = suspended_item {
let wait_json = py_obj_to_json_string(py, &wait_obj)
.map_err(|e| format!("failed to encode suspend payload: {e}"))?;
let wait: WaitRequest =
serde_json::from_str(&wait_json).map_err(|e| format!("invalid suspend payload: {e}"))?;
return Ok(NodeOutcome::Suspended { wait });
}
let update_item = payload_dict
.get_item("update")
.map_err(|e| format!("failed to read update: {e}"))?;
let update = match update_item {
Some(v) if !v.is_none() => Some(v.unbind()),
_ => None,
};
let sends_obj = payload_dict
.get_item("sends")
.map_err(|e| format!("failed to read sends: {e}"))?
.ok_or_else(|| "missing sends".to_string())?;
let sends_list = sends_obj
.downcast::<PyList>()
.map_err(|_| "sends must be a list".to_string())?;
let mut sends = Vec::with_capacity(sends_list.len());
for item in sends_list.iter() {
let send_dict = item
.downcast::<PyDict>()
.map_err(|_| "send item must be a dict".to_string())?;
let node_obj = send_dict
.get_item("node")
.map_err(|e| format!("failed to read send.node: {e}"))?
.ok_or_else(|| "send.node is required".to_string())?;
let node = node_obj
.extract::<String>()
.map_err(|e| format!("send.node must be string: {e}"))?;
let arg: Py<PyAny> = match send_dict.get_item("arg") {
Ok(Some(v)) => v.unbind(),
Ok(None) => Python::with_gil(|py| py.None()),
Err(e) => return Err(format!("failed to read send.arg: {e}")),
};
sends.push(SendPayload { node, arg });
}
Ok(NodeOutcome::Completed(NodeExecResult { update, sends }))
}
#[cfg(feature = "python-bindings")]
fn wrap_resume_arg(arg: &Py<PyAny>, event: &WaitEvent) -> Result<Py<PyAny>, String> {
Python::with_gil(|py| -> Result<Py<PyAny>, String> {
let wrapper = PyDict::new(py);
wrapper
.set_item("__lg_resume_arg__", arg.clone_ref(py))
.map_err(|e| format!("failed to set resume arg: {e}"))?;
let event_json =
serde_json::to_string(event).map_err(|e| format!("failed to encode wait event: {e}"))?;
let event_obj =
json_string_to_py_obj(py, &event_json).map_err(|e| format!("failed to parse event: {e}"))?;
wrapper
.set_item("__lg_resume_event__", event_obj.bind(py))
.map_err(|e| format!("failed to set resume event: {e}"))?;
Ok(wrapper.unbind().into_any())
})
}
#[cfg(feature = "python-bindings")]
fn apply_update_to_state(py: Python<'_>, state: &Py<PyAny>, update: &Py<PyAny>) -> PyResult<()> {
let state_obj = state.bind(py);
let update_obj = update.bind(py);
if update_obj.is_none() {
return Ok(());
}
if state_obj.is_instance_of::<PyDict>() && update_obj.is_instance_of::<PyDict>() {
let state_dict = state_obj.downcast::<PyDict>()?;
let update_dict = update_obj.downcast::<PyDict>()?;
state_dict.call_method1("update", (update_dict,))?;
return Ok(());
}
if let Ok(tuple_like) = update_obj.downcast::<PyList>() {
apply_pair_updates(state_obj, tuple_like)?;
return Ok(());
}
if let Ok(tuple_like) = update_obj.downcast::<PyTuple>() {
let list = PyList::new(py, tuple_like)?;
apply_pair_updates(state_obj, &list)?;
}
Ok(())
}
#[cfg(feature = "python-bindings")]
fn apply_pair_updates(state_obj: &Bound<'_, PyAny>, entries: &Bound<'_, PyList>) -> PyResult<()> {
if !state_obj.is_instance_of::<PyDict>() {
return Ok(());
}
let state_dict = state_obj.downcast::<PyDict>()?;
for entry in entries.iter() {
if let Ok(pair) = entry.downcast::<PyTuple>() {
if pair.len() == 2 {
let key_obj = pair.get_item(0)?;
if let Ok(key) = key_obj.extract::<String>() {
let value_obj = pair.get_item(1)?;
state_dict.set_item(key, value_obj)?;
}
}
}
}
Ok(())
}
#[cfg(feature = "python-bindings")]
fn py_obj_to_json_string(py: Python<'_>, obj: &Bound<'_, PyAny>) -> PyResult<String> {
let json_mod = py.import("json")?;
let dumped = json_mod.call_method1("dumps", (obj,))?;
dumped.extract::<String>()
}
#[cfg(feature = "python-bindings")]
fn json_string_to_py_obj(py: Python<'_>, value: &str) -> PyResult<Py<PyAny>> {
let json_mod = py.import("json")?;
let loaded = json_mod.call_method1("loads", (value,))?;
Ok(loaded.unbind())
}
#[cfg(feature = "python-bindings")]
#[pymodule]
fn langgraph_rust_core(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<PyRustEngine>()?;
Ok(())
}