diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index 72ff74b1a..212b3f8c4 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -26,7 +26,6 @@ pub use continuation_context::MiddlewareContinuationContext; pub(crate) use continuation_context::MiddlewareContinuationLease; pub use global::global_context; pub use llm_execution_context::LlmExecutionContext; -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, @@ -36,6 +35,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_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 d08dead95..a316fd4fb 100644 --- a/crates/core/src/api/runtime/continuation_context.rs +++ b/crates/core/src/api/runtime/continuation_context.rs @@ -9,15 +9,24 @@ use crate::api::optimization::{ LlmOptimizationRecorder, current_llm_optimization_recorder, scope_llm_optimization_recorder, }; use crate::api::runtime::scope_stack::{ - ScopeStackHandle, TASK_SCOPE_STACK, W3cTraceContext, active_event_trace_context, - active_event_uuid, current_context_scope_stack, current_scope_stack, scope_stack_active, - snapshot_scope_stack, with_active_event_trace_context, + AnchoredActiveEvent, ScopeStackHandle, TASK_SCOPE_STACK, W3cTraceContext, + active_event_trace_context, 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, +}; +#[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. @@ -28,7 +37,7 @@ use crate::error::{FlowError, Result}; #[derive(Clone)] pub struct MiddlewareContinuationContext { scope_stack: ScopeStackHandle, - active_event_uuid: Option, + active_event: Option, active_event_trace_context: Option, publication_context: Option, publication_buffer: Option, @@ -42,7 +51,7 @@ impl MiddlewareContinuationContext { pub fn capture() -> Self { Self { scope_stack: current_scope_stack(), - active_event_uuid: active_event_uuid(), + active_event: capture_anchored_active_event(), active_event_trace_context: active_event_trace_context(), publication_context: capture_publication_context(), publication_buffer: capture_nested_publication_buffer(), @@ -66,9 +75,13 @@ 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 + .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(), publication_context: self.publication_context.clone(), publication_buffer: self.publication_buffer.clone(), @@ -76,6 +89,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.clone(), + 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(|| { @@ -96,10 +139,10 @@ 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_trace_context( - uuid, + match self.active_event.clone() { + Some(active_event) => { + with_anchored_active_event( + active_event, self.active_event_trace_context.clone(), published, ) diff --git a/crates/core/src/api/runtime/scope_stack.rs b/crates/core/src/api/runtime/scope_stack.rs index 21f364e84..f8c2f6527 100644 --- a/crates/core/src/api/runtime/scope_stack.rs +++ b/crates/core/src/api/runtime/scope_stack.rs @@ -11,7 +11,7 @@ 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,14 +647,25 @@ impl Default for ScopeStack { /// concurrent readers. pub type ScopeStackHandle = Arc>; +#[derive(Clone)] +pub(crate) struct AnchoredActiveEvent { + event_uuid: Uuid, + // Propagated stacks may share a root UUID, so retain the captured Arc allocation identity. + scope_stack: Weak>, + anchor_scope_uuid: 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 context bound at +/// capture time. #[derive(Clone)] pub struct ThreadScopeStackBinding { stack: ScopeStackHandle, explicit: bool, + active_event: Option, + active_event_trace_context: Option, } impl ThreadScopeStackBinding { @@ -1049,7 +1060,7 @@ 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: AnchoredActiveEvent; /// Exact W3C context of the managed event when one was captured at start. static ACTIVE_EVENT_TRACE_CONTEXT: Option; } @@ -1064,23 +1075,90 @@ pub(crate) async fn with_active_event_trace_context( trace_context: Option, future: impl Future, ) -> T { - ACTIVE_EVENT_UUID + let (scope_stack, anchor_scope_uuid) = scope_stack_identity_and_anchor(); + let active_event = AnchoredActiveEvent { + event_uuid: uuid, + scope_stack, + anchor_scope_uuid, + }; + with_anchored_active_event(active_event, trace_context, future).await +} + +pub(crate) async fn with_anchored_active_event( + active_event: AnchoredActiveEvent, + trace_context: Option, + future: impl Future, +) -> T { + ACTIVE_EVENT .scope( - uuid, + active_event, ACTIVE_EVENT_TRACE_CONTEXT.scope(trace_context, future), ) .await } +pub(crate) fn capture_anchored_active_event() -> Option { + ACTIVE_EVENT + .try_with(Clone::clone) + .ok() + .or_else(thread_active_event) +} + +pub(crate) fn rebind_active_event_to_stack( + active_event: AnchoredActiveEvent, + scope_stack: &ScopeStackHandle, +) -> AnchoredActiveEvent { + let guard = scope_stack + .read() + .unwrap_or_else(|error| error.into_inner()); + AnchoredActiveEvent { + event_uuid: active_event.event_uuid, + 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 { + guard.top().uuid + }, + } +} + pub(crate) fn active_event_uuid() -> Option { - ACTIVE_EVENT_UUID.try_with(|uuid| *uuid).ok() + ACTIVE_EVENT + .try_with(|event| event.event_uuid) + .ok() + .or_else(thread_active_event_uuid) } pub(crate) fn active_event_trace_context() -> Option { - ACTIVE_EVENT_TRACE_CONTEXT - .try_with(Clone::clone) - .ok() - .flatten() + match ACTIVE_EVENT_TRACE_CONTEXT.try_with(Clone::clone) { + Ok(context) => context, + Err(_) => { + thread_active_event()?; + THREAD_ACTIVE_EVENT_TRACE_CONTEXT.with(|context| context.borrow().clone()) + } + } +} + +fn thread_active_event() -> Option { + let mut event = THREAD_ACTIVE_EVENT.with(|active| active.borrow().clone())?; + let stack = current_scope_stack(); + 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()); + if event.anchor_scope_uuid != guard.top().uuid { + if guard.find(&event.anchor_scope_uuid).is_some() { + return None; + } + event.anchor_scope_uuid = guard.top().uuid; + THREAD_ACTIVE_EVENT.with(|active| *active.borrow_mut() = Some(event.clone())); + } + Some(event) +} + +pub(crate) fn thread_active_event_uuid() -> Option { + thread_active_event().map(|event| event.event_uuid) } thread_local! { @@ -1091,6 +1169,10 @@ 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: 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) }; } /// Return the scope stack visible to the current execution context. @@ -1160,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)); } @@ -1171,12 +1254,18 @@ 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 context. 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(|event| event.borrow().clone()), + active_event_trace_context: THREAD_ACTIVE_EVENT_TRACE_CONTEXT + .with(|context| context.borrow().clone()), + } } /// Restore a previously captured thread-local scope stack binding. @@ -1189,6 +1278,26 @@ 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.borrow_mut() = binding.active_event); + THREAD_ACTIVE_EVENT_TRACE_CONTEXT + .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.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); + previous } /// Synchronize the thread-local scope stack without marking it explicit. @@ -1206,9 +1315,40 @@ pub fn restore_thread_scope_stack(binding: ThreadScopeStackBinding) { /// 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 +/// the callback's visible stack before crossing that boundary, so the managed +/// event must be rebound to the snapshot's allocation while retaining its +/// original scope anchor and W3C context. +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 + .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() -> (Weak>, Uuid) { + let stack = current_scope_stack(); + let guard = stack.read().unwrap_or_else(|error| error.into_inner()); + (Arc::downgrade(&stack), 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 ccda2f2e1..41c145378 100644 --- a/crates/core/src/api/shared.rs +++ b/crates/core/src/api/shared.rs @@ -9,7 +9,9 @@ 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::{W3cTraceContext, trace_context_for_managed_span}; +use crate::api::runtime::scope_stack::{ + W3cTraceContext, thread_active_event_uuid, trace_context_for_managed_span, +}; use crate::api::runtime::{ EventSanitizeFn, EventSubscriberFn, NemoRelayContextState, ScopeStackHandle, }; @@ -35,6 +37,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 7617087eb..e2715bba2 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, LlmExecutionContext, LlmExecutionFn, LlmExecutionNextFn, @@ -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_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, @@ -2184,6 +2185,7 @@ async fn invoke_native_async_callback_inner( } 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; @@ -2225,7 +2227,7 @@ async fn invoke_native_async_callback_inner( } }; 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, @@ -2236,38 +2238,41 @@ async fn invoke_native_async_callback_inner( )) 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()); - let state = catch_unwind(AssertUnwindSafe(|| match callback { - NativeAsyncCallback::Middleware(callback) => unsafe { - callback( - user_data.ptr, - invocation as *const NemoRelayNativeString, - next_ref - .map(|next| next as *const NemoRelayNativeAsyncNext) - .unwrap_or(ptr::null()), - completion_ref as *const NemoRelayNativeAsyncCompletion, - ) - }, - NativeAsyncCallback::LlmExecution { callback, .. } => native_context - .as_ref() - .expect("LLM execution callbacks always build a native context") - .with_native_context(|context| unsafe { + 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(|| match callback { + NativeAsyncCallback::Middleware(callback) => unsafe { callback( user_data.ptr, invocation as *const NemoRelayNativeString, - std::ptr::from_ref(&context), next_ref .map(|next| next as *const NemoRelayNativeAsyncNext) .unwrap_or(ptr::null()), completion_ref as *const NemoRelayNativeAsyncCompletion, ) - }), - })); + }, + NativeAsyncCallback::LlmExecution { callback, .. } => native_context + .as_ref() + .expect("LLM execution callbacks always build a native context") + .with_native_context(|context| unsafe { + callback( + user_data.ptr, + invocation as *const NemoRelayNativeString, + std::ptr::from_ref(&context), + next_ref + .map(|next| next as *const NemoRelayNativeAsyncNext) + .unwrap_or(ptr::null()), + completion_ref as *const NemoRelayNativeAsyncCompletion, + ) + }), + })) + }); restore_thread_scope_stack(previous_thread_stack); (state, completion_ref) }; @@ -4085,26 +4090,32 @@ 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()); - let state = catch_unwind(AssertUnwindSafe(|| unsafe { - native_context.with_native_context(|context| { - cb( - user_data.ptr, - invocation, - std::ptr::from_ref(&context), - 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 { + native_context.with_native_context(|context| { + cb( + user_data.ptr, + invocation, + std::ptr::from_ref(&context), + 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/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index e3dc41ae9..0b48e9564 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -1780,7 +1780,7 @@ impl WorkerPluginCallback { registration_name: registration_name.into(), }, )), - ); + )?; guardrail_from_invoke_response(self.invoke_blocking(request)?) } @@ -1790,7 +1790,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(()), @@ -1811,7 +1811,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!( @@ -1832,7 +1832,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!( @@ -1854,7 +1854,7 @@ impl WorkerPluginCallback { surface, continuation_id, Some(invoke_request_payload_tool(tool_name, value, None)), - ); + )?; json_from_invoke_response(self.invoke_async(request).await?) } @@ -1869,7 +1869,7 @@ impl WorkerPluginCallback { RegistrationSurface::ToolConditionalExecutionGuardrail, None, Some(invoke_request_payload_tool(tool_name, value, None)), - ); + )?; guardrail_from_invoke_response(self.invoke_async(request).await?) } @@ -1889,7 +1889,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)) => { @@ -1931,7 +1931,7 @@ impl WorkerPluginCallback { }, ), )), - ); + )?; let capability = context .resolve_codec() .map(|codec| -> FlowResult { @@ -1983,7 +1983,7 @@ impl WorkerPluginCallback { }, ), )), - ); + )?; let capability = context .resolve_codec() .map(|codec| -> FlowResult { @@ -2018,7 +2018,7 @@ impl WorkerPluginCallback { RegistrationSurface::LlmConditionalExecutionGuardrail, None, Some(invoke_request_payload_llm("", Some(request), None, None)), - ); + )?; guardrail_from_invoke_response(self.invoke_async(invoke).await?) } @@ -2039,7 +2039,7 @@ impl WorkerPluginCallback { annotated, None, )), - ); + )?; let response = self.invoke_async(invoke).await?; match response.result { Some(invoke_response_result::Result::LlmRequest(result)) => { @@ -2084,7 +2084,7 @@ impl WorkerPluginCallback { None, None, )), - ); + )?; let codec_capabilities = self.attach_llm_execution_codec_context( &mut invoke, &execution_context, @@ -2115,7 +2115,7 @@ impl WorkerPluginCallback { None, None, )), - ); + )?; let codec_capabilities = self.attach_llm_execution_codec_context( &mut invoke, &execution_context, @@ -2278,12 +2278,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(), @@ -2295,7 +2303,7 @@ impl WorkerPluginCallback { parent_scope_id: String::new(), }), payload, - } + }) } fn invoke_blocking(&self, request: InvokeRequest) -> FlowResult { @@ -2553,6 +2561,7 @@ enum WorkerCodecDirection { struct StoredScopeStack { handle: crate::api::runtime::ScopeStackHandle, publication_buffer: Option, + continuation_context: Option, invocation_base_depth: Option, } @@ -2815,33 +2824,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( @@ -2849,20 +2868,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); @@ -2872,7 +2897,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; }; @@ -2900,11 +2930,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; @@ -2927,6 +2967,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) + }), } } @@ -2980,6 +3026,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")) @@ -2990,6 +3037,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)] @@ -3421,9 +3480,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() }; @@ -3447,6 +3504,7 @@ impl RelayHostRuntime for WorkerHostRuntimeService { StoredScopeStack { handle: crate::api::runtime::create_scope_stack(), publication_buffer: None, + continuation_context: None, invocation_base_depth: None, }, ); @@ -3766,9 +3824,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/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index fd1385191..92be56a2e 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -28,6 +28,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)] @@ -209,6 +228,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) @@ -266,6 +302,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", @@ -356,6 +422,48 @@ impl NativePlugin for FixtureNativePlugin { Ok(stream) }, )?; + ctx.register_llm_execution_intercept("fixture_llm_execution_cancellation", -1, { + let runtime = runtime.clone(); + move |name, request, _context, next| { + let runtime = runtime.clone(); + async move { + if name != "native-fixture-cancelled-unary" { + return next.call(request).await; + } + 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_cancellation", + -1, + { + let runtime = runtime.clone(); + move |name, request, _context, 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(()) } diff --git a/crates/core/tests/fixtures/worker_plugin/src/main.rs b/crates/core/tests/fixtures/worker_plugin/src/main.rs index 0de5ae310..d1903a89b 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,21 @@ fn register_fixture_llm_hooks( ) }, ); + let unary_runtime = runtime.clone(); ctx.register_llm_execution_intercept( "fixture_llm_execution", 0, - |_name, request, _context, next: LlmNext| async move { + move |name, request, _context, 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,27 +396,40 @@ 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, _context, next: LlmStreamNext| async move { - if llm_stream_open_error { - return Err(WorkerSdkError::Callback( - "fixture LLM stream open error requested".into(), - )); + move |name, request, _context, 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, + "worker_plugin_llm_stream_execution_request", + )) + .await?; + let mapped: JsonStream = Box::pin(tokio_stream::StreamExt::map(stream, |chunk| { + chunk.map(|value| mark_json(value, "worker_plugin_llm_stream_execution")) + })); + Ok(mapped) } - let stream = next - .call(mark_llm_request( - request, - "worker_plugin_llm_stream_execution_request", - )) - .await?; - let mapped: JsonStream = Box::pin(tokio_stream::StreamExt::map(stream, |chunk| { - chunk.map(|value| mark_json(value, "worker_plugin_llm_stream_execution")) - })); - Ok(mapped) }, ); } diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index 84c6ffffc..eeccc5f19 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(); @@ -475,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(); @@ -585,6 +646,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 +763,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/integration/worker_plugin_tests.rs b/crates/core/tests/integration/worker_plugin_tests.rs index 78a89096d..48da836da 100644 --- a/crates/core/tests/integration/worker_plugin_tests.rs +++ b/crates/core/tests/integration/worker_plugin_tests.rs @@ -705,6 +705,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", @@ -749,7 +760,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(); } @@ -1691,7 +1722,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/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/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index 98efc9051..e25f7e514 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -1619,12 +1619,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(); @@ -1683,12 +1685,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() @@ -1739,7 +1743,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 }); @@ -1788,6 +1799,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 @@ -1817,6 +1836,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); @@ -1841,9 +1868,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") @@ -1853,7 +1890,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 || { @@ -1867,7 +1904,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() @@ -1895,7 +1932,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 ); } @@ -3331,6 +3372,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, }, ); @@ -3373,6 +3415,188 @@ 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, + timestamp_unix_micros: 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, + timestamp_unix_micros: 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(), diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index 1bd373b76..a418c277e 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; @@ -2440,6 +2441,165 @@ 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_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_for_stack(¤t_scope_stack()); + })); + + 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 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(); + 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_for_stack(¤t_scope_stack()); + 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() @@ -5501,6 +5661,403 @@ 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, + NativeAsyncCodecCapabilities::default(), + ), + ), + )); + 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, + NativeAsyncCodecCapabilities::default(), + ), + invoke_native_async_callback( + interleave_native_callback_scopes, + user_data, + json!({"branch": 1}), + None, + NativeAsyncCodecCapabilities::default(), + ) + ) + }), + ) + .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, + _context: *const NemoRelayNativeLlmExecutionContext, + 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}), + }, + LlmExecutionContext::default(), + downstream.clone(), + ), + wrapped( + "second", + LlmRequest { + headers: Map::new(), + content: json!({"branch": 1}), + }, + LlmExecutionContext::default(), + 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() { diff --git a/crates/plugin/src/async_sdk.rs b/crates/plugin/src/async_sdk.rs index 4d7c75fad..397dbb372 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; @@ -687,20 +688,24 @@ unsafe fn unary_trampoline_impl( let future = catch_unwind(AssertUnwindSafe(|| match (invocation, context) { (Ok(invocation), Ok(context)) => (state.adapter)(invocation, context, next, completion_ref), (Err(error), _) | (_, 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; }) { @@ -713,17 +718,9 @@ unsafe fn unary_trampoline_impl( 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())) @@ -808,6 +805,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> { @@ -852,13 +867,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, + } } } @@ -873,20 +891,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, + } } } @@ -901,12 +928,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, @@ -1097,20 +1130,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.v6.v5.v4.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.v6.v5.v4.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,