Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion crates/core/src/api/runtime.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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;
Expand Down
65 changes: 54 additions & 11 deletions crates/core/src/api/runtime/continuation_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -28,7 +37,7 @@ use crate::error::{FlowError, Result};
#[derive(Clone)]
pub struct MiddlewareContinuationContext {
scope_stack: ScopeStackHandle,
active_event_uuid: Option<uuid::Uuid>,
active_event: Option<AnchoredActiveEvent>,
active_event_trace_context: Option<W3cTraceContext>,
publication_context: Option<PublicationContext>,
publication_buffer: Option<PublicationBuffer>,
Expand All @@ -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(),
Expand All @@ -66,16 +75,50 @@ impl MiddlewareContinuationContext {
/// from the stack captured when the middleware callback began.
#[doc(hidden)]
pub fn isolated_with_scope_stack(&self, scope_stack: &ScopeStackHandle) -> Result<Self> {
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(),
optimization_recorder: self.optimization_recorder.clone(),
})
}

#[cfg(feature = "worker-grpc")]
pub(crate) fn scope_stack(&self) -> ScopeStackHandle {
self.scope_stack.clone()
}

#[cfg(feature = "worker-grpc")]
pub(crate) fn run_sync<T>(&self, callback: impl FnOnce() -> T) -> T {
struct RestoreThreadContext(Option<crate::api::runtime::ThreadScopeStackBinding>);

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<Self> {
let visible_scope_stack = current_context_scope_stack().unwrap_or_else(|| {
Expand All @@ -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,
)
Expand Down
168 changes: 154 additions & 14 deletions crates/core/src/api/runtime/scope_stack.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -647,14 +647,25 @@ impl Default for ScopeStack {
/// concurrent readers.
pub type ScopeStackHandle = Arc<RwLock<ScopeStack>>;

#[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<RwLock<ScopeStack>>,
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<AnchoredActiveEvent>,
active_event_trace_context: Option<W3cTraceContext>,
}

impl ThreadScopeStackBinding {
Expand Down Expand Up @@ -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<W3cTraceContext>;
}
Expand All @@ -1064,23 +1075,90 @@ pub(crate) async fn with_active_event_trace_context<T>(
trace_context: Option<W3cTraceContext>,
future: impl Future<Output = T>,
) -> 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<T>(
active_event: AnchoredActiveEvent,
trace_context: Option<W3cTraceContext>,
future: impl Future<Output = T>,
) -> T {
ACTIVE_EVENT
.scope(
uuid,
active_event,
ACTIVE_EVENT_TRACE_CONTEXT.scope(trace_context, future),
)
.await
}

pub(crate) fn capture_anchored_active_event() -> Option<AnchoredActiveEvent> {
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<Uuid> {
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<W3cTraceContext> {
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<AnchoredActiveEvent> {
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)
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.

pub(crate) fn thread_active_event_uuid() -> Option<Uuid> {
thread_active_event().map(|event| event.event_uuid)
}

thread_local! {
Expand All @@ -1091,6 +1169,10 @@ thread_local! {
static THREAD_SCOPE_STACK: RefCell<ScopeStackHandle> = RefCell::new(create_scope_stack());
/// Whether the current thread explicitly owns a scope stack.
static THREAD_SCOPE_STACK_EXPLICIT: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
/// Managed event propagated into a foreign executor with the thread scope binding.
static THREAD_ACTIVE_EVENT: RefCell<Option<AnchoredActiveEvent>> = const { RefCell::new(None) };
/// Exact W3C context associated with the propagated managed event.
static THREAD_ACTIVE_EVENT_TRACE_CONTEXT: RefCell<Option<W3cTraceContext>> = const { RefCell::new(None) };
}

/// Return the scope stack visible to the current execution context.
Expand Down Expand Up @@ -1160,6 +1242,7 @@ pub fn with_scope_stack<T>(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));
}
Expand All @@ -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.
Expand All @@ -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<AnchoredActiveEvent>,
active_event_trace_context: Option<W3cTraceContext>,
) -> 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.
Expand All @@ -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(&current.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<RwLock<ScopeStack>>, 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
Expand Down
Loading
Loading