From 063db9a78e208262a8ac03ba6f73c9a794df2664 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 29 Sep 2026 11:18:29 -0700 Subject: [PATCH 1/7] fix(plugin): preserve native async callback context Signed-off-by: Alex Fournier --- crates/core/src/api/runtime.rs | 2 +- crates/core/src/api/runtime/scope_stack.rs | 76 +++++++++-- crates/core/src/api/shared.rs | 3 +- crates/core/src/plugin/dynamic/native.rs | 4 +- .../tests/fixtures/native_plugin/src/lib.rs | 92 +++++++++---- .../tests/integration/native_plugin_tests.rs | 125 +++++++++++++++++- crates/core/tests/unit/native_plugin_tests.rs | 90 +++++++++++++ crates/plugin/src/async_sdk.rs | 96 +++++++++----- 8 files changed, 412 insertions(+), 76 deletions(-) diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index f0d085271..97471cf5e 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -24,7 +24,6 @@ pub use continuation_context::MiddlewareContinuationContext; #[cfg(test)] pub(crate) use continuation_context::MiddlewareContinuationLease; pub use global::global_context; -pub(crate) use scope_stack::capture_trace_context; pub use scope_stack::{ PropagationContext, ScopeStack, ScopeStackHandle, TASK_SCOPE_STACK, ThreadScopeStackBinding, capture_propagation_context, capture_propagation_context_with_root, @@ -34,6 +33,7 @@ pub use scope_stack::{ set_thread_scope_stack, sync_thread_scope_stack, task_scope_push, task_scope_remove, task_scope_top, with_active_event_uuid, with_scope_stack, }; +pub(crate) use scope_stack::{capture_trace_context, sync_thread_active_event}; pub use state::NemoRelayContextState; #[doc(hidden)] pub use subscriber_dispatcher::SubscriberDelivery; diff --git a/crates/core/src/api/runtime/scope_stack.rs b/crates/core/src/api/runtime/scope_stack.rs index 861ddf270..866cadf35 100644 --- a/crates/core/src/api/runtime/scope_stack.rs +++ b/crates/core/src/api/runtime/scope_stack.rs @@ -8,7 +8,7 @@ //! can use this module to inspect the active scope chain or propagate scope //! context into worker threads. -use std::cell::RefCell; +use std::cell::{Cell, RefCell}; use std::collections::{HashMap, HashSet}; use std::future::Future; use std::sync::{Arc, RwLock}; @@ -544,14 +544,24 @@ impl Default for ScopeStack { /// concurrent readers. pub type ScopeStackHandle = Arc>; +#[derive(Clone, Copy)] +struct ActiveEventBinding { + uuid: Uuid, + // Propagated stacks may share a root UUID, so identify the captured Arc allocation. + scope_stack_id: usize, + scope_stack_top: Uuid, +} + /// Captured thread-local scope stack binding. /// -/// This preserves both the visible scope stack handle and whether it was -/// explicitly installed on the current thread. +/// This preserves the visible scope stack handle, whether it was explicitly +/// installed on the current thread, and any managed event bound at capture +/// time. #[derive(Clone)] pub struct ThreadScopeStackBinding { stack: ScopeStackHandle, explicit: bool, + active_event: Option, } impl ThreadScopeStackBinding { @@ -659,9 +669,7 @@ pub fn capture_rootless_propagation_context() -> Result { pub fn capture_propagation_context_with_root( root_uuid: Option, ) -> Result { - let parent_uuid = ACTIVE_EVENT_UUID - .try_with(|uuid| *uuid) - .unwrap_or_else(|_| task_scope_top().uuid); + let parent_uuid = active_event_uuid().unwrap_or_else(|| task_scope_top().uuid); let (traceparent, tracestate) = current_scope_stack() .read() .map(|stack| stack.w3c_headers_for_parent(parent_uuid)) @@ -783,16 +791,42 @@ tokio::task_local! { /// Task-local scope stack handle used by async execution contexts. pub static TASK_SCOPE_STACK: ScopeStackHandle; /// Managed tool or LLM event currently executing in this task. - static ACTIVE_EVENT_UUID: Uuid; + static ACTIVE_EVENT: ActiveEventBinding; } /// Run a future with `uuid` as the causally active managed event. pub async fn with_active_event_uuid(uuid: Uuid, future: impl Future) -> T { - ACTIVE_EVENT_UUID.scope(uuid, future).await + let (scope_stack_id, scope_stack_top) = scope_stack_identity_and_top(); + let active_event = ActiveEventBinding { + uuid, + scope_stack_id, + scope_stack_top, + }; + ACTIVE_EVENT.scope(active_event, future).await } pub(crate) fn active_event_uuid() -> Option { - ACTIVE_EVENT_UUID.try_with(|uuid| *uuid).ok() + ACTIVE_EVENT + .try_with(|event| event.uuid) + .ok() + .or_else(thread_active_event_uuid) +} + +pub(crate) fn thread_active_event_uuid() -> Option { + let mut event = THREAD_ACTIVE_EVENT.with(Cell::get)?; + let stack = current_scope_stack(); + if event.scope_stack_id != Arc::as_ptr(&stack) as usize { + return None; + } + let guard = stack.read().unwrap_or_else(|error| error.into_inner()); + if event.scope_stack_top != guard.top().uuid { + if guard.find(&event.scope_stack_top).is_some() { + return None; + } + event.scope_stack_top = guard.top().uuid; + THREAD_ACTIVE_EVENT.with(|active| active.set(Some(event))); + } + Some(event.uuid) } thread_local! { @@ -803,6 +837,8 @@ thread_local! { static THREAD_SCOPE_STACK: RefCell = RefCell::new(create_scope_stack()); /// Whether the current thread explicitly owns a scope stack. static THREAD_SCOPE_STACK_EXPLICIT: std::cell::Cell = const { std::cell::Cell::new(false) }; + /// Managed event propagated into a foreign executor with the thread scope binding. + static THREAD_ACTIVE_EVENT: Cell> = const { Cell::new(None) }; } /// Return the scope stack visible to the current execution context. @@ -883,12 +919,16 @@ pub fn set_thread_scope_stack(handle: ScopeStackHandle) { /// that thread back to their scheduler. /// /// # Returns -/// A [`ThreadScopeStackBinding`] containing the current thread-local stack and -/// explicit-binding flag. +/// A [`ThreadScopeStackBinding`] containing the current thread-local stack, +/// explicit-binding flag, and active managed event. pub fn capture_thread_scope_stack() -> ThreadScopeStackBinding { let stack = THREAD_SCOPE_STACK.with(|stack| stack.borrow().clone()); let explicit = THREAD_SCOPE_STACK_EXPLICIT.with(|flag| flag.get()); - ThreadScopeStackBinding { stack, explicit } + ThreadScopeStackBinding { + stack, + explicit, + active_event: THREAD_ACTIVE_EVENT.with(Cell::get), + } } /// Restore a previously captured thread-local scope stack binding. @@ -901,6 +941,7 @@ pub fn capture_thread_scope_stack() -> ThreadScopeStackBinding { pub fn restore_thread_scope_stack(binding: ThreadScopeStackBinding) { THREAD_SCOPE_STACK.with(|stack| *stack.borrow_mut() = binding.stack); THREAD_SCOPE_STACK_EXPLICIT.with(|flag| flag.set(binding.explicit)); + THREAD_ACTIVE_EVENT.with(|event| event.set(binding.active_event)); } /// Synchronize the thread-local scope stack without marking it explicit. @@ -921,6 +962,17 @@ pub fn sync_thread_scope_stack(handle: ScopeStackHandle) { THREAD_SCOPE_STACK.with(|stack| *stack.borrow_mut() = handle); } +pub(crate) fn sync_thread_active_event() { + let active_event = ACTIVE_EVENT.try_with(|event| *event).ok(); + THREAD_ACTIVE_EVENT.with(|event| event.set(active_event)); +} + +fn scope_stack_identity_and_top() -> (usize, Uuid) { + let stack = current_scope_stack(); + let guard = stack.read().unwrap_or_else(|error| error.into_inner()); + (Arc::as_ptr(&stack) as usize, guard.top().uuid) +} + /// Report whether the current context has an explicitly active scope stack. /// /// This checks task-local state first and otherwise falls back to the diff --git a/crates/core/src/api/shared.rs b/crates/core/src/api/shared.rs index b23366b82..159d68ab7 100644 --- a/crates/core/src/api/shared.rs +++ b/crates/core/src/api/shared.rs @@ -9,7 +9,7 @@ use crate::api::event::{Event, EventSanitizeFields, ScopeCategory}; use crate::api::llm::LlmRequest; use crate::api::registry::{EventMetadataInjector, Guardrail, RuntimeRegistrationKind}; use crate::api::runtime::global_context; -use crate::api::runtime::scope_stack::trace_context_for_llm; +use crate::api::runtime::scope_stack::{thread_active_event_uuid, trace_context_for_llm}; use crate::api::runtime::{ EventSanitizeFn, EventSubscriberFn, NemoRelayContextState, ScopeStackHandle, }; @@ -35,6 +35,7 @@ pub(crate) fn resolve_parent_uuid(parent: Option<&ScopeHandle>) -> Option Some( parent .map(|handle| handle.uuid) + .or_else(thread_active_event_uuid) .unwrap_or_else(|| task_scope_top().uuid), ) } diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 0ace3703e..0b745605f 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -38,7 +38,7 @@ use crate::api::runtime::{ use crate::api::runtime::{ ScopeStackHandle, ThreadScopeStackBinding, capture_thread_scope_stack, create_scope_stack, current_scope_stack, restore_thread_scope_stack, scope_stack_active, set_thread_scope_stack, - sync_thread_scope_stack, with_scope_stack, + sync_thread_active_event, sync_thread_scope_stack, with_scope_stack, }; use crate::api::scope::{ EmitMarkEventParams, PopScopeParams, PushScopeParams, ScopeAttributes, ScopeHandle, ScopeType, @@ -2085,6 +2085,7 @@ async fn invoke_native_async_callback( // SDK can capture it before moving the future to its own executor. let previous_thread_stack = capture_thread_scope_stack(); sync_thread_scope_stack(current_scope_stack()); + sync_thread_active_event(); let state = catch_unwind(AssertUnwindSafe(|| unsafe { cb( user_data.ptr, @@ -3800,6 +3801,7 @@ fn wrap_native_incremental_llm_stream_execution_with_user_data( let stream_ref = Arc::into_raw(stream.clone()); let previous_thread_stack = capture_thread_scope_stack(); sync_thread_scope_stack(current_scope_stack()); + sync_thread_active_event(); let state = catch_unwind(AssertUnwindSafe(|| unsafe { cb( user_data.ptr, diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index fda6a7907..0cb90aa05 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -26,6 +26,25 @@ use tokio::io::{AsyncReadExt, AsyncWriteExt}; struct FixtureNativePlugin; +struct DropScope { + runtime: PluginRuntime, + name: &'static str, +} + +impl Drop for DropScope { + fn drop(&mut self) { + if let Ok(mut scope) = self.runtime.scope( + self.name, + ScopeType::Custom, + None, + None, + None, + ) { + let _ = scope.close(None, None); + } + } +} + static ASYNC_PENDING_ENTERED: AtomicBool = AtomicBool::new(false); #[unsafe(no_mangle)] @@ -325,35 +344,50 @@ impl NativePlugin for FixtureNativePlugin { )) }, )?; - ctx.register_llm_execution_intercept( - "fixture_llm_execution", - 0, - |_name, request, next| async move { - let response = next - .call(mark_llm_request( - request, - "native_plugin_llm_execution_request", - )) - .await?; - Ok(mark_json(response, "native_plugin_llm_execution")) - }, - )?; - ctx.register_llm_stream_execution_intercept( - "fixture_llm_stream_execution", - 0, - |_name, request, next| async move { - let stream = next - .call(mark_llm_request( - request, - "native_plugin_llm_stream_execution_request", - )) - .await?; - let stream: LlmJsonAsyncStream = Box::pin(stream.map(|chunk| { - chunk.map(|chunk| mark_json(chunk, "native_plugin_llm_stream_execution")) - })); - Ok(stream) - }, - )?; + ctx.register_llm_execution_intercept("fixture_llm_execution", 0, { + let runtime = runtime.clone(); + move |name, request, next| { + let drop_scope = (name == "native-fixture-cancelled-unary").then(|| DropScope { + runtime: runtime.clone(), + name: "fixture.native.unary.drop", + }); + async move { + if let Some(_drop_scope) = drop_scope { + ASYNC_PENDING_ENTERED.store(true, Ordering::Release); + return futures::future::pending::>().await; + } + let response = next + .call(mark_llm_request( + request, + "native_plugin_llm_execution_request", + )) + .await?; + Ok(mark_json(response, "native_plugin_llm_execution")) + } + } + })?; + ctx.register_llm_stream_execution_intercept("fixture_llm_stream_execution", 0, { + let runtime = runtime.clone(); + move |name, request, next| { + let drop_scope = (name == "native-fixture-cancelled-stream").then(|| DropScope { + runtime: runtime.clone(), + name: "fixture.native.stream.drop", + }); + async move { + let stream = next + .call(mark_llm_request( + request, + "native_plugin_llm_stream_execution_request", + )) + .await?; + let stream: LlmJsonAsyncStream = Box::pin(stream.map(move |chunk| { + let _ = &drop_scope; + chunk.map(|chunk| mark_json(chunk, "native_plugin_llm_stream_execution")) + })); + Ok(stream) + } + } + })?; Ok(()) } diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index 656d91c07..6457d846d 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -7,8 +7,9 @@ mod plugin_host_test_support; use std::path::{Path, PathBuf}; use std::process::Command; -use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; +use std::task::Poll; use nemo_relay::api::event::{Event, ScopeCategory}; use nemo_relay::api::llm::{ @@ -217,6 +218,15 @@ async fn sdk_cdylib_registers_tool_request_intercept() { manifest_ref: manifest_ref.to_string_lossy().into_owned(), }]) .expect("native plugin should load"); + let fixture_library = unsafe { libloading::Library::new(&fixture.library_path) } + .expect("native fixture should open for synchronization"); + let pending_entered = unsafe { + *fixture_library + .get:: bool>(b"nemo_relay_fixture_async_pending_entered\0") + .expect("native fixture should export its pending-entry signal") + }; + // Clear a signal left by any earlier fixture use in this process. + let _ = unsafe { pending_entered() }; let mut cleanup = NativePluginTestCleanup::new(); let mut plugin_config = PluginConfig::default(); @@ -585,6 +595,59 @@ async fn sdk_cdylib_registers_tool_request_intercept() { ); assert!(llm_end.annotated_response().is_none()); + events.lock().unwrap().clear(); + let cancelled_unary = tokio::spawn(llm_call_execute( + LlmCallExecuteParams::builder() + .name("native-fixture-cancelled-unary") + .request(LlmRequest { + headers: Map::new(), + content: json!({ "prompt": "cancel" }), + }) + .func(Arc::new(|_| Box::pin(std::future::pending()))) + .build(), + )); + tokio::time::timeout(std::time::Duration::from_secs(5), async { + while !unsafe { pending_entered() } { + tokio::task::yield_now().await; + } + }) + .await + .expect("native unary future should start before cancellation"); + cancelled_unary.abort(); + assert!( + cancelled_unary + .await + .expect_err("pending native unary call should be cancelled") + .is_cancelled() + ); + tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + flush_subscribers().expect("cancelled unary events should flush"); + if events.lock().unwrap().iter().any(|event| { + event.name() == "fixture.native.unary.drop" + && event.scope_category() == Some(ScopeCategory::End) + }) { + break; + } + } + }) + .await + .expect("plugin unary future Drop should emit its scope after cancellation"); + let cancelled_unary_events = events.lock().unwrap().clone(); + let cancelled_unary_start = find_event( + &cancelled_unary_events, + "native-fixture-cancelled-unary", + Some(ScopeCategory::Start), + ); + assert_parent( + &cancelled_unary_events, + "fixture.native.unary.drop", + Some(ScopeCategory::End), + Some(cancelled_unary_start.uuid()), + ); + drop(fixture_library); + events.lock().unwrap().clear(); let collected_stream_chunks = Arc::new(Mutex::new(Vec::::new())); let collector_chunks = collected_stream_chunks.clone(); @@ -649,6 +712,66 @@ async fn sdk_cdylib_registers_tool_request_intercept() { true ); + events.lock().unwrap().clear(); + let provider_polled = Arc::new(AtomicBool::new(false)); + let polled = Arc::clone(&provider_polled); + let cancelled_stream = llm_stream_call_execute( + LlmStreamCallExecuteParams::builder() + .name("native-fixture-cancelled-stream") + .request(LlmRequest { + headers: Map::new(), + content: json!({ "prompt": "cancel" }), + }) + .func(Arc::new(move |_request| { + let polled = Arc::clone(&polled); + Box::pin(async move { + Ok(LlmJsonStream::new(futures::stream::poll_fn(move |_| { + polled.store(true, Ordering::SeqCst); + Poll::Pending + }))) + }) + })) + .collector(Box::new(|_| Ok(()))) + .finalizer(Box::new(|| Json::Null)) + .build(), + ) + .await + .expect("pending native stream should open"); + tokio::time::timeout(std::time::Duration::from_secs(5), async { + while !provider_polled.load(Ordering::SeqCst) { + tokio::task::yield_now().await; + } + }) + .await + .expect("provider stream should be polled before cancellation"); + drop(cancelled_stream); + tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + flush_subscribers().expect("cancelled stream events should flush"); + if events.lock().unwrap().iter().any(|event| { + event.name() == "fixture.native.stream.drop" + && event.scope_category() == Some(ScopeCategory::End) + }) { + break; + } + } + }) + .await + .expect("plugin stream Drop should emit its scope after cancellation"); + let cancelled_stream_events = events.lock().unwrap().clone(); + let cancelled_stream_start = find_event( + &cancelled_stream_events, + "native-fixture-cancelled-stream", + Some(ScopeCategory::Start), + ); + assert_parent( + &cancelled_stream_events, + "fixture.native.stream.drop", + Some(ScopeCategory::End), + Some(cancelled_stream_start.uuid()), + ); + events.lock().unwrap().clear(); let llm_request = LlmRequest { headers: Map::new(), diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index ad70be1dd..35e6c970f 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -2357,6 +2357,96 @@ fn native_continuation_context_observation( }) } +#[test] +fn thread_active_event_applies_only_at_the_captured_stack_top() { + let _restore = ThreadScopeStackRestore::capture(); + let callback_stack = create_scope_stack(); + let same_lineage_stack = + crate::api::runtime::scope_stack::snapshot_scope_stack(&callback_stack).unwrap(); + set_thread_scope_stack(callback_stack); + let managed_event_uuid = uuid::Uuid::now_v7(); + let assert_parent = |active_event, expected_parent| { + assert_eq!(active_event_uuid(), active_event); + assert_eq!( + crate::api::shared::resolve_parent_uuid(None), + Some(expected_parent) + ); + assert_eq!( + crate::api::runtime::capture_propagation_context() + .unwrap() + .parent_uuid, + expected_parent + ); + }; + let nested = ScopeHandle::builder() + .name("nested") + .scope_type(ScopeType::Custom) + .build(); + let nested_uuid = nested.uuid; + Runtime::new() + .unwrap() + .block_on(with_active_event_uuid(managed_event_uuid, async { + sync_thread_active_event(); + crate::api::runtime::task_scope_push(nested); + // Re-entering another native callback must not move the managed + // event's anchor past a scope opened by the first callback. + sync_thread_active_event(); + })); + + assert_parent(None, nested_uuid); + crate::api::runtime::task_scope_remove(&nested_uuid).unwrap(); + assert_parent(Some(managed_event_uuid), managed_event_uuid); + + let same_lineage_top_uuid = same_lineage_stack + .read() + .unwrap_or_else(|error| error.into_inner()) + .top() + .uuid; + set_thread_scope_stack(same_lineage_stack); + assert_parent(None, same_lineage_top_uuid); +} + +#[test] +fn restored_thread_active_event_rebases_after_its_stack_anchor_closes() { + let _restore = ThreadScopeStackRestore::capture(); + set_thread_scope_stack(create_scope_stack()); + let callback_parent = ScopeHandle::builder() + .name("callback-parent") + .scope_type(ScopeType::Custom) + .build(); + let callback_parent_uuid = callback_parent.uuid; + crate::api::runtime::task_scope_push(callback_parent); + let managed_event_uuid = uuid::Uuid::now_v7(); + let callback_binding = + Runtime::new() + .unwrap() + .block_on(with_active_event_uuid(managed_event_uuid, async { + sync_thread_active_event(); + capture_thread_scope_stack() + })); + + crate::api::runtime::task_scope_remove(&callback_parent_uuid).unwrap(); + restore_thread_scope_stack(callback_binding); + assert_eq!(active_event_uuid(), Some(managed_event_uuid)); + assert_eq!( + crate::api::shared::resolve_parent_uuid(None), + Some(managed_event_uuid) + ); + + let nested = ScopeHandle::builder() + .name("nested") + .scope_type(ScopeType::Custom) + .build(); + let nested_uuid = nested.uuid; + crate::api::runtime::task_scope_push(nested); + assert_eq!(active_event_uuid(), None); + assert_eq!( + crate::api::shared::resolve_parent_uuid(None), + Some(nested_uuid) + ); + crate::api::runtime::task_scope_remove(&nested_uuid).unwrap(); +} + #[test] fn native_async_next_preserves_runtime_context_for_unary_and_stream_continuations() { let runtime = tokio::runtime::Builder::new_current_thread() diff --git a/crates/plugin/src/async_sdk.rs b/crates/plugin/src/async_sdk.rs index 332ae714b..cf31d2e53 100644 --- a/crates/plugin/src/async_sdk.rs +++ b/crates/plugin/src/async_sdk.rs @@ -7,6 +7,7 @@ use std::collections::BTreeMap; use std::ffi::c_void; use std::future::Future; use std::marker::PhantomData; +use std::mem::ManuallyDrop; use std::panic::{AssertUnwindSafe, catch_unwind}; use std::pin::Pin; use std::ptr; @@ -576,20 +577,24 @@ unsafe extern "C" fn unary_trampoline( let future = catch_unwind(AssertUnwindSafe(|| match invocation { Ok(invocation) => (state.adapter)(invocation, next, completion_ref), Err(error) => Box::pin(async move { Err(error) }) as UnaryFuture, - })); + })) + .unwrap_or_else(|_| { + Box::pin(async move { Err("typed native middleware callback panicked".into()) }) + }); + let future: UnaryFuture = match binding { + Ok(binding) => Box::pin(ScopedFuture::new(future, binding)), + Err(error) => { + completion.reject(&error); + set_last_error(&state.host.0.v3.v1, &error); + return NemoRelayNativeAsyncCallbackState::Pending as u32; + } + }; if let Err(error) = state.executor.ensure_started() { completion.reject(&error); set_last_error(&state.host.0.v3.v1, &error); return NemoRelayNativeAsyncCallbackState::Pending as u32; } - let task = match future { - Ok(future) => drive_unary(future, binding, completion), - Err(_) => drive_unary( - Box::pin(async move { Err("typed native middleware callback panicked".into()) }), - binding, - completion, - ), - }; + let task = drive_unary(future, completion); if let Err(error) = state.executor.spawn(async move { let _ = task.await; }) { @@ -602,17 +607,9 @@ unsafe extern "C" fn unary_trampoline( fn drive_unary( future: UnaryFuture, - binding: Result, completion: Completion, ) -> Pin + Send>> { Box::pin(async move { - let future: UnaryFuture = match binding { - Ok(binding) => Box::pin(ScopedFuture::new(future, binding)), - Err(error) => { - completion.reject(&error); - return; - } - }; let result = tokio::select! { result = AssertUnwindSafe(future).catch_unwind() => { result.unwrap_or_else(|_| Err("typed native middleware future panicked".into())) @@ -697,6 +694,24 @@ impl ScopePollBinding { let status = unsafe { (self.host.scope_stack_restore_thread)(previous) }; status_result(status, "restore executor scope stack") } + + fn drop_bound(&mut self, value: &mut ManuallyDrop) { + let Ok(previous) = self.enter() else { + // Host last-error state is thread-local, so writing it from this + // executor thread would not report the failure to the caller. The + // value must still be reclaimed even when its context cannot be + // installed during teardown. + // SAFETY: `value` is initialized once and only dropped here. + unsafe { ManuallyDrop::drop(value) }; + return; + }; + let mut restore = ScopePollRestore::new(self, previous); + // SAFETY: `value` is initialized once and only dropped here. + unsafe { ManuallyDrop::drop(value) }; + // There is no caller-visible error channel from `Drop`; `exit` still + // makes its best effort to restore the executor's previous context. + let _ = restore.restore(); + } } struct ScopePollRestore<'a> { @@ -741,13 +756,16 @@ impl Drop for ScopePollBinding { } struct ScopedFuture { - future: F, + future: ManuallyDrop, binding: ScopePollBinding, } impl ScopedFuture { fn new(future: F, binding: ScopePollBinding) -> Self { - Self { future, binding } + Self { + future: ManuallyDrop::new(future), + binding, + } } } @@ -762,20 +780,29 @@ impl Future for ScopedFuture { .enter() .unwrap_or_else(|error| panic!("{error}")); let mut restore = ScopePollRestore::new(&mut this.binding, previous); - let result = unsafe { Pin::new_unchecked(&mut this.future) }.poll(cx); + let result = unsafe { Pin::new_unchecked(&mut *this.future) }.poll(cx); restore.restore().unwrap_or_else(|error| panic!("{error}")); result } } +impl Drop for ScopedFuture { + fn drop(&mut self) { + self.binding.drop_bound(&mut self.future); + } +} + struct ScopedStream { - stream: S, + stream: ManuallyDrop, binding: ScopePollBinding, } impl ScopedStream { fn new(stream: S, binding: ScopePollBinding) -> Self { - Self { stream, binding } + Self { + stream: ManuallyDrop::new(stream), + binding, + } } } @@ -790,12 +817,18 @@ impl Stream for ScopedStream { .enter() .unwrap_or_else(|error| panic!("{error}")); let mut restore = ScopePollRestore::new(&mut this.binding, previous); - let result = unsafe { Pin::new_unchecked(&mut this.stream) }.poll_next(cx); + let result = unsafe { Pin::new_unchecked(&mut *this.stream) }.poll_next(cx); restore.restore().unwrap_or_else(|error| panic!("{error}")); result } } +impl Drop for ScopedStream { + fn drop(&mut self) { + self.binding.drop_bound(&mut self.stream); + } +} + #[derive(Deserialize)] struct NameValueInvocation { name: String, @@ -967,20 +1000,21 @@ unsafe extern "C" fn stream_trampoline( let future = future.unwrap_or_else(|_| { Box::pin(async move { Err("typed native stream callback panicked".into()) }) }); + let (future_binding, stream_binding) = match bindings { + Ok(bindings) => bindings, + Err(error) => { + output.reject_once(&error); + set_last_error(&state.host.0.v3.v1, &error); + return NemoRelayNativeAsyncCallbackState::Pending as u32; + } + }; + let future: StreamFuture = Box::pin(ScopedFuture::new(future, future_binding)); if let Err(error) = state.executor.ensure_started() { output.reject_once(&error); set_last_error(&state.host.0.v3.v1, &error); return NemoRelayNativeAsyncCallbackState::Pending as u32; } let task = async move { - let (future_binding, stream_binding) = match bindings { - Ok(bindings) => bindings, - Err(error) => { - output.reject(&error).await; - return; - } - }; - let future: StreamFuture = Box::pin(ScopedFuture::new(future, future_binding)); let stream = tokio::select! { result = AssertUnwindSafe(future).catch_unwind() => match result { Ok(result) => result, From a3304f8f5016651cd3b59de245cb398cddc6a37f Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 29 Sep 2026 18:29:02 -0700 Subject: [PATCH 2/7] fix(plugin): isolate native callback scope stacks Signed-off-by: Alex Fournier --- crates/core/src/api/runtime.rs | 2 +- .../src/api/runtime/continuation_context.rs | 20 +- crates/core/src/api/runtime/scope_stack.rs | 46 +- crates/core/src/plugin/dynamic/native.rs | 71 ++-- .../tests/fixtures/native_plugin/src/lib.rs | 47 ++ .../tests/integration/native_plugin_tests.rs | 51 +++ .../tests/unit/continuation_context_tests.rs | 49 ++- crates/core/tests/unit/native_plugin_tests.rs | 401 +++++++++++++++++- 8 files changed, 639 insertions(+), 48 deletions(-) diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index 97471cf5e..4c3cde70a 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -33,7 +33,7 @@ pub use scope_stack::{ set_thread_scope_stack, sync_thread_scope_stack, task_scope_push, task_scope_remove, task_scope_top, with_active_event_uuid, with_scope_stack, }; -pub(crate) use scope_stack::{capture_trace_context, sync_thread_active_event}; +pub(crate) use scope_stack::{capture_trace_context, sync_thread_active_event_for_stack}; pub use state::NemoRelayContextState; #[doc(hidden)] pub use subscriber_dispatcher::SubscriberDelivery; diff --git a/crates/core/src/api/runtime/continuation_context.rs b/crates/core/src/api/runtime/continuation_context.rs index 7ce3a1c0c..202a3e53d 100644 --- a/crates/core/src/api/runtime/continuation_context.rs +++ b/crates/core/src/api/runtime/continuation_context.rs @@ -9,8 +9,9 @@ use crate::api::optimization::{ LlmOptimizationRecorder, current_llm_optimization_recorder, scope_llm_optimization_recorder, }; use crate::api::runtime::scope_stack::{ - ScopeStackHandle, TASK_SCOPE_STACK, active_event_uuid, current_context_scope_stack, - current_scope_stack, scope_stack_active, snapshot_scope_stack, with_active_event_uuid, + ActiveEventBinding, ScopeStackHandle, TASK_SCOPE_STACK, capture_active_event_binding, + current_context_scope_stack, current_scope_stack, rebase_active_event_binding, + scope_stack_active, snapshot_scope_stack, with_active_event_binding, }; use crate::api::runtime::subscriber_dispatcher::{ PublicationBuffer, PublicationContext, capture_nested_publication_buffer, @@ -27,7 +28,7 @@ use crate::error::{FlowError, Result}; #[derive(Clone)] pub struct MiddlewareContinuationContext { scope_stack: ScopeStackHandle, - active_event_uuid: Option, + active_event: Option, publication_context: Option, publication_buffer: Option, optimization_recorder: Option, @@ -40,7 +41,7 @@ impl MiddlewareContinuationContext { pub fn capture() -> Self { Self { scope_stack: current_scope_stack(), - active_event_uuid: active_event_uuid(), + active_event: capture_active_event_binding(), publication_context: capture_publication_context(), publication_buffer: capture_nested_publication_buffer(), optimization_recorder: current_llm_optimization_recorder(), @@ -63,9 +64,12 @@ impl MiddlewareContinuationContext { /// from the stack captured when the middleware callback began. #[doc(hidden)] pub fn isolated_with_scope_stack(&self, scope_stack: &ScopeStackHandle) -> Result { + let scope_stack = snapshot_scope_stack(scope_stack)?; Ok(Self { - scope_stack: snapshot_scope_stack(scope_stack)?, - active_event_uuid: self.active_event_uuid, + active_event: self + .active_event + .map(|active_event| rebase_active_event_binding(active_event, &scope_stack)), + scope_stack, publication_context: self.publication_context.clone(), publication_buffer: self.publication_buffer.clone(), optimization_recorder: self.optimization_recorder.clone(), @@ -92,8 +96,8 @@ impl MiddlewareContinuationContext { let published = with_task_nested_publication_buffer(self.publication_buffer.clone(), published); let active = async { - match self.active_event_uuid { - Some(uuid) => with_active_event_uuid(uuid, published).await, + match self.active_event { + Some(active_event) => with_active_event_binding(active_event, published).await, None => published.await, } }; diff --git a/crates/core/src/api/runtime/scope_stack.rs b/crates/core/src/api/runtime/scope_stack.rs index 866cadf35..c68313757 100644 --- a/crates/core/src/api/runtime/scope_stack.rs +++ b/crates/core/src/api/runtime/scope_stack.rs @@ -545,7 +545,7 @@ impl Default for ScopeStack { pub type ScopeStackHandle = Arc>; #[derive(Clone, Copy)] -struct ActiveEventBinding { +pub(crate) struct ActiveEventBinding { uuid: Uuid, // Propagated stacks may share a root UUID, so identify the captured Arc allocation. scope_stack_id: usize, @@ -802,9 +802,42 @@ pub async fn with_active_event_uuid(uuid: Uuid, future: impl Future( + active_event: ActiveEventBinding, + future: impl Future, +) -> T { ACTIVE_EVENT.scope(active_event, future).await } +pub(crate) fn capture_active_event_binding() -> Option { + ACTIVE_EVENT.try_with(|event| *event).ok().or_else(|| { + let event = THREAD_ACTIVE_EVENT.with(Cell::get)?; + let stack = current_scope_stack(); + (event.scope_stack_id == Arc::as_ptr(&stack) as usize).then_some(event) + }) +} + +pub(crate) fn rebase_active_event_binding( + active_event: ActiveEventBinding, + scope_stack: &ScopeStackHandle, +) -> ActiveEventBinding { + let guard = scope_stack + .read() + .unwrap_or_else(|error| error.into_inner()); + ActiveEventBinding { + uuid: active_event.uuid, + scope_stack_id: Arc::as_ptr(scope_stack) as usize, + scope_stack_top: if guard.find(&active_event.scope_stack_top).is_some() { + active_event.scope_stack_top + } else { + guard.top().uuid + }, + } +} + pub(crate) fn active_event_uuid() -> Option { ACTIVE_EVENT .try_with(|event| event.uuid) @@ -962,8 +995,15 @@ pub fn sync_thread_scope_stack(handle: ScopeStackHandle) { THREAD_SCOPE_STACK.with(|stack| *stack.borrow_mut() = handle); } -pub(crate) fn sync_thread_active_event() { - let active_event = ACTIVE_EVENT.try_with(|event| *event).ok(); +/// Synchronize the task-local managed event onto an isolated thread stack. +/// +/// Native async callbacks run on a plugin-owned executor. The host snapshots +/// the callback's visible stack before crossing that boundary, so the managed +/// event binding must be rebased to the snapshot's allocation while retaining +/// its original scope anchor. +pub(crate) fn sync_thread_active_event_for_stack(scope_stack: &ScopeStackHandle) { + let active_event = capture_active_event_binding() + .map(|active_event| rebase_active_event_binding(active_event, scope_stack)); THREAD_ACTIVE_EVENT.with(|event| event.set(active_event)); } diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 0b745605f..51bf08d5c 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -27,6 +27,7 @@ use crate::api::registry::{ RuntimeRegistrationKind, deregister_conditional_middleware_guardrail, list_runtime_registrations, register_conditional_middleware_guardrail, }; +use crate::api::runtime::scope_stack::snapshot_scope_stack; use crate::api::runtime::{ ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmConditionalFn, LlmExecutionFn, LlmExecutionNextFn, LlmJsonStream, @@ -38,7 +39,7 @@ use crate::api::runtime::{ use crate::api::runtime::{ ScopeStackHandle, ThreadScopeStackBinding, capture_thread_scope_stack, create_scope_stack, current_scope_stack, restore_thread_scope_stack, scope_stack_active, set_thread_scope_stack, - sync_thread_active_event, sync_thread_scope_stack, with_scope_stack, + sync_thread_active_event_for_stack, sync_thread_scope_stack, with_scope_stack, }; use crate::api::scope::{ EmitMarkEventParams, PopScopeParams, PushScopeParams, ScopeAttributes, ScopeHandle, ScopeType, @@ -2047,6 +2048,7 @@ async fn invoke_native_async_callback( } else { None }; + let callback_scope_stack = snapshot_scope_stack(¤t_scope_stack())?; let invocation = native_string_from_json(&invocation) .ok_or_else(|| FlowError::Internal("failed to allocate native async invocation".into()))? as usize; @@ -2068,7 +2070,7 @@ async fn invoke_native_async_callback( completed: false, }; let completion_ref = Arc::into_raw(completion.clone()) as usize; - let next_ref = match (next, runtime) { + let next_ref = with_scope_stack(callback_scope_stack.clone(), || match (next, runtime) { (Some(inner), Some(runtime)) => Some(Arc::into_raw(Arc::new( NativeAsyncNext::with_completion_owner( inner, @@ -2079,23 +2081,25 @@ async fn invoke_native_async_callback( )) as usize), (None, None) => None, _ => unreachable!("runtime is present exactly for native async intercepts"), - }; + }); // ABI v3 exposes a thread-stack capture operation. Mirror the effective // task-local stack into that slot only while entering plugin code so the // SDK can capture it before moving the future to its own executor. let previous_thread_stack = capture_thread_scope_stack(); - sync_thread_scope_stack(current_scope_stack()); - sync_thread_active_event(); - let state = catch_unwind(AssertUnwindSafe(|| unsafe { - cb( - user_data.ptr, - invocation as *const NemoRelayNativeString, - next_ref - .map(|next| next as *const NemoRelayNativeAsyncNext) - .unwrap_or(ptr::null()), - completion_ref as *const NemoRelayNativeAsyncCompletion, - ) - })); + sync_thread_scope_stack(callback_scope_stack.clone()); + sync_thread_active_event_for_stack(&callback_scope_stack); + let state = with_scope_stack(callback_scope_stack, || { + catch_unwind(AssertUnwindSafe(|| unsafe { + cb( + user_data.ptr, + invocation as *const NemoRelayNativeString, + next_ref + .map(|next| next as *const NemoRelayNativeAsyncNext) + .unwrap_or(ptr::null()), + completion_ref as *const NemoRelayNativeAsyncCompletion, + ) + })) + }); restore_thread_scope_stack(previous_thread_stack); let state = match state { Ok(state) => state, @@ -3792,24 +3796,29 @@ fn wrap_native_incremental_llm_stream_execution_with_user_data( "native async stream intercept requires a Tokio runtime: {error}" )) })?; - let next_ref = Arc::into_raw(Arc::new(NativeAsyncNext::with_stream_owner( - NativeAsyncNextInner::LlmStream(next), - runtime, - Some(user_data.clone()), - &stream, - ))); + let callback_scope_stack = snapshot_scope_stack(¤t_scope_stack())?; + let next_ref = with_scope_stack(callback_scope_stack.clone(), || { + Arc::into_raw(Arc::new(NativeAsyncNext::with_stream_owner( + NativeAsyncNextInner::LlmStream(next), + runtime, + Some(user_data.clone()), + &stream, + ))) + }); let stream_ref = Arc::into_raw(stream.clone()); let previous_thread_stack = capture_thread_scope_stack(); - sync_thread_scope_stack(current_scope_stack()); - sync_thread_active_event(); - let state = catch_unwind(AssertUnwindSafe(|| unsafe { - cb( - user_data.ptr, - invocation, - next_ref as *const NemoRelayNativeAsyncNext, - stream_ref as *const NemoRelayNativeAsyncStream, - ) - })); + sync_thread_scope_stack(callback_scope_stack.clone()); + sync_thread_active_event_for_stack(&callback_scope_stack); + let state = with_scope_stack(callback_scope_stack, || { + catch_unwind(AssertUnwindSafe(|| unsafe { + cb( + user_data.ptr, + invocation, + next_ref as *const NemoRelayNativeAsyncNext, + stream_ref as *const NemoRelayNativeAsyncStream, + ) + })) + }); restore_thread_scope_stack(previous_thread_stack); unsafe { native_string_free(invocation) }; state diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index 0cb90aa05..8f95725eb 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -226,6 +226,23 @@ impl NativePlugin for FixtureNativePlugin { let args = context.args; let args = mark_json(args, "native_plugin_tool_execution_request"); let result = if args + .get("use_scoped_next") + .and_then(Json::as_bool) + .unwrap_or(false) + { + let mut scope = runtime.scope( + "fixture.native.scoped.next", + ScopeType::Custom, + None, + None, + Some(&Json::String("scoped-next-input".into())), + )?; + let call_result = next.call(args).await; + let close_result = + scope.close(Some(&Json::String("scoped-next-output".into())), None); + close_result?; + call_result? + } else if args .get("use_isolated_next") .and_then(Json::as_bool) .unwrap_or(false) @@ -283,6 +300,36 @@ impl NativePlugin for FixtureNativePlugin { } } })?; + ctx.register_tool_execution_intercept("fixture_tool_execution_nested", 1, { + let runtime = runtime.clone(); + move |context, next| { + let runtime = runtime.clone(); + async move { + let args = context.args; + if !args + .get("use_scoped_next") + .and_then(Json::as_bool) + .unwrap_or(false) + { + return next.call(args).await.map(Into::into); + } + let mut scope = runtime.scope( + "fixture.native.scoped.next.downstream", + ScopeType::Custom, + None, + None, + Some(&Json::String("scoped-next-downstream-input".into())), + )?; + let call_result = next.call(args).await; + let close_result = scope.close( + Some(&Json::String("scoped-next-downstream-output".into())), + None, + ); + close_result?; + call_result.map(Into::into) + } + } + })?; ctx.register_llm_sanitize_request_guardrail( "fixture_llm_sanitize_request", diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index 6457d846d..254a0c8c4 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -485,6 +485,57 @@ async fn sdk_cdylib_registers_tool_request_intercept() { "native next callback should use the plugin-selected isolated stack" ); + events.lock().unwrap().clear(); + let result = tool_call_execute( + ToolCallExecuteParams::builder() + .name("native-fixture-tool-scoped-next") + .args(json!({ + "input": "scoped-next", + "use_scoped_next": true + })) + .func(Arc::new(|_args| { + Box::pin(async move { + emit_scope_mark( + EmitMarkEventParams::builder() + .name("native-fixture-tool-scoped-next-callback-mark") + .build(), + )?; + Ok(ToolExecutionResult::new(json!({ "tool_callback": true }))) + }) + })) + .build(), + ) + .await + .expect("native scoped next middleware should run"); + assert_eq!(result.result["tool_callback"], true); + flush_subscribers().expect("scoped next native fixture events should flush"); + let scoped_next_events = events.lock().unwrap().clone(); + let scoped_next_scope = find_event( + &scoped_next_events, + "fixture.native.scoped.next", + Some(ScopeCategory::Start), + ); + let downstream_scoped_next_scope = find_event( + &scoped_next_events, + "fixture.native.scoped.next.downstream", + Some(ScopeCategory::Start), + ); + let scoped_next_callback_mark = find_event( + &scoped_next_events, + "native-fixture-tool-scoped-next-callback-mark", + None, + ); + assert_eq!( + downstream_scoped_next_scope.parent_uuid(), + Some(scoped_next_scope.uuid()), + "the downstream native callback should inherit the scope opened around next.call" + ); + assert_eq!( + scoped_next_callback_mark.parent_uuid(), + Some(downstream_scoped_next_scope.uuid()), + "work below the downstream native callback should remain nested" + ); + events.lock().unwrap().clear(); { let thread_stack = create_scope_stack(); diff --git a/crates/core/tests/unit/continuation_context_tests.rs b/crates/core/tests/unit/continuation_context_tests.rs index 9853f32a3..7bb4afd14 100644 --- a/crates/core/tests/unit/continuation_context_tests.rs +++ b/crates/core/tests/unit/continuation_context_tests.rs @@ -6,9 +6,12 @@ use crate::api::optimization::{ LlmOptimizationRecorder, record_llm_optimization_contribution, scope_llm_optimization_recorder, }; use crate::api::runtime::scope_stack::{ - TASK_SCOPE_STACK, active_event_uuid, create_scope_stack, current_scope_stack, - with_active_event_uuid, + TASK_SCOPE_STACK, active_event_uuid, capture_thread_scope_stack, create_scope_stack, + current_scope_stack, restore_thread_scope_stack, snapshot_scope_stack, + sync_thread_active_event_for_stack, task_scope_push, task_scope_remove, with_active_event_uuid, + with_scope_stack, }; +use crate::api::scope::{ScopeHandle, ScopeType}; use crate::codec::optimization::LlmOptimizationContribution; use crate::error::FlowError; use std::sync::Arc; @@ -83,6 +86,48 @@ fn continuation_context_isolates_each_scope_stack_snapshot() { }); } +#[test] +fn continuation_context_preserves_the_managed_event_anchor_across_nested_scopes() { + let runtime = tokio::runtime::Runtime::new().unwrap(); + runtime.block_on(async { + let scope_stack = create_scope_stack(); + let managed_event_uuid = uuid::Uuid::now_v7(); + TASK_SCOPE_STACK + .scope( + scope_stack, + with_active_event_uuid(managed_event_uuid, async { + let context = MiddlewareContinuationContext::capture(); + let nested = ScopeHandle::builder() + .name("plugin-a") + .scope_type(ScopeType::Custom) + .parent_uuid(managed_event_uuid) + .build(); + let nested_uuid = nested.uuid; + task_scope_push(nested); + let invocation = context.isolated().unwrap(); + + let observed_parent = invocation + .run(async { + let callback_stack = + snapshot_scope_stack(¤t_scope_stack()).unwrap(); + let previous_thread_binding = capture_thread_scope_stack(); + sync_thread_active_event_for_stack(&callback_stack); + let parent = with_scope_stack(callback_stack, || { + crate::api::shared::resolve_parent_uuid(None) + }); + restore_thread_scope_stack(previous_thread_binding); + parent + }) + .await; + + task_scope_remove(&nested_uuid).unwrap(); + assert_eq!(observed_parent, Some(nested_uuid)); + }), + ) + .await; + }); +} + #[test] fn continuation_lease_honors_an_explicit_scope_stack_in_a_spawned_task() { let runtime = tokio::runtime::Runtime::new().unwrap(); diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index 35e6c970f..3c99ae808 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -7,6 +7,7 @@ use super::*; use std::collections::VecDeque; use std::panic::{AssertUnwindSafe, catch_unwind}; +use std::sync::Barrier; use std::sync::atomic::{AtomicUsize, Ordering}; use std::time::Duration; @@ -2386,11 +2387,11 @@ fn thread_active_event_applies_only_at_the_captured_stack_top() { Runtime::new() .unwrap() .block_on(with_active_event_uuid(managed_event_uuid, async { - sync_thread_active_event(); + sync_thread_active_event_for_stack(¤t_scope_stack()); crate::api::runtime::task_scope_push(nested); // Re-entering another native callback must not move the managed // event's anchor past a scope opened by the first callback. - sync_thread_active_event(); + sync_thread_active_event_for_stack(¤t_scope_stack()); })); assert_parent(None, nested_uuid); @@ -2421,7 +2422,7 @@ fn restored_thread_active_event_rebases_after_its_stack_anchor_closes() { Runtime::new() .unwrap() .block_on(with_active_event_uuid(managed_event_uuid, async { - sync_thread_active_event(); + sync_thread_active_event_for_stack(¤t_scope_stack()); capture_thread_scope_stack() })); @@ -5484,6 +5485,400 @@ unsafe extern "C" fn resolve_async_static_json( NemoRelayNativeAsyncCallbackState::Complete as u32 } +#[derive(Debug)] +struct NativeCallbackEntryObservation { + isolated_stack: bool, + scope_local_subscriber_preserved: bool, + managed_event_is_parent: bool, + nested_scope_closed: bool, +} + +struct NativeCallbackEntryProbe { + original_stack: ScopeStackHandle, + expected_subscriber: EventSubscriberFn, + expected_parent: uuid::Uuid, + observation: Mutex>, +} + +unsafe extern "C" fn observe_native_callback_entry( + user_data: *mut c_void, + _invocation_json: *const NemoRelayNativeString, + _next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, +) -> u32 { + let probe = unsafe { &*user_data.cast::() }; + let callback_stack = current_scope_stack(); + let subscribers = callback_stack + .read() + .unwrap_or_else(|error| error.into_inner()) + .collect_scope_local_subscribers(); + let managed_event_is_parent = + crate::api::shared::resolve_parent_uuid(None) == Some(probe.expected_parent); + let nested = ScopeHandle::builder() + .name("callback-entry") + .scope_type(ScopeType::Custom) + .parent_uuid(crate::api::shared::resolve_parent_uuid(None).unwrap()) + .build(); + let nested_uuid = nested.uuid; + crate::api::runtime::task_scope_push(nested); + *probe + .observation + .lock() + .unwrap_or_else(|error| error.into_inner()) = Some(NativeCallbackEntryObservation { + isolated_stack: !Arc::ptr_eq(&callback_stack, &probe.original_stack), + scope_local_subscriber_preserved: subscribers + .iter() + .any(|subscriber| Arc::ptr_eq(subscriber, &probe.expected_subscriber)), + managed_event_is_parent, + nested_scope_closed: crate::api::runtime::task_scope_remove(&nested_uuid).is_ok(), + }); + + let value = native_string_from_json(&Json::Null).unwrap(); + assert_eq!( + unsafe { native_async_completion_resolve_json(completion, value) }, + NemoRelayStatus::Ok + ); + unsafe { native_string_free(value) }; + NemoRelayNativeAsyncCallbackState::Complete as u32 +} + +#[test] +fn native_callback_entry_uses_an_isolated_snapshot() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let original_stack = create_scope_stack(); + let outer = ScopeHandle::builder() + .name("outer") + .scope_type(ScopeType::Custom) + .parent_uuid( + original_stack + .read() + .unwrap_or_else(|error| error.into_inner()) + .root_uuid(), + ) + .build(); + let outer_uuid = outer.uuid; + let subscriber: EventSubscriberFn = Arc::new(|_| {}); + { + let mut stack = original_stack + .write() + .unwrap_or_else(|error| error.into_inner()); + stack.push(outer); + stack + .local_registries_mut(&outer_uuid) + .unwrap() + .event_subscribers + .insert("callback-local".into(), subscriber.clone()); + } + let managed_event_uuid = uuid::Uuid::now_v7(); + let probe = NativeCallbackEntryProbe { + original_stack: original_stack.clone(), + expected_subscriber: subscriber, + expected_parent: managed_event_uuid, + observation: Mutex::new(None), + }; + let user_data = Arc::new(NativeCallbackUserData { + ptr: (&probe as *const NativeCallbackEntryProbe) + .cast_mut() + .cast(), + free_fn: None, + _instance: None, + }); + + let result = runtime.block_on(TASK_SCOPE_STACK.scope( + original_stack.clone(), + with_active_event_uuid( + managed_event_uuid, + invoke_native_async_callback( + observe_native_callback_entry, + user_data, + Json::Null, + None, + None, + ), + ), + )); + assert_eq!(result.unwrap(), Json::Null); + let observation = probe.observation.lock().unwrap().take().unwrap(); + assert!(observation.isolated_stack); + assert!(observation.scope_local_subscriber_preserved); + assert!(observation.managed_event_is_parent); + assert!(observation.nested_scope_closed); + assert_eq!( + original_stack + .read() + .unwrap_or_else(|error| error.into_inner()) + .top() + .uuid, + outer_uuid + ); +} + +struct NativeCallbackInterleave { + original_stack: ScopeStackHandle, + callbacks_entered: Barrier, + first_opened: Barrier, + both_opened: Barrier, + first_closed: Barrier, +} + +impl NativeCallbackInterleave { + fn new(original_stack: ScopeStackHandle) -> Self { + Self { + original_stack, + callbacks_entered: Barrier::new(2), + first_opened: Barrier::new(2), + both_opened: Barrier::new(2), + first_closed: Barrier::new(2), + } + } + + fn run(&self, branch: u8) -> bool { + // Both callbacks must capture their context before either mutates it. + self.callbacks_entered.wait(); + + let scope = ScopeHandle::builder() + .name(format!("callback-{branch}")) + .scope_type(ScopeType::Custom) + .parent_uuid(crate::api::shared::resolve_parent_uuid(None).unwrap()) + .build(); + let scope_uuid = scope.uuid; + if branch == 0 { + crate::api::runtime::task_scope_push(scope); + self.first_opened.wait(); + } else { + self.first_opened.wait(); + crate::api::runtime::task_scope_push(scope); + } + + // Force A-push, B-push, A-close, B-close. A shared LIFO stack rejects + // A's close because B is still on top. + self.both_opened.wait(); + if branch == 0 { + let closed = crate::api::runtime::task_scope_remove(&scope_uuid).is_ok(); + self.first_closed.wait(); + closed + } else { + self.first_closed.wait(); + crate::api::runtime::task_scope_remove(&scope_uuid).is_ok() + } + } +} + +unsafe extern "C" fn free_native_callback_interleave(user_data: *mut c_void) { + drop(unsafe { Box::from_raw(user_data.cast::>()) }); +} + +fn native_callback_interleave_user_data( + state: Arc, +) -> Arc { + Arc::new(NativeCallbackUserData { + ptr: Box::into_raw(Box::new(state)).cast(), + free_fn: Some(free_native_callback_interleave), + _instance: None, + }) +} + +fn native_callback_branch(invocation_json: *const NemoRelayNativeString) -> u8 { + let invocation: Json = serde_json::from_str(&read_native_string(invocation_json).unwrap()) + .expect("native callback invocation should be JSON"); + invocation + .get("branch") + .or_else(|| invocation.pointer("/request/content/branch")) + .and_then(Json::as_u64) + .expect("native callback invocation should identify its branch") as u8 +} + +unsafe extern "C" fn interleave_native_callback_scopes( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + _next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, +) -> u32 { + let branch = native_callback_branch(invocation_json); + let state = unsafe { &*user_data.cast::>() }.clone(); + let isolated_entry = !Arc::ptr_eq(¤t_scope_stack(), &state.original_stack); + let binding = capture_thread_scope_stack(); + let completion = completion as usize; + std::thread::spawn(move || { + let _restore = ThreadScopeStackRestore::capture(); + restore_thread_scope_stack(binding); + let closed = state.run(branch); + let value = native_string_from_json( + &json!({"branch": branch, "closed": closed, "isolated_entry": isolated_entry}), + ) + .unwrap(); + assert_eq!( + unsafe { + native_async_completion_resolve_json( + completion as *const NemoRelayNativeAsyncCompletion, + value, + ) + }, + NemoRelayStatus::Ok + ); + unsafe { + native_string_free(value); + native_async_completion_release(completion as *const NemoRelayNativeAsyncCompletion); + } + }); + NemoRelayNativeAsyncCallbackState::Pending as u32 +} + +#[test] +fn concurrent_native_callbacks_close_scopes_on_isolated_stacks() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let original_stack = create_scope_stack(); + let user_data = native_callback_interleave_user_data(Arc::new(NativeCallbackInterleave::new( + original_stack.clone(), + ))); + + let (first, second) = runtime + .block_on(async { + tokio::time::timeout( + Duration::from_secs(5), + TASK_SCOPE_STACK.scope(original_stack, async { + tokio::join!( + invoke_native_async_callback( + interleave_native_callback_scopes, + user_data.clone(), + json!({"branch": 0}), + None, + None, + ), + invoke_native_async_callback( + interleave_native_callback_scopes, + user_data, + json!({"branch": 1}), + None, + None, + ) + ) + }), + ) + .await + }) + .expect("concurrent native callbacks should settle"); + assert_eq!( + first.unwrap(), + json!({"branch": 0, "closed": true, "isolated_entry": true}) + ); + assert_eq!( + second.unwrap(), + json!({"branch": 1, "closed": true, "isolated_entry": true}) + ); +} + +unsafe extern "C" fn interleave_native_stream_callback_scopes( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const NemoRelayNativeAsyncNext, + stream: *const NemoRelayNativeAsyncStream, +) -> u32 { + let branch = native_callback_branch(invocation_json); + let state = unsafe { &*user_data.cast::>() }.clone(); + let isolated_entry = !Arc::ptr_eq(¤t_scope_stack(), &state.original_stack); + let binding = capture_thread_scope_stack(); + let next = next as usize; + let stream = stream as usize; + std::thread::spawn(move || { + let _restore = ThreadScopeStackRestore::capture(); + restore_thread_scope_stack(binding); + let closed = state.run(branch); + let value = native_string_from_json( + &json!({"branch": branch, "closed": closed, "isolated_entry": isolated_entry}), + ) + .unwrap(); + assert_eq!( + unsafe { + native_async_stream_push_json(stream as *const NemoRelayNativeAsyncStream, value) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { native_async_stream_finish(stream as *const NemoRelayNativeAsyncStream) }, + NemoRelayStatus::Ok + ); + unsafe { + native_string_free(value); + native_async_next_release(next as *const NemoRelayNativeAsyncNext); + native_async_stream_release(stream as *const NemoRelayNativeAsyncStream); + } + }); + NemoRelayNativeAsyncCallbackState::Pending as u32 +} + +#[test] +fn concurrent_native_stream_callbacks_close_scopes_on_isolated_stacks() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let original_stack = create_scope_stack(); + let wrapped = wrap_native_incremental_llm_stream_execution_with_user_data( + interleave_native_stream_callback_scopes, + native_callback_interleave_user_data(Arc::new(NativeCallbackInterleave::new( + original_stack.clone(), + ))), + ); + let downstream: LlmStreamExecutionNextFn = + Arc::new(|_| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })); + + let (first_chunk, second_chunk, first_done, second_done) = runtime + .block_on(async { + tokio::time::timeout(Duration::from_secs(5), async { + let (first, second) = TASK_SCOPE_STACK + .scope(original_stack, async { + tokio::join!( + wrapped( + "first", + LlmRequest { + headers: Map::new(), + content: json!({"branch": 0}), + }, + downstream.clone(), + ), + wrapped( + "second", + LlmRequest { + headers: Map::new(), + content: json!({"branch": 1}), + }, + downstream, + ) + ) + }) + .await; + let mut first = first.unwrap(); + let mut second = second.unwrap(); + let (first_chunk, second_chunk) = tokio::join!(first.next(), second.next()); + ( + first_chunk, + second_chunk, + first.next().await, + second.next().await, + ) + }) + .await + }) + .expect("concurrent native streams should settle"); + assert_eq!( + first_chunk.unwrap().unwrap(), + json!({"branch": 0, "closed": true, "isolated_entry": true}) + ); + assert_eq!( + second_chunk.unwrap().unwrap(), + json!({"branch": 1, "closed": true, "isolated_entry": true}) + ); + assert!(first_done.is_none()); + assert!(second_done.is_none()); +} + #[cfg(unix)] #[tokio::test] async fn native_async_wrappers_validate_callback_result_shapes() { From dcd5de0111033c96de4ff3c7217bef10d9dd85af Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 29 Sep 2026 18:34:26 -0700 Subject: [PATCH 3/7] test(plugin): isolate native cancellation probes Signed-off-by: Alex Fournier --- .../tests/fixtures/native_plugin/src/lib.rs | 99 ++++++++++++------- 1 file changed, 63 insertions(+), 36 deletions(-) diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index 8f95725eb..c3dbc4c15 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -391,50 +391,77 @@ impl NativePlugin for FixtureNativePlugin { )) }, )?; - ctx.register_llm_execution_intercept("fixture_llm_execution", 0, { + ctx.register_llm_execution_intercept( + "fixture_llm_execution", + 0, + |_name, request, next| async move { + let response = next + .call(mark_llm_request( + request, + "native_plugin_llm_execution_request", + )) + .await?; + Ok(mark_json(response, "native_plugin_llm_execution")) + }, + )?; + ctx.register_llm_stream_execution_intercept( + "fixture_llm_stream_execution", + 0, + |_name, request, next| async move { + let stream = next + .call(mark_llm_request( + request, + "native_plugin_llm_stream_execution_request", + )) + .await?; + let stream: LlmJsonAsyncStream = Box::pin(stream.map(|chunk| { + chunk.map(|chunk| mark_json(chunk, "native_plugin_llm_stream_execution")) + })); + Ok(stream) + }, + )?; + ctx.register_llm_execution_intercept("fixture_llm_execution_cancellation", -1, { let runtime = runtime.clone(); move |name, request, next| { - let drop_scope = (name == "native-fixture-cancelled-unary").then(|| DropScope { - runtime: runtime.clone(), - name: "fixture.native.unary.drop", - }); + let runtime = runtime.clone(); async move { - if let Some(_drop_scope) = drop_scope { - ASYNC_PENDING_ENTERED.store(true, Ordering::Release); - return futures::future::pending::>().await; + if name != "native-fixture-cancelled-unary" { + return next.call(request).await; } - let response = next - .call(mark_llm_request( - request, - "native_plugin_llm_execution_request", - )) - .await?; - Ok(mark_json(response, "native_plugin_llm_execution")) + let _drop_scope = DropScope { + runtime, + name: "fixture.native.unary.drop", + }; + ASYNC_PENDING_ENTERED.store(true, Ordering::Release); + futures::future::pending::>().await } } })?; - ctx.register_llm_stream_execution_intercept("fixture_llm_stream_execution", 0, { - let runtime = runtime.clone(); - move |name, request, next| { - let drop_scope = (name == "native-fixture-cancelled-stream").then(|| DropScope { - runtime: runtime.clone(), - name: "fixture.native.stream.drop", - }); - async move { - let stream = next - .call(mark_llm_request( - request, - "native_plugin_llm_stream_execution_request", - )) - .await?; - let stream: LlmJsonAsyncStream = Box::pin(stream.map(move |chunk| { - let _ = &drop_scope; - chunk.map(|chunk| mark_json(chunk, "native_plugin_llm_stream_execution")) - })); - Ok(stream) + ctx.register_llm_stream_execution_intercept( + "fixture_llm_stream_execution_cancellation", + -1, + { + let runtime = runtime.clone(); + move |name, request, next| { + let runtime = runtime.clone(); + async move { + if name != "native-fixture-cancelled-stream" { + return next.call(request).await; + } + let drop_scope = DropScope { + runtime, + name: "fixture.native.stream.drop", + }; + let stream = next.call(request).await?; + let stream: LlmJsonAsyncStream = Box::pin(stream.map(move |chunk| { + let _ = &drop_scope; + chunk + })); + Ok(stream) + } } - } - })?; + }, + )?; Ok(()) } From eba3b9d6db2b81f915c95d422a5779e3bfbcac65 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 30 Sep 2026 08:47:40 -0700 Subject: [PATCH 4/7] refactor(runtime): clarify anchored event state Signed-off-by: Alex Fournier --- .../src/api/runtime/continuation_context.rs | 14 ++-- crates/core/src/api/runtime/scope_stack.rs | 72 +++++++++---------- 2 files changed, 43 insertions(+), 43 deletions(-) diff --git a/crates/core/src/api/runtime/continuation_context.rs b/crates/core/src/api/runtime/continuation_context.rs index 202a3e53d..068653fcc 100644 --- a/crates/core/src/api/runtime/continuation_context.rs +++ b/crates/core/src/api/runtime/continuation_context.rs @@ -9,9 +9,9 @@ use crate::api::optimization::{ LlmOptimizationRecorder, current_llm_optimization_recorder, scope_llm_optimization_recorder, }; use crate::api::runtime::scope_stack::{ - ActiveEventBinding, ScopeStackHandle, TASK_SCOPE_STACK, capture_active_event_binding, - current_context_scope_stack, current_scope_stack, rebase_active_event_binding, - scope_stack_active, snapshot_scope_stack, with_active_event_binding, + AnchoredActiveEvent, ScopeStackHandle, TASK_SCOPE_STACK, capture_anchored_active_event, + current_context_scope_stack, current_scope_stack, rebind_active_event_to_stack, + scope_stack_active, snapshot_scope_stack, with_anchored_active_event, }; use crate::api::runtime::subscriber_dispatcher::{ PublicationBuffer, PublicationContext, capture_nested_publication_buffer, @@ -28,7 +28,7 @@ use crate::error::{FlowError, Result}; #[derive(Clone)] pub struct MiddlewareContinuationContext { scope_stack: ScopeStackHandle, - active_event: Option, + active_event: Option, publication_context: Option, publication_buffer: Option, optimization_recorder: Option, @@ -41,7 +41,7 @@ impl MiddlewareContinuationContext { pub fn capture() -> Self { Self { scope_stack: current_scope_stack(), - active_event: capture_active_event_binding(), + active_event: capture_anchored_active_event(), publication_context: capture_publication_context(), publication_buffer: capture_nested_publication_buffer(), optimization_recorder: current_llm_optimization_recorder(), @@ -68,7 +68,7 @@ impl MiddlewareContinuationContext { Ok(Self { active_event: self .active_event - .map(|active_event| rebase_active_event_binding(active_event, &scope_stack)), + .map(|active_event| rebind_active_event_to_stack(active_event, &scope_stack)), scope_stack, publication_context: self.publication_context.clone(), publication_buffer: self.publication_buffer.clone(), @@ -97,7 +97,7 @@ impl MiddlewareContinuationContext { with_task_nested_publication_buffer(self.publication_buffer.clone(), published); let active = async { match self.active_event { - Some(active_event) => with_active_event_binding(active_event, published).await, + Some(active_event) => with_anchored_active_event(active_event, published).await, None => published.await, } }; diff --git a/crates/core/src/api/runtime/scope_stack.rs b/crates/core/src/api/runtime/scope_stack.rs index c68313757..ef8848a49 100644 --- a/crates/core/src/api/runtime/scope_stack.rs +++ b/crates/core/src/api/runtime/scope_stack.rs @@ -545,11 +545,11 @@ impl Default for ScopeStack { pub type ScopeStackHandle = Arc>; #[derive(Clone, Copy)] -pub(crate) struct ActiveEventBinding { - uuid: Uuid, +pub(crate) struct AnchoredActiveEvent { + event_uuid: Uuid, // Propagated stacks may share a root UUID, so identify the captured Arc allocation. - scope_stack_id: usize, - scope_stack_top: Uuid, + stack_identity: usize, + anchor_scope_uuid: Uuid, } /// Captured thread-local scope stack binding. @@ -561,7 +561,7 @@ pub(crate) struct ActiveEventBinding { pub struct ThreadScopeStackBinding { stack: ScopeStackHandle, explicit: bool, - active_event: Option, + active_event: Option, } impl ThreadScopeStackBinding { @@ -791,47 +791,47 @@ tokio::task_local! { /// Task-local scope stack handle used by async execution contexts. pub static TASK_SCOPE_STACK: ScopeStackHandle; /// Managed tool or LLM event currently executing in this task. - static ACTIVE_EVENT: ActiveEventBinding; + static ACTIVE_EVENT: AnchoredActiveEvent; } /// Run a future with `uuid` as the causally active managed event. pub async fn with_active_event_uuid(uuid: Uuid, future: impl Future) -> T { - let (scope_stack_id, scope_stack_top) = scope_stack_identity_and_top(); - let active_event = ActiveEventBinding { - uuid, - scope_stack_id, - scope_stack_top, + let (stack_identity, anchor_scope_uuid) = scope_stack_identity_and_anchor(); + let active_event = AnchoredActiveEvent { + event_uuid: uuid, + stack_identity, + anchor_scope_uuid, }; - with_active_event_binding(active_event, future).await + with_anchored_active_event(active_event, future).await } -pub(crate) async fn with_active_event_binding( - active_event: ActiveEventBinding, +pub(crate) async fn with_anchored_active_event( + active_event: AnchoredActiveEvent, future: impl Future, ) -> T { ACTIVE_EVENT.scope(active_event, future).await } -pub(crate) fn capture_active_event_binding() -> Option { +pub(crate) fn capture_anchored_active_event() -> Option { ACTIVE_EVENT.try_with(|event| *event).ok().or_else(|| { let event = THREAD_ACTIVE_EVENT.with(Cell::get)?; let stack = current_scope_stack(); - (event.scope_stack_id == Arc::as_ptr(&stack) as usize).then_some(event) + (event.stack_identity == Arc::as_ptr(&stack) as usize).then_some(event) }) } -pub(crate) fn rebase_active_event_binding( - active_event: ActiveEventBinding, +pub(crate) fn rebind_active_event_to_stack( + active_event: AnchoredActiveEvent, scope_stack: &ScopeStackHandle, -) -> ActiveEventBinding { +) -> AnchoredActiveEvent { let guard = scope_stack .read() .unwrap_or_else(|error| error.into_inner()); - ActiveEventBinding { - uuid: active_event.uuid, - scope_stack_id: Arc::as_ptr(scope_stack) as usize, - scope_stack_top: if guard.find(&active_event.scope_stack_top).is_some() { - active_event.scope_stack_top + AnchoredActiveEvent { + event_uuid: active_event.event_uuid, + stack_identity: Arc::as_ptr(scope_stack) as usize, + anchor_scope_uuid: if guard.find(&active_event.anchor_scope_uuid).is_some() { + active_event.anchor_scope_uuid } else { guard.top().uuid }, @@ -840,7 +840,7 @@ pub(crate) fn rebase_active_event_binding( pub(crate) fn active_event_uuid() -> Option { ACTIVE_EVENT - .try_with(|event| event.uuid) + .try_with(|event| event.event_uuid) .ok() .or_else(thread_active_event_uuid) } @@ -848,18 +848,18 @@ pub(crate) fn active_event_uuid() -> Option { pub(crate) fn thread_active_event_uuid() -> Option { let mut event = THREAD_ACTIVE_EVENT.with(Cell::get)?; let stack = current_scope_stack(); - if event.scope_stack_id != Arc::as_ptr(&stack) as usize { + if event.stack_identity != Arc::as_ptr(&stack) as usize { return None; } let guard = stack.read().unwrap_or_else(|error| error.into_inner()); - if event.scope_stack_top != guard.top().uuid { - if guard.find(&event.scope_stack_top).is_some() { + if event.anchor_scope_uuid != guard.top().uuid { + if guard.find(&event.anchor_scope_uuid).is_some() { return None; } - event.scope_stack_top = guard.top().uuid; + event.anchor_scope_uuid = guard.top().uuid; THREAD_ACTIVE_EVENT.with(|active| active.set(Some(event))); } - Some(event.uuid) + Some(event.event_uuid) } thread_local! { @@ -871,7 +871,7 @@ thread_local! { /// Whether the current thread explicitly owns a scope stack. static THREAD_SCOPE_STACK_EXPLICIT: std::cell::Cell = const { std::cell::Cell::new(false) }; /// Managed event propagated into a foreign executor with the thread scope binding. - static THREAD_ACTIVE_EVENT: Cell> = const { Cell::new(None) }; + static THREAD_ACTIVE_EVENT: Cell> = const { Cell::new(None) }; } /// Return the scope stack visible to the current execution context. @@ -999,15 +999,15 @@ pub fn sync_thread_scope_stack(handle: ScopeStackHandle) { /// /// Native async callbacks run on a plugin-owned executor. The host snapshots /// the callback's visible stack before crossing that boundary, so the managed -/// event binding must be rebased to the snapshot's allocation while retaining -/// its original scope anchor. +/// event must be rebound to the snapshot's allocation while retaining its +/// original scope anchor. pub(crate) fn sync_thread_active_event_for_stack(scope_stack: &ScopeStackHandle) { - let active_event = capture_active_event_binding() - .map(|active_event| rebase_active_event_binding(active_event, scope_stack)); + let active_event = capture_anchored_active_event() + .map(|active_event| rebind_active_event_to_stack(active_event, scope_stack)); THREAD_ACTIVE_EVENT.with(|event| event.set(active_event)); } -fn scope_stack_identity_and_top() -> (usize, Uuid) { +fn scope_stack_identity_and_anchor() -> (usize, Uuid) { let stack = current_scope_stack(); let guard = stack.read().unwrap_or_else(|error| error.into_inner()); (Arc::as_ptr(&stack) as usize, guard.top().uuid) From d4ca84874464acd618d5d45ed6d67670199b9d74 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Thu, 1 Oct 2026 11:32:58 -0700 Subject: [PATCH 5/7] fix(plugin): isolate worker callback scope context Signed-off-by: Alex Fournier --- .../src/api/runtime/continuation_context.rs | 38 +++ crates/core/src/api/runtime/scope_stack.rs | 16 + crates/core/src/plugin/dynamic/worker.rs | 158 ++++++---- .../tests/fixtures/worker_plugin/src/main.rs | 38 ++- .../tests/integration/worker_plugin_tests.rs | 46 ++- .../core/tests/unit/dynamic_worker_tests.rs | 282 ++++++++++++++++-- 6 files changed, 499 insertions(+), 79 deletions(-) diff --git a/crates/core/src/api/runtime/continuation_context.rs b/crates/core/src/api/runtime/continuation_context.rs index 79ff62005..a0762fd7e 100644 --- a/crates/core/src/api/runtime/continuation_context.rs +++ b/crates/core/src/api/runtime/continuation_context.rs @@ -14,11 +14,19 @@ use crate::api::runtime::scope_stack::{ current_scope_stack, rebind_active_event_to_stack, scope_stack_active, snapshot_scope_stack, with_anchored_active_event, }; +#[cfg(feature = "worker-grpc")] +use crate::api::runtime::scope_stack::{ + install_thread_continuation_context, restore_thread_scope_stack, with_scope_stack, +}; use crate::api::runtime::subscriber_dispatcher::{ PublicationBuffer, PublicationContext, capture_nested_publication_buffer, capture_publication_context, with_task_nested_publication_buffer, with_task_publication_context, }; +#[cfg(feature = "worker-grpc")] +use crate::api::runtime::subscriber_dispatcher::{ + with_nested_publication_buffer, with_publication_context, +}; use crate::error::{FlowError, Result}; /// Opaque Relay task context captured for a middleware `next` continuation. @@ -80,6 +88,36 @@ impl MiddlewareContinuationContext { }) } + #[cfg(feature = "worker-grpc")] + pub(crate) fn scope_stack(&self) -> ScopeStackHandle { + self.scope_stack.clone() + } + + #[cfg(feature = "worker-grpc")] + pub(crate) fn run_sync(&self, callback: impl FnOnce() -> T) -> T { + struct RestoreThreadContext(Option); + + impl Drop for RestoreThreadContext { + fn drop(&mut self) { + if let Some(previous) = self.0.take() { + restore_thread_scope_stack(previous); + } + } + } + + let previous = install_thread_continuation_context( + &self.scope_stack, + self.active_event, + self.active_event_trace_context.clone(), + ); + let _restore = RestoreThreadContext(Some(previous)); + with_publication_context(self.publication_context.clone(), || { + with_nested_publication_buffer(self.publication_buffer.clone(), || { + with_scope_stack(self.scope_stack.clone(), callback) + }) + }) + } + /// Clone this context with the scope selection visible to a continuation call. pub(crate) fn isolated_for_current_invocation(&self) -> Result { let visible_scope_stack = current_context_scope_stack().unwrap_or_else(|| { diff --git a/crates/core/src/api/runtime/scope_stack.rs b/crates/core/src/api/runtime/scope_stack.rs index 97558d536..87078ac90 100644 --- a/crates/core/src/api/runtime/scope_stack.rs +++ b/crates/core/src/api/runtime/scope_stack.rs @@ -1281,6 +1281,22 @@ pub fn restore_thread_scope_stack(binding: ThreadScopeStackBinding) { .with(|context| *context.borrow_mut() = binding.active_event_trace_context); } +#[cfg(feature = "worker-grpc")] +pub(crate) fn install_thread_continuation_context( + scope_stack: &ScopeStackHandle, + active_event: Option, + active_event_trace_context: Option, +) -> ThreadScopeStackBinding { + let previous = capture_thread_scope_stack(); + sync_thread_scope_stack(scope_stack.clone()); + THREAD_ACTIVE_EVENT.with(|event| { + event.set(active_event.map(|event| rebind_active_event_to_stack(event, scope_stack))) + }); + THREAD_ACTIVE_EVENT_TRACE_CONTEXT + .with(|context| *context.borrow_mut() = active_event_trace_context); + previous +} + /// Synchronize the thread-local scope stack without marking it explicit. /// /// This updates the thread-local slot used by native runtime code while diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index 622a379d9..138d71f64 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -1697,7 +1697,7 @@ impl WorkerPluginCallback { registration_name: registration_name.into(), }, )), - ); + )?; guardrail_from_invoke_response(self.invoke_blocking(request)?) } @@ -1707,7 +1707,7 @@ impl WorkerPluginCallback { RegistrationSurface::Subscriber, None, Some(invoke_request_payload_event(event)), - ); + )?; let response = self.invoke_blocking(request)?; match response.result { Some(invoke_response_result::Result::Empty(_)) | None => Ok(()), @@ -1728,7 +1728,7 @@ impl WorkerPluginCallback { RegistrationSurface::EventMetadataInjector, None, Some(invoke_request_payload_event(event)), - ); + )?; let value = json_from_invoke_response(self.invoke_async(request).await?)?; let additions = serde_json::from_value::>(value).map_err(|err| { FlowError::Internal(format!( @@ -1749,7 +1749,7 @@ impl WorkerPluginCallback { surface, None, Some(invoke_request_payload_event(event)), - ); + )?; let value = json_from_invoke_response(self.invoke_async(request).await?)?; serde_json::from_value(value).map_err(|err| { FlowError::Internal(format!( @@ -1771,7 +1771,7 @@ impl WorkerPluginCallback { surface, continuation_id, Some(invoke_request_payload_tool(tool_name, value, None)), - ); + )?; json_from_invoke_response(self.invoke_async(request).await?) } @@ -1786,7 +1786,7 @@ impl WorkerPluginCallback { RegistrationSurface::ToolConditionalExecutionGuardrail, None, Some(invoke_request_payload_tool(tool_name, value, None)), - ); + )?; guardrail_from_invoke_response(self.invoke_async(request).await?) } @@ -1806,7 +1806,7 @@ impl WorkerPluginCallback { RegistrationSurface::ToolExecutionIntercept, Some(continuation_id), Some(invoke_request_payload_tool(tool_name, value, tool_call_id)), - ); + )?; let response = self.invoke_async(request).await?; match response.result { Some(invoke_response_result::Result::ToolExecution(result)) => { @@ -1848,7 +1848,7 @@ impl WorkerPluginCallback { }, ), )), - ); + )?; let capability_id = context.resolve_codec().map(|codec| { let capability_id = self .host_state @@ -1898,7 +1898,7 @@ impl WorkerPluginCallback { }, ), )), - ); + )?; let capability_id = context.resolve_codec().map(|codec| { let capability_id = self .host_state @@ -1931,7 +1931,7 @@ impl WorkerPluginCallback { RegistrationSurface::LlmConditionalExecutionGuardrail, None, Some(invoke_request_payload_llm("", Some(request), None, None)), - ); + )?; guardrail_from_invoke_response(self.invoke_async(invoke).await?) } @@ -1952,7 +1952,7 @@ impl WorkerPluginCallback { annotated, None, )), - ); + )?; let response = self.invoke_async(invoke).await?; match response.result { Some(invoke_response_result::Result::LlmRequest(result)) => { @@ -1996,7 +1996,7 @@ impl WorkerPluginCallback { None, None, )), - ); + )?; json_from_invoke_response(self.invoke_async(invoke).await?) } @@ -2020,7 +2020,7 @@ impl WorkerPluginCallback { None, None, )), - ); + )?; let mut client = self.client.clone(); let mut guard = WorkerInvocationGuard::new(self, &invoke); let (tx, rx) = mpsc::channel(16); @@ -2094,12 +2094,20 @@ impl WorkerPluginCallback { surface: RegistrationSurface, continuation_id: Option, payload: Option, - ) -> InvokeRequest { - let scope_stack_id = self.host_state.insert_invocation_scope_stack( + ) -> FlowResult { + let scope_stack_id = match self.host_state.insert_invocation_scope_stack( current_scope_stack(), capture_nested_publication_buffer(), - ); - InvokeRequest { + ) { + Ok(scope_stack_id) => scope_stack_id, + Err(error) => { + if let Some(continuation_id) = continuation_id.as_deref() { + self.host_state.remove_continuation(continuation_id); + } + return Err(error); + } + }; + Ok(InvokeRequest { activation_id: self.activation_id.clone(), auth_token: self.host_state.auth_token.clone(), invocation_id: Uuid::now_v7().to_string(), @@ -2111,7 +2119,7 @@ impl WorkerPluginCallback { parent_scope_id: String::new(), }), payload, - } + }) } fn invoke_blocking(&self, request: InvokeRequest) -> FlowResult { @@ -2361,6 +2369,7 @@ enum WorkerCodecDirection { struct StoredScopeStack { handle: crate::api::runtime::ScopeStackHandle, publication_buffer: Option, + continuation_context: Option, invocation_base_depth: Option, } @@ -2610,33 +2619,43 @@ impl WorkerHostRuntimeState { fn insert_invocation_scope_stack( &self, - stack: crate::api::runtime::ScopeStackHandle, + source_stack: crate::api::runtime::ScopeStackHandle, publication_buffer: Option, - ) -> String { - let id = format!("invoke-{}", Uuid::now_v7()); - let Ok(mut stacks) = self.scope_stacks.lock() else { - return id; - }; + ) -> FlowResult { + let mut stacks = self + .scope_stacks + .lock() + .map_err(|error| FlowError::Internal(format!("scope stack lock poisoned: {error}")))?; loop { - let Ok(cleanups) = self.scope_stack_cleanups.lock() else { - return id; - }; - if !cleanups.iter().any(|handle| Arc::ptr_eq(handle, &stack)) { + let cleanups = self.scope_stack_cleanups.lock().map_err(|error| { + FlowError::Internal(format!("scope cleanup lock poisoned: {error}")) + })?; + if !cleanups + .iter() + .any(|handle| Arc::ptr_eq(handle, &source_stack)) + { break; } drop(stacks); - let Ok(guard) = self.scope_stack_cleanup_complete.wait(cleanups) else { - return id; - }; + let guard = self + .scope_stack_cleanup_complete + .wait(cleanups) + .map_err(|error| { + FlowError::Internal(format!("scope cleanup lock poisoned: {error}")) + })?; drop(guard); - let Ok(guard) = self.scope_stacks.lock() else { - return id; - }; - stacks = guard; + stacks = self.scope_stacks.lock().map_err(|error| { + FlowError::Internal(format!("scope stack lock poisoned: {error}")) + })?; } - let Ok(stack_guard) = stack.read() else { - return id; - }; + + let continuation_context = + MiddlewareContinuationContext::capture().isolated_with_scope_stack(&source_stack)?; + let stack = continuation_context.scope_stack(); + let id = format!("invoke-{}", Uuid::now_v7()); + let stack_guard = stack + .read() + .map_err(|error| FlowError::Internal(format!("scope stack lock poisoned: {error}")))?; let invocation_base_depth = stack_guard.scopes().len(); drop(stack_guard); stacks.insert( @@ -2644,20 +2663,26 @@ impl WorkerHostRuntimeState { StoredScopeStack { handle: stack, publication_buffer, + continuation_context: Some(continuation_context), invocation_base_depth: Some(invocation_base_depth), }, ); - id + Ok(id) } fn cleanup_invocation_scope_stack(&self, id: &str) { let unwind = self.take_invocation_scope_cleanup(id); - if let Some((handle, base_depth)) = unwind { + if let Some((handle, base_depth, continuation_context, publication_buffer)) = unwind { let _cleanup = ScopeStackCleanupGuard { state: self, handle: handle.clone(), }; - Self::unwind_scope_stack(&handle, base_depth); + Self::unwind_scope_stack( + &handle, + base_depth, + continuation_context, + publication_buffer, + ); } if let Ok(mut handles) = self.scope_handles.lock() { handles.retain(|_, handle| handle.scope_stack_id != id); @@ -2667,7 +2692,12 @@ impl WorkerHostRuntimeState { fn take_invocation_scope_cleanup( &self, id: &str, - ) -> Option<(crate::api::runtime::ScopeStackHandle, usize)> { + ) -> Option<( + crate::api::runtime::ScopeStackHandle, + usize, + Option, + Option, + )> { let Ok(mut stacks) = self.scope_stacks.lock() else { return None; }; @@ -2695,11 +2725,21 @@ impl WorkerHostRuntimeState { return None; }; cleanups.push(stored.handle.clone()); - Some((stored.handle, base_depth)) + Some(( + stored.handle, + base_depth, + stored.continuation_context, + stored.publication_buffer, + )) } - fn unwind_scope_stack(stack: &crate::api::runtime::ScopeStackHandle, base_depth: usize) { - loop { + fn unwind_scope_stack( + stack: &crate::api::runtime::ScopeStackHandle, + base_depth: usize, + continuation_context: Option, + publication_buffer: Option, + ) { + let unwind = || loop { let top_uuid = { let Ok(stack) = stack.read() else { return; @@ -2722,6 +2762,12 @@ impl WorkerHostRuntimeState { if stack.remove(&top_uuid).is_err() { return; } + }; + match continuation_context { + Some(context) => context.run_sync(unwind), + None => with_nested_publication_buffer(publication_buffer, || { + with_scope_stack(stack.clone(), unwind) + }), } } @@ -2775,6 +2821,7 @@ impl WorkerHostRuntimeState { .map(|stored| StoredInvocationContext { scope_stack: stored.handle.clone(), publication_buffer: stored.publication_buffer.clone(), + continuation_context: stored.continuation_context.clone(), }) .map(Some) .ok_or_else(|| Status::not_found("scope stack not found")) @@ -2785,6 +2832,18 @@ impl WorkerHostRuntimeState { struct StoredInvocationContext { scope_stack: crate::api::runtime::ScopeStackHandle, publication_buffer: Option, + continuation_context: Option, +} + +impl StoredInvocationContext { + fn run(&self, callback: impl FnOnce() -> T) -> T { + match &self.continuation_context { + Some(context) => context.run_sync(callback), + None => with_nested_publication_buffer(self.publication_buffer.clone(), || { + with_scope_stack(self.scope_stack.clone(), callback) + }), + } + } } #[derive(Clone)] @@ -3210,9 +3269,7 @@ impl RelayHostRuntime for WorkerHostRuntimeService { let result = if handle.scope_stack_id.is_empty() { pop() } else if let Some(context) = self.state.invocation_context(&handle.scope_stack_id)? { - with_nested_publication_buffer(context.publication_buffer, || { - with_scope_stack(context.scope_stack, pop) - }) + context.run(pop) } else { pop() }; @@ -3236,6 +3293,7 @@ impl RelayHostRuntime for WorkerHostRuntimeService { StoredScopeStack { handle: crate::api::runtime::create_scope_stack(), publication_buffer: None, + continuation_context: None, invocation_base_depth: None, }, ); @@ -3555,9 +3613,7 @@ impl WorkerHostRuntimeService { else { return f(); }; - with_nested_publication_buffer(context.publication_buffer, || { - with_scope_stack(context.scope_stack, f) - }) + context.run(f) } } diff --git a/crates/core/tests/fixtures/worker_plugin/src/main.rs b/crates/core/tests/fixtures/worker_plugin/src/main.rs index 9089b515d..7c0870f8c 100644 --- a/crates/core/tests/fixtures/worker_plugin/src/main.rs +++ b/crates/core/tests/fixtures/worker_plugin/src/main.rs @@ -175,7 +175,7 @@ impl WorkerPlugin for FixtureWorkerPlugin { ); register_fixture_tool_hooks( ctx, - runtime, + runtime.clone(), fixture_flag(config, "block_tool"), fixture_flag(config, "tool_request_error"), fixture_flag(config, "exit_in_tool_request"), @@ -183,6 +183,7 @@ impl WorkerPlugin for FixtureWorkerPlugin { ); register_fixture_llm_hooks( ctx, + runtime, fixture_flag(config, "llm_request_error"), fixture_flag(config, "llm_stream_open_error"), ); @@ -310,6 +311,7 @@ fn register_fixture_tool_hooks( fn register_fixture_llm_hooks( ctx: &mut PluginContext, + runtime: nemo_relay_worker::PluginRuntime, llm_request_error: bool, llm_stream_open_error: bool, ) { @@ -372,10 +374,18 @@ fn register_fixture_llm_hooks( ) }, ); - ctx.register_llm_execution_intercept( - "fixture_llm_execution", - 0, - |_name, request, next: LlmNext| async move { + let unary_runtime = runtime.clone(); + ctx.register_llm_execution_intercept("fixture_llm_execution", 0, move |name, request, next: LlmNext| { + let runtime = unary_runtime.clone(); + let name = name.to_owned(); + async move { + runtime + .emit_mark( + "fixture.worker.llm_execution.runtime.mark", + None, + Some(json!({ "name": name })), + ) + .await?; let response = next .call(mark_llm_request( request, @@ -383,17 +393,28 @@ fn register_fixture_llm_hooks( )) .await?; Ok(mark_json(response, "worker_plugin_llm_execution")) - }, - ); + } + }); + let stream_runtime = runtime; ctx.register_llm_stream_execution_intercept( "fixture_llm_stream_execution", 0, - move |_name, request, next: LlmStreamNext| async move { + move |name, request, next: LlmStreamNext| { + let runtime = stream_runtime.clone(); + let name = name.to_owned(); + async move { if llm_stream_open_error { return Err(WorkerSdkError::Callback( "fixture LLM stream open error requested".into(), )); } + runtime + .emit_mark( + "fixture.worker.llm_stream_execution.runtime.mark", + None, + Some(json!({ "name": name })), + ) + .await?; let stream = next .call(mark_llm_request( request, @@ -404,6 +425,7 @@ fn register_fixture_llm_hooks( chunk.map(|value| mark_json(value, "worker_plugin_llm_stream_execution")) })); Ok(mapped) + } }, ); } diff --git a/crates/core/tests/integration/worker_plugin_tests.rs b/crates/core/tests/integration/worker_plugin_tests.rs index c5a842d0f..80c4eecf3 100644 --- a/crates/core/tests/integration/worker_plugin_tests.rs +++ b/crates/core/tests/integration/worker_plugin_tests.rs @@ -702,6 +702,17 @@ async fn rust_worker_registers_and_invokes_all_current_surfaces() { ); assert_eq!(pending_mark.metadata().unwrap()["fixture"], true); assert_eq!(pending_mark.metadata().unwrap()["worker_plugin_mark"], true); + let runtime_mark = find_event( + &captured_events, + "fixture.worker.llm_execution.runtime.mark", + None, + ); + assert_eq!(runtime_mark.parent_uuid(), Some(llm_start.uuid())); + assert!(runtime_mark.propagation_traceparent().is_some()); + assert_eq!( + runtime_mark.metadata().unwrap()["name"], + "worker-fixture-llm-execute" + ); let llm_end = find_event( &captured_events, "worker-fixture-llm-execute", @@ -746,7 +757,27 @@ async fn rust_worker_registers_and_invokes_all_current_surfaces() { stream_value["request"]["worker_plugin_llm_stream_execution_request"], true ); + flush_subscribers().expect("worker fixture streaming events should flush"); + let captured_events = events.lock().unwrap().clone(); + let stream_start = find_event( + &captured_events, + "worker-fixture-llm-stream", + Some(ScopeCategory::Start), + ); + let stream_runtime_mark = find_event( + &captured_events, + "fixture.worker.llm_stream_execution.runtime.mark", + None, + ); + assert_eq!(stream_runtime_mark.parent_uuid(), Some(stream_start.uuid())); + assert!(stream_runtime_mark.propagation_traceparent().is_some()); + assert_eq!( + stream_runtime_mark.metadata().unwrap()["name"], + "worker-fixture-llm-stream" + ); + deregister_subscriber("worker_plugin_fixture_events") + .expect("worker fixture subscriber should deregister"); loaded.clear(); } @@ -1676,7 +1707,20 @@ async fn python_worker_host_runtime_mark_and_mutated_request_round_trip() { flush_subscribers().expect("Python callback mark should flush"); let captured_events = events.lock().unwrap(); - find_event(&captured_events, "example.python_worker.tool_request", None); + let tool_start = find_event( + &captured_events, + "python-worker-tool", + Some(ScopeCategory::Start), + ); + let runtime_scope = find_event( + &captured_events, + "example.python_worker.request", + Some(ScopeCategory::Start), + ); + assert_eq!(runtime_scope.parent_uuid(), Some(tool_start.uuid())); + assert!(runtime_scope.propagation_traceparent().is_some()); + let runtime_mark = find_event(&captured_events, "example.python_worker.tool_request", None); + assert_eq!(runtime_mark.parent_uuid(), Some(runtime_scope.uuid())); let tool_mark = find_event( &captured_events, "example.python_worker.tool_execution", diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index a2e9b99d3..a4566c080 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -1377,12 +1377,14 @@ async fn callback_timeout_sends_explicit_worker_cancellation() { |_| Box::pin(tokio_stream::empty()), ) .await; - let request = callback.base_request( - "timeout", - RegistrationSurface::ToolRequestIntercept, - None, - Some(invoke_request_payload_tool("tool", json!({}), None)), - ); + let request = callback + .base_request( + "timeout", + RegistrationSurface::ToolRequestIntercept, + None, + Some(invoke_request_payload_tool("tool", json!({}), None)), + ) + .expect("worker invocation request should build"); let invocation_id = request.invocation_id.clone(); let callback_task = callback.clone(); @@ -1441,12 +1443,14 @@ async fn dropping_callback_future_cancels_worker_and_cleans_host_state() { Box::pin(async move { Ok(ToolExecutionResult::new(value)) }) }))) .expect("continuation should insert"); - let request = callback.base_request( - "cancel", - RegistrationSurface::ToolExecutionIntercept, - Some(continuation_id), - Some(invoke_request_payload_tool("tool", json!({}), None)), - ); + let request = callback + .base_request( + "cancel", + RegistrationSurface::ToolExecutionIntercept, + Some(continuation_id), + Some(invoke_request_payload_tool("tool", json!({}), None)), + ) + .expect("worker invocation request should build"); let scope_stack_id = request .scope .as_ref() @@ -1496,7 +1500,14 @@ async fn dropping_callback_future_cancels_worker_and_cleans_host_state() { ); let overlapping_scope_stack_id = callback .host_state - .insert_invocation_scope_stack(invocation_stack.clone(), None); + .insert_invocation_scope_stack(invocation_stack.clone(), None) + .expect("overlapping invocation scope stack should insert"); + let overlapping_stack = callback + .host_state + .stack(&overlapping_scope_stack_id) + .expect("overlapping scope stack lookup should succeed") + .expect("overlapping scope stack should exist"); + assert!(!Arc::ptr_eq(&invocation_stack, &overlapping_stack)); let invocation_id = request.invocation_id.clone(); let callback_task = callback.clone(); let task = tokio::spawn(async move { callback_task.invoke_async(request).await }); @@ -1545,6 +1556,14 @@ async fn dropping_callback_future_cancels_worker_and_cleans_host_state() { .expect("invocation scope stack lock") .scopes() .len(), + baseline_depth + ); + assert_eq!( + overlapping_stack + .read() + .expect("overlapping scope stack lock") + .scopes() + .len(), baseline_depth + 2 ); callback @@ -1574,6 +1593,14 @@ async fn dropping_callback_future_cancels_worker_and_cleans_host_state() { .len(), baseline_depth ); + assert_eq!( + overlapping_stack + .read() + .expect("overlapping scope stack lock") + .scopes() + .len(), + baseline_depth + 2 + ); callback .host_state .cleanup_invocation_scope_stack(&scope_stack_id); @@ -1598,9 +1625,19 @@ fn invocation_cleanup_releases_host_state_locks_before_unwinding() { AUTH_TOKEN.into(), )); let stack = crate::api::runtime::create_scope_stack(); - let baseline_depth = stack.read().expect("scope stack lock").scopes().len(); - let scope_stack_id = state.insert_invocation_scope_stack(stack.clone(), None); - with_scope_stack(stack.clone(), || { + let scope_stack_id = state + .insert_invocation_scope_stack(stack.clone(), None) + .expect("invocation scope stack should insert"); + let invocation_stack = state + .stack(&scope_stack_id) + .expect("invocation scope stack lookup should succeed") + .expect("invocation scope stack should exist"); + let baseline_depth = invocation_stack + .read() + .expect("scope stack lock") + .scopes() + .len(); + with_scope_stack(invocation_stack.clone(), || { push_scope( PushScopeParams::builder() .name("cleanup-lock-test") @@ -1610,7 +1647,7 @@ fn invocation_cleanup_releases_host_state_locks_before_unwinding() { }) .expect("worker scope should push"); - let stack_guard = stack.write().expect("scope stack lock"); + let stack_guard = invocation_stack.write().expect("scope stack lock"); let (done_tx, done_rx) = std::sync::mpsc::channel(); let cleanup_state = state.clone(); let cleanup = std::thread::spawn(move || { @@ -1624,7 +1661,7 @@ fn invocation_cleanup_releases_host_state_locks_before_unwinding() { .lock() .expect("scope cleanup lock") .iter() - .any(|handle| Arc::ptr_eq(handle, &stack)); + .any(|handle| Arc::ptr_eq(handle, &invocation_stack)); if cleanup_registered && state.scope_stacks.try_lock().is_ok() && state.pending_scope_cleanups.try_lock().is_ok() @@ -1652,7 +1689,11 @@ fn invocation_cleanup_releases_host_state_locks_before_unwinding() { .is_empty() ); assert_eq!( - stack.read().expect("scope stack lock").scopes().len(), + invocation_stack + .read() + .expect("scope stack lock") + .scopes() + .len(), baseline_depth ); } @@ -1711,6 +1752,28 @@ async fn dropping_host_stream_sends_explicit_worker_cancellation() { .expect("host should cancel abandoned stream") .expect("cancellation channel should remain open"); assert!(cancellation.reason.contains("stopped consuming")); + tokio::time::timeout(std::time::Duration::from_secs(1), async { + loop { + let continuations_empty = callback + .host_state + .continuations + .lock() + .expect("continuation lock") + .is_empty(); + let stacks_empty = callback + .host_state + .scope_stacks + .lock() + .expect("scope stack lock") + .is_empty(); + if continuations_empty && stacks_empty { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("abandoned stream host state should be cleaned up"); } #[tokio::test(flavor = "multi_thread")] @@ -2813,6 +2876,7 @@ async fn worker_continuations_use_the_scope_stack_selected_for_each_call() { StoredScopeStack { handle: stack, publication_buffer: None, + continuation_context: None, invocation_base_depth: None, }, ); @@ -2855,6 +2919,186 @@ async fn worker_continuations_use_the_scope_stack_selected_for_each_call() { assert_eq!(decode(second.unwrap()), json!(expected[1])); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn overlapping_worker_invocations_isolate_scope_stack_mutations() { + let state = Arc::new(WorkerHostRuntimeState::new( + ACTIVATION_ID.into(), + AUTH_TOKEN.into(), + )); + let service = WorkerHostRuntimeService { + state: state.clone(), + }; + let source_stack = crate::api::runtime::create_scope_stack(); + let first_stack_id = state + .insert_invocation_scope_stack(source_stack.clone(), None) + .expect("first invocation scope stack should insert"); + let second_stack_id = state + .insert_invocation_scope_stack(source_stack, None) + .expect("second invocation scope stack should insert"); + let first_stack = state + .stack(&first_stack_id) + .expect("first invocation scope stack lookup should succeed") + .expect("first invocation scope stack should exist"); + let second_stack = state + .stack(&second_stack_id) + .expect("second invocation scope stack lookup should succeed") + .expect("second invocation scope stack should exist"); + assert!(!Arc::ptr_eq(&first_stack, &second_stack)); + + let next_barrier = Arc::new(tokio::sync::Barrier::new(2)); + let continuation_id = state + .insert_continuation(Continuation::tool(Arc::new(move |_| { + let next_barrier = next_barrier.clone(); + Box::pin(async move { + next_barrier.wait().await; + tokio::task::yield_now().await; + Ok(ToolExecutionResult::new(json!( + crate::api::runtime::task_scope_top().uuid.to_string() + ))) + }) + }))) + .expect("tool continuation should insert"); + let service = Arc::new(service); + let invoke = |scope_stack_id: String, name: &'static str| { + let service = service.clone(); + let continuation_id = continuation_id.clone(); + tokio::spawn(async move { + let response = service + .push_scope(Request::new(PushScopeRequest { + activation_id: ACTIVATION_ID.into(), + auth_token: AUTH_TOKEN.into(), + scope: Some(ScopeContext { + scope_stack_id: scope_stack_id.clone(), + parent_scope_id: String::new(), + }), + name: name.into(), + scope_type: ProtoScopeType::Custom as i32, + data: None, + metadata: None, + input: None, + })) + .await + .expect("worker scope should push") + .into_inner(); + assert!(response.error.is_none(), "{:?}", response.error); + let expected = response + .scope_handle_id + .strip_prefix("scope-") + .expect("worker scope handle should use the documented prefix") + .to_owned(); + let response = service + .tool_next(Request::new(ToolNextRequest { + activation_id: ACTIVATION_ID.into(), + auth_token: AUTH_TOKEN.into(), + continuation_id, + value: Some(json_envelope(JSON_SCHEMA, &json!({})).expect("json envelope")), + scope: Some(ScopeContext { + scope_stack_id, + parent_scope_id: String::new(), + }), + })) + .await + .expect("worker continuation should run") + .into_inner(); + let value = response.value.expect("tool next should return a value"); + let observed = + decode_json_value::(value.result.as_ref().expect("tool next result")) + .expect("tool next result should decode"); + (expected, observed) + }) + }; + + let (first, second) = tokio::join!( + invoke(first_stack_id.clone(), "worker-overlap-first"), + invoke(second_stack_id.clone(), "worker-overlap-second"), + ); + let (first_expected, first_observed) = first.expect("first invocation task should finish"); + let (second_expected, second_observed) = second.expect("second invocation task should finish"); + assert_ne!(first_expected, second_expected); + assert_eq!(first_observed, json!(first_expected)); + assert_eq!(second_observed, json!(second_expected)); + + state.remove_continuation(&continuation_id); + state.cleanup_invocation_scope_stack(&first_stack_id); + state.cleanup_invocation_scope_stack(&second_stack_id); +} + +#[tokio::test] +async fn worker_runtime_scope_calls_restore_managed_parent_and_trace_context() { + let state = Arc::new(WorkerHostRuntimeState::new( + ACTIVATION_ID.into(), + AUTH_TOKEN.into(), + )); + let service = WorkerHostRuntimeService { + state: state.clone(), + }; + let source_stack = crate::api::runtime::create_scope_stack(); + let managed_event_uuid = uuid::Uuid::now_v7(); + let traceparent = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"; + let trace_context = crate::api::runtime::scope_stack::W3cTraceContext::new( + traceparent, + Some("vendor=value".into()), + ) + .expect("trace context should be valid"); + let scope_stack_id = crate::api::runtime::scope_stack::TASK_SCOPE_STACK + .scope( + source_stack.clone(), + crate::api::runtime::scope_stack::with_active_event_trace_context( + managed_event_uuid, + Some(trace_context), + async { + state + .insert_invocation_scope_stack(source_stack, None) + .expect("invocation scope stack should insert") + }, + ), + ) + .await; + let context = state + .invocation_context(&scope_stack_id) + .expect("invocation context lookup should succeed") + .expect("invocation context should exist"); + let observed_traceparent = context.run(|| { + crate::api::runtime::scope_stack::active_event_trace_context() + .map(|context| context.traceparent().to_owned()) + }); + assert_eq!(observed_traceparent.as_deref(), Some(traceparent)); + + let parent_response = service + .push_scope(Request::new(PushScopeRequest { + activation_id: ACTIVATION_ID.into(), + auth_token: AUTH_TOKEN.into(), + scope: Some(ScopeContext { + scope_stack_id: scope_stack_id.clone(), + parent_scope_id: String::new(), + }), + name: "worker-managed-parent".into(), + scope_type: ProtoScopeType::Custom as i32, + data: None, + metadata: None, + input: None, + })) + .await + .expect("managed-parent scope should push") + .into_inner(); + assert!( + parent_response.error.is_none(), + "{:?}", + parent_response.error + ); + let parent_handle = state + .scope_handles + .lock() + .expect("scope handle lock") + .get(&parent_response.scope_handle_id) + .expect("parent scope handle should be retained") + .handle + .clone(); + assert_eq!(parent_handle.parent_uuid, Some(managed_event_uuid)); + + state.cleanup_invocation_scope_stack(&scope_stack_id); +} + fn valid_llm_request() -> LlmRequest { LlmRequest { headers: serde_json::Map::new(), From c22b25a4e4a9d1aa8be4e54ba15dd361f9dd8adb Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Thu, 1 Oct 2026 13:05:15 -0700 Subject: [PATCH 6/7] fix(runtime): invalidate stale thread event context Signed-off-by: Alex Fournier --- .../src/api/runtime/continuation_context.rs | 5 +- crates/core/src/api/runtime/scope_stack.rs | 54 +++++++++------ crates/core/tests/unit/native_plugin_tests.rs | 69 +++++++++++++++++++ 3 files changed, 106 insertions(+), 22 deletions(-) diff --git a/crates/core/src/api/runtime/continuation_context.rs b/crates/core/src/api/runtime/continuation_context.rs index a0762fd7e..a316fd4fb 100644 --- a/crates/core/src/api/runtime/continuation_context.rs +++ b/crates/core/src/api/runtime/continuation_context.rs @@ -79,6 +79,7 @@ impl MiddlewareContinuationContext { Ok(Self { active_event: self .active_event + .clone() .map(|active_event| rebind_active_event_to_stack(active_event, &scope_stack)), scope_stack, active_event_trace_context: self.active_event_trace_context.clone(), @@ -107,7 +108,7 @@ impl MiddlewareContinuationContext { let previous = install_thread_continuation_context( &self.scope_stack, - self.active_event, + self.active_event.clone(), self.active_event_trace_context.clone(), ); let _restore = RestoreThreadContext(Some(previous)); @@ -138,7 +139,7 @@ impl MiddlewareContinuationContext { let published = with_task_nested_publication_buffer(self.publication_buffer.clone(), published); let active = async { - match self.active_event { + match self.active_event.clone() { Some(active_event) => { with_anchored_active_event( active_event, diff --git a/crates/core/src/api/runtime/scope_stack.rs b/crates/core/src/api/runtime/scope_stack.rs index 87078ac90..f8c2f6527 100644 --- a/crates/core/src/api/runtime/scope_stack.rs +++ b/crates/core/src/api/runtime/scope_stack.rs @@ -8,10 +8,10 @@ //! can use this module to inspect the active scope chain or propagate scope //! context into worker threads. -use std::cell::{Cell, RefCell}; +use std::cell::RefCell; use std::collections::{HashMap, HashSet}; use std::future::Future; -use std::sync::{Arc, RwLock}; +use std::sync::{Arc, RwLock, Weak}; use opentelemetry::propagation::TextMapPropagator; use opentelemetry::trace::{SpanContext, TraceContextExt, TraceFlags, TraceState}; @@ -647,11 +647,11 @@ impl Default for ScopeStack { /// concurrent readers. pub type ScopeStackHandle = Arc>; -#[derive(Clone, Copy)] +#[derive(Clone)] pub(crate) struct AnchoredActiveEvent { event_uuid: Uuid, - // Propagated stacks may share a root UUID, so identify the captured Arc allocation. - stack_identity: usize, + // Propagated stacks may share a root UUID, so retain the captured Arc allocation identity. + scope_stack: Weak>, anchor_scope_uuid: Uuid, } @@ -1075,10 +1075,10 @@ pub(crate) async fn with_active_event_trace_context( trace_context: Option, future: impl Future, ) -> T { - let (stack_identity, anchor_scope_uuid) = scope_stack_identity_and_anchor(); + let (scope_stack, anchor_scope_uuid) = scope_stack_identity_and_anchor(); let active_event = AnchoredActiveEvent { event_uuid: uuid, - stack_identity, + scope_stack, anchor_scope_uuid, }; with_anchored_active_event(active_event, trace_context, future).await @@ -1099,7 +1099,7 @@ pub(crate) async fn with_anchored_active_event( pub(crate) fn capture_anchored_active_event() -> Option { ACTIVE_EVENT - .try_with(|event| *event) + .try_with(Clone::clone) .ok() .or_else(thread_active_event) } @@ -1113,7 +1113,7 @@ pub(crate) fn rebind_active_event_to_stack( .unwrap_or_else(|error| error.into_inner()); AnchoredActiveEvent { event_uuid: active_event.event_uuid, - stack_identity: Arc::as_ptr(scope_stack) as usize, + scope_stack: Arc::downgrade(scope_stack), anchor_scope_uuid: if guard.find(&active_event.anchor_scope_uuid).is_some() { active_event.anchor_scope_uuid } else { @@ -1140,9 +1140,10 @@ pub(crate) fn active_event_trace_context() -> Option { } fn thread_active_event() -> Option { - let mut event = THREAD_ACTIVE_EVENT.with(Cell::get)?; + let mut event = THREAD_ACTIVE_EVENT.with(|active| active.borrow().clone())?; let stack = current_scope_stack(); - if event.stack_identity != Arc::as_ptr(&stack) as usize { + let event_stack = event.scope_stack.upgrade()?; + if !Arc::ptr_eq(&event_stack, &stack) { return None; } let guard = stack.read().unwrap_or_else(|error| error.into_inner()); @@ -1151,7 +1152,7 @@ fn thread_active_event() -> Option { return None; } event.anchor_scope_uuid = guard.top().uuid; - THREAD_ACTIVE_EVENT.with(|active| active.set(Some(event))); + THREAD_ACTIVE_EVENT.with(|active| *active.borrow_mut() = Some(event.clone())); } Some(event) } @@ -1169,7 +1170,7 @@ thread_local! { /// Whether the current thread explicitly owns a scope stack. static THREAD_SCOPE_STACK_EXPLICIT: std::cell::Cell = const { std::cell::Cell::new(false) }; /// Managed event propagated into a foreign executor with the thread scope binding. - static THREAD_ACTIVE_EVENT: Cell> = const { Cell::new(None) }; + static THREAD_ACTIVE_EVENT: RefCell> = const { RefCell::new(None) }; /// Exact W3C context associated with the propagated managed event. static THREAD_ACTIVE_EVENT_TRACE_CONTEXT: RefCell> = const { RefCell::new(None) }; } @@ -1241,6 +1242,7 @@ pub fn with_scope_stack(handle: ScopeStackHandle, f: impl FnOnce() -> T) -> T /// # Notes /// Use this when propagating an existing scope stack into worker threads. pub fn set_thread_scope_stack(handle: ScopeStackHandle) { + clear_thread_active_event_for_stack_change(&handle); THREAD_SCOPE_STACK.with(|stack| *stack.borrow_mut() = handle); THREAD_SCOPE_STACK_EXPLICIT.with(|flag| flag.set(true)); } @@ -1260,7 +1262,7 @@ pub fn capture_thread_scope_stack() -> ThreadScopeStackBinding { ThreadScopeStackBinding { stack, explicit, - active_event: THREAD_ACTIVE_EVENT.with(Cell::get), + active_event: THREAD_ACTIVE_EVENT.with(|event| event.borrow().clone()), active_event_trace_context: THREAD_ACTIVE_EVENT_TRACE_CONTEXT .with(|context| context.borrow().clone()), } @@ -1276,7 +1278,7 @@ pub fn capture_thread_scope_stack() -> ThreadScopeStackBinding { pub fn restore_thread_scope_stack(binding: ThreadScopeStackBinding) { THREAD_SCOPE_STACK.with(|stack| *stack.borrow_mut() = binding.stack); THREAD_SCOPE_STACK_EXPLICIT.with(|flag| flag.set(binding.explicit)); - THREAD_ACTIVE_EVENT.with(|event| event.set(binding.active_event)); + THREAD_ACTIVE_EVENT.with(|event| *event.borrow_mut() = binding.active_event); THREAD_ACTIVE_EVENT_TRACE_CONTEXT .with(|context| *context.borrow_mut() = binding.active_event_trace_context); } @@ -1290,7 +1292,8 @@ pub(crate) fn install_thread_continuation_context( let previous = capture_thread_scope_stack(); sync_thread_scope_stack(scope_stack.clone()); THREAD_ACTIVE_EVENT.with(|event| { - event.set(active_event.map(|event| rebind_active_event_to_stack(event, scope_stack))) + *event.borrow_mut() = + active_event.map(|event| rebind_active_event_to_stack(event, scope_stack)); }); THREAD_ACTIVE_EVENT_TRACE_CONTEXT .with(|context| *context.borrow_mut() = active_event_trace_context); @@ -1312,9 +1315,18 @@ pub(crate) fn install_thread_continuation_context( /// Python bindings use this to mirror `ContextVar` state into Rust without /// forcing `scope_stack_active()` to become `true` for the thread. pub fn sync_thread_scope_stack(handle: ScopeStackHandle) { + clear_thread_active_event_for_stack_change(&handle); THREAD_SCOPE_STACK.with(|stack| *stack.borrow_mut() = handle); } +fn clear_thread_active_event_for_stack_change(handle: &ScopeStackHandle) { + let stack_changed = THREAD_SCOPE_STACK.with(|current| !Arc::ptr_eq(¤t.borrow(), handle)); + if stack_changed { + THREAD_ACTIVE_EVENT.with(|event| *event.borrow_mut() = None); + THREAD_ACTIVE_EVENT_TRACE_CONTEXT.with(|context| *context.borrow_mut() = None); + } +} + /// Synchronize the task-local managed event onto an isolated thread stack. /// /// Native async callbacks run on a plugin-owned executor. The host snapshots @@ -1324,15 +1336,17 @@ pub fn sync_thread_scope_stack(handle: ScopeStackHandle) { pub(crate) fn sync_thread_active_event_for_stack(scope_stack: &ScopeStackHandle) { let active_event = capture_anchored_active_event() .map(|active_event| rebind_active_event_to_stack(active_event, scope_stack)); - let trace_context = active_event.and_then(|_| active_event_trace_context()); - THREAD_ACTIVE_EVENT.with(|event| event.set(active_event)); + let trace_context = active_event + .as_ref() + .and_then(|_| active_event_trace_context()); + THREAD_ACTIVE_EVENT.with(|event| *event.borrow_mut() = active_event); THREAD_ACTIVE_EVENT_TRACE_CONTEXT.with(|context| *context.borrow_mut() = trace_context); } -fn scope_stack_identity_and_anchor() -> (usize, Uuid) { +fn scope_stack_identity_and_anchor() -> (Weak>, Uuid) { let stack = current_scope_stack(); let guard = stack.read().unwrap_or_else(|error| error.into_inner()); - (Arc::as_ptr(&stack) as usize, guard.top().uuid) + (Arc::downgrade(&stack), guard.top().uuid) } /// Report whether the current context has an explicitly active scope stack. diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index 176d061d7..a418c277e 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -2490,6 +2490,75 @@ fn thread_active_event_applies_only_at_the_captured_stack_top() { assert_parent(None, same_lineage_top_uuid); } +#[test] +fn thread_stack_setters_clear_context_only_when_the_stack_allocation_changes() { + let _restore = ThreadScopeStackRestore::capture(); + let runtime = Runtime::new().unwrap(); + let traceparent = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"; + let install_event = |event_uuid| { + let trace_context = crate::api::runtime::scope_stack::W3cTraceContext::new( + traceparent, + Some("vendor=value".into()), + ) + .expect("trace context should be valid"); + runtime.block_on( + crate::api::runtime::scope_stack::with_active_event_trace_context( + event_uuid, + Some(trace_context), + async { + sync_thread_active_event_for_stack(¤t_scope_stack()); + }, + ), + ); + }; + let assert_event = |expected| { + assert_eq!(active_event_uuid(), expected); + assert_eq!( + crate::api::runtime::scope_stack::active_event_trace_context() + .map(|context| context.traceparent().to_owned()), + expected.map(|_| traceparent.to_string()) + ); + }; + + let explicit_stack = create_scope_stack(); + set_thread_scope_stack(explicit_stack.clone()); + let explicit_event = uuid::Uuid::now_v7(); + install_event(explicit_event); + set_thread_scope_stack(explicit_stack); + assert_event(Some(explicit_event)); + + let synchronized_stack = create_scope_stack(); + let synchronized_top = synchronized_stack + .read() + .unwrap_or_else(|error| error.into_inner()) + .top() + .uuid; + set_thread_scope_stack(synchronized_stack.clone()); + assert_event(None); + assert_eq!( + crate::api::shared::resolve_parent_uuid(None), + Some(synchronized_top) + ); + + let synchronized_event = uuid::Uuid::now_v7(); + install_event(synchronized_event); + sync_thread_scope_stack(synchronized_stack); + assert_event(Some(synchronized_event)); + + let replacement = create_scope_stack(); + let replacement_top = replacement + .read() + .unwrap_or_else(|error| error.into_inner()) + .top() + .uuid; + sync_thread_scope_stack(replacement); + assert_event(None); + assert_eq!( + crate::api::shared::resolve_parent_uuid(None), + Some(replacement_top) + ); +} + #[test] fn restored_thread_active_event_rebases_after_its_stack_anchor_closes() { let _restore = ThreadScopeStackRestore::capture(); From e2484dc608846873fe63aa8cc2089de1320b38ee Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Thu, 1 Oct 2026 13:32:23 -0700 Subject: [PATCH 7/7] test(worker): include scope timestamp field Signed-off-by: Alex Fournier --- crates/core/tests/unit/dynamic_worker_tests.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index 3911d611a..e25f7e514 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -3472,6 +3472,7 @@ async fn overlapping_worker_invocations_isolate_scope_stack_mutations() { data: None, metadata: None, input: None, + timestamp_unix_micros: None, })) .await .expect("worker scope should push") @@ -3573,6 +3574,7 @@ async fn worker_runtime_scope_calls_restore_managed_parent_and_trace_context() { data: None, metadata: None, input: None, + timestamp_unix_micros: None, })) .await .expect("managed-parent scope should push")