mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-22 23:52:23 +02:00
simplify
This commit is contained in:
@@ -100,3 +100,10 @@ dmypy.json
|
||||
.turbo
|
||||
.editorconfig
|
||||
.scratch
|
||||
|
||||
# macOS debug symbol bundles generated during local Rust builds
|
||||
saf-python-sdk/python/saf_python_sdk/*.dSYM/
|
||||
|
||||
# Local PyO3 extension artifacts for saf-python-sdk
|
||||
saf-python-sdk/python/saf_python_sdk/langgraph_rust_core*.so
|
||||
saf-python-sdk/python/saf_python_sdk/langgraph_rust_core*.pyd
|
||||
|
||||
@@ -116,8 +116,9 @@ func (c *Context) SendCustomStreamEvent(value any) error {
|
||||
}
|
||||
|
||||
type Handler[StateT any] struct {
|
||||
engine *RustEngine
|
||||
done chan resultOrErr[StateT]
|
||||
engine *RustEngine
|
||||
done chan resultOrErr[StateT]
|
||||
streamReadyC chan struct{}
|
||||
}
|
||||
|
||||
type resultOrErr[StateT any] struct {
|
||||
@@ -135,6 +136,7 @@ func (h *Handler[StateT]) WaitForResult() (StateT, error) {
|
||||
}
|
||||
|
||||
func (h *Handler[StateT]) ReceiveStream() (any, error) {
|
||||
<-h.streamReadyC
|
||||
event, hasEvent, err := h.engine.ReceiveStream()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -146,6 +148,7 @@ func (h *Handler[StateT]) ReceiveStream() (any, error) {
|
||||
}
|
||||
|
||||
func (h *Handler[StateT]) CloseStream() error {
|
||||
<-h.streamReadyC
|
||||
return h.engine.CloseStream()
|
||||
}
|
||||
|
||||
@@ -165,15 +168,27 @@ func (g *CompiledGraph[StateT]) Start(initialInput any, initialState StateT, str
|
||||
}
|
||||
|
||||
handler := &Handler[StateT]{
|
||||
engine: engine,
|
||||
done: make(chan resultOrErr[StateT], 1),
|
||||
engine: engine,
|
||||
done: make(chan resultOrErr[StateT], 1),
|
||||
streamReadyC: make(chan struct{}),
|
||||
}
|
||||
go func() {
|
||||
defer engine.Close()
|
||||
streamModeForRun := resolvedStreamMode
|
||||
if resolvedStreamMode != "" {
|
||||
if err := engine.StartStream(resolvedStreamMode); err != nil {
|
||||
close(handler.streamReadyC)
|
||||
handler.done <- resultOrErr[StateT]{err: err}
|
||||
close(handler.done)
|
||||
return
|
||||
}
|
||||
streamModeForRun = ""
|
||||
}
|
||||
close(handler.streamReadyC)
|
||||
rawState, err := engine.RunGraph(
|
||||
g.entryPoint,
|
||||
g.finishPoint,
|
||||
resolvedStreamMode,
|
||||
streamModeForRun,
|
||||
initialState,
|
||||
initialInput,
|
||||
func(node string, nodeInput any, fallbackState map[string]any) (Command, error) {
|
||||
|
||||
@@ -491,3 +491,253 @@ impl Engine {
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct CallbackSendPayloadJson {
|
||||
node: String,
|
||||
#[serde(default)]
|
||||
arg: Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct CallbackNodeExecResultJsonWire {
|
||||
update: Option<Value>,
|
||||
#[serde(default)]
|
||||
sends: Vec<CallbackSendPayloadJson>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct CallbackEnvelopeIn {
|
||||
ok: bool,
|
||||
#[serde(default)]
|
||||
payload: Option<CallbackNodeExecResultJsonWire>,
|
||||
#[serde(default)]
|
||||
suspend: Option<WaitRequest>,
|
||||
#[serde(default)]
|
||||
error: Option<String>,
|
||||
}
|
||||
|
||||
enum SchedulerEvent<U, A> {
|
||||
Node(Result<NodeExecution<U, A>, String>),
|
||||
Resume {
|
||||
node: String,
|
||||
arg: A,
|
||||
event: WaitEvent,
|
||||
},
|
||||
WaitError(String),
|
||||
}
|
||||
|
||||
struct NodeExecution<U, A> {
|
||||
node: String,
|
||||
arg: A,
|
||||
outcome: NodeOutcome<U, A>,
|
||||
}
|
||||
|
||||
fn spawn_node_task<State, U, A, F>(
|
||||
node: String,
|
||||
arg: A,
|
||||
state_snapshot: State,
|
||||
tx: tokio_mpsc::UnboundedSender<SchedulerEvent<U, A>>,
|
||||
callback: Arc<F>,
|
||||
) -> Result<(), String>
|
||||
where
|
||||
State: Send + 'static,
|
||||
U: Send + 'static,
|
||||
A: Clone + Send + 'static,
|
||||
F: Fn(String, A, State) -> Result<NodeOutcome<U, A>, String> + Send + Sync + 'static,
|
||||
{
|
||||
node_pool_execute(move || {
|
||||
let node_for_result = node.clone();
|
||||
let arg_for_result = arg.clone();
|
||||
let result = callback(node, arg, state_snapshot).map(|outcome| NodeExecution {
|
||||
node: node_for_result,
|
||||
arg: arg_for_result,
|
||||
outcome,
|
||||
});
|
||||
let _ = tx.send(SchedulerEvent::Node(result));
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn run_graph_with_callback<State, U, A, FCallback, FMerge, FWrap>(
|
||||
entry_point: String,
|
||||
finish_point: String,
|
||||
initial_state: State,
|
||||
initial_input: A,
|
||||
engine: Engine,
|
||||
callback: FCallback,
|
||||
merge_update: FMerge,
|
||||
wrap_resume_arg: FWrap,
|
||||
) -> Result<State, String>
|
||||
where
|
||||
State: Clone + Send + 'static,
|
||||
U: Send + 'static,
|
||||
A: Clone + Send + 'static,
|
||||
FCallback: Fn(String, A, State) -> Result<NodeOutcome<U, A>, String> + Send + Sync + 'static,
|
||||
FMerge: Fn(&mut State, Option<U>) -> Result<(), String> + Send + Sync + 'static,
|
||||
FWrap: Fn(A, WaitEvent) -> Result<A, String> + Send + Sync + 'static,
|
||||
{
|
||||
let callback = Arc::new(callback);
|
||||
let merge_update = Arc::new(merge_update);
|
||||
let wrap_resume_arg = Arc::new(wrap_resume_arg);
|
||||
|
||||
let (tx, mut rx) = tokio_mpsc::unbounded_channel::<SchedulerEvent<U, A>>();
|
||||
let state = Arc::new(StdMutex::new(initial_state));
|
||||
let tx_for_spawn = tx.clone();
|
||||
let state_for_spawn = Arc::clone(&state);
|
||||
let mut active: usize = 1;
|
||||
let mut waiting: usize = 0;
|
||||
|
||||
spawn_node_task(
|
||||
entry_point,
|
||||
initial_input,
|
||||
state_for_spawn
|
||||
.lock()
|
||||
.expect("state mutex poisoned")
|
||||
.clone(),
|
||||
tx_for_spawn.clone(),
|
||||
Arc::clone(&callback),
|
||||
)?;
|
||||
|
||||
while active > 0 || waiting > 0 {
|
||||
let evt = rx
|
||||
.recv()
|
||||
.await
|
||||
.ok_or_else(|| "scheduler event channel closed".to_string())?;
|
||||
match evt {
|
||||
SchedulerEvent::Node(result) => {
|
||||
active = active.saturating_sub(1);
|
||||
let exec = result?;
|
||||
match exec.outcome {
|
||||
NodeOutcome::Completed(node_result) => {
|
||||
let mut guard = state.lock().expect("state mutex poisoned");
|
||||
merge_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_node_task(
|
||||
send.node,
|
||||
send.arg,
|
||||
snapshot,
|
||||
tx_for_spawn.clone(),
|
||||
Arc::clone(&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(SchedulerEvent::Resume { node, arg, event });
|
||||
}
|
||||
Err(e) => {
|
||||
let _ = tx_wait.send(SchedulerEvent::WaitError(e));
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
SchedulerEvent::Resume { node, arg, event } => {
|
||||
waiting = waiting.saturating_sub(1);
|
||||
active += 1;
|
||||
let snapshot = state_for_spawn
|
||||
.lock()
|
||||
.expect("state mutex poisoned")
|
||||
.clone();
|
||||
let resume_arg = wrap_resume_arg(arg, event)?;
|
||||
spawn_node_task(
|
||||
node,
|
||||
resume_arg,
|
||||
snapshot,
|
||||
tx_for_spawn.clone(),
|
||||
Arc::clone(&callback),
|
||||
)?;
|
||||
}
|
||||
SchedulerEvent::WaitError(e) => return Err(e),
|
||||
}
|
||||
}
|
||||
|
||||
let final_state = state.lock().expect("state mutex poisoned").clone();
|
||||
Ok(final_state)
|
||||
}
|
||||
|
||||
pub async fn run_graph_json_with_callback<F>(
|
||||
entry_point: String,
|
||||
finish_point: String,
|
||||
initial_state: Value,
|
||||
initial_input: Value,
|
||||
engine: Engine,
|
||||
callback: F,
|
||||
) -> Result<Value, String>
|
||||
where
|
||||
F: Fn(String, Value, Value) -> Result<NodeOutcome<Value, Value>, String>
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
{
|
||||
run_graph_with_callback(
|
||||
entry_point,
|
||||
finish_point,
|
||||
initial_state,
|
||||
initial_input,
|
||||
engine,
|
||||
callback,
|
||||
|state: &mut Value, update: Option<Value>| {
|
||||
merge_json_update(state, update);
|
||||
Ok(())
|
||||
},
|
||||
|arg: Value, event: WaitEvent| {
|
||||
Ok(serde_json::json!({
|
||||
"__lg_resume_arg__": arg,
|
||||
"__lg_resume_event__": event,
|
||||
}))
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub fn parse_callback_envelope_json(
|
||||
raw: &str,
|
||||
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,
|
||||
}))
|
||||
}
|
||||
|
||||
+42
-250
@@ -1,39 +1,11 @@
|
||||
use crate::engine::{
|
||||
merge_json_update, node_pool_execute, run_loop_block_on, run_loop_spawn, AnyOfCondition,
|
||||
Engine, NodeExecResult, NodeOutcome, SendPayload, WaitEvent, WaitRequest,
|
||||
parse_callback_envelope_json, run_graph_json_with_callback, run_loop_block_on, run_loop_spawn,
|
||||
AnyOfCondition, Engine, NodeOutcome,
|
||||
};
|
||||
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,
|
||||
@@ -42,9 +14,6 @@ type CNodeCallback = unsafe extern "C" fn(
|
||||
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());
|
||||
@@ -63,217 +32,6 @@ fn into_c_ptr(s: String) -> *mut c_char {
|
||||
}
|
||||
}
|
||||
|
||||
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()))
|
||||
@@ -528,22 +286,56 @@ pub unsafe extern "C" fn rc_run_graph_json(
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(e) = (*ptr).start_stream(stream_mode.as_deref()) {
|
||||
return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}"));
|
||||
if let Some(mode) = stream_mode.as_deref() {
|
||||
if let Err(e) = (*ptr).start_stream(Some(mode)) {
|
||||
return into_c_ptr(format!("{{\"ok\":false,\"error\":\"{e}\"}}"));
|
||||
}
|
||||
}
|
||||
|
||||
let (tx, rx) = mpsc::channel::<Result<Value, String>>();
|
||||
let user_data = CUserData(user_data);
|
||||
let user_data_bits = user_data;
|
||||
let run_engine = (*ptr).clone();
|
||||
let submit = run_loop_spawn(async move {
|
||||
let out = run_graph_scheduler_json(
|
||||
let callback_wrapper = move |node: String,
|
||||
arg: Value,
|
||||
state_snapshot: Value|
|
||||
-> Result<NodeOutcome<Value, Value>, 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());
|
||||
}
|
||||
parse_callback_envelope_json(&out_raw, &node)
|
||||
};
|
||||
let out = run_graph_json_with_callback(
|
||||
entry_point,
|
||||
finish_point,
|
||||
initial_state,
|
||||
initial_input,
|
||||
run_engine.clone(),
|
||||
user_data,
|
||||
callback,
|
||||
callback_wrapper,
|
||||
)
|
||||
.await;
|
||||
run_engine.close_stream();
|
||||
|
||||
+66
-186
@@ -1,6 +1,6 @@
|
||||
#[cfg(feature = "python-bindings")]
|
||||
use crate::engine::{
|
||||
node_pool_execute, run_loop_block_on, run_loop_spawn, AnyOfCondition, Engine, NodeExecResult,
|
||||
run_graph_with_callback, run_loop_block_on, AnyOfCondition, Engine, NodeExecResult,
|
||||
NodeOutcome, SendPayload, WaitCondition, WaitEvent, WaitRequest,
|
||||
};
|
||||
#[cfg(feature = "python-bindings")]
|
||||
@@ -14,11 +14,7 @@ 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]
|
||||
@@ -50,11 +46,7 @@ impl PyRustEngine {
|
||||
|
||||
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)
|
||||
self.publish_json(channel, &value_json)
|
||||
}
|
||||
|
||||
fn wait_any_of_json(&self, any_of_json: &str) -> PyResult<String> {
|
||||
@@ -89,12 +81,7 @@ impl PyRustEngine {
|
||||
|
||||
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}")))?;
|
||||
let event_json = self.wait_any_of_json(&payload_json)?;
|
||||
json_string_to_py_obj(py, &event_json)
|
||||
}
|
||||
|
||||
@@ -107,6 +94,7 @@ impl PyRustEngine {
|
||||
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))
|
||||
}
|
||||
|
||||
#[pyo3(signature = (stream_mode=None))]
|
||||
fn start_stream(&self, stream_mode: Option<&str>) -> PyResult<()> {
|
||||
self.inner
|
||||
.start_stream(stream_mode)
|
||||
@@ -114,7 +102,8 @@ impl PyRustEngine {
|
||||
}
|
||||
|
||||
fn receive_stream_obj(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
|
||||
match run_loop_block_on(self.inner.receive_stream_async()) {
|
||||
let event = py.allow_threads(|| run_loop_block_on(self.inner.receive_stream_async()));
|
||||
match event {
|
||||
Some(value) => {
|
||||
let event_json = serde_json::to_string(&value)
|
||||
.map_err(|e| PyValueError::new_err(format!("Serialize event failed: {e}")))?;
|
||||
@@ -136,6 +125,7 @@ impl PyRustEngine {
|
||||
self.inner.close_stream();
|
||||
}
|
||||
|
||||
#[pyo3(signature = (entry_point, finish_point, initial_state, callback, stream_mode=None))]
|
||||
fn run_graph_py(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
@@ -145,184 +135,74 @@ impl PyRustEngine {
|
||||
callback: Py<PyAny>,
|
||||
stream_mode: Option<&str>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let state = Arc::new(initial_state);
|
||||
if let Some(mode) = stream_mode {
|
||||
self.inner
|
||||
.start_stream(Some(mode))
|
||||
.map_err(PyValueError::new_err)?;
|
||||
}
|
||||
|
||||
let engine = self.inner.clone();
|
||||
let callback = Arc::new(callback);
|
||||
let entry_point = entry_point.to_string();
|
||||
let finish_point = finish_point.to_string();
|
||||
self.inner
|
||||
.start_stream(stream_mode)
|
||||
.map_err(PyValueError::new_err)?;
|
||||
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.clone(),
|
||||
)
|
||||
.await;
|
||||
engine.close_stream();
|
||||
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))
|
||||
}
|
||||
}
|
||||
let initial_state = Arc::new(initial_state);
|
||||
let initial_input = Arc::clone(&initial_state);
|
||||
|
||||
#[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) => {
|
||||
let run_result =
|
||||
py.allow_threads(move || {
|
||||
run_loop_block_on(run_graph_with_callback(
|
||||
entry_point,
|
||||
finish_point,
|
||||
initial_state,
|
||||
initial_input,
|
||||
engine.clone(),
|
||||
{
|
||||
let callback = Arc::clone(&callback);
|
||||
move |node: String,
|
||||
arg: Arc<Py<PyAny>>,
|
||||
state_snapshot: Arc<Py<PyAny>>|
|
||||
-> Result<NodeOutcome<Arc<Py<PyAny>>, Arc<Py<PyAny>>>, String> {
|
||||
Python::with_gil(|py| -> Result<NodeOutcome<Arc<Py<PyAny>>, Arc<Py<PyAny>>>, String> {
|
||||
let callback_bound = callback.as_ref().bind(py);
|
||||
let payload_obj = callback_bound
|
||||
.call1((
|
||||
node.as_str(),
|
||||
arg.as_ref().clone_ref(py),
|
||||
state_snapshot.as_ref().clone_ref(py),
|
||||
))
|
||||
.map_err(|e| format!("callback failed for node `{node}`: {e}"))?;
|
||||
parse_node_outcome_arc(py, &payload_obj).map_err(|e| {
|
||||
format!("invalid callback payload for `{node}`: {e}")
|
||||
})
|
||||
})
|
||||
}
|
||||
},
|
||||
|state: &mut Arc<Py<PyAny>>, update: Option<Arc<Py<PyAny>>>| -> Result<(), String> {
|
||||
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)
|
||||
})?;
|
||||
if let Some(update) = update {
|
||||
apply_update_to_state(py, state.as_ref(), update.as_ref())
|
||||
.map_err(|e| format!("state merge failed: {e}"))?;
|
||||
}
|
||||
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),
|
||||
}
|
||||
})
|
||||
},
|
||||
|arg: Arc<Py<PyAny>>, event: WaitEvent| -> Result<Arc<Py<PyAny>>, String> {
|
||||
wrap_resume_arg(arg.as_ref(), &event).map(Arc::new)
|
||||
},
|
||||
))
|
||||
});
|
||||
|
||||
self.inner.close_stream();
|
||||
let out = run_result.map_err(PyValueError::new_err)?;
|
||||
Ok(out.as_ref().clone_ref(py))
|
||||
}
|
||||
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(
|
||||
fn parse_node_outcome_arc(
|
||||
py: Python<'_>,
|
||||
payload_obj: &Bound<'_, PyAny>,
|
||||
) -> Result<NodeOutcome<Py<PyAny>, Py<PyAny>>, String> {
|
||||
) -> Result<NodeOutcome<Arc<Py<PyAny>>, Arc<Py<PyAny>>>, String> {
|
||||
let payload_dict = payload_obj
|
||||
.downcast::<PyDict>()
|
||||
.map_err(|_| "payload must be a dict".to_string())?;
|
||||
@@ -341,7 +221,7 @@ fn parse_node_outcome(
|
||||
.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()),
|
||||
Some(v) if !v.is_none() => Some(Arc::new(v.unbind())),
|
||||
_ => None,
|
||||
};
|
||||
|
||||
@@ -365,9 +245,9 @@ fn parse_node_outcome(
|
||||
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()),
|
||||
let arg: Arc<Py<PyAny>> = match send_dict.get_item("arg") {
|
||||
Ok(Some(v)) => Arc::new(v.unbind()),
|
||||
Ok(None) => Arc::new(Python::with_gil(|py| py.None())),
|
||||
Err(e) => return Err(format!("failed to read send.arg: {e}")),
|
||||
};
|
||||
sends.push(SendPayload { node, arg });
|
||||
|
||||
@@ -211,11 +211,9 @@ class GraphRunHandler(Generic[StateT]):
|
||||
await self._run.publish(channel, value)
|
||||
|
||||
async def receive_stream(self) -> Any | None:
|
||||
loop = asyncio.get_running_loop()
|
||||
return await loop.run_in_executor(
|
||||
_advanced_graph_executor(),
|
||||
self._run.receive_stream_sync,
|
||||
)
|
||||
# Use a separate thread pool from graph execution to avoid deadlock
|
||||
# when LANGGRAPH_ADVANCED_GRAPH_PY_THREADS is configured to 1.
|
||||
return await asyncio.to_thread(self._run.receive_stream_sync)
|
||||
|
||||
def close_stream(self) -> None:
|
||||
self._run.close_stream_sync()
|
||||
@@ -244,6 +242,9 @@ class _GraphEngineRun:
|
||||
self._rust_engine = PyRustEngine()
|
||||
for name in async_channel_specs:
|
||||
self._rust_engine.add_async_channel(name)
|
||||
self._stream_ready = threading.Event()
|
||||
if self._stream_mode is None:
|
||||
self._stream_ready.set()
|
||||
self._tasks: set[asyncio.Task[list[Send]]] = set()
|
||||
self._finished = False
|
||||
self._state: Any = None
|
||||
@@ -252,16 +253,24 @@ class _GraphEngineRun:
|
||||
|
||||
async def run(self, initial_state: StateT) -> StateT:
|
||||
finish_point = self._finish_point or ""
|
||||
if self._stream_mode is not None:
|
||||
try:
|
||||
self._rust_engine.start_stream(self._stream_mode)
|
||||
finally:
|
||||
self._stream_ready.set()
|
||||
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._stream_mode,
|
||||
)
|
||||
try:
|
||||
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,
|
||||
None,
|
||||
)
|
||||
finally:
|
||||
self._stream_ready.set()
|
||||
self._state = result_obj
|
||||
return cast(StateT, self._state)
|
||||
|
||||
@@ -328,6 +337,7 @@ class _GraphEngineRun:
|
||||
self._rust_engine.send_custom_stream_event_obj(value)
|
||||
|
||||
def receive_stream_sync(self) -> Any | None:
|
||||
self._stream_ready.wait()
|
||||
return self._rust_engine.receive_stream_obj()
|
||||
|
||||
def close_stream_sync(self) -> None:
|
||||
|
||||
Binary file not shown.
Generated
+1
-1
@@ -4,5 +4,5 @@ requires-python = ">=3.10"
|
||||
|
||||
[[package]]
|
||||
name = "saf-python-sdk"
|
||||
version = "0.1.1"
|
||||
version = "0.1.2"
|
||||
source = { editable = "." }
|
||||
|
||||
Reference in New Issue
Block a user