From 229944337cd6e286762902b8c46615333cae661c Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 13:31:57 -0400 Subject: [PATCH 01/52] feat: queue scope and mark event publication Signed-off-by: Will Killian --- .../src/api/runtime/subscriber_dispatcher.rs | 69 ++++++++++++++++++- crates/core/src/api/scope.rs | 59 ++++++++++------ crates/core/src/api/shared.rs | 27 ++++++-- .../subscriber_dispatcher_tests.rs | 65 +++++++++++++++++ docs/reference/event-sanitizers.mdx | 13 ++++ 5 files changed, 205 insertions(+), 28 deletions(-) diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index ffb4ef3b0..4ac440724 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -4,7 +4,10 @@ //! Asynchronous subscriber delivery for native targets. use crate::api::event::Event; -use crate::api::runtime::EventSubscriberFn; +use crate::api::registry::Guardrail; +use crate::api::runtime::{ + EventSanitizeFn, EventSubscriberFn, NemoRelayContextState, ScopeStackHandle, +}; use crate::error::Result; mod native { @@ -24,6 +27,7 @@ mod native { enum DispatcherMessage { Deliver { event: Box, + sanitizers: Vec>, subscribers: Vec, scope_stack: ScopeStackHandle, }, @@ -46,6 +50,7 @@ mod native { } let message = DispatcherMessage::Deliver { event: Box::new(event.clone()), + sanitizers: Vec::new(), subscribers: subscribers.to_vec(), scope_stack: current_scope_stack(), }; @@ -76,6 +81,41 @@ mod native { } } + pub(super) fn dispatch_sanitized_event( + event: Event, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, + ) -> bool { + let message = DispatcherMessage::Deliver { + event: Box::new(event), + sanitizers, + subscribers: subscribers.to_vec(), + scope_stack, + }; + match dispatcher_sender() { + Ok(sender) if sender.send(message).is_ok() => true, + Ok(_) => { + log::warn!( + target: "nemo_relay.runtime", + event = "subscriber_event_dropped", + reason = "dispatcher_disconnected"; + "Subscriber event was dropped because the dispatcher stopped" + ); + false + } + Err(error) if !DISPATCHER_FAILURE_LOGGED.swap(true, Ordering::AcqRel) => { + log::error!( + target: "nemo_relay.runtime", + event = "subscriber_dispatcher_failed"; + "Subscriber dispatcher failed to start: {error}" + ); + false + } + Err(_) => false, + } + } + pub(super) fn flush_subscribers() -> Result<()> { if IN_DISPATCHER.with(Cell::get) { return Ok(()); @@ -149,9 +189,10 @@ mod native { match message { DispatcherMessage::Deliver { event, + sanitizers, subscribers, scope_stack, - } => deliver_event(event, subscribers, scope_stack), + } => deliver_event(event, sanitizers, subscribers, scope_stack), DispatcherMessage::Flush { done } => { let _ = done.send(()); } @@ -160,12 +201,25 @@ mod native { fn deliver_event( event: Box, + sanitizers: Vec>, subscribers: Vec, scope_stack: ScopeStackHandle, ) { let previous_scope_stack = capture_thread_scope_stack(); set_thread_scope_stack(scope_stack); IN_DISPATCHER.with(|flag| flag.set(true)); + let original = (*event).clone(); + let event = catch_unwind(AssertUnwindSafe(|| { + NemoRelayContextState::event_sanitize_snapshot_chain(*event, &sanitizers) + })) + .unwrap_or_else(|_| { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_panicked"; + "Event sanitizer panicked; publishing the original event snapshot" + ); + original + }); for subscriber in subscribers { if catch_unwind(AssertUnwindSafe(|| subscriber(&event))).is_err() { log::error!( @@ -185,6 +239,17 @@ pub(crate) fn dispatch_event(event: &Event, subscribers: &[EventSubscriberFn]) - native::dispatch_event(event, subscribers) } +/// Queue a snapshot for serial event sanitization followed by subscriber +/// delivery. Used by synchronous scope and mark APIs. +pub(crate) fn dispatch_sanitized_event( + event: Event, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, +) -> bool { + native::dispatch_sanitized_event(event, sanitizers, subscribers, scope_stack) +} + /// Wait for all queued subscriber callbacks submitted before this call. pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() diff --git a/crates/core/src/api/scope.rs b/crates/core/src/api/scope.rs index 60a1aa53d..8cc0297e1 100644 --- a/crates/core/src/api/scope.rs +++ b/crates/core/src/api/scope.rs @@ -2,13 +2,14 @@ // SPDX-License-Identifier: Apache-2.0 use crate::api::event::{BaseEvent, CategoryProfile, DataSchema, EventCategory, MarkEvent}; -use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::global_context; +use crate::api::runtime::subscriber_dispatcher; use crate::api::runtime::{ current_scope_stack, task_scope_push, task_scope_remove, task_scope_top, }; use crate::api::shared::{ - ensure_runtime_owner, resolve_parent_uuid, sanitize_event, snapshot_event_subscribers, + ensure_runtime_owner, resolve_parent_uuid, snapshot_event_sanitizers, + snapshot_event_subscribers, }; use crate::error::{FlowError, Result}; use crate::json::Json; @@ -216,12 +217,13 @@ pub fn get_handle() -> Result { /// cannot be read safely. /// /// # Notes -/// Scope-local subscribers attached to ancestor scopes observe the emitted -/// start event before the function returns. +/// The event and its visible middleware/subscriber chains are snapshotted +/// before this function returns. Sanitization and subscriber delivery happen +/// later on the serial publication dispatcher. pub fn push_scope(params: PushScopeParams<'_>) -> Result { ensure_runtime_owner()?; let parent_uuid = resolve_parent_uuid(params.parent); - let (handle, event, subscribers) = { + let (handle, event, subscribers, emission_scope_stack) = { let scope_stack = current_scope_stack(); let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); let scope_subscribers = scope_guard.collect_scope_local_subscribers(); @@ -241,12 +243,16 @@ pub fn push_scope(params: PushScopeParams<'_>) -> Result { .build(); let handle = state.create_scope_handle(handle_params); let event = state.build_scope_start_event(&handle, params.input); - (handle, event, subscribers) + (handle, event, subscribers, scope_stack.clone()) }; - let event = sanitize_event(event); task_scope_push(handle.clone()); - if let Some(event) = event { - NemoRelayContextState::emit_event(&event, &subscribers); + if let Some(sanitizers) = snapshot_event_sanitizers(&event, &emission_scope_stack) { + let _ = subscriber_dispatcher::dispatch_sanitized_event( + event, + sanitizers, + &subscribers, + emission_scope_stack, + ); } Ok(handle) } @@ -276,7 +282,7 @@ pub fn push_scope(params: PushScopeParams<'_>) -> Result { pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { ensure_runtime_owner()?; let scope_stack = current_scope_stack(); - let (scope, event, subscribers) = { + let (scope, event, subscribers, emission_scope_stack) = { let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); let top = scope_guard.top(); if top.uuid != *params.handle_uuid { @@ -302,13 +308,20 @@ pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { .metadata_opt(params.metadata) .build(), ); - (scope, event, subscribers) + (scope, event, subscribers, scope_stack.clone()) }; - let event = sanitize_event(event); + // Snapshot scope-local middleware before removing its owner. Publication + // happens later, but cleanup must not change the chain visible at emission. + let sanitizers = snapshot_event_sanitizers(&event, &emission_scope_stack); let removed = task_scope_remove(params.handle_uuid)?; debug_assert_eq!(removed.uuid, scope.uuid); - if let Some(event) = event { - NemoRelayContextState::emit_event(&event, &subscribers); + if let Some(sanitizers) = sanitizers { + let _ = subscriber_dispatcher::dispatch_sanitized_event( + event, + sanitizers, + &subscribers, + emission_scope_stack, + ); } Ok(()) } @@ -335,13 +348,14 @@ pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { /// cannot be read safely. /// /// # Notes -/// Scope-local subscribers attached to ancestor scopes observe the emitted -/// mark event just like scope, tool, and LLM lifecycle events. +/// The event and its visible middleware/subscriber chains are snapshotted +/// before this function returns. Sanitization and subscriber delivery happen +/// later on the serial publication dispatcher. pub fn event(params: EmitMarkEventParams<'_>) -> Result<()> { ensure_runtime_owner()?; let parent_uuid = resolve_parent_uuid(params.parent); let scope_stack = current_scope_stack(); - let (event, subscribers) = { + let (event, subscribers, emission_scope_stack) = { let subscribers = if params.name == COMPACTION_EVENT_NAME { let mut scope_guard = scope_stack.write().expect("scope stack lock poisoned"); let subscribers = @@ -368,10 +382,15 @@ pub fn event(params: EmitMarkEventParams<'_>) -> Result<()> { params.category, params.category_profile, )); - (event, subscribers) + (event, subscribers, scope_stack.clone()) }; - if let Some(event) = sanitize_event(event) { - NemoRelayContextState::emit_event(&event, &subscribers); + if let Some(sanitizers) = snapshot_event_sanitizers(&event, &emission_scope_stack) { + let _ = subscriber_dispatcher::dispatch_sanitized_event( + event, + sanitizers, + &subscribers, + emission_scope_stack, + ); } Ok(()) } diff --git a/crates/core/src/api/shared.rs b/crates/core/src/api/shared.rs index a3baf051d..92aebeed1 100644 --- a/crates/core/src/api/shared.rs +++ b/crates/core/src/api/shared.rs @@ -7,8 +7,11 @@ use uuid::Uuid; use crate::api::event::{Event, ScopeCategory}; use crate::api::llm::LlmRequest; +use crate::api::registry::Guardrail; use crate::api::runtime::global_context; -use crate::api::runtime::{EventSubscriberFn, NemoRelayContextState, ScopeStackHandle}; +use crate::api::runtime::{ + EventSanitizeFn, EventSubscriberFn, NemoRelayContextState, ScopeStackHandle, +}; use crate::api::runtime::{current_scope_stack, task_scope_top}; use crate::api::scope::ScopeHandle; use crate::api::scope::ScopeType; @@ -51,7 +54,22 @@ pub(crate) fn sanitize_event_with_scope_stack( event: Event, scope_stack: &ScopeStackHandle, ) -> Option { - let entries = { + let entries = snapshot_event_sanitizers(&event, scope_stack)?; + Some(NemoRelayContextState::event_sanitize_snapshot_chain( + event, &entries, + )) +} + +/// Snapshot the event sanitizer chain visible on a captured scope stack. +/// +/// The snapshot remains valid after the emitting scope is removed, allowing +/// synchronous scope and mark APIs to enqueue publication without changing +/// which scope-local middleware observes the event. +pub(crate) fn snapshot_event_sanitizers( + event: &Event, + scope_stack: &ScopeStackHandle, +) -> Option>> { + Some({ let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); let context = global_context(); let state = match context.read() { @@ -87,10 +105,7 @@ pub(crate) fn sanitize_event_with_scope_stack( ) } } - }; - Some(NemoRelayContextState::event_sanitize_snapshot_chain( - event, &entries, - )) + }) } pub(crate) fn ensure_runtime_owner() -> Result<()> { diff --git a/crates/core/tests/integration/subscriber_dispatcher_tests.rs b/crates/core/tests/integration/subscriber_dispatcher_tests.rs index e83fc7d6d..70fb5c87d 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -6,11 +6,16 @@ use std::sync::{Arc, Mutex, mpsc}; use std::time::Duration; +use nemo_relay::api::event::Event; +use nemo_relay::api::registry::{ + deregister_mark_sanitize_guardrail, register_mark_sanitize_guardrail, +}; use nemo_relay::api::runtime::{ NemoRelayContextState, create_scope_stack, global_context, set_thread_scope_stack, }; use nemo_relay::api::scope::{EmitMarkEventParams, event}; use nemo_relay::api::subscriber::{deregister_subscriber, flush_subscribers, register_subscriber}; +use serde_json::json; static TEST_MUTEX: Mutex<()> = Mutex::new(()); @@ -96,6 +101,66 @@ fn dispatcher_preserves_event_order() { assert_eq!(observed.lock().unwrap().as_slice(), ["one", "two"]); } +#[test] +fn mark_emission_snapshots_sanitizers_and_returns_before_they_finish() { + let _lock = TEST_MUTEX.lock().unwrap(); + flush_subscribers().unwrap(); + reset_global(); + setup_isolated_thread(); + + let (sanitizer_started_tx, sanitizer_started_rx) = mpsc::channel(); + let (release_tx, release_rx) = mpsc::channel(); + let release_rx = Arc::new(Mutex::new(release_rx)); + register_mark_sanitize_guardrail( + "blocking-mark-sanitizer", + 10, + Arc::new(move |_, mut fields| { + sanitizer_started_tx.send(()).unwrap(); + release_rx.lock().unwrap().recv().unwrap(); + fields.data = Some(json!({"sanitized": true})); + fields + }), + ) + .unwrap(); + + let observed = Arc::new(Mutex::new(Vec::::new())); + let observed_events = Arc::clone(&observed); + register_subscriber( + "sanitized-mark-subscriber", + Arc::new(move |event| observed_events.lock().unwrap().push(event.clone())), + ) + .unwrap(); + + let (returned_tx, returned_rx) = mpsc::channel(); + let event_thread = std::thread::spawn(move || { + emit_mark("queued-sanitizer"); + returned_tx.send(()).unwrap(); + }); + + sanitizer_started_rx + .recv_timeout(Duration::from_secs(1)) + .expect("sanitizer should start on the dispatcher thread"); + returned_rx + .recv_timeout(Duration::from_secs(1)) + .expect("mark emission should return while its sanitizer is blocked"); + + // Removing the global registration cannot affect the already-snapshotted + // publication chain. + deregister_mark_sanitize_guardrail("blocking-mark-sanitizer").unwrap(); + release_tx.send(()).unwrap(); + event_thread.join().unwrap(); + flush_subscribers().unwrap(); + + let events = observed.lock().unwrap(); + assert_eq!(events.len(), 1); + assert_eq!( + events[0].sanitize_fields().data, + Some(json!({"sanitized": true})) + ); + drop(events); + deregister_subscriber("sanitized-mark-subscriber").unwrap(); +} + #[test] fn dispatcher_continues_after_subscriber_panic() { let _lock = TEST_MUTEX.lock().unwrap(); diff --git a/docs/reference/event-sanitizers.mdx b/docs/reference/event-sanitizers.mdx index 8ded00de5..8a4689662 100644 --- a/docs/reference/event-sanitizers.mdx +++ b/docs/reference/event-sanitizers.mdx @@ -53,6 +53,19 @@ binding callback results fail open and preserve the current fields. In Node.js, a synchronous sanitizer callback that throws also fails open; Relay records the error for `getLastCallbackError()`. +## Publication Semantics + +Scope and mark emission APIs remain synchronous. They snapshot the event, +visible sanitizer chain, and subscribers, then enqueue that snapshot for +sanitization and publication on a serial background dispatcher. Subscribers +and exporters therefore receive the sanitized event after the emission call +returns. + +The dispatcher processes snapshots in FIFO order, preserving scope start/end +and mark ordering. Closing a scope or deregistering middleware after emission +does not alter an already-snapshotted publication chain. Use the binding's +subscriber flush API when a test or shutdown path must wait for queued delivery. + ## Registration Lifetimes Where you register a sanitizer determines how long it stays active. From 7eeb5235e57cdc98c28537b9361c6b001484534f Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 14:39:58 -0400 Subject: [PATCH 02/52] test: flush queued FFI event sanitizers Signed-off-by: Will Killian --- crates/ffi/tests/unit/api/registry_tests.rs | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/crates/ffi/tests/unit/api/registry_tests.rs b/crates/ffi/tests/unit/api/registry_tests.rs index 185cadd79..97452897c 100644 --- a/crates/ffi/tests/unit/api/registry_tests.rs +++ b/crates/ffi/tests/unit/api/registry_tests.rs @@ -243,6 +243,8 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { nemo_relay_deregister_mark_sanitize_guardrail(invalid_guard.as_ptr()), NemoRelayStatus::Ok ); + // The queued event retains its sanitizer snapshot after deregistration. + assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); assert_eq!(*lock_unpoisoned(plugin_frees()), 4); let mut owner = ptr::null_mut(); @@ -343,6 +345,8 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::Ok ); + // Scope removal does not alter the sanitizer snapshots already queued. + assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); assert_eq!(*lock_unpoisoned(plugin_frees()), 7); let invalid_uuid = cstring("not-a-uuid"); From fb4cc1af312b7e9f28060bedb9af36bad5f53f3e Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 15:32:21 -0400 Subject: [PATCH 03/52] perf: skip queued sanitizers without subscribers Signed-off-by: Will Killian --- .../src/api/runtime/subscriber_dispatcher.rs | 3 +++ .../subscriber_dispatcher_tests.rs | 27 +++++++++++++++++++ 2 files changed, 30 insertions(+) diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 4ac440724..39dea23fc 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -87,6 +87,9 @@ mod native { subscribers: &[EventSubscriberFn], scope_stack: ScopeStackHandle, ) -> bool { + if subscribers.is_empty() { + return true; + } let message = DispatcherMessage::Deliver { event: Box::new(event), sanitizers, diff --git a/crates/core/tests/integration/subscriber_dispatcher_tests.rs b/crates/core/tests/integration/subscriber_dispatcher_tests.rs index 70fb5c87d..037c05cc7 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -3,6 +3,7 @@ //! Integration tests for native subscriber dispatch behavior. +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, mpsc}; use std::time::Duration; @@ -161,6 +162,32 @@ fn mark_emission_snapshots_sanitizers_and_returns_before_they_finish() { deregister_subscriber("sanitized-mark-subscriber").unwrap(); } +#[test] +fn mark_emission_skips_sanitizers_without_subscribers() { + let _lock = TEST_MUTEX.lock().unwrap(); + flush_subscribers().unwrap(); + reset_global(); + setup_isolated_thread(); + + let sanitizer_called = Arc::new(AtomicBool::new(false)); + let called = Arc::clone(&sanitizer_called); + register_mark_sanitize_guardrail( + "unused-mark-sanitizer", + 10, + Arc::new(move |_, fields| { + called.store(true, Ordering::Release); + fields + }), + ) + .unwrap(); + + emit_mark("no-subscribers"); + flush_subscribers().unwrap(); + deregister_mark_sanitize_guardrail("unused-mark-sanitizer").unwrap(); + + assert!(!sanitizer_called.load(Ordering::Acquire)); +} + #[test] fn dispatcher_continues_after_subscriber_panic() { let _lock = TEST_MUTEX.lock().unwrap(); From ec5b6c53a7f83e5fef4b334422c3a86a8c4eeba0 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 16:11:47 -0400 Subject: [PATCH 04/52] perf: avoid unsanitized event clones Signed-off-by: Will Killian --- .../src/api/runtime/subscriber_dispatcher.rs | 28 ++++++------ .../subscriber_dispatcher_tests.rs | 43 +++++++++++++++++++ 2 files changed, 59 insertions(+), 12 deletions(-) diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 39dea23fc..29a18a525 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -211,18 +211,22 @@ mod native { let previous_scope_stack = capture_thread_scope_stack(); set_thread_scope_stack(scope_stack); IN_DISPATCHER.with(|flag| flag.set(true)); - let original = (*event).clone(); - let event = catch_unwind(AssertUnwindSafe(|| { - NemoRelayContextState::event_sanitize_snapshot_chain(*event, &sanitizers) - })) - .unwrap_or_else(|_| { - log::error!( - target: "nemo_relay.runtime", - event = "event_sanitizer_panicked"; - "Event sanitizer panicked; publishing the original event snapshot" - ); - original - }); + let event = if sanitizers.is_empty() { + *event + } else { + let original = (*event).clone(); + catch_unwind(AssertUnwindSafe(|| { + NemoRelayContextState::event_sanitize_snapshot_chain(*event, &sanitizers) + })) + .unwrap_or_else(|_| { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_panicked"; + "Event sanitizer panicked; publishing the original event snapshot" + ); + original + }) + }; for subscriber in subscribers { if catch_unwind(AssertUnwindSafe(|| subscriber(&event))).is_err() { log::error!( diff --git a/crates/core/tests/integration/subscriber_dispatcher_tests.rs b/crates/core/tests/integration/subscriber_dispatcher_tests.rs index 037c05cc7..ead76744b 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -188,6 +188,49 @@ fn mark_emission_skips_sanitizers_without_subscribers() { assert!(!sanitizer_called.load(Ordering::Acquire)); } +#[test] +fn sanitizer_panic_publishes_the_original_event() { + let _lock = TEST_MUTEX.lock().unwrap(); + flush_subscribers().unwrap(); + reset_global(); + setup_isolated_thread(); + + register_mark_sanitize_guardrail( + "panicking-mark-sanitizer", + 10, + Arc::new(move |_, _| panic!("sanitizer failed")), + ) + .unwrap(); + + let observed = Arc::new(Mutex::new(Vec::::new())); + let observed_events = Arc::clone(&observed); + register_subscriber( + "panic-fallback-subscriber", + Arc::new(move |event| observed_events.lock().unwrap().push(event.clone())), + ) + .unwrap(); + + event( + EmitMarkEventParams::builder() + .name("panic-fallback") + .data(json!({"original": true})) + .build(), + ) + .unwrap(); + flush_subscribers().unwrap(); + + deregister_mark_sanitize_guardrail("panicking-mark-sanitizer").unwrap(); + deregister_subscriber("panic-fallback-subscriber").unwrap(); + + let events = observed.lock().unwrap(); + assert_eq!(events.len(), 1); + assert_eq!(events[0].name(), "panic-fallback"); + assert_eq!( + events[0].sanitize_fields().data, + Some(json!({"original": true})) + ); +} + #[test] fn dispatcher_continues_after_subscriber_panic() { let _lock = TEST_MUTEX.lock().unwrap(); From 7254dca3f978bcd9f4c410e17324c680e01e24f5 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 16:16:43 -0400 Subject: [PATCH 05/52] fix: preserve sanitized snapshots after callback panics Signed-off-by: Will Killian --- crates/core/src/api/runtime/state.rs | 16 ++++++++++++++-- .../src/api/runtime/subscriber_dispatcher.rs | 17 +---------------- .../integration/subscriber_dispatcher_tests.rs | 14 ++++++++++++-- 3 files changed, 27 insertions(+), 20 deletions(-) diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index eb3dbb950..fa45397a4 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -10,6 +10,7 @@ use std::any::Any; use std::collections::HashMap; +use std::panic::{AssertUnwindSafe, catch_unwind}; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; @@ -638,8 +639,19 @@ impl NemoRelayContextState { entries: &[Guardrail], ) -> Event { for entry in entries { - let fields = (entry.payload)(&event, event.sanitize_fields()); - event.apply_sanitize_fields(fields); + if catch_unwind(AssertUnwindSafe(|| { + let fields = (entry.payload)(&event, event.sanitize_fields()); + event.apply_sanitize_fields(fields); + })) + .is_err() + { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_panicked", + guardrail = entry.name.as_str(); + "Event sanitizer panicked; publishing the latest valid event snapshot" + ); + } } event } diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 29a18a525..b050a6f45 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -211,22 +211,7 @@ mod native { let previous_scope_stack = capture_thread_scope_stack(); set_thread_scope_stack(scope_stack); IN_DISPATCHER.with(|flag| flag.set(true)); - let event = if sanitizers.is_empty() { - *event - } else { - let original = (*event).clone(); - catch_unwind(AssertUnwindSafe(|| { - NemoRelayContextState::event_sanitize_snapshot_chain(*event, &sanitizers) - })) - .unwrap_or_else(|_| { - log::error!( - target: "nemo_relay.runtime", - event = "event_sanitizer_panicked"; - "Event sanitizer panicked; publishing the original event snapshot" - ); - original - }) - }; + let event = NemoRelayContextState::event_sanitize_snapshot_chain(*event, &sanitizers); for subscriber in subscribers { if catch_unwind(AssertUnwindSafe(|| subscriber(&event))).is_err() { log::error!( diff --git a/crates/core/tests/integration/subscriber_dispatcher_tests.rs b/crates/core/tests/integration/subscriber_dispatcher_tests.rs index ead76744b..2276dee6c 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -189,12 +189,21 @@ fn mark_emission_skips_sanitizers_without_subscribers() { } #[test] -fn sanitizer_panic_publishes_the_original_event() { +fn sanitizer_panic_publishes_the_latest_valid_event() { let _lock = TEST_MUTEX.lock().unwrap(); flush_subscribers().unwrap(); reset_global(); setup_isolated_thread(); + register_mark_sanitize_guardrail( + "successful-mark-sanitizer", + 0, + Arc::new(move |_, mut fields| { + fields.data = Some(json!({"redacted": true})); + fields + }), + ) + .unwrap(); register_mark_sanitize_guardrail( "panicking-mark-sanitizer", 10, @@ -219,6 +228,7 @@ fn sanitizer_panic_publishes_the_original_event() { .unwrap(); flush_subscribers().unwrap(); + deregister_mark_sanitize_guardrail("successful-mark-sanitizer").unwrap(); deregister_mark_sanitize_guardrail("panicking-mark-sanitizer").unwrap(); deregister_subscriber("panic-fallback-subscriber").unwrap(); @@ -227,7 +237,7 @@ fn sanitizer_panic_publishes_the_original_event() { assert_eq!(events[0].name(), "panic-fallback"); assert_eq!( events[0].sanitize_fields().data, - Some(json!({"original": true})) + Some(json!({"redacted": true})) ); } From e701c9b5e13a35df1fcc1e88ff3d5c97c71f747a Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:08:05 -0400 Subject: [PATCH 06/52] docs: clarify deferred scope-end publication Signed-off-by: Will Killian --- crates/core/src/api/scope.rs | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/crates/core/src/api/scope.rs b/crates/core/src/api/scope.rs index 8cc0297e1..27763f836 100644 --- a/crates/core/src/api/scope.rs +++ b/crates/core/src/api/scope.rs @@ -279,6 +279,11 @@ pub fn push_scope(params: PushScopeParams<'_>) -> Result { /// /// # Notes /// The implicit root scope cannot be removed. +/// +/// Scope-end emission snapshots the visible scope-local sanitizers before +/// removing the scope. Publication is then queued after removal using that +/// snapshot, so cleanup does not change the middleware applied to the emitted +/// event. pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { ensure_runtime_owner()?; let scope_stack = current_scope_stack(); From d51bb0fdc3f04da91cbbd66c9a42d0bb898f3fe1 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 19:28:33 -0400 Subject: [PATCH 07/52] fix(node): avoid subscriber flush deadlock Signed-off-by: Will Killian --- crates/node/README.md | 9 ++++---- crates/node/src/api/mod.rs | 14 ++++++++----- crates/node/tests/event_sanitizers_tests.mjs | 22 ++++++++++---------- crates/node/tests/llm_tests.mjs | 2 +- crates/node/tests/scope_tests.mjs | 7 +++---- crates/node/tests/tools_tests.mjs | 2 +- 6 files changed, 30 insertions(+), 26 deletions(-) diff --git a/crates/node/README.md b/crates/node/README.md index 13d0ccda9..ea6486462 100644 --- a/crates/node/README.md +++ b/crates/node/README.md @@ -88,7 +88,7 @@ async function main() { event("initialized", handle, { binding: "node" }, null); }); - flushSubscribers(); + await flushSubscribers(); await new Promise((resolve) => setImmediate(resolve)); deregisterSubscriber("printer"); } @@ -99,9 +99,10 @@ main().catch((error) => { }); ``` -Native subscriber delivery is asynchronous. `flushSubscribers()` drains the -native dispatcher. The extra event-loop turn lets queued JavaScript callback -side effects complete before deregistration or exit. +Native subscriber delivery is asynchronous. Awaiting `flushSubscribers()` drains +the native dispatcher without blocking the Node.js event loop. The extra +event-loop turn lets queued JavaScript callback side effects complete before +deregistration or exit. The main runtime API is exported from `nemo-relay-node`. Additional entry points are available at `nemo-relay-node/typed`, `nemo-relay-node/plugin`, diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index 24bc87b80..26838ee41 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -3237,17 +3237,21 @@ pub fn deregister_subscriber(name: String) -> Result { core_subscriber_api::deregister_subscriber(&name).map_err(to_napi_err) } -/// Wait for native subscriber callbacks queued before this call to finish. +/// Return a Promise that resolves when native subscriber callbacks queued +/// before this call finish. /// /// Call this function outside native subscriber callbacks. A re-entrant call returns without /// waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can /// still run. /// -/// JavaScript subscribers are queued through Node's `ThreadsafeFunction`; callers that -/// need JS callback side effects should await an event-loop tick after this returns. +/// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this +/// Promise does not block the Node event loop while event sanitizers settle. #[napi] -pub fn flush_subscribers() -> Result<()> { - core_subscriber_api::flush_subscribers().map_err(to_napi_err) +pub async fn flush_subscribers() -> Result<()> { + tokio::task::spawn_blocking(core_subscriber_api::flush_subscribers) + .await + .map_err(|error| to_napi_err(FlowError::Internal(error.to_string())))? + .map_err(to_napi_err) } // --------------------------------------------------------------------------- diff --git a/crates/node/tests/event_sanitizers_tests.mjs b/crates/node/tests/event_sanitizers_tests.mjs index 812a2a0ef..3679b52cf 100644 --- a/crates/node/tests/event_sanitizers_tests.mjs +++ b/crates/node/tests/event_sanitizers_tests.mjs @@ -57,7 +57,7 @@ describe('event sanitizer registries', () => { }); try { lib.event('checkpoint', null, { secret: 'raw' }, { secret: 'raw' }); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, 1); } finally { lib.deregisterMarkSanitizeGuardrail('node-event-first'); @@ -93,7 +93,7 @@ describe('event sanitizer registries', () => { { secret: 'input' }, ); lib.popScope(handle, { secret: 'output' }, null, { secret: 'end' }); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, 2); } finally { lib.deregisterScopeSanitizeStartGuardrail('node-scope-start'); @@ -129,14 +129,14 @@ describe('event sanitizer registries', () => { lib.registerMarkSanitizeGuardrail(name, 0, sanitizer); try { lib.event(name, null, { kept: kind }, { kept: kind }); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, Object.keys(invalidResults).indexOf(kind) + 1); } finally { lib.deregisterMarkSanitizeGuardrail(seedName); lib.deregisterMarkSanitizeGuardrail(name); } assertSanitizerFieldsCleared(events.at(-1)); - assert.match(lib.getLastCallbackError(), /event sanitizer callback failed/); + assert.match(lib.getLastCallbackError(), /invalid JS event sanitizer result/); } } finally { lib.deregisterSubscriber('node-event-sanitize-invalid-sub'); @@ -151,7 +151,7 @@ describe('event sanitizer registries', () => { })); try { await lib.toolCallExecute('background-tool', { raw: true }, (args) => args); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, 2); } finally { lib.deregisterScopeSanitizeStartGuardrail('node-background-start'); @@ -184,7 +184,7 @@ describe('event sanitizer registries', () => { lib.registerScopeSanitizeStartGuardrail(name, 0, sanitizer); try { await lib.toolCallExecute(name, { kept: kind }, (args) => args); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, (Object.keys(invalidResults).indexOf(kind) + 1) * 2); } finally { lib.deregisterScopeSanitizeStartGuardrail(seedName); @@ -215,7 +215,7 @@ describe('event sanitizer registries', () => { }); try { await lib.toolCallExecute('background-throw-tool', { kept: true }, (args) => args); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, 2); const start = events.find( (event) => event.kind === 'scope' && event.name === 'background-throw-tool' && event.scope_category === 'start', @@ -243,7 +243,7 @@ describe('event sanitizer registries', () => { lib.popScope(child); lib.popScope(owner); lib.event('outside', null, { raw: true }); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, 3); lib.deregisterSubscriber('node-event-sanitize-local-sub'); const marks = Object.fromEntries( @@ -271,11 +271,11 @@ describe('event sanitizer registries', () => { components: [plugin.ComponentSpec(kind)], }); lib.event('configured', null, { raw: true }); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, 1); plugin.clear(); lib.event('cleared', null, { raw: true }); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, 2); } finally { plugin.clear(); @@ -312,7 +312,7 @@ describe('event sanitizer registries', () => { components: [plugin.ComponentSpec(kind)], }); lib.event('plugin-throw', null, { raw: true }, { raw: true }); - lib.flushSubscribers(); + await lib.flushSubscribers(); await waitFor(events, 1); assertSanitizerFieldsCleared(events.at(-1)); assert.match(lib.getLastCallbackError() ?? '', /plugin sanitizer boom/i); diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index 3115b2136..61eabb89b 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -51,7 +51,7 @@ function rejectWith(value) { } async function flushSubscriberCallbacks() { - flushSubscribers(); + await flushSubscribers(); for (let i = 0; i < 10; i += 1) { await new Promise((resolve) => setImmediate(resolve)); } diff --git a/crates/node/tests/scope_tests.mjs b/crates/node/tests/scope_tests.mjs index a4cb41aed..cb50f1ade 100644 --- a/crates/node/tests/scope_tests.mjs +++ b/crates/node/tests/scope_tests.mjs @@ -30,7 +30,7 @@ function rejectWithPrimitive(value) { } async function flushSubscriberCallbacks() { - flushSubscribers(); + await flushSubscribers(); for (let i = 0; i < 10; i += 1) { await new Promise((resolve) => setImmediate(resolve)); } @@ -362,13 +362,12 @@ describe('Subscribers', () => { } }); - it('flushSubscribers is a native barrier before JS event-loop delivery', async () => { + it('flushSubscribers asynchronously drains the native dispatcher', async () => { const events = []; registerSubscriber('node_flush_collector', (e) => events.push(e)); try { event('node_flush_mark', null, null, null); - flushSubscribers(); - assert.equal(events.length, 0); + await flushSubscribers(); await new Promise((resolve) => setImmediate(resolve)); assert.ok(events.some((e) => e.kind === 'mark' && e.name === 'node_flush_mark')); } finally { diff --git a/crates/node/tests/tools_tests.mjs b/crates/node/tests/tools_tests.mjs index 125475561..3d20628b6 100644 --- a/crates/node/tests/tools_tests.mjs +++ b/crates/node/tests/tools_tests.mjs @@ -48,7 +48,7 @@ function sparseArray() { } async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { - flushSubscribers(); + await flushSubscribers(); // flushSubscribers() waits for Relay's Rust subscriber dispatcher, but JS // subscriber callbacks are queued onto Node's event loop through N-API // ThreadsafeFunction. Yield event-loop turns until the observed JS-side From 64ab12af98b07a2384eb4198049ee7f4218280b3 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 19:56:31 -0400 Subject: [PATCH 08/52] fix(node): await subscriber flush in OpenClaw Signed-off-by: Will Killian --- crates/node/src/api/mod.rs | 3 +++ integrations/openclaw/src/hooks-backend.ts | 6 +++--- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index 26838ee41..24df871fb 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -3246,6 +3246,9 @@ pub fn deregister_subscriber(name: String) -> Result { /// /// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this /// Promise does not block the Node event loop while event sanitizers settle. +/// +/// The Promise rejects if the blocking task fails or the core subscriber flush returns an error. +/// Callers should handle errors when awaiting it. #[napi] pub async fn flush_subscribers() -> Result<()> { tokio::task::spawn_blocking(core_subscriber_api::flush_subscribers) diff --git a/integrations/openclaw/src/hooks-backend.ts b/integrations/openclaw/src/hooks-backend.ts index e72f5c6cb..41bcca99b 100644 --- a/integrations/openclaw/src/hooks-backend.ts +++ b/integrations/openclaw/src/hooks-backend.ts @@ -426,7 +426,7 @@ export class HookReplayBackend { this.materializeDeferredSessionRoot(session); drainSession(this.sessionManager(), session); closeSessionRoot(this.sessionManager(), session, summary, session.finalOutput ?? summary, metadata); - this.flushSubscriberDelivery('session_close'); + await this.flushSubscriberDelivery('session_close'); this.forgetPendingSubagentLineage(session); deleteSession(this.stateValue, session); } @@ -466,9 +466,9 @@ export class HookReplayBackend { } /** Wait for native subscriber/exporter delivery after a replay closure boundary. */ - private flushSubscriberDelivery(label: string): void { + private async flushSubscriberDelivery(label: string): Promise { try { - this.nf.flushSubscribers?.(); + await this.nf.flushSubscribers?.(); } catch (error) { this.logBoundedWarn( `flush-subscribers:${label}`, From e048edc205723c00b82e99c91e63c4c09cc1f373 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 14:00:50 -0400 Subject: [PATCH 09/52] feat!: make middleware async across primary bindings Signed-off-by: Will Killian --- crates/adaptive/src/acg_component.rs | 46 +- .../adaptive/src/adaptive_hints_intercept.rs | 49 +- crates/adaptive/src/lib.rs | 4 + .../integration/runtime_integration_tests.rs | 12 +- crates/adaptive/tests/support/mod.rs | 12 + .../tests/unit/acg_component_tests.rs | 13 +- .../unit/adaptive_hints_intercept_tests.rs | 19 +- .../tests/unit/plugin_component_tests.rs | 1 + .../tests/unit/runtime_features_tests.rs | 25 +- crates/adaptive/tests/unit/runtime_tests.rs | 1 + crates/cli/src/sessions/mod.rs | 26 +- .../cli/tests/coverage/shared/server_tests.rs | 122 ++- crates/core/src/api/llm.rs | 465 +++++++-- crates/core/src/api/runtime/callbacks.rs | 54 +- crates/core/src/api/runtime/state.rs | 123 ++- .../src/api/runtime/subscriber_dispatcher.rs | 198 +++- crates/core/src/api/scope.rs | 33 +- crates/core/src/api/shared.rs | 58 +- crates/core/src/api/tool.rs | 229 ++++- crates/core/src/logging/rotation.rs | 4 + crates/core/src/plugin/dynamic/native.rs | 969 +++++++++++++++--- crates/core/src/plugin/dynamic/worker.rs | 172 ++-- crates/core/src/stream.rs | 233 +++-- .../tests/coverage/logging_rotation_tests.rs | 34 + .../core/tests/coverage/logging_sink_tests.rs | 67 +- .../tests/fixtures/native_plugin/src/lib.rs | 291 +++++- .../tests/integration/api_surface_tests.rs | 228 +++-- .../tests/integration/middleware_tests.rs | 383 ++++--- .../tests/integration/native_plugin_tests.rs | 128 ++- .../core/tests/integration/pipeline_tests.rs | 55 +- .../tests/integration/scope_local_tests.rs | 37 +- .../subscriber_dispatcher_tests.rs | 173 ++-- crates/core/tests/integration/test_support.rs | 18 + .../tests/integration/worker_plugin_tests.rs | 6 + crates/core/tests/unit/context_tests.rs | 25 +- .../core/tests/unit/dynamic_worker_tests.rs | 36 +- crates/core/tests/unit/llm_api_tests.rs | 33 +- crates/core/tests/unit/native_plugin_tests.rs | 127 ++- crates/core/tests/unit/plugin_tests.rs | 267 +++-- crates/core/tests/unit/shared_tests.rs | 67 +- crates/ffi/src/api/mod.rs | 17 +- crates/ffi/src/callable.rs | 351 ++++--- .../tests/integration/callable_extra_tests.rs | 61 +- crates/ffi/tests/unit/callable_tests.rs | 93 +- crates/node/src/api/mod.rs | 333 +++--- crates/node/src/callable.rs | 816 ++++++++++----- crates/node/src/callback_factory.rs | 18 +- crates/node/src/promise_call.rs | 67 +- crates/node/tests/callback_error_tests.mjs | 4 +- crates/node/tests/event_sanitizers_tests.mjs | 50 +- crates/node/tests/llm_tests.mjs | 102 +- crates/node/tests/scope_tests.mjs | 24 +- crates/node/tests/tools_tests.mjs | 62 ++ crates/pii-redaction/src/builtin.rs | 242 ++--- .../tests/unit/component_tests.rs | 192 ++-- crates/plugin/README.md | 8 +- crates/plugin/src/lib.rs | 178 +++- crates/plugin/tests/typed_callbacks.rs | 4 +- crates/python/src/py_api/mod.rs | 146 ++- crates/python/src/py_callable.rs | 434 ++++---- .../python/tests/coverage/coverage_tests.rs | 69 +- .../tests/coverage/py_api_coverage_tests.rs | 51 +- .../coverage/py_callable_coverage_tests.rs | 71 +- docs/about-nemo-relay/concepts/middleware.mdx | 26 + .../dynamic-plugins/native-dynamic/about.mdx | 39 +- docs/reference/event-sanitizers.mdx | 53 +- docs/reference/migration-guides.mdx | 114 ++- integrations/openclaw/test/live-smoke.test.ts | 14 +- python/nemo_relay/__init__.py | 23 +- python/nemo_relay/__init__.pyi | 25 +- python/nemo_relay/_native.pyi | 47 +- python/tests/test_adaptive.py | 2 +- python/tests/test_builtin_codecs.py | 9 +- python/tests/test_context_isolation.py | 8 +- python/tests/test_event_sanitizers.py | 4 +- python/tests/test_llm.py | 31 +- python/tests/test_tools.py | 8 +- 77 files changed, 6186 insertions(+), 2453 deletions(-) create mode 100644 crates/adaptive/tests/support/mod.rs create mode 100644 crates/core/tests/coverage/logging_rotation_tests.rs create mode 100644 crates/core/tests/integration/test_support.rs diff --git a/crates/adaptive/src/acg_component.rs b/crates/adaptive/src/acg_component.rs index 443e376ff..c71ec811d 100644 --- a/crates/adaptive/src/acg_component.rs +++ b/crates/adaptive/src/acg_component.rs @@ -582,26 +582,32 @@ pub(crate) fn create_acg_llm_request_intercept( provider: String, plugin: Arc, ) -> LlmRequestInterceptFn { - Arc::new(move |_name: &str, request: LlmRequest, annotated| { - let input_content = request.content.clone(); - let translated = - translate_request(&request, &agent_id, &provider, plugin.as_ref(), &hot_cache) - .unwrap_or(request); - if annotated.is_some() && translated.content != input_content { - let translated_annotated = build_semantic_request_view(&translated) - .map_err(|error| nemo_relay::error::FlowError::Internal(error.to_string()))? - .annotated_request; - return Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - LlmRequest { - headers: translated.headers, - content: input_content, - }, - Some(translated_annotated), - )); - } - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - translated, annotated, - )) + Arc::new(move |_name: String, request: LlmRequest, annotated| { + let hot_cache = hot_cache.clone(); + let agent_id = agent_id.clone(); + let provider = provider.clone(); + let plugin = plugin.clone(); + Box::pin(async move { + let input_content = request.content.clone(); + let translated = + translate_request(&request, &agent_id, &provider, plugin.as_ref(), &hot_cache) + .unwrap_or(request); + if annotated.is_some() && translated.content != input_content { + let translated_annotated = build_semantic_request_view(&translated) + .map_err(|error| nemo_relay::error::FlowError::Internal(error.to_string()))? + .annotated_request; + return Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + LlmRequest { + headers: translated.headers, + content: input_content, + }, + Some(translated_annotated), + )); + } + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + translated, annotated, + )) + }) }) } diff --git a/crates/adaptive/src/adaptive_hints_intercept.rs b/crates/adaptive/src/adaptive_hints_intercept.rs index c9f245505..b4b649264 100644 --- a/crates/adaptive/src/adaptive_hints_intercept.rs +++ b/crates/adaptive/src/adaptive_hints_intercept.rs @@ -174,31 +174,34 @@ impl AdaptiveHintsIntercept { pub fn into_request_fn(self) -> LlmRequestInterceptFn { let this = Arc::new(self); Arc::new( - move |_name: &str, + move |_name: String, mut request: LlmRequest, mut annotated: Option| { - let scope_path = extract_scope_path(); - let manual_ls = read_manual_latency_sensitivity(); - let scope_depth = scope_path.len(); - let call_index = this.call_counter.fetch_add(1, Ordering::Relaxed); - - let effective_agent_id = this.effective_agent_id(); - let cached_hints = - this.load_hints(&scope_path, &effective_agent_id, call_index, scope_depth); - let final_hints = apply_manual_latency_override( - cached_hints, - manual_ls, - &effective_agent_id, - scope_depth, - ); - - if let Some(hints) = final_hints { - inject_agent_hints(&mut request, &mut annotated, &hints); - } - - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + let this = this.clone(); + Box::pin(async move { + let scope_path = extract_scope_path(); + let manual_ls = read_manual_latency_sensitivity(); + let scope_depth = scope_path.len(); + let call_index = this.call_counter.fetch_add(1, Ordering::Relaxed); + + let effective_agent_id = this.effective_agent_id(); + let cached_hints = + this.load_hints(&scope_path, &effective_agent_id, call_index, scope_depth); + let final_hints = apply_manual_latency_override( + cached_hints, + manual_ls, + &effective_agent_id, + scope_depth, + ); + + if let Some(hints) = final_hints { + inject_agent_hints(&mut request, &mut annotated, &hints); + } + + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }, ) } diff --git a/crates/adaptive/src/lib.rs b/crates/adaptive/src/lib.rs index ebe78d534..74ec930b7 100644 --- a/crates/adaptive/src/lib.rs +++ b/crates/adaptive/src/lib.rs @@ -14,6 +14,10 @@ pub(crate) static TEST_GLOBAL_CONTEXT_MUTEX: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); +#[cfg(test)] +#[path = "../tests/support/mod.rs"] +pub(crate) mod test_support; + pub mod acg; pub mod acg_component; pub mod acg_learner; diff --git a/crates/adaptive/tests/integration/runtime_integration_tests.rs b/crates/adaptive/tests/integration/runtime_integration_tests.rs index dd65079ea..f0e64fb1a 100644 --- a/crates/adaptive/tests/integration/runtime_integration_tests.rs +++ b/crates/adaptive/tests/integration/runtime_integration_tests.rs @@ -605,6 +605,7 @@ async fn test_adaptive_plugin_registers_and_passes_calls_through() { content: json!({"messages": []}), }, ) + .await .unwrap(); assert_eq!(request.request.content["messages"], json!([])); @@ -739,9 +740,11 @@ impl Plugin for HeaderPlugin { false, Arc::new(|_name, mut request, annotated| { request.headers.insert("x-plugin".into(), json!("set")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + Box::pin(async move { + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }), )?; ctx.register_tool_request_intercept( @@ -752,7 +755,7 @@ impl Plugin for HeaderPlugin { if let Json::Object(ref mut map) = args { map.insert("x-tool-plugin".into(), json!(true)); } - Ok(args) + Box::pin(async move { Ok(args) }) }), )?; ctx.register_llm_execution_intercept( @@ -823,6 +826,7 @@ async fn test_top_level_plugin_registers_request_and_execution_intercepts() { content: json!({"messages": []}), }, ) + .await .unwrap(); assert_eq!(request.request.headers.get("x-plugin"), Some(&json!("set"))); diff --git a/crates/adaptive/tests/support/mod.rs b/crates/adaptive/tests/support/mod.rs new file mode 100644 index 000000000..c9f4775d0 --- /dev/null +++ b/crates/adaptive/tests/support/mod.rs @@ -0,0 +1,12 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::future::Future; + +pub(crate) fn block_on(future: F) -> F::Output { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .expect("test runtime should build") + .block_on(future) +} diff --git a/crates/adaptive/tests/unit/acg_component_tests.rs b/crates/adaptive/tests/unit/acg_component_tests.rs index 51477bb3b..69004ee72 100644 --- a/crates/adaptive/tests/unit/acg_component_tests.rs +++ b/crates/adaptive/tests/unit/acg_component_tests.rs @@ -1087,11 +1087,11 @@ fn acg_component_request_intercept_passes_original_request_and_annotation_when_t plugin, ); - let outcome = intercept( - "anthropic", + let outcome = crate::test_support::block_on(intercept( + "anthropic".to_string(), invalid_request.clone(), Some(annotated.clone()), - ) + )) .expect("request intercept should pass through"); let translated = outcome.request; let returned_annotated = outcome.annotated_request; @@ -1335,7 +1335,12 @@ fn acg_component_request_intercept_rewrites_annotation_without_mutating_provider plugin, ); - let outcome = intercept("anthropic", request, Some(original_annotation.clone())).unwrap(); + let outcome = crate::test_support::block_on(intercept( + "anthropic".to_string(), + request, + Some(original_annotation.clone()), + )) + .unwrap(); assert_eq!(outcome.request.content, original_content); let annotation = outcome diff --git a/crates/adaptive/tests/unit/adaptive_hints_intercept_tests.rs b/crates/adaptive/tests/unit/adaptive_hints_intercept_tests.rs index e3ef485b5..19296c842 100644 --- a/crates/adaptive/tests/unit/adaptive_hints_intercept_tests.rs +++ b/crates/adaptive/tests/unit/adaptive_hints_intercept_tests.rs @@ -4,6 +4,7 @@ //! Unit tests for adaptive hints intercept in the NeMo Relay adaptive crate. use super::*; + use std::sync::{Mutex, OnceLock}; use crate::trie::data_models::{LlmCallPrediction, PredictionMetrics}; @@ -196,14 +197,14 @@ fn test_adaptive_hints_intercept_injects_prediction_hints_and_manual_override() stream: None, extra: serde_json::Map::new(), }; - let outcome = req_fn( - "model", + let outcome = crate::test_support::block_on(req_fn( + "model".to_string(), LlmRequest { headers: serde_json::Map::new(), content: serde_json::json!({}), }, Some(annotated.clone()), - ) + )) .unwrap(); let request = outcome.request; let returned_annotated = outcome.annotated_request; @@ -266,14 +267,14 @@ fn test_adaptive_hints_intercept_uses_defaults_and_ignores_poisoned_cache() { })); let req_fn = AdaptiveHintsIntercept::new(hot_cache, "fallback-agent".to_string()).into_request_fn(); - let outcome = req_fn( - "model", + let outcome = crate::test_support::block_on(req_fn( + "model".to_string(), LlmRequest { headers: serde_json::Map::new(), content: serde_json::json!({}), }, None, - ) + )) .unwrap(); let request = outcome.request; let annotated = outcome.annotated_request; @@ -305,14 +306,14 @@ fn test_adaptive_hints_intercept_uses_defaults_and_ignores_poisoned_cache() { }); let poisoned_req_fn = AdaptiveHintsIntercept::new(poisoned_cache, "fallback-agent".to_string()).into_request_fn(); - let poisoned_outcome = poisoned_req_fn( - "model", + let poisoned_outcome = crate::test_support::block_on(poisoned_req_fn( + "model".to_string(), LlmRequest { headers: serde_json::Map::new(), content: serde_json::json!({"existing": true}), }, None, - ) + )) .unwrap(); let poisoned_request = poisoned_outcome.request; assert!( diff --git a/crates/adaptive/tests/unit/plugin_component_tests.rs b/crates/adaptive/tests/unit/plugin_component_tests.rs index 044504d19..5a49feb25 100644 --- a/crates/adaptive/tests/unit/plugin_component_tests.rs +++ b/crates/adaptive/tests/unit/plugin_component_tests.rs @@ -364,6 +364,7 @@ async fn adaptive_plugin_registers_runtime_and_rolls_back_registration() { content: json!({}), }, ) + .await .unwrap(); assert!(request.request.headers.is_empty()); diff --git a/crates/adaptive/tests/unit/runtime_features_tests.rs b/crates/adaptive/tests/unit/runtime_features_tests.rs index 5a53c7734..efdb60541 100644 --- a/crates/adaptive/tests/unit/runtime_features_tests.rs +++ b/crates/adaptive/tests/unit/runtime_features_tests.rs @@ -139,9 +139,11 @@ fn assert_llm_request_intercept_registered(name: &str) { i32::MAX, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + Box::pin(async move { + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }), ), name, @@ -154,9 +156,11 @@ fn assert_llm_request_intercept_absent(name: &str) { i32::MAX, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + Box::pin(async move { + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }), ) .unwrap(); @@ -565,6 +569,7 @@ async fn adaptive_hints_feature_registers_request_intercept() { content: json!({}), }, ) + .await .unwrap(); assert!(request.request.headers.contains_key(AGENT_HINTS_HEADER_KEY)); @@ -730,9 +735,11 @@ async fn registration_context_registers_all_supported_callback_types() { 5, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + Box::pin(async move { + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }), ) .unwrap(); diff --git a/crates/adaptive/tests/unit/runtime_tests.rs b/crates/adaptive/tests/unit/runtime_tests.rs index f14d83fb0..cfad601bd 100644 --- a/crates/adaptive/tests/unit/runtime_tests.rs +++ b/crates/adaptive/tests/unit/runtime_tests.rs @@ -637,6 +637,7 @@ async fn adaptive_runtime_bind_scope_requires_registration_and_passes_through_wi }; let translated = llm_request_intercepts("anthropic", request.clone()) + .await .expect("request intercept chain should pass through when no hot-cache state exists"); assert_eq!(translated.request.content, request.content); diff --git a/crates/cli/src/sessions/mod.rs b/crates/cli/src/sessions/mod.rs index 732f9c31e..5f5510966 100644 --- a/crates/cli/src/sessions/mod.rs +++ b/crates/cli/src/sessions/mod.rs @@ -582,6 +582,10 @@ impl SessionManager { .map_err(CliError::from) }) .await?; + // Manual lifecycle events publish on the serial dispatcher. This + // test-only seam returns after the matching end event is observable so + // a subsequent synthetic provider call cannot overtake it. + nemo_relay::api::subscriber::flush_subscribers().map_err(CliError::from)?; let mut sessions = self.inner.lock().await; if let Some(session) = sessions.get_mut(&session_id) { session.record_completed_llm_response(response_for_hints, owner_subagent_id); @@ -784,9 +788,9 @@ impl Session { NormalizedEvent::SubagentStarted(event) => self.start_subagent(event).await, NormalizedEvent::SubagentEnded(event) => self.end_subagent(event).await, NormalizedEvent::LlmHint(event) => self.add_llm_hint(event), - NormalizedEvent::LlmStarted(event) => self.start_hook_llm(event), - NormalizedEvent::LlmEnded(event) => self.end_hook_llm(event), - NormalizedEvent::ToolStarted(event) => self.start_tool(event), + NormalizedEvent::LlmStarted(event) => self.start_hook_llm(event).await, + NormalizedEvent::LlmEnded(event) => self.end_hook_llm(event).await, + NormalizedEvent::ToolStarted(event) => self.start_tool(event).await, NormalizedEvent::ToolEnded(event) => self.end_tool(event).await, NormalizedEvent::PromptSubmitted(event) => self.start_turn(event).await, NormalizedEvent::Compaction(event) => self.mark("compaction", event), @@ -1142,8 +1146,8 @@ impl Session { if self.turn_scope.is_none() { return Ok(Vec::new()); } - self.close_active_llms(reason)?; - self.close_active_tools(reason)?; + self.close_active_llms(reason).await?; + self.close_active_tools(reason).await?; let closed_subagents = self.close_active_subagents(reason).await?; let output = self.last_turn_llm_output.take().unwrap_or(output); self.clear_correlation_state(); @@ -1182,7 +1186,7 @@ impl Session { } // Ends all active hook-observed LLM calls before closing their containing scopes. - fn close_active_llms(&mut self, reason: &str) -> Result<(), CliError> { + async fn close_active_llms(&mut self, reason: &str) -> Result<(), CliError> { let active_llms: Vec<_> = self.llms.drain().map(|(_, handle)| handle).collect(); for handle in active_llms { llm_call_end( @@ -1198,7 +1202,7 @@ impl Session { // Ends all active tool calls with a synthetic close result before ending their containing scopes. // Draining first avoids holding mutable map state while the runtime emits lifecycle events. - fn close_active_tools(&mut self, reason: &str) -> Result<(), CliError> { + async fn close_active_tools(&mut self, reason: &str) -> Result<(), CliError> { let active_tools: Vec<_> = self .tools .drain() @@ -1428,7 +1432,7 @@ impl Session { // ignored so repeated pre hooks do not create parallel handles for one provider call. Aliased // child-session LLMs carry their subagent owner in metadata and are resolved by // `hook_llm_owner`. - fn start_hook_llm(&mut self, event: LlmEvent) -> Result<(), CliError> { + async fn start_hook_llm(&mut self, event: LlmEvent) -> Result<(), CliError> { self.ensure_turn_started(event.metadata.clone())?; if self.llms.contains_key(&event.api_call_id) { return Ok(()); @@ -1454,7 +1458,7 @@ impl Session { // Ends a hook-observed LLM call, synthesizing a start if only the post hook arrives. The same // alias metadata recovery used by `start_hook_llm` keeps post-only aliased child LLMs under the // subagent instead of falling back to the root agent. - fn end_hook_llm(&mut self, event: LlmEvent) -> Result<(), CliError> { + async fn end_hook_llm(&mut self, event: LlmEvent) -> Result<(), CliError> { self.ensure_turn_started(event.metadata.clone())?; let (parent, metadata) = self.hook_llm_owner(event.metadata); let handle = match self.llms.remove(&event.api_call_id) { @@ -1511,7 +1515,7 @@ impl Session { // Starts a tool call under an explicit subagent when available, otherwise under the turn // scope. Duplicate tool IDs are ignored so repeated pre-tool hooks do not create parallel // handles for one agent tool invocation. - fn start_tool(&mut self, event: ToolEvent) -> Result<(), CliError> { + async fn start_tool(&mut self, event: ToolEvent) -> Result<(), CliError> { self.ensure_turn_started(event.metadata.clone())?; if self.tools.contains_key(&event.tool_call_id) { return Ok(()); @@ -1529,7 +1533,7 @@ impl Session { let active_tool_arguments = arguments.clone(); let active_tool_name = event.tool_name.clone(); let active_tool_owner_subagent_id = owner.subagent_id.clone(); - tool_conditional_execution(event.tool_name.as_str(), &arguments)?; + tool_conditional_execution(event.tool_name.as_str(), &arguments).await?; let metadata = tool_correlation_metadata( self.event_identity_metadata(event.metadata), owner.status, diff --git a/crates/cli/tests/coverage/shared/server_tests.rs b/crates/cli/tests/coverage/shared/server_tests.rs index e6c4b52cf..0f7ef777c 100644 --- a/crates/cli/tests/coverage/shared/server_tests.rs +++ b/crates/cli/tests/coverage/shared/server_tests.rs @@ -2339,7 +2339,9 @@ async fn pre_tool_hook_rejects_when_conditional_guardrail_blocks() { "cli-pre-tool-blocker", 1, Arc::new(|name, _args| { - Ok((name == BLOCKED_TEST_TOOL).then(|| "blocked by policy".to_string())) + Box::pin(async move { + Ok((name == BLOCKED_TEST_TOOL).then(|| "blocked by policy".to_string())) + }) }), ) .unwrap(); @@ -2565,21 +2567,24 @@ async fn gateway_request_codec_exposes_annotations_and_applies_buffered_edits() 1, false, Arc::new(move |_name, mut request, annotated| { - if request.headers.get("x-codec-test").and_then(Value::as_str) != Some("buffered") { - return Ok(LlmRequestInterceptOutcome::new(request, annotated)); - } - let mut annotated = annotated.expect("gateway generation route must have a codec"); - *captured.lock().unwrap() = Some(serde_json::to_value(&annotated).unwrap()); - let nemo_relay::codec::request::Message::User { content, .. } = - &mut annotated.messages[0] - else { - panic!("expected portable Responses string input"); - }; - *content = nemo_relay::codec::request::MessageContent::Text("edited".into()); - request - .headers - .insert("x-test-intercept".into(), json!("visible")); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + let captured = captured.clone(); + Box::pin(async move { + if request.headers.get("x-codec-test").and_then(Value::as_str) != Some("buffered") { + return Ok(LlmRequestInterceptOutcome::new(request, annotated)); + } + let mut annotated = annotated.expect("gateway generation route must have a codec"); + *captured.lock().unwrap() = Some(serde_json::to_value(&annotated).unwrap()); + let nemo_relay::codec::request::Message::User { content, .. } = + &mut annotated.messages[0] + else { + panic!("expected portable Responses string input"); + }; + *content = nemo_relay::codec::request::MessageContent::Text("edited".into()); + request + .headers + .insert("x-test-intercept".into(), json!("visible")); + Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + }) }), ) .unwrap(); @@ -2622,10 +2627,12 @@ async fn gateway_request_codec_rejects_raw_body_mutation_before_upstream() { 1, false, Arc::new(|_name, mut request, annotated| { - if request.headers.get("x-codec-test").and_then(Value::as_str) == Some("raw") { - request.content["input"] = json!("forbidden raw edit"); - } - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + Box::pin(async move { + if request.headers.get("x-codec-test").and_then(Value::as_str) == Some("raw") { + request.content["input"] = json!("forbidden raw edit"); + } + Ok(LlmRequestInterceptOutcome::new(request, annotated)) + }) }), ) .unwrap(); @@ -2688,17 +2695,20 @@ async fn gateway_request_codec_rejects_stream_mode_changes_before_upstream() { 1, false, Arc::new(|_name, request, annotated| { - if request - .headers - .get("x-codec-stream-toggle") - .and_then(Value::as_str) - == Some("true") - { - let mut annotated = annotated.expect("generation route must expose an annotation"); - annotated.stream = Some(!annotated.stream.unwrap_or(false)); - return Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))); - } - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + Box::pin(async move { + if request + .headers + .get("x-codec-stream-toggle") + .and_then(Value::as_str) + == Some("true") + { + let mut annotated = + annotated.expect("generation route must expose an annotation"); + annotated.stream = Some(!annotated.stream.unwrap_or(false)); + return Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))); + } + Ok(LlmRequestInterceptOutcome::new(request, annotated)) + }) }), ) .unwrap(); @@ -2746,29 +2756,33 @@ async fn gateway_request_codecs_apply_buffered_and_streaming_edits_on_all_genera 1, false, Arc::new(move |_name, mut request, annotated| { - let Some(marker) = request - .headers - .get("x-codec-matrix") - .and_then(Value::as_str) - .map(str::to_string) - else { - return Ok(LlmRequestInterceptOutcome::new(request, annotated)); - }; - let mut annotated = annotated.expect("generation route must expose an annotation"); - captured_annotations.lock().unwrap().push(json!({ - "marker": marker, - "annotation": annotated, - })); - let nemo_relay::codec::request::Message::User { content, .. } = - &mut annotated.messages[0] - else { - panic!("expected the first request item to be a portable user message"); - }; - *content = nemo_relay::codec::request::MessageContent::Text(format!("edited-{marker}")); - request - .headers - .insert("x-codec-edited".into(), json!(marker)); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + let captured_annotations = captured_annotations.clone(); + Box::pin(async move { + let Some(marker) = request + .headers + .get("x-codec-matrix") + .and_then(Value::as_str) + .map(str::to_string) + else { + return Ok(LlmRequestInterceptOutcome::new(request, annotated)); + }; + let mut annotated = annotated.expect("generation route must expose an annotation"); + captured_annotations.lock().unwrap().push(json!({ + "marker": marker, + "annotation": annotated, + })); + let nemo_relay::codec::request::Message::User { content, .. } = + &mut annotated.messages[0] + else { + panic!("expected the first request item to be a portable user message"); + }; + *content = + nemo_relay::codec::request::MessageContent::Text(format!("edited-{marker}")); + request + .headers + .insert("x-codec-edited".into(), json!(marker)); + Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + }) }), ) .unwrap(); diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index c34829241..06ac5119b 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -1,6 +1,7 @@ // SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 +use std::future::Future; use std::sync::Arc; use chrono::{DateTime, TimeDelta, Utc}; @@ -19,6 +20,9 @@ use crate::api::optimization::{ use crate::api::runtime::LlmCodecIdentity; use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::global_context; +use crate::api::runtime::subscriber_dispatcher::{ + dispatch_sanitized_event, dispatch_transformed_event, +}; use crate::api::runtime::{ EventSubscriberFn, LlmCollectorFn, LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn, @@ -30,7 +34,7 @@ use crate::api::scope::{EmitMarkEventParams, ScopeHandle}; use crate::api::shared::{ ensure_runtime_owner, inject_dynamo_session_ids, metadata_with_otel_status, resolve_parent_uuid, run_request_intercepts_with_codec_and_recorder, - sanitize_event_with_scope_stack, snapshot_event_subscribers, + sanitize_event_with_scope_stack, snapshot_event_sanitizers, snapshot_event_subscribers, }; use crate::codec::request::{AnnotatedLlmRequest, Message}; use crate::codec::response::{AnnotatedLlmResponse, attach_estimated_cost_for_provider}; @@ -400,28 +404,7 @@ fn limit_annotated_request_history_to_current_user_turn( ) } -fn emit_llm_start( - handle: &LlmHandle, - request: &LlmRequest, - annotated_request: Option>, - request_codec: Option>, -) -> Result<()> { - ensure_runtime_owner()?; - let subscribers = { - let scope_stack = handle.captured_scope_stack(); - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); - snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())? - }; - emit_llm_start_with_subscribers( - handle, - request, - annotated_request, - request_codec, - &subscribers, - ) -} - -fn emit_llm_start_with_subscribers( +async fn emit_llm_start_with_subscribers( handle: &LlmHandle, request: &LlmRequest, annotated_request: Option>, @@ -446,7 +429,8 @@ fn emit_llm_start_with_subscribers( observable_request.clone(), LlmSanitizeRequestContext::for_request_codec(request_codec.clone()), &entries, - ); + ) + .await; let request_changed = sanitized_request .as_ref() .is_some_and(|sanitized_request| sanitized_request != &observable_request); @@ -480,7 +464,7 @@ fn emit_llm_start_with_subscribers( .map_err(|error| FlowError::Internal(error.to_string()))?; state.build_llm_start_event(handle, input, annotated_request) }; - if let Some(event) = sanitize_event_with_scope_stack(event, scope_stack) { + if let Some(event) = sanitize_event_with_scope_stack(event, scope_stack).await { NemoRelayContextState::emit_event(&event, subscribers); } Ok(()) @@ -495,7 +479,33 @@ fn remove_observability_credential_headers(mut request: LlmRequest) -> LlmReques request } -fn emit_pending_request_marks( +/// Synchronous test seam retained for lifecycle unit tests. Public manual +/// lifecycle emission is synchronous too, but its work is queued; this helper +/// exercises the managed start-event transformation directly. +#[cfg(test)] +fn emit_llm_start( + handle: &LlmHandle, + request: &LlmRequest, + annotated_request: Option>, + request_codec: Option>, +) -> Result<()> { + let subscribers = { + let scope_stack = handle.captured_scope_stack(); + let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())? + }; + tokio::runtime::Runtime::new() + .map_err(|error| FlowError::Internal(error.to_string()))? + .block_on(emit_llm_start_with_subscribers( + handle, + request, + annotated_request, + request_codec, + &subscribers, + )) +} + +async fn emit_pending_request_marks( handle: &LlmHandle, marks: Vec, subscribers: &[EventSubscriberFn], @@ -517,28 +527,58 @@ fn emit_pending_request_marks( mark.category, mark.category_profile, )); - if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) { + if let Some(event) = + sanitize_event_with_scope_stack(event, handle.captured_scope_stack()).await + { NemoRelayContextState::emit_event(&event, subscribers); } } Ok(()) } -pub(crate) fn emit_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) { - emit_optimization_marks_with( +pub(crate) async fn emit_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) { + emit_optimization_marks_with_async( handle, subscribers, |event| sanitize_event_with_scope_stack(event, handle.captured_scope_stack()), |event, subscribers| NemoRelayContextState::try_emit_event(event, subscribers), - ); + ) + .await; } -fn emit_optimization_marks_with( +/// Queue optimization marks from a synchronous lifecycle API. +/// +/// The public manual lifecycle APIs must not await middleware. Capture each +/// event's sanitizer chain now and enqueue the immutable snapshots ahead of +/// the corresponding end event, preserving publication order. +fn enqueue_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) { + let contributions = handle.optimization_recorder.unemitted_with_timestamps(); + if contributions.is_empty() || ensure_runtime_owner().is_err() { + return; + } + let scope_stack = handle.captured_scope_stack().clone(); + for (contribution, recorded_at) in contributions { + let event = optimization_mark_event(handle, &contribution, recorded_at); + let Some(sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) else { + break; + }; + if dispatch_sanitized_event(event, sanitizers, subscribers, scope_stack.clone()) { + handle.optimization_recorder.mark_emitted(1); + } else { + break; + } + } +} + +async fn emit_optimization_marks_with_async( handle: &LlmHandle, subscribers: &[EventSubscriberFn], - mut sanitize: impl FnMut(Event) -> Option, + mut sanitize: F, mut enqueue: impl FnMut(&Event, &[EventSubscriberFn]) -> bool, -) { +) where + F: FnMut(Event) -> Fut, + Fut: Future>, +{ let contributions = handle.optimization_recorder.unemitted_with_timestamps(); if contributions.is_empty() { return; @@ -554,30 +594,8 @@ fn emit_optimization_marks_with( return; } for (contribution, recorded_at) in contributions { - let offset = contribution.sequence.unwrap_or(0).saturating_add(2); - let offset = i64::try_from(offset).unwrap_or(i64::MAX); - let request_ordered_timestamp = handle.started_at + TimeDelta::microseconds(offset); - let timestamp = recorded_at.max(request_ordered_timestamp); - let data = serde_json::to_value(&contribution).unwrap_or(Json::Null); - let event = Event::Mark(MarkEvent::new( - BaseEvent::builder() - .name("nemo_relay.llm.optimization") - .parent_uuid(handle.uuid) - .timestamp(timestamp) - .data(data) - .data_schema(DataSchema { - name: "nemo.relay.llm_optimization_contribution".to_string(), - version: "1".to_string(), - }) - .build(), - Some(EventCategory::custom()), - Some( - CategoryProfile::builder() - .subtype("nemo_relay.llm.optimization") - .build(), - ), - )); - let Some(event) = sanitize(event) else { + let event = optimization_mark_event(handle, &contribution, recorded_at); + let Some(event) = sanitize(event).await else { // Sanitizers currently rewrite fields rather than intentionally // dropping events. `None` means the sanitizer context was // unavailable, so preserve this ordered suffix for a later retry. @@ -594,6 +612,63 @@ fn emit_optimization_marks_with( } } +fn optimization_mark_event( + handle: &LlmHandle, + contribution: &crate::codec::optimization::LlmOptimizationContribution, + recorded_at: DateTime, +) -> Event { + let offset = contribution.sequence.unwrap_or(0).saturating_add(2); + let offset = i64::try_from(offset).unwrap_or(i64::MAX); + let request_ordered_timestamp = handle.started_at + TimeDelta::microseconds(offset); + Event::Mark(MarkEvent::new( + BaseEvent::builder() + .name("nemo_relay.llm.optimization") + .parent_uuid(handle.uuid) + .timestamp(recorded_at.max(request_ordered_timestamp)) + .data(serde_json::to_value(contribution).unwrap_or(Json::Null)) + .data_schema(DataSchema { + name: "nemo.relay.llm_optimization_contribution".to_string(), + version: "1".to_string(), + }) + .build(), + Some(EventCategory::custom()), + Some( + CategoryProfile::builder() + .subtype("nemo_relay.llm.optimization") + .build(), + ), + )) +} + +/// Synchronous test seam for optimization-mark accounting. Production paths +/// always use [`emit_optimization_marks_with_async`]; unit tests use this seam +/// to isolate cursor behavior from asynchronous event publication. +#[cfg(test)] +fn emit_optimization_marks_with( + handle: &LlmHandle, + subscribers: &[EventSubscriberFn], + mut sanitize: F, + mut enqueue: impl FnMut(&Event, &[EventSubscriberFn]) -> bool, +) where + F: FnMut(Event) -> Option, +{ + let contributions = handle.optimization_recorder.unemitted_with_timestamps(); + if contributions.is_empty() || ensure_runtime_owner().is_err() { + return; + } + for (contribution, recorded_at) in contributions { + let event = optimization_mark_event(handle, &contribution, recorded_at); + let Some(event) = sanitize(event) else { + break; + }; + if enqueue(&event, subscribers) { + handle.optimization_recorder.mark_emitted(1); + } else { + break; + } + } +} + /// Start a manual LLM lifecycle span. /// /// This emits an LLM-start event after applying sanitize-request guardrails to @@ -641,7 +716,77 @@ pub fn llm_call(params: LlmCallParams<'_>) -> Result { .timestamp_opt(params.timestamp) .build(); let handle = create_llm_handle(handle_params)?; - emit_llm_start(&handle, params.request, params.annotated_request, None)?; + let scope_stack = handle.captured_scope_stack().clone(); + let (entries, subscribers, agent_is_fresh) = { + let mut scope_guard = scope_stack.write().expect("scope stack lock poisoned"); + let scope_locals = scope_guard.collect_scope_local_registries(|registries| { + ®istries.llm_sanitize_request_guardrails + }); + let subscribers = + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + let entries = state.llm_sanitize_request_entries(&scope_locals); + drop(state); + let agent_is_fresh = scope_guard.take_agent_freshness(handle.parent_uuid); + (entries, subscribers, agent_is_fresh) + }; + // Middleware and event publication only observe a credential-free copy. + // Keep `params.request` untouched: it remains the caller/provider request. + let request = remove_observability_credential_headers(params.request.clone()); + let annotated_request = params.annotated_request; + let event = { + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + state.build_llm_start_event(&handle, None, None) + }; + let queued_handle = handle.clone(); + if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { + dispatch_transformed_event( + event, + Box::new(move |event| { + Box::pin(async move { + let mut sanitized_request = + NemoRelayContextState::llm_sanitize_request_snapshot_chain( + request.clone(), + LlmSanitizeRequestContext::default(), + &entries, + ) + .await; + let request_changed = sanitized_request + .as_ref() + .is_some_and(|sanitized| sanitized != &request); + let mut annotation = if sanitized_request.is_none() || request_changed { + None + } else { + annotated_request + }; + if !agent_is_fresh && let Some(sanitized_request) = sanitized_request.as_mut() { + project_llm_request_to_current_user_turn( + sanitized_request, + &mut annotation, + None, + ); + } + let input = sanitized_request + .as_ref() + .and_then(|request| serde_json::to_value(request).ok()); + let context = global_context(); + match context.read() { + Ok(state) => state.build_llm_start_event(&queued_handle, input, annotation), + Err(_) => event, + } + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); + } Ok(handle) } @@ -682,17 +827,137 @@ struct LlmCallEndBehavior { /// Sanitize-response guardrails affect only the emitted end-event payload, not /// the caller-owned `response` value. pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> { - llm_call_end_with_behavior( - params, - LlmCallEndBehavior { - response_codec_errors_fatal: true, - attach_estimated_cost: false, - }, - None, - ) + ensure_runtime_owner()?; + let scope_stack = params.handle.captured_scope_stack().clone(); + let (entries, subscribers) = { + let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_locals = scope_guard.collect_scope_local_registries(|registries| { + ®istries.llm_sanitize_response_guardrails + }); + let subscribers = + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + ( + state.llm_sanitize_response_entries(&scope_locals), + subscribers, + ) + }; + let response = if params.response.is_null() { + params.data.unwrap_or(params.response) + } else { + params.response + }; + let response_was_null_without_fallback = response.is_null(); + let handle = params.handle.clone(); + let metadata = params.metadata; + let timestamp = params.timestamp; + let annotated_response = params.annotated_response; + let response_codec = params.response_codec; + handle.optimization_recorder.close_for_finalization(None); + enqueue_optimization_marks(&handle, &subscribers); + let event = { + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + state.build_llm_end_event( + EndLlmHandleParams::builder() + .handle(&handle) + .data(Json::Null) + .metadata_opt(metadata.clone()) + .annotated_response_opt(annotated_response.clone()) + .timestamp_opt(timestamp) + .build(), + ) + }; + if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { + dispatch_transformed_event( + event, + Box::new(move |event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( + response.clone(), + LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), + &entries, + ) + .await; + let changed = sanitized + .as_ref() + .is_some_and(|sanitized| sanitized != &response); + let data = match sanitized { + Some(response) + if response_was_null_without_fallback && response.is_null() => + { + None + } + response => response, + }; + let annotation_omitted = data.as_ref().is_none_or(Json::is_null); + let (mut annotation, decode_error) = if annotation_omitted { + (None, None) + } else { + resolve_llm_end_annotation( + (!changed).then_some(annotated_response).flatten(), + response_codec, + data.as_ref(), + &LlmCallEndBehavior { + response_codec_errors_fatal: false, + attach_estimated_cost: false, + }, + &handle.name, + ) + }; + if let Some(error) = decode_error { + log::error!( + target: "nemo_relay.runtime", + event = "manual_llm_response_codec_failed"; + "Manual LLM response annotation failed during queued publication: {error}" + ); + } + let pricing = crate::codec::response::active_pricing_resolver(); + let summary = finalize_optimization_summary( + &handle.optimization_recorder, + annotation.as_mut(), + handle.model_name.as_deref(), + &pricing, + ); + if !annotation_omitted + && annotation.is_none() + && let Some(summary) = summary + { + annotation = Some(AnnotatedLlmResponse { + optimization_summary: Some(summary), + ..AnnotatedLlmResponse::default() + }); + } + let context = global_context(); + let Ok(state) = context.read() else { + return event; + }; + let end_metadata = metadata_with_otel_status(metadata, "OK", None); + state.build_llm_end_event( + EndLlmHandleParams::builder() + .handle(&handle) + .data_opt(data) + .metadata_opt(end_metadata) + .annotated_response_opt(annotation.map(Arc::new)) + .timestamp_opt(timestamp) + .build(), + ) + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); + } + Ok(()) } -fn llm_call_end_with_behavior( +async fn llm_call_end_with_behavior( params: LlmCallEndParams<'_>, behavior: LlmCallEndBehavior, lifecycle_subscribers: Option<&[EventSubscriberFn]>, @@ -735,7 +1000,8 @@ fn llm_call_end_with_behavior( response.clone(), LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), &entries, - ); + ) + .await; let response_changed = sanitized_response .as_ref() .is_some_and(|sanitized_response| sanitized_response != &response); @@ -756,7 +1022,7 @@ fn llm_call_end_with_behavior( ) }; handle.optimization_recorder.close_for_finalization(None); - emit_optimization_marks(handle, &subscribers); + emit_optimization_marks(handle, &subscribers).await; let pricing = crate::codec::response::active_pricing_resolver(); let summary = finalize_optimization_summary( &handle.optimization_recorder, @@ -790,7 +1056,8 @@ fn llm_call_end_with_behavior( .build(), ) }; - if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) { + if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()).await + { NemoRelayContextState::emit_event(&event, &subscribers); } if let Some(error) = decode_error @@ -842,7 +1109,7 @@ fn resolve_llm_end_annotation( } } -fn emit_llm_end_without_output( +async fn emit_llm_end_without_output( handle: &LlmHandle, metadata: Option, response_codec: Option>, @@ -868,17 +1135,20 @@ fn emit_llm_end_without_output( (entries, subscribers) }; let had_fallback_data = handle.data.is_some(); - let data = handle.data.clone().and_then(|data| { + let data = if let Some(data) = handle.data.clone() { NemoRelayContextState::llm_sanitize_response_snapshot_chain( data, LlmSanitizeResponseContext::for_response_codec(response_codec), &entries, ) - }); + .await + } else { + None + }; let annotation_omitted = (had_fallback_data && data.is_none()) || data.as_ref().is_some_and(Json::is_null); handle.optimization_recorder.close_for_finalization(None); - emit_optimization_marks(handle, &subscribers); + emit_optimization_marks(handle, &subscribers).await; let pricing = crate::codec::response::active_pricing_resolver(); let annotated_response = (!annotation_omitted) .then(|| { @@ -903,7 +1173,8 @@ fn emit_llm_end_without_output( .map_err(|error| FlowError::Internal(error.to_string()))?; state.end_llm_handle(handle, data, metadata, annotated_response) }; - if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) { + if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()).await + { NemoRelayContextState::emit_event(&event, &subscribers); } Ok(()) @@ -990,7 +1261,9 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { &subscribers, parent_uuid, guardrail_metadata, - )? { + ) + .await? + { let mut rejection_data = json!({}); if let Some(object) = rejection_data.as_object_mut() { object.insert("rejected".into(), json!(true)); @@ -1018,6 +1291,7 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { codec, &optimization_recorder, ) + .await }) .await?; @@ -1043,12 +1317,13 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { annotated_request.clone(), request_codec.clone(), &lifecycle_subscribers, - )?; - emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers)?; + ) + .await?; + emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers).await?; handle .optimization_recorder .record_all(optimization_contributions); - emit_optimization_marks(&handle, &lifecycle_subscribers); + emit_optimization_marks(&handle, &lifecycle_subscribers).await; let execution_name = name.clone(); let event_uuid = handle.uuid; @@ -1087,7 +1362,8 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { attach_estimated_cost: true, }, Some(&lifecycle_subscribers), - )?; + ) + .await?; Ok(response) } Err(error) => { @@ -1098,7 +1374,8 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { end_metadata, response_codec, Some(&lifecycle_subscribers), - ); + ) + .await; Err(error) } } @@ -1186,7 +1463,9 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu &subscribers, parent_uuid, guardrail_metadata, - )? { + ) + .await? + { let mut rejection_data = json!({}); if let Some(object) = rejection_data.as_object_mut() { object.insert("rejected".into(), json!(true)); @@ -1214,6 +1493,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu codec, &optimization_recorder, ) + .await }) .await?; @@ -1239,12 +1519,13 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu annotated_request, request_codec.clone(), &lifecycle_subscribers, - )?; - emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers)?; + ) + .await?; + emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers).await?; handle .optimization_recorder .record_all(optimization_contributions); - emit_optimization_marks(&handle, &lifecycle_subscribers); + emit_optimization_marks(&handle, &lifecycle_subscribers).await; let execution_name = name.clone(); let event_uuid = handle.uuid; @@ -1289,7 +1570,8 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu end_metadata, response_codec, Some(&lifecycle_subscribers), - ); + ) + .await; Err(error) } } @@ -1318,7 +1600,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu /// /// This helper does not emit the returned marks because it does not own an LLM /// lifecycle. Callers must attach them to the lifecycle they own. -pub fn llm_request_intercepts( +pub async fn llm_request_intercepts( name: &str, request: LlmRequest, ) -> Result { @@ -1336,7 +1618,8 @@ pub fn llm_request_intercepts( }; let mut outcome = NemoRelayContextState::llm_request_intercepts_snapshot_chain( name, request, None, &entries, false, - )?; + ) + .await?; inject_dynamo_session_ids(&mut outcome.request); Ok(outcome) } @@ -1361,7 +1644,7 @@ pub fn llm_request_intercepts( /// This helper is useful for preflight checks when the caller needs the /// rejection result without starting an LLM span. Guardrail scopes are still /// emitted for the conditional checks themselves. -pub fn llm_conditional_execution(request: &LlmRequest) -> Result<()> { +pub async fn llm_conditional_execution(request: &LlmRequest) -> Result<()> { ensure_runtime_owner()?; let (entries, subscribers, parent_uuid) = { let scope_stack = current_scope_stack(); @@ -1384,7 +1667,9 @@ pub fn llm_conditional_execution(request: &LlmRequest) -> Result<()> { &subscribers, parent_uuid, None, - )? { + ) + .await? + { return Err(FlowError::GuardrailRejected(error)); } Ok(()) diff --git a/crates/core/src/api/runtime/callbacks.rs b/crates/core/src/api/runtime/callbacks.rs index e40558da1..b07a82fde 100644 --- a/crates/core/src/api/runtime/callbacks.rs +++ b/crates/core/src/api/runtime/callbacks.rs @@ -27,8 +27,14 @@ use crate::json::Json; /// /// The callback receives the current event as immutable context and the fields /// it may replace. Later callbacks observe fields returned by earlier entries. -pub type EventSanitizeFn = - Arc EventSanitizeFields + Send + Sync>; +pub type EventSanitizeFn = Arc< + dyn Fn( + Event, + EventSanitizeFields, + ) -> Pin> + Send>> + + Send + + Sync, +>; /// Sanitize a tool request payload before the runtime records it. /// @@ -42,7 +48,8 @@ pub type EventSanitizeFn = /// /// # Returns /// Sanitized JSON payload for the emitted event. -pub type ToolSanitizeFn = Arc Json + Send + Sync>; +pub type ToolSanitizeFn = + Arc Pin> + Send>> + Send + Sync>; /// Decide whether a tool call is allowed to continue. /// /// The callback receives the tool name and the current argument payload. It can @@ -64,7 +71,11 @@ pub type ToolSanitizeFn = Arc Json + Send + Sync>; /// # Errors /// The callback can return any [`FlowError`](crate::error::FlowError) to abort /// guardrail evaluation. -pub type ToolConditionalFn = Arc Result> + Send + Sync>; +pub type ToolConditionalFn = Arc< + dyn Fn(String, Json) -> Pin>> + Send>> + + Send + + Sync, +>; /// Rewrite tool arguments before execution. /// /// Tool request intercepts run in priority order and can transform the JSON @@ -80,7 +91,8 @@ pub type ToolConditionalFn = Arc Result> + /// # Errors /// The callback can return any [`FlowError`](crate::error::FlowError) to abort /// the request-intercept chain. -pub type ToolInterceptFn = Arc Result + Send + Sync>; +pub type ToolInterceptFn = + Arc Pin> + Send>> + Send + Sync>; /// Continuation type invoked by tool execution intercepts. /// /// Execution intercepts receive this callable as their `next` continuation and @@ -308,8 +320,14 @@ impl LlmSanitizeResponseContext { /// /// The context is always supplied and distinguishes no codec, built-in codecs, /// runtime-registered codecs, and opaque active codecs. -pub type LlmSanitizeRequestFn = - Arc Option + Send + Sync>; +pub type LlmSanitizeRequestFn = Arc< + dyn Fn( + LlmRequest, + LlmSanitizeRequestContext, + ) -> Pin>> + Send>> + + Send + + Sync, +>; /// Sanitize an LLM response before the runtime records it. /// /// These callbacks rewrite the JSON response payload captured on LLM-end @@ -325,8 +343,14 @@ pub type LlmSanitizeRequestFn = /// /// The context is always supplied and distinguishes no codec, built-in codecs, /// runtime-registered codecs, and opaque active codecs. -pub type LlmSanitizeResponseFn = - Arc Option + Send + Sync>; +pub type LlmSanitizeResponseFn = Arc< + dyn Fn( + Json, + LlmSanitizeResponseContext, + ) -> Pin>> + Send>> + + Send + + Sync, +>; /// Decide whether an LLM call is allowed to continue. /// /// The callback receives the current [`LlmRequest`] and can allow execution, @@ -346,7 +370,11 @@ pub type LlmSanitizeResponseFn = /// # Errors /// The callback can return any [`FlowError`](crate::error::FlowError) to abort /// guardrail evaluation. -pub type LlmConditionalFn = Arc Result> + Send + Sync>; +pub type LlmConditionalFn = Arc< + dyn Fn(LlmRequest) -> Pin>> + Send>> + + Send + + Sync, +>; /// Rewrite or annotate an LLM request before execution. /// /// Request intercepts can transform the wire request, attach or replace a @@ -368,7 +396,11 @@ pub type LlmConditionalFn = Arc Result> + /// The callback can return any [`FlowError`](crate::error::FlowError) to abort /// the request-intercept chain. pub type LlmRequestInterceptFn = Arc< - dyn Fn(&str, LlmRequest, Option) -> Result + dyn Fn( + String, + LlmRequest, + Option, + ) -> Pin> + Send>> + Send + Sync, >; diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index fa45397a4..e2098b9a7 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -10,7 +10,6 @@ use std::any::Any; use std::collections::HashMap; -use std::panic::{AssertUnwindSafe, catch_unwind}; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; @@ -565,7 +564,7 @@ impl NemoRelayContextState { )) } - fn emit_guardrail_scope_start( + async fn emit_guardrail_scope_start( name: &str, parent_uuid: Option, metadata: Option, @@ -592,13 +591,13 @@ impl NemoRelayContextState { EventCategory::from(handle.scope_type), None, )); - if let Some(event) = sanitize_event(event) { + if let Some(event) = sanitize_event(event).await { Self::emit_event(&event, subscribers); } handle } - fn emit_guardrail_scope_end( + async fn emit_guardrail_scope_end( handle: &ScopeHandle, output: Json, subscribers: &[EventSubscriberFn], @@ -617,7 +616,7 @@ impl NemoRelayContextState { EventCategory::from(handle.scope_type), None, )); - if let Some(event) = sanitize_event(event) { + if let Some(event) = sanitize_event(event).await { Self::emit_event(&event, subscribers); } } @@ -634,23 +633,21 @@ impl NemoRelayContextState { } /// Apply an event sanitizer snapshot to the mutable observability fields. - pub(crate) fn event_sanitize_snapshot_chain( + pub(crate) async fn event_sanitize_snapshot_chain( mut event: Event, entries: &[Guardrail], ) -> Event { for entry in entries { - if catch_unwind(AssertUnwindSafe(|| { - let fields = (entry.payload)(&event, event.sanitize_fields()); - event.apply_sanitize_fields(fields); - })) - .is_err() - { - log::error!( + let fields = event.sanitize_fields(); + match (entry.payload)(event.clone(), fields).await { + Ok(fields) => event.apply_sanitize_fields(fields), + Err(error) => log::error!( target: "nemo_relay.runtime", - event = "event_sanitizer_panicked", - guardrail = entry.name.as_str(); - "Event sanitizer panicked; publishing the latest valid event snapshot" - ); + event = "event_sanitizer_failed", + sanitizer = entry.name.as_str(), + event_name = event.name(); + "Event sanitizer failed; preserving the last valid event snapshot: {error}" + ), } } event @@ -684,14 +681,23 @@ impl NemoRelayContextState { /// /// # Returns /// The sanitized JSON payload after every provided guardrail has run. - pub(crate) fn tool_sanitize_request_snapshot_chain( + pub(crate) async fn tool_sanitize_request_snapshot_chain( name: &str, args: Json, entries: &[Guardrail], ) -> Json { let mut value = args; for entry in entries { - value = (entry.payload)(name, value); + match (entry.payload)(name.to_string(), value.clone()).await { + Ok(next) => value = next, + Err(error) => log::error!( + target: "nemo_relay.runtime", + event = "tool_request_sanitizer_failed", + sanitizer = entry.name.as_str(), + tool_name = name; + "Tool request sanitizer failed; preserving the last valid payload: {error}" + ), + } } value } @@ -724,14 +730,23 @@ impl NemoRelayContextState { /// /// # Returns /// The sanitized JSON payload after every provided guardrail has run. - pub(crate) fn tool_sanitize_response_snapshot_chain( + pub(crate) async fn tool_sanitize_response_snapshot_chain( name: &str, result: Json, entries: &[Guardrail], ) -> Json { let mut value = result; for entry in entries { - value = (entry.payload)(name, value); + match (entry.payload)(name.to_string(), value.clone()).await { + Ok(next) => value = next, + Err(error) => log::error!( + target: "nemo_relay.runtime", + event = "tool_response_sanitizer_failed", + sanitizer = entry.name.as_str(), + tool_name = name; + "Tool response sanitizer failed; preserving the last valid payload: {error}" + ), + } } value } @@ -781,7 +796,7 @@ impl NemoRelayContextState { /// # Errors /// Propagates any error returned by a guardrail callback after emitting the /// corresponding guardrail scope end event. - pub(crate) fn tool_conditional_execution_snapshot_chain( + pub(crate) async fn tool_conditional_execution_snapshot_chain( name: &str, args: &Json, entries: &[Guardrail], @@ -799,8 +814,9 @@ impl NemoRelayContextState { "target_name": name, }), subscribers, - ); - let result = (entry.payload)(name, args); + ) + .await; + let result = (entry.payload)(name.to_string(), args.clone()).await; let output = match &result { Ok(Some(reason)) => json!({ "allowed": false, @@ -816,7 +832,7 @@ impl NemoRelayContextState { "error": error.to_string(), }), }; - Self::emit_guardrail_scope_end(&handle, output, subscribers); + Self::emit_guardrail_scope_end(&handle, output, subscribers).await; if let Some(error) = result? { return Ok(Some(error)); } @@ -859,14 +875,14 @@ impl NemoRelayContextState { /// # Notes /// If an intercept entry has `break_chain` enabled, later intercepts are /// skipped after that entry runs. - pub(crate) fn tool_request_intercepts_snapshot_chain( + pub(crate) async fn tool_request_intercepts_snapshot_chain( name: &str, args: Json, entries: &[Intercept], ) -> crate::error::Result { let mut value = args; for entry in entries { - value = (entry.payload.callable)(name, value)?; + value = (entry.payload.callable)(name.to_string(), value).await?; if entry.payload.break_chain { break; } @@ -976,14 +992,28 @@ impl NemoRelayContextState { /// /// # Returns /// The sanitized [`LlmRequest`] after every provided guardrail has run. - pub(crate) fn llm_sanitize_request_snapshot_chain( + pub(crate) async fn llm_sanitize_request_snapshot_chain( request: LlmRequest, context: LlmSanitizeRequestContext, entries: &[Guardrail], ) -> Option { let mut value = Some(request); for entry in entries { - value = value.and_then(|value| (entry.payload)(value, context.clone())); + if let Some(current) = value.take() { + match (entry.payload)(current.clone(), context.clone()).await { + Ok(next) => value = next, + Err(error) => { + log::error!( + target: "nemo_relay.runtime", + event = "llm_request_sanitizer_failed", + sanitizer = entry.name.as_str(), + preserved_value = "unsanitized_request"; + "LLM request sanitizer failed; preserving the last valid unsanitized request: {error}" + ); + value = Some(current); + } + } + } } value } @@ -1015,14 +1045,28 @@ impl NemoRelayContextState { /// /// # Returns /// The sanitized response payload after every provided guardrail has run. - pub(crate) fn llm_sanitize_response_snapshot_chain( + pub(crate) async fn llm_sanitize_response_snapshot_chain( response: Json, context: LlmSanitizeResponseContext, entries: &[Guardrail], ) -> Option { let mut value = Some(response); for entry in entries { - value = value.and_then(|value| (entry.payload)(value, context.clone())); + if let Some(current) = value.take() { + match (entry.payload)(current.clone(), context.clone()).await { + Ok(next) => value = next, + Err(error) => { + log::error!( + target: "nemo_relay.runtime", + event = "llm_response_sanitizer_failed", + sanitizer = entry.name.as_str(), + preserved_value = "unsanitized_response"; + "LLM response sanitizer failed; preserving the last valid unsanitized response: {error}" + ); + value = Some(current); + } + } + } } value } @@ -1071,7 +1115,7 @@ impl NemoRelayContextState { /// # Errors /// Propagates any error returned by a guardrail callback after emitting the /// corresponding guardrail scope end event. - pub(crate) fn llm_conditional_execution_snapshot_chain( + pub(crate) async fn llm_conditional_execution_snapshot_chain( request: &LlmRequest, entries: &[Guardrail], subscribers: &[EventSubscriberFn], @@ -1087,8 +1131,9 @@ impl NemoRelayContextState { "kind": "llm_conditional_execution", }), subscribers, - ); - let result = (entry.payload)(request); + ) + .await; + let result = (entry.payload)(request.clone()).await; let output = match &result { Ok(Some(reason)) => json!({ "allowed": false, @@ -1104,7 +1149,7 @@ impl NemoRelayContextState { "error": error.to_string(), }), }; - Self::emit_guardrail_scope_end(&handle, output, subscribers); + Self::emit_guardrail_scope_end(&handle, output, subscribers).await; if let Some(error) = result? { return Ok(Some(error)); } @@ -1151,7 +1196,7 @@ impl NemoRelayContextState { /// # Notes /// If an intercept entry has `break_chain` enabled, later intercepts are /// skipped after that entry runs. - pub(crate) fn llm_request_intercepts_snapshot_chain( + pub(crate) async fn llm_request_intercepts_snapshot_chain( name: &str, request: LlmRequest, annotated: Option, @@ -1166,11 +1211,12 @@ impl NemoRelayContextState { codec_active, None, ) + .await } /// Run a request-intercept snapshot while ingesting optimization evidence /// directly into the managed call's bounded accumulator. - pub(crate) fn llm_request_intercepts_snapshot_chain_with_recorder( + pub(crate) async fn llm_request_intercepts_snapshot_chain_with_recorder( name: &str, request: LlmRequest, annotated: Option, @@ -1184,7 +1230,8 @@ impl NemoRelayContextState { let mut optimization_contributions = Vec::new(); for entry in entries { let input_content = request_value.content.clone(); - let outcome = (entry.payload.callable)(name, request_value, annotated_value)?; + let outcome = + (entry.payload.callable)(name.to_string(), request_value, annotated_value).await?; if codec_active && outcome.request.content != input_content { return Err(crate::error::FlowError::InvalidArgument(format!( "LLM request intercept '{}' changed request.content while a request codec is active; modify annotated_request instead", diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index b050a6f45..bdd48d20f 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -9,6 +9,12 @@ use crate::api::runtime::{ EventSanitizeFn, EventSubscriberFn, NemoRelayContextState, ScopeStackHandle, }; use crate::error::Result; +use std::future::Future; +use std::pin::Pin; + +pub(crate) type EventTransformFn = Box< + dyn FnOnce(Event) -> Pin + Send + 'static>> + Send + 'static, +>; mod native { use std::cell::Cell; @@ -27,6 +33,7 @@ mod native { enum DispatcherMessage { Deliver { event: Box, + transform: Option, sanitizers: Vec>, subscribers: Vec, scope_stack: ScopeStackHandle, @@ -34,11 +41,17 @@ mod native { Flush { done: Sender<()>, }, + Barrier { + done: Receiver<()>, + }, } static DISPATCHER: OnceLock, String>> = OnceLock::new(); + static SANITIZER_RUNTIME: OnceLock> = + OnceLock::new(); static DISPATCHER_FAILURE_LOGGED: AtomicBool = AtomicBool::new(false); + static SANITIZER_RUNTIME_FAILURE_LOGGED: AtomicBool = AtomicBool::new(false); thread_local! { static IN_DISPATCHER: Cell = const { Cell::new(false) }; @@ -50,6 +63,7 @@ mod native { } let message = DispatcherMessage::Deliver { event: Box::new(event.clone()), + transform: None, sanitizers: Vec::new(), subscribers: subscribers.to_vec(), scope_stack: current_scope_stack(), @@ -92,31 +106,41 @@ mod native { } let message = DispatcherMessage::Deliver { event: Box::new(event), + transform: None, sanitizers, subscribers: subscribers.to_vec(), scope_stack, }; - match dispatcher_sender() { - Ok(sender) if sender.send(message).is_ok() => true, - Ok(_) => { - log::warn!( - target: "nemo_relay.runtime", - event = "subscriber_event_dropped", - reason = "dispatcher_disconnected"; - "Subscriber event was dropped because the dispatcher stopped" - ); - false - } - Err(error) if !DISPATCHER_FAILURE_LOGGED.swap(true, Ordering::AcqRel) => { - log::error!( - target: "nemo_relay.runtime", - event = "subscriber_dispatcher_failed"; - "Subscriber dispatcher failed to start: {error}" - ); - false - } - Err(_) => false, - } + send_dispatch_message(message) + } + + pub(super) fn dispatch_transformed_event( + event: Event, + transform: EventTransformFn, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, + ) -> bool { + let message = DispatcherMessage::Deliver { + event: Box::new(event), + transform: Some(transform), + sanitizers, + subscribers: subscribers.to_vec(), + scope_stack, + }; + send_dispatch_message(message) + } + + /// Insert a FIFO barrier for work that will enqueue a publication from an + /// async task. A later flush waits for the task to signal completion, then + /// drains the event it queued before acknowledging the flush. + pub(super) fn register_async_publication() -> Option> { + let sender = dispatcher_sender().ok()?; + let (done_tx, done_rx) = mpsc::channel(); + sender + .send(DispatcherMessage::Barrier { done: done_rx }) + .ok() + .map(|_| done_tx) } pub(super) fn flush_subscribers() -> Result<()> { @@ -145,6 +169,30 @@ mod native { DISPATCHER.get_or_init(start_dispatcher).clone() } + fn send_dispatch_message(message: DispatcherMessage) -> bool { + match dispatcher_sender() { + Ok(sender) if sender.send(message).is_ok() => true, + Ok(_) => { + log::warn!( + target: "nemo_relay.runtime", + event = "subscriber_event_dropped", + reason = "dispatcher_disconnected"; + "Subscriber event was dropped because the dispatcher stopped" + ); + false + } + Err(error) if !DISPATCHER_FAILURE_LOGGED.swap(true, Ordering::AcqRel) => { + log::error!( + target: "nemo_relay.runtime", + event = "subscriber_dispatcher_failed"; + "Subscriber dispatcher failed to start: {error}" + ); + false + } + Err(_) => false, + } + } + fn start_dispatcher() -> std::result::Result, String> { let (tx, rx) = mpsc::channel::(); let sender = std::thread::Builder::new() @@ -172,6 +220,9 @@ mod native { let _ = pending.send(()); } } + DispatcherMessage::Barrier { done } => { + let _ = done.recv(); + } message => handle_message(message), } } @@ -182,6 +233,9 @@ mod native { while let Ok(message) = rx.try_recv() { match message { DispatcherMessage::Flush { done } => pending_flushes.push(done), + DispatcherMessage::Barrier { done } => { + let _ = done.recv(); + } message => handle_message(message), } } @@ -192,18 +246,23 @@ mod native { match message { DispatcherMessage::Deliver { event, + transform, sanitizers, subscribers, scope_stack, - } => deliver_event(event, sanitizers, subscribers, scope_stack), + } => deliver_event(event, transform, sanitizers, subscribers, scope_stack), DispatcherMessage::Flush { done } => { let _ = done.send(()); } + DispatcherMessage::Barrier { done } => { + let _ = done.recv(); + } } } fn deliver_event( event: Box, + transform: Option, sanitizers: Vec>, subscribers: Vec, scope_stack: ScopeStackHandle, @@ -211,7 +270,11 @@ mod native { let previous_scope_stack = capture_thread_scope_stack(); set_thread_scope_stack(scope_stack); IN_DISPATCHER.with(|flag| flag.set(true)); - let event = NemoRelayContextState::event_sanitize_snapshot_chain(*event, &sanitizers); + let Some(event) = sanitize_event_snapshot(*event, transform, sanitizers) else { + IN_DISPATCHER.with(|flag| flag.set(false)); + restore_thread_scope_stack(previous_scope_stack); + return; + }; for subscriber in subscribers { if catch_unwind(AssertUnwindSafe(|| subscriber(&event))).is_err() { log::error!( @@ -224,6 +287,75 @@ mod native { IN_DISPATCHER.with(|flag| flag.set(false)); restore_thread_scope_stack(previous_scope_stack); } + + /// Apply a transform and sanitizers on the dispatcher thread. A transform + /// failure drops the event because it may be responsible for inserting the + /// sanitized payload. A sanitizer failure retains the transformed snapshot + /// and continues publication (fail open). + fn sanitize_event_snapshot( + event: Event, + transform: Option, + sanitizers: Vec>, + ) -> Option { + let runtime = match SANITIZER_RUNTIME.get_or_init(|| { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| error.to_string()) + }) { + Ok(runtime) => runtime, + Err(error) => { + if !SANITIZER_RUNTIME_FAILURE_LOGGED.swap(true, Ordering::AcqRel) { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_runtime_failed"; + "Event sanitizer runtime failed; dropping events: {error}" + ); + } + return None; + } + }; + let transformed = match catch_unwind(AssertUnwindSafe(|| { + runtime.block_on(async move { + match transform { + Some(transform) => transform(event).await, + None => event, + } + }) + })) { + Ok(event) => event, + Err(_) => { + log::error!( + target: "nemo_relay.runtime", + event = "event_transform_panicked"; + "Event transform panicked; dropping the event" + ); + return None; + } + }; + if sanitizers.is_empty() { + return Some(transformed); + } + let original = transformed.clone(); + Some( + match catch_unwind(AssertUnwindSafe(|| { + runtime.block_on(NemoRelayContextState::event_sanitize_snapshot_chain( + transformed, + &sanitizers, + )) + })) { + Ok(event) => event, + Err(_) => { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_panicked"; + "Event sanitizer panicked; publishing the transformed event snapshot" + ); + original + } + }, + ) + } } /// Queue an event for subscriber delivery. @@ -242,6 +374,26 @@ pub(crate) fn dispatch_sanitized_event( native::dispatch_sanitized_event(event, sanitizers, subscribers, scope_stack) } +/// Queue a snapshot for a middleware-specific asynchronous transformation, +/// followed by event sanitization and subscriber delivery. +pub(crate) fn dispatch_transformed_event( + event: Event, + transform: EventTransformFn, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, +) -> bool { + native::dispatch_transformed_event(event, transform, sanitizers, subscribers, scope_stack) +} + +/// Register a FIFO barrier for async work that will queue a subscriber event. +/// +/// Dropping the returned sender releases the barrier, so error paths cannot +/// leave the dispatcher blocked. +pub(crate) fn register_async_publication() -> Option> { + native::register_async_publication() +} + /// Wait for all queued subscriber callbacks submitted before this call. pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() diff --git a/crates/core/src/api/scope.rs b/crates/core/src/api/scope.rs index 27763f836..24e5d8cc3 100644 --- a/crates/core/src/api/scope.rs +++ b/crates/core/src/api/scope.rs @@ -217,9 +217,8 @@ pub fn get_handle() -> Result { /// cannot be read safely. /// /// # Notes -/// The event and its visible middleware/subscriber chains are snapshotted -/// before this function returns. Sanitization and subscriber delivery happen -/// later on the serial publication dispatcher. +/// Scope-local subscribers attached to ancestor scopes observe the emitted +/// start event before the function returns. pub fn push_scope(params: PushScopeParams<'_>) -> Result { ensure_runtime_owner()?; let parent_uuid = resolve_parent_uuid(params.parent); @@ -315,8 +314,9 @@ pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { ); (scope, event, subscribers, scope_stack.clone()) }; - // Snapshot scope-local middleware before removing its owner. Publication - // happens later, but cleanup must not change the chain visible at emission. + // Capture the scope-local chain before removing its owner. The event is + // published later, but scope cleanup must not change the middleware that + // was visible when the end event was emitted. let sanitizers = snapshot_event_sanitizers(&event, &emission_scope_stack); let removed = task_scope_remove(params.handle_uuid)?; debug_assert_eq!(removed.uuid, scope.uuid); @@ -353,22 +353,35 @@ pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { /// cannot be read safely. /// /// # Notes -/// The event and its visible middleware/subscriber chains are snapshotted -/// before this function returns. Sanitization and subscriber delivery happen -/// later on the serial publication dispatcher. +/// Scope-local subscribers attached to ancestor scopes observe the emitted +/// mark event just like scope, tool, and LLM lifecycle events. pub fn event(params: EmitMarkEventParams<'_>) -> Result<()> { ensure_runtime_owner()?; let parent_uuid = resolve_parent_uuid(params.parent); let scope_stack = current_scope_stack(); let (event, subscribers, emission_scope_stack) = { let subscribers = if params.name == COMPACTION_EVENT_NAME { - let mut scope_guard = scope_stack.write().expect("scope stack lock poisoned"); + let mut scope_guard = scope_stack.write().map_err(|error| { + log::error!( + target: "nemo_relay.runtime", + event = "mark_event_scope_stack_unavailable"; + "Mark event was dropped because the scope stack lock is poisoned: {error}" + ); + FlowError::Internal(error.to_string()) + })?; let subscribers = snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; scope_guard.mark_agent_fresh(parent_uuid); subscribers } else { - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_guard = scope_stack.read().map_err(|error| { + log::error!( + target: "nemo_relay.runtime", + event = "mark_event_scope_stack_unavailable"; + "Mark event was dropped because the scope stack lock is poisoned: {error}" + ); + FlowError::Internal(error.to_string()) + })?; snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())? }; let context = global_context(); diff --git a/crates/core/src/api/shared.rs b/crates/core/src/api/shared.rs index 92aebeed1..b97c147bb 100644 --- a/crates/core/src/api/shared.rs +++ b/crates/core/src/api/shared.rs @@ -45,36 +45,52 @@ pub(crate) fn snapshot_event_subscribers( } /// Apply the event sanitizer chain visible on the current scope stack. -pub(crate) fn sanitize_event(event: Event) -> Option { - sanitize_event_with_scope_stack(event, ¤t_scope_stack()) +pub(crate) async fn sanitize_event(event: Event) -> Option { + sanitize_event_with_scope_stack(event, ¤t_scope_stack()).await } /// Apply the event sanitizer chain visible on a captured scope stack. -pub(crate) fn sanitize_event_with_scope_stack( +pub(crate) async fn sanitize_event_with_scope_stack( event: Event, scope_stack: &ScopeStackHandle, ) -> Option { let entries = snapshot_event_sanitizers(&event, scope_stack)?; - Some(NemoRelayContextState::event_sanitize_snapshot_chain( - event, &entries, - )) + Some(NemoRelayContextState::event_sanitize_snapshot_chain(event, &entries).await) } -/// Snapshot the event sanitizer chain visible on a captured scope stack. +/// Snapshot the event sanitizers visible to an event without invoking them. /// -/// The snapshot remains valid after the emitting scope is removed, allowing -/// synchronous scope and mark APIs to enqueue publication without changing -/// which scope-local middleware observes the event. +/// Scope and mark emission use this to capture middleware ownership while the +/// scope is still active, then let the serial dispatcher sanitize and publish +/// the immutable event snapshot later. This keeps public scope APIs +/// synchronous while ensuring scope removal cannot affect queued work. pub(crate) fn snapshot_event_sanitizers( event: &Event, scope_stack: &ScopeStackHandle, ) -> Option>> { - Some({ - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let entries = { + let scope_guard = match scope_stack.read() { + Ok(guard) => guard, + Err(error) => { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_snapshot_failed"; + "Event was dropped because the scope stack lock is poisoned: {error}" + ); + return None; + } + }; let context = global_context(); let state = match context.read() { Ok(state) => state, - Err(_) => return None, + Err(error) => { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_snapshot_failed"; + "Event was dropped because the runtime context lock is poisoned: {error}" + ); + return None; + } }; match &event { Event::Mark(_) => { @@ -105,7 +121,8 @@ pub(crate) fn snapshot_event_sanitizers( ) } } - }) + }; + Some(entries) } pub(crate) fn ensure_runtime_owner() -> Result<()> { @@ -211,26 +228,26 @@ pub(crate) type InterceptedLlmRequest = ( ); #[cfg(test)] -pub(crate) fn run_request_intercepts_with_codec( +pub(crate) async fn run_request_intercepts_with_codec( name: &str, request: LlmRequest, codec: Option>, ) -> Result { - run_request_intercepts_with_codec_inner(name, request, codec, None) + run_request_intercepts_with_codec_inner(name, request, codec, None).await } /// Run request intercepts and record optimization contributions directly into /// the managed call's bounded accumulator as each intercept completes. -pub(crate) fn run_request_intercepts_with_codec_and_recorder( +pub(crate) async fn run_request_intercepts_with_codec_and_recorder( name: &str, request: LlmRequest, codec: Option>, recorder: &crate::api::optimization::LlmOptimizationRecorder, ) -> Result { - run_request_intercepts_with_codec_inner(name, request, codec, Some(recorder)) + run_request_intercepts_with_codec_inner(name, request, codec, Some(recorder)).await } -fn run_request_intercepts_with_codec_inner( +async fn run_request_intercepts_with_codec_inner( name: &str, request: LlmRequest, codec: Option>, @@ -261,7 +278,8 @@ fn run_request_intercepts_with_codec_inner( &entries, codec.is_some(), recorder, - )?; + ) + .await?; let mut request = outcome.request; inject_dynamo_session_ids(&mut request); let pending_marks = outcome.pending_marks; diff --git a/crates/core/src/api/tool.rs b/crates/core/src/api/tool.rs index 6fbe6cd70..7d6d10f71 100644 --- a/crates/core/src/api/tool.rs +++ b/crates/core/src/api/tool.rs @@ -7,12 +7,15 @@ use crate::api::event::{BaseEvent, Event, MarkEvent, PendingMarkSpec}; use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::current_scope_stack; use crate::api::runtime::global_context; +use crate::api::runtime::subscriber_dispatcher::{ + dispatch_sanitized_event, dispatch_transformed_event, +}; use crate::api::runtime::{EventSubscriberFn, ToolExecutionNextFn, with_active_event_uuid}; use crate::api::scope::event; use crate::api::scope::{EmitMarkEventParams, ScopeHandle}; use crate::api::shared::{ ensure_runtime_owner, metadata_with_otel_status, resolve_parent_uuid, sanitize_event, - snapshot_event_subscribers, + snapshot_event_sanitizers, snapshot_event_subscribers, }; use crate::api::skill_load; use crate::error::{FlowError, Result}; @@ -206,11 +209,107 @@ pub struct ToolCallEndParams<'a> { /// Sanitize-request guardrails affect only the emitted start-event payload, not /// the caller-owned `args` value. pub fn tool_call(params: ToolCallParams<'_>) -> Result { - let (handle, _) = tool_call_with_subscriber_snapshot(params)?; + ensure_runtime_owner()?; + let scope_stack = current_scope_stack(); + let (entries, subscribers) = { + let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_locals = scope_guard.collect_scope_local_registries(|registries| { + ®istries.tool_sanitize_request_guardrails + }); + let subscribers = + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + ( + state.tool_sanitize_request_entries(&scope_locals), + subscribers, + ) + }; + let handled_skill_loads = params + .metadata + .as_ref() + .and_then(Json::as_object) + .and_then(|metadata| metadata.get(skill_load::HANDLED_METADATA_KEY)) + .and_then(Json::as_bool) + .is_some_and(|handled| handled); + let skill_loads = if handled_skill_loads { + Vec::new() + } else if let Some(skill_loads) = skill_load::precomputed(params.metadata.as_ref()) { + skill_loads + } else { + skill_load::detect(params.name, ¶ms.args) + }; + let raw_args = params.args; + let (handle, event, marks) = { + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + let handle = state.create_tool_handle( + CreateToolHandleParams::builder() + .name(params.name) + .parent_uuid_opt(resolve_parent_uuid(params.parent)) + .attributes(params.attributes) + .data_opt(params.data) + .metadata_opt(params.metadata) + .tool_call_id_opt(params.tool_call_id) + .timestamp_opt(params.timestamp) + .build(), + ); + let event = state.build_tool_start_event(&handle, None); + let marks = skill_loads + .into_iter() + .map(|skill_load| { + state.create_event(MarkEvent::new( + BaseEvent::builder() + .name("skill.load") + .parent_uuid(handle.uuid) + .timestamp(handle.started_at) + .data(json!({"skill_name": skill_load.name})) + .metadata(json!({ + "skill_load_source": <&str>::from(skill_load.source), + "tool_name": handle.name, + })) + .build(), + None, + None, + )) + }) + .collect::>(); + (handle, event, marks) + }; + let tool_name = handle.name.clone(); + if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { + dispatch_transformed_event( + event, + Box::new(move |mut event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::tool_sanitize_request_snapshot_chain( + &tool_name, raw_args, &entries, + ) + .await; + let mut fields = event.sanitize_fields(); + fields.data = Some(sanitized); + event.apply_sanitize_fields(fields); + event + }) + }), + event_sanitizers, + &subscribers, + scope_stack.clone(), + ); + } + for mark in marks { + if let Some(sanitizers) = snapshot_event_sanitizers(&mark, &scope_stack) { + dispatch_sanitized_event(mark, sanitizers, &subscribers, scope_stack.clone()); + } + } Ok(handle) } -fn tool_call_with_subscriber_snapshot( +async fn tool_call_with_subscriber_snapshot( params: ToolCallParams<'_>, ) -> Result<(ToolHandle, Vec)> { ensure_runtime_owner()?; @@ -248,7 +347,8 @@ fn tool_call_with_subscriber_snapshot( params.name, params.args, &entries, - ); + ) + .await; let (handle, event, marks) = { let context = global_context(); let state = context @@ -286,14 +386,16 @@ fn tool_call_with_subscriber_snapshot( .collect::>(); (handle, event, marks) }; - let marks = marks - .into_iter() - .filter_map(sanitize_event) - .collect::>(); - if let Some(event) = sanitize_event(event) { + let mut sanitized_marks = Vec::with_capacity(marks.len()); + for mark in marks { + if let Some(mark) = sanitize_event(mark).await { + sanitized_marks.push(mark); + } + } + if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, &subscribers); } - for mark in marks { + for mark in sanitized_marks { NemoRelayContextState::emit_event(&mark, &subscribers); } Ok((handle, subscribers)) @@ -326,10 +428,69 @@ fn tool_call_with_subscriber_snapshot( /// Sanitize-response guardrails affect only the emitted end-event payload, not /// the caller-owned `result` value. pub fn tool_call_end(params: ToolCallEndParams<'_>) -> Result<()> { - tool_call_end_with_pending_marks(params, Vec::new(), None) + ensure_runtime_owner()?; + let scope_stack = current_scope_stack(); + let (entries, subscribers) = { + let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_locals = scope_guard.collect_scope_local_registries(|registries| { + ®istries.tool_sanitize_response_guardrails + }); + let subscribers = + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + ( + state.tool_sanitize_response_entries(&scope_locals), + subscribers, + ) + }; + let result = params.result; + let fallback = params.data; + let event = { + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + state.build_tool_end_event( + EndToolHandleParams::builder() + .handle(params.handle) + .data(Json::Null) + .metadata_opt(params.metadata) + .timestamp_opt(params.timestamp) + .build(), + ) + }; + let tool_name = params.handle.name.clone(); + if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { + dispatch_transformed_event( + event, + Box::new(move |mut event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::tool_sanitize_response_snapshot_chain( + &tool_name, result, &entries, + ) + .await; + let mut fields = event.sanitize_fields(); + fields.data = if sanitized.is_null() { + fallback + } else { + Some(sanitized) + }; + event.apply_sanitize_fields(fields); + event + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); + } + Ok(()) } -fn tool_call_end_with_pending_marks( +async fn tool_call_end_with_pending_marks( params: ToolCallEndParams<'_>, pending_marks: Vec, lifecycle_subscribers: Option<&[EventSubscriberFn]>, @@ -358,7 +519,8 @@ fn tool_call_end_with_pending_marks( ¶ms.handle.name, params.result, &entries, - ); + ) + .await; let data = if sanitized_result.is_null() { params.data } else { @@ -396,18 +558,23 @@ fn tool_call_end_with_pending_marks( mark.category_profile, )) }) - .filter_map(sanitize_event) .collect::>(); - if let Some(event) = sanitize_event(event) { + let mut sanitized_marks = Vec::with_capacity(marks.len()); + for mark in marks { + if let Some(mark) = sanitize_event(mark).await { + sanitized_marks.push(mark); + } + } + if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, subscribers); } - for mark in marks { + for mark in sanitized_marks { NemoRelayContextState::emit_event(&mark, subscribers); } Ok(()) } -fn emit_tool_end_without_output( +async fn emit_tool_end_without_output( handle: &ToolHandle, metadata: Option, lifecycle_subscribers: &[EventSubscriberFn], @@ -420,7 +587,7 @@ fn emit_tool_end_without_output( .map_err(|error| FlowError::Internal(error.to_string()))?; state.end_tool_handle(handle, handle.data.clone(), metadata) }; - if let Some(event) = sanitize_event(event) { + if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, lifecycle_subscribers); } Ok(()) @@ -493,7 +660,9 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { &subscribers, parent_uuid, guardrail_metadata, - )? { + ) + .await? + { let mut rejection_data = json!({}); if let Some(object) = rejection_data.as_object_mut() { object.insert("rejected".into(), json!(true)); @@ -526,7 +695,8 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { &name, args, &intercept_entries, - )?; + ) + .await?; let (handle, lifecycle_subscribers) = tool_call_with_subscriber_snapshot( ToolCallParams::builder() @@ -537,7 +707,8 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { .data_opt(data.clone()) .metadata_opt(metadata.clone()) .build(), - )?; + ) + .await?; let execution = { let scope_stack = current_scope_stack(); @@ -567,13 +738,15 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { .build(), pending_marks, Some(&lifecycle_subscribers), - )?; + ) + .await?; Ok(result) } Err(error) => { let end_metadata = metadata_with_otel_status(metadata, "ERROR", Some(error.to_string())); - let _ = emit_tool_end_without_output(&handle, end_metadata, &lifecycle_subscribers); + let _ = + emit_tool_end_without_output(&handle, end_metadata, &lifecycle_subscribers).await; Err(error) } } @@ -596,7 +769,7 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { /// /// # Notes /// Conditional guardrails and execution intercepts are not run by this helper. -pub fn tool_request_intercepts(name: &str, args: Json) -> Result { +pub async fn tool_request_intercepts(name: &str, args: Json) -> Result { ensure_runtime_owner()?; let entries = { let scope_stack = current_scope_stack(); @@ -609,7 +782,7 @@ pub fn tool_request_intercepts(name: &str, args: Json) -> Result { .map_err(|error| FlowError::Internal(error.to_string()))?; state.tool_request_intercept_entries(&scope_locals) }; - NemoRelayContextState::tool_request_intercepts_snapshot_chain(name, args, &entries) + NemoRelayContextState::tool_request_intercepts_snapshot_chain(name, args, &entries).await } /// Run only the tool conditional-execution guardrail chain. @@ -633,7 +806,7 @@ pub fn tool_request_intercepts(name: &str, args: Json) -> Result { /// This helper is useful for preflight checks when the caller needs the /// rejection result without starting a tool span. Guardrail scopes are still /// emitted for the conditional checks themselves. -pub fn tool_conditional_execution(name: &str, args: &Json) -> Result<()> { +pub async fn tool_conditional_execution(name: &str, args: &Json) -> Result<()> { ensure_runtime_owner()?; let (entries, subscribers, parent_uuid) = { let scope_stack = current_scope_stack(); @@ -657,7 +830,9 @@ pub fn tool_conditional_execution(name: &str, args: &Json) -> Result<()> { &subscribers, parent_uuid, None, - )? { + ) + .await? + { return Err(FlowError::GuardrailRejected(error)); } Ok(()) diff --git a/crates/core/src/logging/rotation.rs b/crates/core/src/logging/rotation.rs index 3314377c1..fa55780bb 100644 --- a/crates/core/src/logging/rotation.rs +++ b/crates/core/src/logging/rotation.rs @@ -146,3 +146,7 @@ pub(crate) fn rotated_log_path(base_path: &Path, index: usize) -> PathBuf { } base_path.with_file_name(file_name) } + +#[cfg(test)] +#[path = "../../tests/coverage/logging_rotation_tests.rs"] +mod tests; diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index f1217c746..7481d86ee 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -14,24 +14,27 @@ use std::panic::{AssertUnwindSafe, catch_unwind}; use std::path::{Path, PathBuf}; use std::pin::Pin; use std::ptr; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, OnceLock}; use std::task::{Context, Poll}; use chrono::{DateTime, Utc}; use libloading::{Library, Symbol}; use nemo_relay_plugin::{ - NEMO_RELAY_NATIVE_ABI_VERSION, NemoRelayNativeEventSanitizeCb, - NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, - NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, - NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, - NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, - NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, - NemoRelayNativeLlmStreamV1, NemoRelayNativePluginContext, NemoRelayNativePluginEntry, - NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, - NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, NemoRelayNativeString, - NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, NemoRelayNativeToolJsonCb, - NemoRelayNativeWithScopeStackCb, NemoRelayStatus, + NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, + NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncCompletion, + NemoRelayNativeAsyncMiddlewareCb, NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, + NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, + NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeLlmCodecKind, + NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmRequestCodec, + NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, + NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, + NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, + NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativePluginContext, + NemoRelayNativePluginEntry, NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, + NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, + NemoRelayNativeString, NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, + NemoRelayNativeToolJsonCb, NemoRelayNativeWithScopeStackCb, NemoRelayStatus, }; use semver::{Version, VersionReq}; use serde_json::{Map, Value as Json}; @@ -374,7 +377,14 @@ fn load_one_native_plugin( library_path.display() )) })?; - let status = entry(native_host_api(), &mut plugin); + let mut status = entry(native_host_api(), &mut plugin); + // SDKs compiled against ABI v2 correctly reject a v3 table. Retry + // their entry point with the frozen v2 prefix instead of making a + // runtime upgrade a breaking change for installed native plugins. + if status == NemoRelayStatus::InvalidArg { + drop_native_plugin_descriptor(&mut plugin); + status = entry(native_host_api_legacy(), &mut plugin); + } if status != NemoRelayStatus::Ok { drop_native_plugin_descriptor(&mut plugin); return Err(PluginError::RegistrationFailed(format!( @@ -784,10 +794,19 @@ unsafe extern "C" fn native_llm_response_codec_decode( } fn native_host_api() -> *const NemoRelayNativeHostApiV1 { + static HOST_API: OnceLock = OnceLock::new(); + &HOST_API.get_or_init(build_native_host_api_v3).v1 as *const NemoRelayNativeHostApiV1 +} + +fn native_host_api_legacy() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); + HOST_API.get_or_init(build_native_host_api_legacy) as *const _ +} + +fn build_native_host_api_legacy() -> NemoRelayNativeHostApiV1 { static RELAY_VERSION: &[u8] = concat!(env!("CARGO_PKG_VERSION"), "\0").as_bytes(); - HOST_API.get_or_init(|| NemoRelayNativeHostApiV1 { - abi_version: NEMO_RELAY_NATIVE_ABI_VERSION, + NemoRelayNativeHostApiV1 { + abi_version: NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, struct_size: std::mem::size_of::(), relay_version: RELAY_VERSION.as_ptr().cast(), string_new: native_string_new, @@ -841,7 +860,23 @@ fn native_host_api() -> *const NemoRelayNativeHostApiV1 { native_plugin_context_register_scope_sanitize_start_guardrail, plugin_context_register_scope_sanitize_end_guardrail: native_plugin_context_register_scope_sanitize_end_guardrail, - }) as *const _ + } +} + +fn build_native_host_api_v3() -> NemoRelayNativeHostApiV3 { + let mut v1 = build_native_host_api_legacy(); + v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION; + v1.struct_size = std::mem::size_of::(); + NemoRelayNativeHostApiV3 { + v1, + async_completion_resolve_json: native_async_completion_resolve_json, + async_completion_reject: native_async_completion_reject, + async_completion_is_cancelled: native_async_completion_is_cancelled, + async_completion_release: native_async_completion_release, + async_next_invoke: native_async_next_invoke, + async_next_release: native_async_next_release, + plugin_context_register_async_middleware: native_plugin_context_register_async_middleware, + } } fn read_native_string(value: *const NemoRelayNativeString) -> crate::plugin::Result { @@ -1305,6 +1340,680 @@ fn make_user_data( }) } +/// One-shot state retained by a v3 native async callback. +struct NativeAsyncCompletion { + sender: Mutex>>>, + cancelled: AtomicBool, + // A pending native callback can continue running after its completion + // wakes the awaiting task. Keep the callback's dynamic-library instance + // alive until native code explicitly releases this handle. + _callback_user_data: Option>, +} + +struct NativeAsyncWait { + completion: Arc, + receiver: tokio::sync::oneshot::Receiver>, +} + +impl Drop for NativeAsyncWait { + fn drop(&mut self) { + self.completion.cancelled.store(true, Ordering::Release); + } +} + +enum NativeAsyncNextInner { + Tool(ToolExecutionNextFn), + Llm(LlmExecutionNextFn), + LlmStream(LlmStreamExecutionNextFn), +} + +struct NativeAsyncNext { + inner: NativeAsyncNextInner, + runtime: tokio::runtime::Handle, + // The native callback owns this handle independently of its completion. + // Retaining the library here prevents an unload while it still uses `next`. + _callback_user_data: Option>, +} + +async fn invoke_native_async_callback( + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: Arc, + invocation: Json, + next: Option, +) -> FlowResult { + let runtime = if next.is_some() { + Some(tokio::runtime::Handle::try_current().map_err(|error| { + FlowError::Internal(format!( + "native async intercept requires a Tokio runtime: {error}" + )) + })?) + } else { + None + }; + let invocation = native_string_from_json(&invocation) + .ok_or_else(|| FlowError::Internal("failed to allocate native async invocation".into()))? + as usize; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + _callback_user_data: Some(user_data.clone()), + }); + let completion_ref = Arc::into_raw(completion.clone()) as usize; + let next_ref = match (next, runtime) { + (Some(inner), Some(runtime)) => Some(Arc::into_raw(Arc::new(NativeAsyncNext { + inner, + runtime, + _callback_user_data: Some(user_data.clone()), + })) as usize), + (None, None) => None, + _ => unreachable!("runtime is present exactly for native async intercepts"), + }; + let state = match catch_unwind(AssertUnwindSafe(|| unsafe { + cb( + user_data.ptr, + invocation as *const NemoRelayNativeString, + next_ref + .map(|next| next as *const NemoRelayNativeAsyncNext) + .unwrap_or(ptr::null()), + completion_ref as *const NemoRelayNativeAsyncCompletion, + ) + })) { + Ok(state) => state, + Err(_) => { + unsafe { + drop(Arc::from_raw( + completion_ref as *const NativeAsyncCompletion, + )); + if let Some(next_ref) = next_ref { + drop(Arc::from_raw(next_ref as *const NativeAsyncNext)); + } + native_string_free(invocation as *mut NemoRelayNativeString); + } + return Err(FlowError::Internal("native async callback panicked".into())); + } + }; + unsafe { native_string_free(invocation as *mut NemoRelayNativeString) }; + if state == NemoRelayNativeAsyncCallbackState::Complete { + unsafe { + drop(Arc::from_raw( + completion_ref as *const NativeAsyncCompletion, + )); + if let Some(next_ref) = next_ref { + drop(Arc::from_raw(next_ref as *const NativeAsyncNext)); + } + } + if completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .is_some() + { + return Err(FlowError::Internal( + "native async callback returned Complete without settling".into(), + )); + } + } + let mut wait = NativeAsyncWait { + completion, + receiver, + }; + (&mut wait.receiver) + .await + .map_err(|_| FlowError::Internal("native async callback dropped without settling".into()))? +} + +unsafe extern "C" fn native_async_completion_resolve_json( + completion: *const NemoRelayNativeAsyncCompletion, + value_json: *const NemoRelayNativeString, +) -> NemoRelayStatus { + let Some(completion) = (unsafe { (completion as *const NativeAsyncCompletion).as_ref() }) + else { + return NemoRelayStatus::NullPointer; + }; + if completion.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let value = match parse_json_arg(value_json, "native async completion result") { + Ok(value) => value, + Err(status) => return status, + }; + let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + else { + return NemoRelayStatus::InvalidArg; + }; + let _ = sender.send(Ok(value)); + NemoRelayStatus::Ok +} + +unsafe extern "C" fn native_async_completion_reject( + completion: *const NemoRelayNativeAsyncCompletion, + message: *const NemoRelayNativeString, +) -> NemoRelayStatus { + let Some(completion) = (unsafe { (completion as *const NativeAsyncCompletion).as_ref() }) + else { + return NemoRelayStatus::NullPointer; + }; + if completion.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let message = if message.is_null() { + "native async callback rejected".to_string() + } else { + match read_native_string(message) { + Ok(message) => message, + Err(error) => { + set_native_last_error(error.to_string()); + return NemoRelayStatus::InvalidArg; + } + } + }; + let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + else { + return NemoRelayStatus::InvalidArg; + }; + let _ = sender.send(Err(FlowError::Internal(message))); + NemoRelayStatus::Ok +} + +unsafe extern "C" fn native_async_completion_is_cancelled( + completion: *const NemoRelayNativeAsyncCompletion, +) -> bool { + unsafe { (completion as *const NativeAsyncCompletion).as_ref() } + .is_none_or(|completion| completion.cancelled.load(Ordering::Acquire)) +} + +unsafe extern "C" fn native_async_completion_release( + completion: *const NemoRelayNativeAsyncCompletion, +) { + if !completion.is_null() { + unsafe { drop(Arc::from_raw(completion as *const NativeAsyncCompletion)) }; + } +} + +unsafe extern "C" fn native_async_next_release(next: *const NemoRelayNativeAsyncNext) { + if !next.is_null() { + unsafe { drop(Arc::from_raw(next as *const NativeAsyncNext)) }; + } +} + +/// Invokes the runtime continuation without blocking the calling native thread. +unsafe extern "C" fn native_async_next_invoke( + next: *const NemoRelayNativeAsyncNext, + invocation_json: *const NemoRelayNativeString, + completion: *const NemoRelayNativeAsyncCompletion, +) -> NemoRelayStatus { + let Some(next) = (unsafe { (next as *const NativeAsyncNext).as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if completion.is_null() { + return NemoRelayStatus::NullPointer; + } + let invocation = match parse_json_arg(invocation_json, "native async next invocation") { + Ok(value) => value, + Err(status) => return status, + }; + unsafe { Arc::increment_strong_count(completion as *const NativeAsyncCompletion) }; + let completion = unsafe { Arc::from_raw(completion as *const NativeAsyncCompletion) }; + let future: Pin> + Send>> = match &next.inner { + NativeAsyncNextInner::Tool(next) => { + let next = next.clone(); + Box::pin(async move { + serde_json::to_value(ToolExecutionInterceptOutcome::new(next(invocation).await?)) + .map_err(|error| { + FlowError::Internal(format!( + "failed to serialize native async tool outcome: {error}" + )) + }) + }) + } + NativeAsyncNextInner::Llm(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(error) => { + let _ = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); + return NemoRelayStatus::InvalidArg; + } + }; + let next = next.clone(); + Box::pin(async move { next(request).await }) + } + NativeAsyncNextInner::LlmStream(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(error) => { + let _ = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); + return NemoRelayStatus::InvalidArg; + } + }; + let next = next.clone(); + Box::pin(async move { + let mut stream = next(request).await?; + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk?); + } + Ok(Json::Array(chunks)) + }) + } + }; + next.runtime.spawn(async move { + let result = future.await; + if let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + { + let _ = sender.send(result); + } + }); + NemoRelayStatus::Ok +} + +fn wrap_native_async_tool_json( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> ToolSanitizeFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, value| { + let user_data = user_data.clone(); + Box::pin(async move { + let value = invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "value": value}), + None, + ) + .await?; + Ok(value) + }) + }) +} + +fn wrap_native_async_tool_conditional( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> ToolConditionalFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, value| { + let user_data = user_data.clone(); + Box::pin(async move { + match invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "value": value}), + None, + ) + .await? + { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => Err(FlowError::Internal(format!( + "native async tool conditional callback returned {other}; expected string or null" + ))), + } + }) + }) +} + +fn wrap_native_async_llm_conditional( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmConditionalFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |request| { + let user_data = user_data.clone(); + Box::pin(async move { + match invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"request": request}), + None, + ) + .await? + { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => Err(FlowError::Internal(format!( + "native async LLM conditional callback returned {other}; expected string or null" + ))), + } + }) + }) +} + +fn wrap_native_async_llm_sanitize_request( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmSanitizeRequestFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |request, context| { + let user_data = user_data.clone(); + let codec = format!("{:?}", context.codec()); + Box::pin(async move { + let value = invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"request": request, "context": {"codec": codec}}), + None, + ) + .await?; + if value.is_null() { + Ok(None) + } else { + serde_json::from_value(value) + .map(Some) + .map_err(|error| FlowError::Internal(error.to_string())) + } + }) + }) +} + +fn wrap_native_async_llm_sanitize_response( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmSanitizeResponseFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |response, context| { + let user_data = user_data.clone(); + let codec = format!("{:?}", context.codec()); + Box::pin(async move { + let value = invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"response": response, "context": {"codec": codec}}), + None, + ) + .await?; + Ok((!value.is_null()).then_some(value)) + }) + }) +} + +fn wrap_native_async_llm_request_intercept( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmRequestInterceptFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, request, annotated| { + let user_data = user_data.clone(); + Box::pin(async move { + serde_json::from_value( + invoke_native_async_callback( + cb, + user_data, + serde_json::json!({ + "name": name, + "request": request, + "annotated": annotated, + }), + None, + ) + .await?, + ) + .map_err(|error| { + FlowError::Internal(format!( + "invalid native async LLM intercept outcome: {error}" + )) + }) + }) + }) +} + +fn wrap_native_async_event_sanitize( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> EventSanitizeFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |event, fields| { + let user_data = user_data.clone(); + Box::pin(async move { + serde_json::from_value( + invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"event": event, "fields": fields}), + None, + ) + .await?, + ) + .map_err(|error| { + FlowError::Internal(format!("invalid native async event fields: {error}")) + }) + }) + }) +} + +fn wrap_native_async_tool_execution( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> ToolExecutionFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, args, next| { + let user_data = user_data.clone(); + let invocation = serde_json::json!({"name": name, "value": args}); + Box::pin(async move { + serde_json::from_value( + invoke_native_async_callback( + cb, + user_data, + invocation, + Some(NativeAsyncNextInner::Tool(next)), + ) + .await?, + ) + .map_err(|error| { + FlowError::Internal(format!("invalid native async tool outcome: {error}")) + }) + }) + }) +} + +fn wrap_native_async_llm_execution( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmExecutionFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, request, next| { + let user_data = user_data.clone(); + Box::pin(invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "request": request}), + Some(NativeAsyncNextInner::Llm(next)), + )) + }) +} + +fn wrap_native_async_llm_stream_execution( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmStreamExecutionFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, request, next| { + let user_data = user_data.clone(); + let name = name.to_owned(); + Box::pin(async move { + let value = invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "request": request}), + Some(NativeAsyncNextInner::LlmStream(next)), + ) + .await?; + let chunks = value.as_array().cloned().ok_or_else(|| { + FlowError::Internal( + "native async LLM stream intercept must resolve to an array".into(), + ) + })?; + Ok(LlmJsonStream::new(tokio_stream::iter( + chunks.into_iter().map(Ok), + ))) + }) + }) +} + +unsafe extern "C" fn native_plugin_context_register_async_middleware( + ctx: *mut NemoRelayNativePluginContext, + kind: NemoRelayNativeAsyncMiddlewareKind, + name: *const NemoRelayNativeString, + priority: i32, + break_chain: bool, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> NemoRelayStatus { + clear_native_last_error(); + let host_ctx = match host_ctx_mut(ctx) { + Ok(ctx) => ctx, + Err(status) => return status, + }; + let instance = host_ctx.instance.clone(); + let name = match read_name(name) { + Ok(name) => name, + Err(status) => return status, + }; + let context = unsafe { &mut *host_ctx.ctx }; + let registration = match kind { + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest => context + .register_tool_sanitize_request_guardrail( + &name, + priority, + wrap_native_async_tool_json(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeResponse => context + .register_tool_sanitize_response_guardrail( + &name, + priority, + wrap_native_async_tool_json(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ToolConditionalExecution => context + .register_tool_conditional_execution_guardrail( + &name, + priority, + wrap_native_async_tool_conditional(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ToolRequestIntercept => context + .register_tool_request_intercept( + &name, + priority, + break_chain, + wrap_native_async_tool_json(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ToolExecutionIntercept => context + .register_tool_execution_intercept( + &name, + priority, + wrap_native_async_tool_execution(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeRequest => context + .register_llm_sanitize_request_guardrail( + &name, + priority, + wrap_native_async_llm_sanitize_request(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeResponse => context + .register_llm_sanitize_response_guardrail( + &name, + priority, + wrap_native_async_llm_sanitize_response(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmConditionalExecution => context + .register_llm_conditional_execution_guardrail( + &name, + priority, + wrap_native_async_llm_conditional(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept => { + if let Err(error) = validate_annotated_request_consumer_compatibility( + &instance.relay_compat, + &instance.plugin_kind, + ) { + return status_from_plugin_error(error); + } + context.register_llm_request_intercept( + &name, + priority, + break_chain, + wrap_native_async_llm_request_intercept(instance, cb, user_data, free_fn), + ) + } + NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept => context + .register_llm_execution_intercept( + &name, + priority, + wrap_native_async_llm_execution(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept => context + .register_llm_stream_execution_intercept( + &name, + priority, + wrap_native_async_llm_stream_execution(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::MarkSanitize => context + .register_mark_sanitize_guardrail( + &name, + priority, + wrap_native_async_event_sanitize(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeStart => context + .register_scope_sanitize_start_guardrail( + &name, + priority, + wrap_native_async_event_sanitize(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeEnd => context + .register_scope_sanitize_end_guardrail( + &name, + priority, + wrap_native_async_event_sanitize(instance, cb, user_data, free_fn), + ), + }; + match registration { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_plugin_error(error), + } +} + fn host_ctx_mut<'a>( ctx: *mut NemoRelayNativePluginContext, ) -> Result<&'a mut NativeHostPluginContext, NemoRelayStatus> { @@ -1740,7 +2449,8 @@ fn wrap_event_sanitize_fn( ) -> EventSanitizeFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |event, fields| { - call_event_sanitize_callback(cb, user_data.ptr, event, &fields).unwrap_or_default() + let user_data = user_data.clone(); + Box::pin(async move { call_event_sanitize_callback(cb, user_data.ptr, &event, &fields) }) }) } @@ -1795,7 +2505,8 @@ fn wrap_tool_json_fn( ) -> ToolSanitizeFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |name, payload| { - call_tool_json_callback(cb, user_data.ptr, name, &payload).unwrap_or(Json::Null) + let user_data = user_data.clone(); + Box::pin(async move { call_tool_json_callback(cb, user_data.ptr, &name, &payload) }) }) } @@ -1806,7 +2517,10 @@ fn wrap_tool_intercept_fn( free_fn: NemoRelayNativeFreeFn, ) -> ToolInterceptFn { let user_data = make_user_data(instance, user_data, free_fn); - Arc::new(move |name, payload| call_tool_json_callback(cb, user_data.ptr, name, &payload)) + Arc::new(move |name, payload| { + let user_data = user_data.clone(); + Box::pin(async move { call_tool_json_callback(cb, user_data.ptr, &name, &payload) }) + }) } fn call_tool_json_callback( @@ -1846,32 +2560,35 @@ fn wrap_tool_conditional_fn( ) -> ToolConditionalFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |name, args| { - clear_native_last_error(); - let name_string = native_string_from_str(name) - .ok_or_else(|| FlowError::Internal("failed to allocate native name".into()))?; - let args_string = native_string_from_json(args) - .ok_or_else(|| FlowError::Internal("failed to allocate native args".into()))?; - let mut out = ptr::null_mut(); - let status = unsafe { cb(user_data.ptr, name_string, args_string, &mut out) }; - unsafe { - native_string_free(name_string); - native_string_free(args_string); - } - if status != NemoRelayStatus::Ok { - if !out.is_null() { - unsafe { native_string_free(out) }; + let user_data = user_data.clone(); + Box::pin(async move { + clear_native_last_error(); + let name_string = native_string_from_str(&name) + .ok_or_else(|| FlowError::Internal("failed to allocate native name".into()))?; + let args_string = native_string_from_json(&args) + .ok_or_else(|| FlowError::Internal("failed to allocate native args".into()))?; + let mut out = ptr::null_mut(); + let status = unsafe { cb(user_data.ptr, name_string, args_string, &mut out) }; + unsafe { + native_string_free(name_string); + native_string_free(args_string); } - return Err(flow_error_from_status( - status, - "native tool conditional failed", - )); - } - if out.is_null() { - Ok(None) - } else { - let reason = take_native_string(out)?; - Ok(Some(reason)) - } + if status != NemoRelayStatus::Ok { + if !out.is_null() { + unsafe { native_string_free(out) }; + } + return Err(flow_error_from_status( + status, + "native tool conditional failed", + )); + } + if out.is_null() { + Ok(None) + } else { + let reason = take_native_string(out)?; + Ok(Some(reason)) + } + }) }) } @@ -1961,9 +2678,10 @@ fn wrap_llm_sanitize_request_fn( ) -> LlmSanitizeRequestFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |request, context| { - call_llm_sanitize_request_callback(cb, user_data.ptr, &request, context) - .ok() - .flatten() + let user_data = user_data.clone(); + Box::pin( + async move { call_llm_sanitize_request_callback(cb, user_data.ptr, &request, context) }, + ) }) } @@ -1975,9 +2693,10 @@ fn wrap_llm_sanitize_response_fn( ) -> LlmSanitizeResponseFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |payload, context| { - call_llm_sanitize_response_callback(cb, user_data.ptr, &payload, context) - .ok() - .flatten() + let user_data = user_data.clone(); + Box::pin(async move { + call_llm_sanitize_response_callback(cb, user_data.ptr, &payload, context) + }) }) } @@ -2118,30 +2837,34 @@ fn wrap_llm_conditional_fn( ) -> LlmConditionalFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |request| { - clear_native_last_error(); - let request_json = serde_json::to_value(request).map_err(|err| { - FlowError::Internal(format!("failed to serialize LLM request: {err}")) - })?; - let request_string = native_string_from_json(&request_json) - .ok_or_else(|| FlowError::Internal("failed to allocate native LLM request".into()))?; - let mut out = ptr::null_mut(); - let status = unsafe { cb(user_data.ptr, request_string, &mut out) }; - unsafe { native_string_free(request_string) }; - if status != NemoRelayStatus::Ok { - if !out.is_null() { - unsafe { native_string_free(out) }; + let user_data = user_data.clone(); + Box::pin(async move { + clear_native_last_error(); + let request_json = serde_json::to_value(request).map_err(|err| { + FlowError::Internal(format!("failed to serialize LLM request: {err}")) + })?; + let request_string = native_string_from_json(&request_json).ok_or_else(|| { + FlowError::Internal("failed to allocate native LLM request".into()) + })?; + let mut out = ptr::null_mut(); + let status = unsafe { cb(user_data.ptr, request_string, &mut out) }; + unsafe { native_string_free(request_string) }; + if status != NemoRelayStatus::Ok { + if !out.is_null() { + unsafe { native_string_free(out) }; + } + return Err(flow_error_from_status( + status, + "native LLM conditional failed", + )); } - return Err(flow_error_from_status( - status, - "native LLM conditional failed", - )); - } - if out.is_null() { - Ok(None) - } else { - let reason = take_native_string(out)?; - Ok(Some(reason)) - } + if out.is_null() { + Ok(None) + } else { + let reason = take_native_string(out)?; + Ok(Some(reason)) + } + }) }) } @@ -2153,58 +2876,62 @@ fn wrap_llm_request_intercept_fn( ) -> LlmRequestInterceptFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |name, request, annotated| { - clear_native_last_error(); - let name_string = native_string_from_str(name) - .ok_or_else(|| FlowError::Internal("failed to allocate native name".into()))?; - let request_json = serde_json::to_value(&request).map_err(|err| { - FlowError::Internal(format!("failed to serialize LLM request: {err}")) - })?; - let request_string = native_string_from_json(&request_json) - .ok_or_else(|| FlowError::Internal("failed to allocate native LLM request".into()))?; - let annotated_string = match &annotated { - Some(annotated) => { - let value = serde_json::to_value(annotated).map_err(|err| { - FlowError::Internal(format!("failed to serialize annotated request: {err}")) - })?; - native_string_from_json(&value).ok_or_else(|| { - FlowError::Internal("failed to allocate annotated request".into()) - })? + let user_data = user_data.clone(); + Box::pin(async move { + clear_native_last_error(); + let name_string = native_string_from_str(&name) + .ok_or_else(|| FlowError::Internal("failed to allocate native name".into()))?; + let request_json = serde_json::to_value(&request).map_err(|err| { + FlowError::Internal(format!("failed to serialize LLM request: {err}")) + })?; + let request_string = native_string_from_json(&request_json).ok_or_else(|| { + FlowError::Internal("failed to allocate native LLM request".into()) + })?; + let annotated_string = match &annotated { + Some(annotated) => { + let value = serde_json::to_value(annotated).map_err(|err| { + FlowError::Internal(format!("failed to serialize annotated request: {err}")) + })?; + native_string_from_json(&value).ok_or_else(|| { + FlowError::Internal("failed to allocate annotated request".into()) + })? + } + None => ptr::null_mut(), + }; + let mut out_outcome = ptr::null_mut(); + let status = unsafe { + cb( + user_data.ptr, + name_string, + request_string, + annotated_string, + &mut out_outcome, + ) + }; + unsafe { + native_string_free(name_string); + native_string_free(request_string); + native_string_free(annotated_string); } - None => ptr::null_mut(), - }; - let mut out_outcome = ptr::null_mut(); - let status = unsafe { - cb( - user_data.ptr, - name_string, - request_string, - annotated_string, - &mut out_outcome, - ) - }; - unsafe { - native_string_free(name_string); - native_string_free(request_string); - native_string_free(annotated_string); - } - if status != NemoRelayStatus::Ok { + if status != NemoRelayStatus::Ok { + unsafe { + native_string_free(out_outcome); + } + return Err(flow_error_from_status( + status, + "native LLM request intercept failed", + )); + } + let outcome_json = json_from_native_string( + out_outcome, + "native LLM request intercept returned null outcome", + ); unsafe { native_string_free(out_outcome); } - return Err(flow_error_from_status( - status, - "native LLM request intercept failed", - )); - } - let outcome_json = json_from_native_string( - out_outcome, - "native LLM request intercept returned null outcome", - ); - unsafe { - native_string_free(out_outcome); - } - serde_json::from_value::(outcome_json?).map_err(|err| { - FlowError::Internal(format!("invalid LLM request intercept outcome JSON: {err}")) + serde_json::from_value::(outcome_json?).map_err(|err| { + FlowError::Internal(format!("invalid LLM request intercept outcome JSON: {err}")) + }) }) }) } diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index 79e2e70ce..5262dd784 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -59,9 +59,9 @@ use tower::service_fn; use crate::api::event::{Event, EventSanitizeFields}; use crate::api::llm::{LLM_REQUEST_INTERCEPT_OUTCOME_SCHEMA, LlmRequest}; use crate::api::runtime::{ - LlmCodecIdentity, LlmExecutionNextFn, LlmJsonStream, LlmSanitizeRequestContext, - LlmSanitizeResponseContext, LlmStreamExecutionNextFn, ToolExecutionNextFn, current_scope_stack, - with_scope_stack, + EventSanitizeFn, LlmCodecIdentity, LlmExecutionNextFn, LlmJsonStream, + LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn, + ToolExecutionNextFn, current_scope_stack, with_scope_stack, }; use crate::api::scope::{ EmitMarkEventParams, PopScopeParams, PushScopeParams, ScopeAttributes, ScopeHandle, ScopeType, @@ -1119,14 +1119,16 @@ impl WorkerPluginInstance { ) -> crate::plugin::Result<()> { let instance = Arc::new(self.clone_for_callback()); let callback_name = name.to_owned(); - let callback = Arc::new(move |event: &Event, _fields: EventSanitizeFields| { - instance - .invoke_event_sanitize(&callback_name, surface, event) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - EventSanitizeFields::default() + let callback: EventSanitizeFn = + Arc::new(move |event: Event, _fields: EventSanitizeFields| { + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_event_sanitize(&callback_name, surface, &event) + .await }) - }); + }); match surface { RegistrationSurface::MarkSanitizeGuardrail => { ctx.register_mark_sanitize_guardrail(name, priority, callback) @@ -1157,18 +1159,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |tool_name, value| { - instance - .invoke_tool_json( - &callback_name, - surface, - tool_name, - value.clone(), - None, - ) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - value - }) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_tool_json(&callback_name, surface, &tool_name, value, None) + .await + }) }), ), RegistrationSurface::ToolSanitizeResponseGuardrail => ctx @@ -1176,18 +1173,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |tool_name, value| { - instance - .invoke_tool_json( - &callback_name, - surface, - tool_name, - value.clone(), - None, - ) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - value - }) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_tool_json(&callback_name, surface, &tool_name, value, None) + .await + }) }), ), RegistrationSurface::ToolConditionalExecutionGuardrail => ctx @@ -1195,7 +1187,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |tool_name, value| { - instance.invoke_tool_guardrail(&callback_name, tool_name, value.clone()) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_tool_guardrail(&callback_name, &tool_name, value) + .await + }) }), ), RegistrationSurface::ToolRequestIntercept => ctx.register_tool_request_intercept( @@ -1203,7 +1201,13 @@ impl WorkerPluginInstance { priority, registration.break_chain, Arc::new(move |tool_name, value| { - instance.invoke_tool_json(&callback_name, surface, tool_name, value, None) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_tool_json(&callback_name, surface, &tool_name, value, None) + .await + }) }), ), RegistrationSurface::ToolExecutionIntercept => ctx.register_tool_execution_intercept( @@ -1244,12 +1248,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |request, context| { - instance - .invoke_llm_sanitize_request(&callback_name, request.clone(), context) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - None - }) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_llm_sanitize_request(&callback_name, request, context) + .await + }) }), ), RegistrationSurface::LlmSanitizeResponseGuardrail => ctx @@ -1257,12 +1262,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |value, context| { - instance - .invoke_llm_sanitize_response(&callback_name, value.clone(), context) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - None - }) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_llm_sanitize_response(&callback_name, value, context) + .await + }) }), ), RegistrationSurface::LlmConditionalExecutionGuardrail => ctx @@ -1270,7 +1276,11 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |request| { - instance.invoke_llm_guardrail(&callback_name, request.clone()) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance.invoke_llm_guardrail(&callback_name, request).await + }) }), ), RegistrationSurface::LlmRequestIntercept => ctx.register_llm_request_intercept( @@ -1278,12 +1288,18 @@ impl WorkerPluginInstance { priority, registration.break_chain, Arc::new(move |model_name, request, annotated| { - instance.invoke_llm_request_intercept( - &callback_name, - model_name, - request, - annotated, - ) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_llm_request_intercept( + &callback_name, + &model_name, + request, + annotated, + ) + .await + }) }), ), RegistrationSurface::LlmExecutionIntercept => ctx.register_llm_execution_intercept( @@ -1476,7 +1492,7 @@ impl WorkerPluginCallback { } } - fn invoke_event_sanitize( + async fn invoke_event_sanitize( &self, registration_name: &str, surface: RegistrationSurface, @@ -1488,7 +1504,7 @@ impl WorkerPluginCallback { None, Some(invoke_request_payload_event(event)), ); - let value = json_from_invoke_response(self.invoke_blocking(request)?)?; + let value = json_from_invoke_response(self.invoke_async(request).await?)?; serde_json::from_value(value).map_err(|err| { FlowError::Internal(format!( "worker returned invalid event sanitize fields: {err}" @@ -1496,7 +1512,7 @@ impl WorkerPluginCallback { }) } - fn invoke_tool_json( + async fn invoke_tool_json( &self, registration_name: &str, surface: RegistrationSurface, @@ -1510,10 +1526,10 @@ impl WorkerPluginCallback { continuation_id, Some(invoke_request_payload_tool(tool_name, value)), ); - json_from_invoke_response(self.invoke_blocking(request)?) + json_from_invoke_response(self.invoke_async(request).await?) } - fn invoke_tool_guardrail( + async fn invoke_tool_guardrail( &self, registration_name: &str, tool_name: &str, @@ -1525,7 +1541,7 @@ impl WorkerPluginCallback { None, Some(invoke_request_payload_tool(tool_name, value)), ); - guardrail_from_invoke_response(self.invoke_blocking(request)?) + guardrail_from_invoke_response(self.invoke_async(request).await?) } async fn invoke_tool_execution( @@ -1568,7 +1584,7 @@ impl WorkerPluginCallback { } } - fn invoke_llm_sanitize_request( + async fn invoke_llm_sanitize_request( &self, registration_name: &str, request: LlmRequest, @@ -1606,7 +1622,7 @@ impl WorkerPluginCallback { context.codec_capability_id = Some(capability_id.clone()); capability_id }); - let response = self.invoke_blocking(invoke); + let response = self.invoke_async(invoke).await; if let Some(capability_id) = capability_id { self.host_state.remove_codec(&capability_id); } @@ -1618,7 +1634,7 @@ impl WorkerPluginCallback { }) } - fn invoke_llm_sanitize_response( + async fn invoke_llm_sanitize_response( &self, registration_name: &str, response: Json, @@ -1656,14 +1672,14 @@ impl WorkerPluginCallback { context.codec_capability_id = Some(capability_id.clone()); capability_id }); - let response = self.invoke_blocking(invoke); + let response = self.invoke_async(invoke).await; if let Some(capability_id) = capability_id { self.host_state.remove_codec(&capability_id); } optional_json_from_invoke_response(response?) } - fn invoke_llm_guardrail( + async fn invoke_llm_guardrail( &self, registration_name: &str, request: LlmRequest, @@ -1674,10 +1690,10 @@ impl WorkerPluginCallback { None, Some(invoke_request_payload_llm("", Some(request), None, None)), ); - guardrail_from_invoke_response(self.invoke_blocking(invoke)?) + guardrail_from_invoke_response(self.invoke_async(invoke).await?) } - fn invoke_llm_request_intercept( + async fn invoke_llm_request_intercept( &self, registration_name: &str, model_name: &str, @@ -1695,7 +1711,7 @@ impl WorkerPluginCallback { None, )), ); - let response = self.invoke_blocking(invoke)?; + let response = self.invoke_async(invoke).await?; match response.result { Some(invoke_response_result::Result::LlmRequest(result)) => { let outcome = required_envelope(result.outcome, "llm request intercept outcome")?; @@ -1853,8 +1869,22 @@ impl WorkerPluginCallback { } async fn invoke_async(&self, request: InvokeRequest) -> FlowResult { - self.invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) - .await + let callback_name = request.registration_name.clone(); + let surface = request.surface; + let result = self + .invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) + .await; + if let Err(error) = &result { + log::warn!( + target: "nemo_relay.worker", + event = "worker_callback_failed", + plugin_id = self.plugin_kind.as_str(), + callback = callback_name.as_str(), + surface; + "Worker plugin callback failed: {error}" + ); + } + result } async fn invoke_async_with_timeout( diff --git a/crates/core/src/stream.rs b/crates/core/src/stream.rs index 3fee8dd3e..756f066a7 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -39,14 +39,16 @@ use crate::api::optimization::finalize_optimization_summary; use crate::api::runtime::LlmSanitizeResponseContext; use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::global_context; +use crate::api::runtime::subscriber_dispatcher; use crate::api::runtime::{ EventSubscriberFn, LlmJsonStream, LlmStreamInner, ScopeStackHandle, current_scope_stack, }; -use crate::api::shared::metadata_with_otel_status; -use crate::api::shared::sanitize_event_with_scope_stack; +use crate::api::shared::{ + metadata_with_otel_status, sanitize_event_with_scope_stack, snapshot_event_sanitizers, +}; use crate::codec::response::{AnnotatedLlmResponse, attach_estimated_cost_for_provider}; use crate::codec::traits::LlmResponseCodec; -use crate::error::Result; +use crate::error::{FlowError, Result}; use crate::json::Json; use serde_json::Map; @@ -78,6 +80,8 @@ pub struct LlmStreamWrapper { chunk_index: u64, ended: bool, close_result: Option>, + finalization: Option>, + terminal_result: Option>, } impl LlmStreamWrapper { @@ -157,6 +161,8 @@ impl LlmStreamWrapper { chunk_index: 0, ended: false, close_result: None, + finalization: None, + terminal_result: None, } } @@ -182,7 +188,12 @@ impl LlmStreamWrapper { "ERROR", Some("stream dropped before clean completion".to_string()), ); - self.emit_end_event(metadata, true); + // Drop cannot await the async finalizer. Close the recorder before + // spawning it so late optimization evidence is rejected immediately. + self.handle + .optimization_recorder + .close_for_finalization(Some("stream_interrupted")); + self.finalization = self.emit_end_event(metadata, true, true); } fn finish_with_status( @@ -197,14 +208,23 @@ impl LlmStreamWrapper { self.ended = true; let metadata = metadata_with_otel_status(self.metadata.clone(), status_code, status_message); - self.emit_end_event(metadata, interrupted); + self.finalization = self.emit_end_event(metadata, interrupted, false); } /// Emit the LLM END event with aggregated response data. /// /// Calls the finalizer to produce the aggregated response, runs sanitize /// response guardrails, and emits the END event. - fn emit_end_event(&mut self, metadata: Option, interrupted: bool) { + fn emit_end_event( + &mut self, + metadata: Option, + interrupted: bool, + background_thread: bool, + ) -> Option> { + // The finalizer below runs on the caller's Tokio runtime. Register a + // dispatcher barrier before spawning it so a synchronous subscriber + // flush after this stream is dropped cannot overtake the END event. + let publication_barrier = subscriber_dispatcher::register_async_publication(); let aggregated = match self.finalizer.take() { Some(finalizer) => finalizer(), None => Json::Null, @@ -230,68 +250,108 @@ impl LlmStreamWrapper { Err(_) => None, } }; - let Some(entries) = snapshot else { - return; - }; - let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( - response, - self.sanitize_context.clone(), - &entries, - ); - let data = match sanitized { - Some(response) if response_was_null_without_fallback && response.is_null() => None, - response => response, - }; - let annotation_omitted = data.as_ref().is_none_or(Json::is_null); - let mut annotated_response: Option = (!annotation_omitted) - .then(|| { - data.as_ref().and_then(|response| { - self.response_codec.as_ref().and_then(|codec| { - let mut decoded = codec.decode_response(response).ok()?; - attach_estimated_cost_for_provider(&mut decoded, Some(&self.handle.name)); - Some(decoded) + let entries = snapshot?; + let handle = self.handle.clone(); + let scope_stack = self.scope_stack.clone(); + let subscribers = self.subscribers.clone(); + let response_codec = self.response_codec.clone(); + let sanitize_context = self.sanitize_context.clone(); + let finalize = async move { + let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( + response, + sanitize_context, + &entries, + ) + .await; + let data = match sanitized { + Some(response) if response_was_null_without_fallback && response.is_null() => None, + response => response, + }; + let annotation_omitted = data.as_ref().is_none_or(Json::is_null); + let mut annotated_response: Option = (!annotation_omitted) + .then(|| { + data.as_ref().and_then(|response| { + response_codec.as_ref().and_then(|codec| { + let mut decoded = codec.decode_response(response).ok()?; + attach_estimated_cost_for_provider(&mut decoded, Some(&handle.name)); + Some(decoded) + }) }) }) - }) - .flatten(); - let interruption = (interrupted - && !has_authoritative_final_usage(annotated_response.as_ref())) - .then_some("stream_interrupted"); - self.handle - .optimization_recorder - .close_for_finalization(interruption); - emit_optimization_marks(&self.handle, &self.subscribers); - let pricing = crate::codec::response::active_pricing_resolver(); - let summary = finalize_optimization_summary( - &self.handle.optimization_recorder, - annotated_response.as_mut(), - self.handle.model_name.as_deref(), - &pricing, - ); - if !annotation_omitted - && annotated_response.is_none() - && let Some(summary) = summary - { - annotated_response = Some(AnnotatedLlmResponse { - optimization_summary: Some(summary), - ..AnnotatedLlmResponse::default() + .flatten(); + let interruption = (interrupted + && !has_authoritative_final_usage(annotated_response.as_ref())) + .then_some("stream_interrupted"); + handle + .optimization_recorder + .close_for_finalization(interruption); + emit_optimization_marks(&handle, &subscribers).await; + let pricing = crate::codec::response::active_pricing_resolver(); + let summary = finalize_optimization_summary( + &handle.optimization_recorder, + annotated_response.as_mut(), + handle.model_name.as_deref(), + &pricing, + ); + if !annotation_omitted + && annotated_response.is_none() + && let Some(summary) = summary + { + annotated_response = Some(AnnotatedLlmResponse { + optimization_summary: Some(summary), + ..AnnotatedLlmResponse::default() + }); + } + let annotated_response = annotated_response.map(Arc::new); + let event_snapshot = { + let ctx = global_context(); + let state = ctx.read(); + match state { + Ok(state) => { + Some(state.end_llm_handle(&handle, data, metadata, annotated_response)) + } + Err(_) => None, + } + }; + if let Some(event) = event_snapshot + && let Some(event) = sanitize_event_with_scope_stack(event, &scope_stack).await + { + let _ = subscriber_dispatcher::dispatch_sanitized_event( + event, + Vec::new(), + &subscribers, + scope_stack.clone(), + ); + } + if let Some(done) = publication_barrier { + let _ = done.send(()); + } + }; + if background_thread { + // `Drop` can run while the current-thread Tokio executor is + // synchronously flushing subscribers. Use a dedicated runtime so + // the FIFO publication barrier can still be released. + std::thread::spawn(move || { + if let Ok(runtime) = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + { + runtime.block_on(finalize); + } }); + return None; } - let annotated_response = annotated_response.map(Arc::new); - let event_snapshot = { - let ctx = global_context(); - let state = ctx.read(); - match state { - Ok(state) => { - Some(state.end_llm_handle(&self.handle, data, metadata, annotated_response)) + match tokio::runtime::Handle::try_current() { + Ok(handle) => Some(handle.spawn(finalize)), + Err(_) => { + if let Ok(runtime) = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + { + runtime.block_on(finalize); } - Err(_) => None, + None } - }; - if let Some(event) = event_snapshot - && let Some(event) = sanitize_event_with_scope_stack(event, &self.scope_stack) - { - NemoRelayContextState::emit_event(&event, &self.subscribers); } } @@ -318,9 +378,14 @@ impl LlmStreamWrapper { } }; if let Some(event) = event_snapshot - && let Some(event) = sanitize_event_with_scope_stack(event, &self.scope_stack) + && let Some(sanitizers) = snapshot_event_sanitizers(&event, &self.scope_stack) { - NemoRelayContextState::emit_event(&event, &self.subscribers); + let _ = subscriber_dispatcher::dispatch_sanitized_event( + event, + sanitizers, + &self.subscribers, + self.scope_stack.clone(), + ); } } } @@ -328,8 +393,31 @@ impl LlmStreamWrapper { impl Stream for LlmStreamWrapper { type Item = Result; - fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - let this = self.get_mut(); + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.as_mut().get_mut(); + + // The END event runs async because response and event sanitizers may + // await. Do not expose stream termination until that work has queued + // the event: callers commonly flush subscribers immediately after + // exhausting a stream, and that flush must include its END event. + if let Some(finalization) = this.finalization.as_mut() { + return match Pin::new(finalization).poll(cx) { + Poll::Pending => Poll::Pending, + Poll::Ready(Ok(())) => { + this.finalization = None; + match this.terminal_result.take() { + Some(result) => Poll::Ready(Some(result)), + None => Poll::Ready(None), + } + } + Poll::Ready(Err(error)) => { + this.finalization = None; + Poll::Ready(Some(Err(FlowError::Internal(format!( + "stream finalization task failed: {error}" + ))))) + } + }; + } if this.ended { return Poll::Ready(None); @@ -346,19 +434,21 @@ impl Stream for LlmStreamWrapper { Ok(()) => Poll::Ready(Some(Ok(raw_chunk))), Err(e) => { let message = e.to_string(); + this.terminal_result = Some(Err(e)); this.finish_with_status("ERROR", Some(message), true); - Poll::Ready(Some(Err(e))) + self.poll_next(cx) } } } Poll::Ready(Some(Err(e))) => { let message = e.to_string(); + this.terminal_result = Some(Err(e)); this.finish_with_status("ERROR", Some(message), true); - Poll::Ready(Some(Err(e))) + self.poll_next(cx) } Poll::Ready(None) => { this.finish_with_status("OK", None, false); - Poll::Ready(None) + self.poll_next(cx) } Poll::Pending => Poll::Pending, } @@ -374,6 +464,11 @@ impl LlmStreamInner for LlmStreamWrapper { } let result = this.inner.close().await; this.finish(); + if let Some(finalization) = this.finalization.take() { + finalization.await.map_err(|error| { + FlowError::Internal(format!("stream finalization task failed: {error}")) + })?; + } this.close_result = Some(result.clone()); this.close_result .as_ref() diff --git a/crates/core/tests/coverage/logging_rotation_tests.rs b/crates/core/tests/coverage/logging_rotation_tests.rs new file mode 100644 index 000000000..842fa7927 --- /dev/null +++ b/crates/core/tests/coverage/logging_rotation_tests.rs @@ -0,0 +1,34 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use super::*; + +#[test] +fn rotating_writer_rotates_retains_and_reports_closed_file_errors() { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("nested").join("relay.log"); + let mut writer = SizeRotatingFileWriter::new(path.clone(), 4, 2).unwrap(); + assert_eq!(writer.write(b"abcd").unwrap(), 4); + writer.flush().unwrap(); + assert_eq!(writer.write(b"e").unwrap(), 1); + writer.flush().unwrap(); + + assert_eq!(std::fs::read(rotated_log_path(&path, 1)).unwrap(), b"abcd"); + assert_eq!(std::fs::read(&path).unwrap(), b"e"); + + writer.file = None; + assert!(writer.write(b"x").is_err()); + assert!(writer.flush().is_err()); +} + +#[test] +fn rotation_helpers_handle_empty_and_missing_files() { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("missing.log"); + rotate_files(&path, 2).unwrap(); + assert_eq!( + rotated_log_path(&path, 2), + directory.path().join("missing.2.log") + ); + create_parent_directory(std::path::Path::new("plain.log")).unwrap(); +} diff --git a/crates/core/tests/coverage/logging_sink_tests.rs b/crates/core/tests/coverage/logging_sink_tests.rs index 044ce2b47..88f142655 100644 --- a/crates/core/tests/coverage/logging_sink_tests.rs +++ b/crates/core/tests/coverage/logging_sink_tests.rs @@ -2,10 +2,16 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - DROP_REPORT_INTERVAL_MILLIS, DropNoticeRateLimiter, dropped_record_error_handler, - log_level_filter, now_millis, spdlog_level, stderr_error_handler, + DROP_REPORT_INTERVAL_MILLIS, DropNoticeRateLimiter, build_logger, dropped_record_error_handler, + log_level_filter, now_millis, reserved_sink_paths, resolve_log_path, rotated_log_path, + spdlog_level, stderr_error_handler, }; use crate::logging::LogLevel; +use crate::logging::{ + FileLogRotationConfig, FileLogSinkConfig, LogFormat, LogSinkConfig, LoggingConfig, + MAX_FILE_SINK_QUEUE_ENTRIES, +}; +use std::path::PathBuf; #[test] fn drop_notice_rate_limiter_reports_immediately_then_once_per_interval() { @@ -33,3 +39,60 @@ fn sink_helpers_cover_boundary_levels_time_and_emergency_handlers() { "expected test error", ))); } + +#[test] +fn logger_builder_rejects_duplicate_conflicting_and_invalid_file_sinks() { + let directory = tempfile::tempdir().unwrap(); + let log_path = directory.path().join("relay.log"); + let file_sink = |path: PathBuf, rotation| { + LogSinkConfig::File(FileLogSinkConfig { + path, + level: LogLevel::Info, + format: LogFormat::Jsonl, + queue_capacity: 8, + rotation, + }) + }; + + assert!(resolve_log_path(std::path::Path::new("")).is_err()); + let rotation = FileLogRotationConfig::new(32, 1).unwrap(); + assert_eq!(reserved_sink_paths(&log_path, Some(rotation)).len(), 2); + + let duplicate = LoggingConfig { + sinks: vec![ + file_sink(log_path.clone(), None), + file_sink(log_path.clone(), None), + ], + ..LoggingConfig::default() + }; + let error = match build_logger(&duplicate, "root".into()) { + Ok(_) => panic!("duplicate file sinks must be rejected"), + Err(error) => error, + }; + assert!(error.to_string().contains("duplicate logging sink path")); + + let conflict = LoggingConfig { + sinks: vec![ + file_sink(log_path.clone(), Some(rotation)), + file_sink(rotated_log_path(&log_path, 1), None), + ], + ..LoggingConfig::default() + }; + let error = match build_logger(&conflict, "root".into()) { + Ok(_) => panic!("active and rotated file paths must not overlap"), + Err(error) => error, + }; + assert!(error.to_string().contains("conflicts")); + + let mut invalid_capacity = LoggingConfig { + sinks: vec![file_sink(log_path, None)], + ..LoggingConfig::default() + }; + let LogSinkConfig::File(file_sink) = &mut invalid_capacity.sinks[0]; + file_sink.queue_capacity = MAX_FILE_SINK_QUEUE_ENTRIES + 1; + let error = match build_logger(&invalid_capacity, "root".into()) { + Ok(_) => panic!("oversized async queues must be rejected"), + Err(error) => error, + }; + assert!(error.to_string().contains("queue_capacity")); +} diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index 350cf1414..acade0d0c 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -7,7 +7,9 @@ use std::ptr; use nemo_relay_plugin::{ CategoryProfile, ConfigDiagnostic, DiagnosticLevel, Event, EventCategory, EventSanitizeFields, Json, LlmJsonStream, LlmRequest, LlmRequestInterceptOutcome, NemoRelayNativeHostApiV1, - NemoRelayNativePluginContext, NemoRelayNativePluginV1, NemoRelayNativeString, NemoRelayStatus, + NemoRelayNativeHostApiV3, NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncMiddlewareCb, + NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativePluginContext, NemoRelayNativePluginV1, + NemoRelayNativeString, NemoRelayStatus, NemoRelayNativeToolNextFn, NativePlugin, PendingMarkSpec, PluginContext, PluginRuntime, ScopeCategory, ScopeType, ToolExecutionInterceptOutcome, }; @@ -280,6 +282,33 @@ fn mark_json(mut value: Json, key: &str) -> Json { nemo_relay_plugin::nemo_relay_plugin!(nemo_relay_fixture_native_plugin, || FixtureNativePlugin); +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_fixture_async_entry( + host: *const NemoRelayNativeHostApiV1, + out: *mut NemoRelayNativePluginV1, +) -> NemoRelayStatus { + if host.is_null() || out.is_null() { + return NemoRelayStatus::NullPointer; + } + let host_v1 = unsafe { &*host }; + if host_v1.abi_version < 3 + || host_v1.struct_size < std::mem::size_of::() + { + return NemoRelayStatus::InvalidArg; + } + let host_v2 = unsafe { &*(host as *const NemoRelayNativeHostApiV3) }; + let mut plugin = NemoRelayNativePluginV1::default(); + plugin.plugin_kind = unsafe { raw_host_string(&host_v2.v1, "fixture_async") }; + if plugin.plugin_kind.is_null() { + return NemoRelayStatus::Internal; + } + plugin.user_data = Box::into_raw(Box::new(*host_v2)).cast(); + plugin.register = Some(raw_register_async_tool_request); + plugin.drop = Some(raw_drop_async_host); + unsafe { *out = plugin }; + NemoRelayStatus::Ok +} + #[unsafe(no_mangle)] pub unsafe extern "C" fn nemo_relay_fixture_observability_collision( host: *const NemoRelayNativeHostApiV1, @@ -570,6 +599,260 @@ unsafe extern "C" fn raw_register_event_sanitize_errors( status } +unsafe extern "C" fn raw_register_async_tool_request( + user_data: *mut c_void, + _plugin_config_json: *const NemoRelayNativeString, + ctx: *mut NemoRelayNativePluginContext, +) -> NemoRelayStatus { + if user_data.is_null() { + return NemoRelayStatus::NullPointer; + } + let host = unsafe { &*(user_data as *const NemoRelayNativeHostApiV3) }; + let registrations: [ + (NemoRelayNativeAsyncMiddlewareKind, &str, NemoRelayNativeAsyncMiddlewareCb); + 14 + ] = [ + (NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest, "fixture_async_tool_sanitize_request", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeResponse, "fixture_async_tool_sanitize_response", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::ToolConditionalExecution, "fixture_async_tool_conditional", raw_async_allow_callback), + (NemoRelayNativeAsyncMiddlewareKind::ToolRequestIntercept, "fixture_async_request", raw_async_tool_request_callback), + (NemoRelayNativeAsyncMiddlewareKind::ToolExecutionIntercept, "fixture_async_execution", raw_async_tool_execution_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeRequest, "fixture_async_llm_sanitize_request", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeResponse, "fixture_async_llm_sanitize_response", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmConditionalExecution, "fixture_async_llm_conditional", raw_async_allow_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept, "fixture_async_llm_request", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept, "fixture_async_llm_execution", raw_async_tool_execution_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept, "fixture_async_llm_stream", raw_async_tool_execution_callback), + (NemoRelayNativeAsyncMiddlewareKind::MarkSanitize, "fixture_async_mark", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeStart, "fixture_async_scope_start", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeEnd, "fixture_async_scope_end", raw_async_passthrough_callback), + ]; + for (kind, registration_name, callback) in registrations { + let name = unsafe { raw_host_string(&host.v1, registration_name) }; + if name.is_null() { + return NemoRelayStatus::Internal; + } + let status = unsafe { + (host.plugin_context_register_async_middleware)( + ctx, kind, name, 0, false, callback, user_data, None, + ) + }; + unsafe { (host.v1.string_free)(name) }; + if status != NemoRelayStatus::Ok { + return status; + } + } + NemoRelayStatus::Ok +} + +unsafe extern "C" fn raw_async_allow_callback( + user_data: *mut c_void, + _invocation_json: *const NemoRelayNativeString, + _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> NemoRelayNativeAsyncCallbackState { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let result = unsafe { raw_host_string(&host.v1, "null") }; + if result.is_null() { + unsafe { reject_async_completion(host, completion, "failed to allocate async allow result") }; + } else { + unsafe { + (host.async_completion_resolve_json)(completion, result); + (host.v1.string_free)(result); + } + } + NemoRelayNativeAsyncCallbackState::Complete +} + +unsafe extern "C" fn raw_async_passthrough_callback( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> NemoRelayNativeAsyncCallbackState { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let result = unsafe { raw_host_string_value(&host.v1, invocation_json) } + .and_then(|value| serde_json::from_str::(&value).ok()) + .and_then(|invocation| { + invocation.get("annotated").map(|annotated| { + json!({ + "request": invocation["request"], + "annotated_request": annotated, + "pending_marks": [], + "optimization_contributions": [], + }) + }).or_else(|| { + ["value", "request", "response", "fields"] + .into_iter() + .find_map(|key| invocation.get(key).cloned()) + }) + }) + .and_then(|value| serde_json::to_string(&value).ok()); + let Some(result) = result else { + unsafe { reject_async_completion(host, completion, "invalid async passthrough invocation") }; + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let result = unsafe { raw_host_string(&host.v1, &result) }; + if result.is_null() { + unsafe { reject_async_completion(host, completion, "failed to allocate async passthrough result") }; + } else { + unsafe { + (host.async_completion_resolve_json)(completion, result); + (host.v1.string_free)(result); + } + } + NemoRelayNativeAsyncCallbackState::Complete +} + +unsafe extern "C" fn raw_async_tool_request_callback( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> NemoRelayNativeAsyncCallbackState { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let invocation = unsafe { raw_host_string_value(&host.v1, invocation_json) } + .and_then(|json| serde_json::from_str::(&json).ok()) + .and_then(|mut invocation| { + let pending = invocation["name"].as_str() == Some("async-pending"); + let duplicate = invocation["name"].as_str() == Some("async-double"); + invocation + .get_mut("value") + .and_then(Json::as_object_mut) + .map(|value| { + value.insert("native_async".into(), json!(true)); + (Json::Object(value.clone()), pending, duplicate) + }) + }) + .and_then(|(value, pending, duplicate)| { + serde_json::to_string(&value) + .ok() + .map(|value| (value, pending, duplicate)) + }); + let Some((result, pending, duplicate)) = invocation else { + unsafe { reject_async_completion(host, completion, "invalid async tool request invocation") }; + return NemoRelayNativeAsyncCallbackState::Complete; + }; + if pending { + let host = *host; + let completion = completion as usize; + std::thread::spawn(move || { + std::thread::sleep(std::time::Duration::from_millis(10)); + let result = unsafe { raw_host_string(&host.v1, &result) }; + if !result.is_null() { + unsafe { + (host.async_completion_resolve_json)( + completion as *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + result, + ); + (host.v1.string_free)(result); + (host.async_completion_release)( + completion as *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + ); + } + } else { + unsafe { + reject_async_completion( + &host, + completion as *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + "failed to allocate async tool request result", + ); + (host.async_completion_release)( + completion as *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + ); + } + } + }); + return NemoRelayNativeAsyncCallbackState::Pending; + } + let result = unsafe { raw_host_string(&host.v1, &result) }; + if !result.is_null() { + unsafe { + (host.async_completion_resolve_json)(completion, result); + if duplicate { + let _ = (host.async_completion_resolve_json)(completion, result); + } + (host.v1.string_free)(result); + } + } else { + unsafe { reject_async_completion(host, completion, "failed to allocate async tool request result") }; + } + NemoRelayNativeAsyncCallbackState::Complete +} + +unsafe extern "C" fn raw_async_tool_execution_callback( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> NemoRelayNativeAsyncCallbackState { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete; + }; + if next.is_null() || completion.is_null() { + unsafe { reject_async_completion(host, completion, "async tool execution requires next and completion") }; + return NemoRelayNativeAsyncCallbackState::Complete; + } + let value = unsafe { raw_host_string_value(&host.v1, invocation_json) } + .and_then(|json| serde_json::from_str::(&json).ok()) + .and_then(|mut invocation| { + if let Some(value) = invocation.get_mut("value").and_then(Json::as_object_mut) { + value.insert("native_async_execution".into(), json!(true)); + Some(Json::Object(value.clone())) + } else { + invocation.get("request").cloned() + } + }) + .and_then(|value| serde_json::to_string(&value).ok()); + let Some(value) = value else { + unsafe { reject_async_completion(host, completion, "invalid async tool execution invocation") }; + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let value = unsafe { raw_host_string(&host.v1, &value) }; + if value.is_null() { + unsafe { reject_async_completion(host, completion, "failed to allocate async tool execution invocation") }; + return NemoRelayNativeAsyncCallbackState::Complete; + } + let status = unsafe { (host.async_next_invoke)(next, value, completion) }; + unsafe { + (host.v1.string_free)(value); + } + if status == NemoRelayStatus::Ok { + unsafe { + (host.async_next_release)(next); + (host.async_completion_release)(completion); + } + NemoRelayNativeAsyncCallbackState::Pending + } else { + unsafe { reject_async_completion(host, completion, "failed to invoke async tool execution next") }; + NemoRelayNativeAsyncCallbackState::Complete + } +} + +unsafe fn reject_async_completion( + host: &NemoRelayNativeHostApiV3, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + message: &str, +) { + if completion.is_null() { + return; + } + let message = unsafe { raw_host_string(&host.v1, message) }; + if message.is_null() { + return; + } + unsafe { + let _ = (host.async_completion_reject)(completion, message); + (host.v1.string_free)(message); + } +} + unsafe extern "C" fn raw_tool_outcome_callback( user_data: *mut c_void, name: *const NemoRelayNativeString, @@ -638,6 +921,12 @@ unsafe extern "C" fn raw_drop_host(user_data: *mut c_void) { } } +unsafe extern "C" fn raw_drop_async_host(user_data: *mut c_void) { + if !user_data.is_null() { + drop(unsafe { Box::from_raw(user_data as *mut NemoRelayNativeHostApiV3) }); + } +} + unsafe fn raw_host_from_user_data<'a>( user_data: *mut c_void, ) -> Option<&'a NemoRelayNativeHostApiV1> { diff --git a/crates/core/tests/integration/api_surface_tests.rs b/crates/core/tests/integration/api_surface_tests.rs index 7a6045ac0..61f56251f 100644 --- a/crates/core/tests/integration/api_surface_tests.rs +++ b/crates/core/tests/integration/api_surface_tests.rs @@ -7,6 +7,9 @@ use std::sync::{Arc, Mutex}; +mod test_support; +use test_support::ready; + use chrono::{DateTime, TimeDelta, Utc}; use futures::StreamExt; use nemo_relay::api::event::{CategoryProfile, Event, ScopeCategory}; @@ -65,6 +68,7 @@ use nemo_relay::api::tool::{ tool_call, tool_call_end, tool_call_execute, tool_conditional_execution, tool_request_intercepts, }; +use nemo_relay::codec::optimization::LlmOptimizationContribution; use nemo_relay::error::{FlowError, Result}; use nemo_relay::json::Json; use serde_json::{Map, json}; @@ -90,7 +94,7 @@ fn event_sanitizers_rewrite_only_observability_fields_in_priority_order() { Arc::new(|event, mut fields| { assert_eq!(event.data().unwrap()["order"], json!(["early"])); fields.data = Some(json!({"order": ["early", "late"]})); - fields + ready(fields) }), ) .unwrap(); @@ -101,7 +105,7 @@ fn event_sanitizers_rewrite_only_observability_fields_in_priority_order() { fields.data = Some(json!({"order": ["early"]})); fields.metadata = Some(json!({"redacted": true})); fields.category_profile = Some(CategoryProfile::builder().subtype("sanitized").build()); - fields + ready(fields) }), ) .unwrap(); @@ -151,7 +155,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 20, Arc::new(|_, mut fields| { fields.data = Some(json!({"mark": "global"})); - fields + ready(fields) }), ) .unwrap(); @@ -160,7 +164,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 20, Arc::new(|_, mut fields| { fields.metadata = Some(json!({"scope_end": true})); - fields + ready(fields) }), ) .unwrap(); @@ -178,7 +182,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 10, Arc::new(|_, mut fields| { fields.data = Some(json!({"mark": "local"})); - fields + ready(fields) }), ) .unwrap(); @@ -188,7 +192,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 10, Arc::new(|_, mut fields| { fields.metadata = Some(json!({"scope_start": true})); - fields + ready(fields) }), ) .unwrap(); @@ -198,7 +202,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 10, Arc::new(|_, mut fields| { fields.data = Some(json!({"scope_end": "local"})); - fields + ready(fields) }), ) .unwrap(); @@ -512,7 +516,7 @@ fn skill_load_detection_uses_original_arguments_before_observability_sanitizatio register_tool_sanitize_request_guardrail( "strip-skill-path", 1, - Arc::new(|_name, _args| json!({"path": "[redacted]"})), + Arc::new(|_name, _args| ready(json!({"path": "[redacted]"}))), ) .unwrap(); let events = capture_events("sanitized-skill-load-api-events"); @@ -575,7 +579,7 @@ async fn managed_skill_load_marks_survive_failures_repeat_per_call_and_skip_bloc register_tool_conditional_execution_guardrail( "block-skill-load", 1, - Arc::new(|_name, _args| Ok(Some("blocked before start".into()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("blocked before start".into())) })), ) .unwrap(); let blocked = tool_call_execute( @@ -703,6 +707,56 @@ fn test_manual_lifecycle_timestamp_overrides() { deregister_subscriber("timestamp-api-events").unwrap(); } +#[test] +fn test_manual_llm_end_queues_optimization_marks_before_end_event() { + let _lock = TEST_MUTEX.lock().unwrap(); + reset_global(); + setup_isolated_thread(); + + let events = capture_events("manual-optimization-events"); + let request = make_llm_request(json!({"messages": []})); + let handle = llm_call( + LlmCallParams::builder() + .name("manual-optimized-llm") + .request(&request) + .build(), + ) + .unwrap(); + assert!( + handle + .optimization_recorder + .record(LlmOptimizationContribution::new( + "test.manual", + "test_manual_kind", + )) + ); + + llm_call_end( + nemo_relay::api::llm::LlmCallEndParams::builder() + .handle(&handle) + .response(json!({"ok": true})) + .build(), + ) + .unwrap(); + + let names = captured_events_snapshot(&events) + .into_iter() + .filter(|event| { + event.name() == "manual-optimized-llm" || event.name() == "nemo_relay.llm.optimization" + }) + .map(|event| event.name().to_owned()) + .collect::>(); + assert_eq!( + names, + [ + "manual-optimized-llm", + "nemo_relay.llm.optimization", + "manual-optimized-llm", + ] + ); + deregister_subscriber("manual-optimization-events").unwrap(); +} + #[test] fn test_manual_lifecycle_default_end_timestamps_follow_explicit_starts() { let _lock = TEST_MUTEX.lock().unwrap(); @@ -834,10 +888,19 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { reset_global(); setup_isolated_thread(); - register_mark_sanitize_guardrail("mark-sanitize", 1, Arc::new(|_, fields| fields)).unwrap(); + register_mark_sanitize_guardrail( + "mark-sanitize", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); expect_already_exists( - register_mark_sanitize_guardrail("mark-sanitize", 1, Arc::new(|_, fields| fields)) - .unwrap_err(), + register_mark_sanitize_guardrail( + "mark-sanitize", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap_err(), "mark-sanitize", ); assert!(deregister_mark_sanitize_guardrail("mark-sanitize").unwrap()); @@ -846,28 +909,32 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_scope_sanitize_start_guardrail( "scope-start-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); assert!(deregister_scope_sanitize_start_guardrail("scope-start-sanitize").unwrap()); assert!(!deregister_scope_sanitize_start_guardrail("scope-start-sanitize").unwrap()); - register_scope_sanitize_end_guardrail("scope-end-sanitize", 1, Arc::new(|_, fields| fields)) - .unwrap(); + register_scope_sanitize_end_guardrail( + "scope-end-sanitize", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); assert!(deregister_scope_sanitize_end_guardrail("scope-end-sanitize").unwrap()); assert!(!deregister_scope_sanitize_end_guardrail("scope-end-sanitize").unwrap()); register_tool_sanitize_request_guardrail( "tool-sanitize-request", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); expect_already_exists( register_tool_sanitize_request_guardrail( "tool-sanitize-request", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap_err(), "tool-sanitize-request", @@ -878,7 +945,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_tool_sanitize_response_guardrail( "tool-sanitize-response", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); assert!(deregister_tool_sanitize_response_guardrail("tool-sanitize-response").unwrap()); @@ -886,13 +953,18 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_tool_conditional_execution_guardrail( "tool-conditional", 1, - Arc::new(|_name, _args| Ok(None)), + Arc::new(|_name, _args| Box::pin(async { Ok(None) })), ) .unwrap(); assert!(deregister_tool_conditional_execution_guardrail("tool-conditional").unwrap()); - register_tool_request_intercept("tool-request", 1, false, Arc::new(|_name, args| Ok(args))) - .unwrap(); + register_tool_request_intercept( + "tool-request", + 1, + false, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ) + .unwrap(); assert!(deregister_tool_request_intercept("tool-request").unwrap()); register_tool_execution_intercept( @@ -906,7 +978,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_llm_sanitize_request_guardrail( "llm-sanitize-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); assert!(deregister_llm_sanitize_request_guardrail("llm-sanitize-request").unwrap()); @@ -914,7 +986,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_llm_sanitize_response_guardrail( "llm-sanitize-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); assert!(deregister_llm_sanitize_response_guardrail("llm-sanitize-response").unwrap()); @@ -922,7 +994,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_llm_conditional_execution_guardrail( "llm-conditional", 1, - Arc::new(|_request| Ok(None)), + Arc::new(|_request| Box::pin(async { Ok(None) })), ) .unwrap(); assert!(deregister_llm_conditional_execution_guardrail("llm-conditional").unwrap()); @@ -932,7 +1004,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { 1, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -1023,7 +1095,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "mark-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); expect_already_exists( @@ -1031,7 +1103,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "mark-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap_err(), "mark-sanitize", @@ -1043,7 +1115,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "scope-start-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); assert!( @@ -1059,7 +1131,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "scope-end-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); assert!( @@ -1073,7 +1145,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "tool-sanitize-request", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); expect_already_exists( @@ -1081,7 +1153,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "tool-sanitize-request", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap_err(), "tool-sanitize-request", @@ -1095,7 +1167,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "tool-sanitize-response", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); assert!( @@ -1107,7 +1179,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "tool-conditional", 1, - Arc::new(|_name, _args| Ok(None)), + Arc::new(|_name, _args| Box::pin(async { Ok(None) })), ) .unwrap(); assert!( @@ -1120,7 +1192,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss "tool-request", 1, false, - Arc::new(|_name, args| Ok(args)), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); assert!(scope_deregister_tool_request_intercept(&scope.uuid, "tool-request").unwrap()); @@ -1138,7 +1210,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "llm-sanitize-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); assert!( @@ -1150,7 +1222,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "llm-sanitize-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); assert!( @@ -1162,7 +1234,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "llm-conditional", 1, - Arc::new(|_request| Ok(None)), + Arc::new(|_request| Box::pin(async { Ok(None) })), ) .unwrap(); assert!( @@ -1176,7 +1248,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss 1, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -1229,7 +1301,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "missing-mark-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap_err(), "scope", @@ -1239,7 +1311,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "missing-scope-start-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap_err(), "scope", @@ -1249,7 +1321,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "missing-scope-end-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap_err(), "scope", @@ -1259,7 +1331,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "missing-tool-sanitize", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap_err(), "scope", @@ -1270,7 +1342,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss "missing-tool-request", 1, false, - Arc::new(|_name, args| Ok(args)), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap_err(), "scope", @@ -1307,7 +1379,7 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { args.as_object_mut() .unwrap() .insert("sanitized_request".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -1319,7 +1391,7 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { .as_object_mut() .unwrap() .insert("sanitized_response".into(), json!(true)); - result + ready(result) }), ) .unwrap(); @@ -1331,7 +1403,7 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { assert_eq!(event.input().unwrap()["sanitized_request"], true); fields.metadata = Some(json!({"generic_start": true})); } - fields + ready(fields) }), ) .unwrap(); @@ -1343,7 +1415,7 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { assert_eq!(event.output().unwrap()["sanitized_response"], true); fields.metadata = Some(json!({"generic_end": true})); } - fields + ready(fields) }), ) .unwrap(); @@ -1401,12 +1473,14 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { args.as_object_mut() .unwrap() .insert("intercepted".into(), json!(true)); - Ok(args) + ready(args) }), ) .unwrap(); assert_eq!( - tool_request_intercepts("tool-api", json!({"value": 2})).unwrap()["intercepted"], + tool_request_intercepts("tool-api", json!({"value": 2})) + .await + .unwrap()["intercepted"], json!(true) ); deregister_tool_request_intercept("tool-request").unwrap(); @@ -1414,11 +1488,11 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { register_tool_conditional_execution_guardrail( "tool-reject", 1, - Arc::new(|_name, _args| Ok(Some("tool denied".into()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("tool denied".into())) })), ) .unwrap(); assert!(matches!( - tool_conditional_execution("tool-api", &json!({"value": 3})), + tool_conditional_execution("tool-api", &json!({"value": 3})).await, Err(FlowError::GuardrailRejected(reason)) if reason == "tool denied" )); assert!(matches!( @@ -1492,7 +1566,7 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { 1, Arc::new(|mut request, _context| { request.headers.insert("x-sanitized".into(), json!(true)); - Some(request) + ready(Some(request)) }), ) .unwrap(); @@ -1504,7 +1578,7 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { .as_object_mut() .unwrap() .insert("sanitized_response".into(), json!(true)); - Some(response) + ready(Some(response)) }), ) .unwrap(); @@ -1557,7 +1631,7 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { false, Arc::new(|_name, mut request, annotated| { request.headers.insert("x-intercepted".into(), json!(true)); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -1567,6 +1641,7 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { "llm-api", make_llm_request(json!({"messages": [{"role": "user", "content": "hello"}]})), ) + .await .unwrap(); assert_eq!( intercepted.request.headers.get("x-intercepted"), @@ -1577,11 +1652,11 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { register_llm_conditional_execution_guardrail( "llm-reject", 1, - Arc::new(|_request| Ok(Some("llm denied".into()))), + Arc::new(|_request| Box::pin(async { Ok(Some("llm denied".into())) })), ) .unwrap(); assert!(matches!( - llm_conditional_execution(&make_llm_request(json!({"messages": []}))), + llm_conditional_execution(&make_llm_request(json!({"messages": []}))).await, Err(FlowError::GuardrailRejected(reason)) if reason == "llm denied" )); assert!(matches!( @@ -1661,7 +1736,7 @@ async fn test_llm_stream_chunk_marks_track_successful_chunks() { Arc::new(|event, mut fields| { assert_eq!(event.name(), "llm.chunk"); fields.metadata = Some(json!({"sanitized": true})); - fields + ready(fields) }), ) .unwrap(); @@ -1703,6 +1778,7 @@ async fn test_llm_stream_chunk_marks_track_successful_chunks() { yielded.push(item.unwrap()); } assert_eq!(yielded, raw_chunks); + stream.close().await.unwrap(); let captured = captured_events_snapshot(&events); assert_eq!(captured.len(), 4); @@ -1775,6 +1851,7 @@ async fn test_llm_stream_chunk_mark_survives_collector_failure() { Err(FlowError::Internal(message)) if message == "collector failed" )); assert!(stream.next().await.is_none()); + stream.close().await.unwrap(); let captured = captured_events_snapshot(&events); assert_eq!(captured.len(), 3); @@ -1841,31 +1918,34 @@ async fn test_llm_stream_api_covers_success_rejection_and_execution_error_paths( chunks, vec![json!({"messages": [{"role": "user", "content": "hello"}]})] ); + stream.close().await.unwrap(); let success_events = captured_events_snapshot(&events); - assert_eq!(success_events[0].kind(), "scope"); - assert_eq!( - success_events[0].scope_category(), - Some(ScopeCategory::Start) - ); - assert_eq!(success_events[0].category().unwrap().as_str(), "llm"); - assert_eq!(success_events.last().unwrap().kind(), "scope"); - assert_eq!( - success_events.last().unwrap().scope_category(), - Some(ScopeCategory::End) - ); - assert_eq!( - success_events.last().unwrap().category().unwrap().as_str(), - "llm" - ); + let success_start = success_events + .iter() + .find(|event| { + event.kind() == "scope" && event.scope_category() == Some(ScopeCategory::Start) + }) + .expect("stream start event"); + let success_end = success_events + .iter() + .rev() + .find(|event| event.kind() == "scope" && event.scope_category() == Some(ScopeCategory::End)) + .expect("stream end event"); + assert_eq!(success_start.kind(), "scope"); + assert_eq!(success_start.scope_category(), Some(ScopeCategory::Start)); + assert_eq!(success_start.category().unwrap().as_str(), "llm"); + assert_eq!(success_end.kind(), "scope"); + assert_eq!(success_end.scope_category(), Some(ScopeCategory::End)); + assert_eq!(success_end.category().unwrap().as_str(), "llm"); assert_eq!( - success_events.last().unwrap().output().unwrap(), + success_end.output().unwrap(), &json!([{"messages": [{"role": "user", "content": "hello"}]}]) ); register_llm_conditional_execution_guardrail( "llm-stream-reject", 1, - Arc::new(|_request| Ok(Some("stream denied".into()))), + Arc::new(|_request| Box::pin(async { Ok(Some("stream denied".into())) })), ) .unwrap(); let reject_collector: Box Result<()> + Send> = Box::new(|_chunk| Ok(())); diff --git a/crates/core/tests/integration/middleware_tests.rs b/crates/core/tests/integration/middleware_tests.rs index c5062b469..d448e7867 100644 --- a/crates/core/tests/integration/middleware_tests.rs +++ b/crates/core/tests/integration/middleware_tests.rs @@ -13,6 +13,9 @@ use std::sync::atomic::{AtomicBool, AtomicU32, Ordering}; use std::sync::{Arc, Mutex}; +mod test_support; +use test_support::{ready, ready_result}; + use futures::StreamExt; use nemo_relay::api::event::{ CategoryProfile, DataSchema, Event, EventCategory, PendingMarkSpec, ScopeCategory, @@ -143,8 +146,8 @@ fn assert_middleware_callback_labels( /// Register 3 tool sanitize request guardrails at priorities 1, 3, 2; /// verify execution order is 1, 2, 3. -#[test] -fn test_sanitize_guardrail_priority_ordering() { +#[tokio::test] +async fn test_sanitize_guardrail_priority_ordering() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -158,7 +161,7 @@ fn test_sanitize_guardrail_priority_ordering() { 1, Arc::new(move |_name, args| { o1.lock().unwrap().push(1); - args + ready(args) }), ) .unwrap(); @@ -170,7 +173,7 @@ fn test_sanitize_guardrail_priority_ordering() { 3, Arc::new(move |_name, args| { o3.lock().unwrap().push(3); - args + ready(args) }), ) .unwrap(); @@ -182,7 +185,7 @@ fn test_sanitize_guardrail_priority_ordering() { 2, Arc::new(move |_name, args| { o2.lock().unwrap().push(2); - args + ready(args) }), ) .unwrap(); @@ -195,6 +198,7 @@ fn test_sanitize_guardrail_priority_ordering() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); let recorded = order.lock().unwrap(); assert_eq!( @@ -211,8 +215,8 @@ fn test_sanitize_guardrail_priority_ordering() { /// Register 3 tool request intercepts at priorities 1, 3, 2; /// verify execution order is 1, 2, 3. -#[test] -fn test_request_intercept_priority_ordering() { +#[tokio::test] +async fn test_request_intercept_priority_ordering() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -226,7 +230,7 @@ fn test_request_intercept_priority_ordering() { false, Arc::new(move |_name, args| { o1.lock().unwrap().push(1); - Ok(args) + ready(args) }), ) .unwrap(); @@ -238,7 +242,7 @@ fn test_request_intercept_priority_ordering() { false, Arc::new(move |_name, args| { o3.lock().unwrap().push(3); - Ok(args) + ready(args) }), ) .unwrap(); @@ -250,13 +254,15 @@ fn test_request_intercept_priority_ordering() { false, Arc::new(move |_name, args| { o2.lock().unwrap().push(2); - Ok(args) + ready(args) }), ) .unwrap(); // Use the standalone intercept chain function - let _result = tool_request_intercepts("test_tool", json!({})).unwrap(); + let _result = tool_request_intercepts("test_tool", json!({})) + .await + .unwrap(); let recorded = order.lock().unwrap(); assert_eq!( @@ -272,8 +278,8 @@ fn test_request_intercept_priority_ordering() { } /// Verify that deregistering and re-registering at a different priority re-sorts. -#[test] -fn test_re_registration_at_different_priority_re_sorts() { +#[tokio::test] +async fn test_re_registration_at_different_priority_re_sorts() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -287,7 +293,7 @@ fn test_re_registration_at_different_priority_re_sorts() { false, Arc::new(move |_name, args| { o_a.lock().unwrap().push("a_p10".into()); - Ok(args) + ready(args) }), ) .unwrap(); @@ -299,13 +305,13 @@ fn test_re_registration_at_different_priority_re_sorts() { false, Arc::new(move |_name, args| { o_b.lock().unwrap().push("b_p20".into()); - Ok(args) + ready(args) }), ) .unwrap(); // First call: a runs before b - let _ = tool_request_intercepts("test", json!({})).unwrap(); + let _ = tool_request_intercepts("test", json!({})).await.unwrap(); { let recorded = order.lock().unwrap(); assert_eq!(*recorded, vec!["a_p10", "b_p20"]); @@ -320,14 +326,14 @@ fn test_re_registration_at_different_priority_re_sorts() { false, Arc::new(move |_name, args| { o_a2.lock().unwrap().push("a_p30".into()); - Ok(args) + ready(args) }), ) .unwrap(); // Clear and re-run order.lock().unwrap().clear(); - let _ = tool_request_intercepts("test", json!({})).unwrap(); + let _ = tool_request_intercepts("test", json!({})).await.unwrap(); { let recorded = order.lock().unwrap(); assert_eq!( @@ -348,8 +354,8 @@ fn test_re_registration_at_different_priority_re_sorts() { /// Register 2 request intercepts, first with break_chain=true. /// Verify second intercept is NOT called and the result from the first is used. -#[test] -fn test_break_chain_stops_subsequent_intercepts() { +#[tokio::test] +async fn test_break_chain_stops_subsequent_intercepts() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -364,7 +370,7 @@ fn test_break_chain_stops_subsequent_intercepts() { args.as_object_mut() .unwrap() .insert("breaker_ran".into(), json!(true)); - Ok(args) + ready(args) }), ) .unwrap(); @@ -379,12 +385,12 @@ fn test_break_chain_stops_subsequent_intercepts() { args.as_object_mut() .unwrap() .insert("after_ran".into(), json!(true)); - Ok(args) + ready(args) }), ) .unwrap(); - let result = tool_request_intercepts("tool", json!({})).unwrap(); + let result = tool_request_intercepts("tool", json!({})).await.unwrap(); // First intercept's transformation should be applied assert_eq!(result["breaker_ran"], true); @@ -404,8 +410,8 @@ fn test_break_chain_stops_subsequent_intercepts() { } /// With break_chain=false on all intercepts, both should be called. -#[test] -fn test_no_break_chain_runs_all_intercepts() { +#[tokio::test] +async fn test_no_break_chain_runs_all_intercepts() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -419,7 +425,7 @@ fn test_no_break_chain_runs_all_intercepts() { false, Arc::new(move |_name, args| { c1.fetch_add(1, Ordering::SeqCst); - Ok(args) + ready(args) }), ) .unwrap(); @@ -431,12 +437,12 @@ fn test_no_break_chain_runs_all_intercepts() { false, Arc::new(move |_name, args| { c2.fetch_add(1, Ordering::SeqCst); - Ok(args) + ready(args) }), ) .unwrap(); - let _ = tool_request_intercepts("tool", json!({})).unwrap(); + let _ = tool_request_intercepts("tool", json!({})).await.unwrap(); assert_eq!( call_count.load(Ordering::SeqCst), @@ -690,7 +696,7 @@ async fn test_tool_execution_outcome_marks_follow_end_with_tool_parentage() { let mut metadata = fields.metadata.unwrap_or_else(|| json!({})); metadata["sanitized"] = json!(true); fields.metadata = Some(metadata); - fields + ready(fields) }), ) .unwrap(); @@ -1329,7 +1335,7 @@ async fn test_conditional_guardrail_rejects() { register_tool_conditional_execution_guardrail( "rejector", 1, - Arc::new(|_name, _args| Ok(Some("not allowed".to_string()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("not allowed".to_string())) })), ) .unwrap(); @@ -1363,8 +1369,12 @@ async fn test_conditional_guardrail_allows() { reset_global(); setup_isolated_thread(); - register_tool_conditional_execution_guardrail("allower", 1, Arc::new(|_name, _args| Ok(None))) - .unwrap(); + register_tool_conditional_execution_guardrail( + "allower", + 1, + Arc::new(|_name, _args| Box::pin(async { Ok(None) })), + ) + .unwrap(); let func: ToolExecutionNextFn = Arc::new(|args| Box::pin(async move { Ok(args) })); @@ -1402,12 +1412,16 @@ async fn test_tool_conditional_guardrail_emits_guardrail_scope() { ) .unwrap(); - register_tool_conditional_execution_guardrail("tool_scope_allow", 1, Arc::new(|_, _| Ok(None))) - .unwrap(); + register_tool_conditional_execution_guardrail( + "tool_scope_allow", + 1, + Arc::new(|_, _| ready(None)), + ) + .unwrap(); register_tool_conditional_execution_guardrail( "tool_scope_reject", 2, - Arc::new(|_, _| Ok(Some("blocked by tool guardrail".to_string()))), + Arc::new(|_, _| ready(Some("blocked by tool guardrail".to_string()))), ) .unwrap(); @@ -1487,13 +1501,17 @@ async fn test_conditional_guardrail_first_rejection_wins() { reset_global(); setup_isolated_thread(); - register_tool_conditional_execution_guardrail("allows", 1, Arc::new(|_name, _args| Ok(None))) - .unwrap(); + register_tool_conditional_execution_guardrail( + "allows", + 1, + Arc::new(|_name, _args| Box::pin(async { Ok(None) })), + ) + .unwrap(); register_tool_conditional_execution_guardrail( "rejects", 2, - Arc::new(|_name, _args| Ok(Some("blocked by second".to_string()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("blocked by second".to_string())) })), ) .unwrap(); @@ -1533,9 +1551,9 @@ async fn test_conditional_guardrail_tool_name_filtering() { 1, Arc::new(|name, _args| { if name == "dangerous_tool" { - Ok(Some("dangerous_tool is forbidden".to_string())) + ready(Some("dangerous_tool is forbidden".to_string())) } else { - Ok(None) + ready(None) } }), ) @@ -1575,8 +1593,8 @@ async fn test_conditional_guardrail_tool_name_filtering() { /// Push scope, register scope-local guardrail, verify it applies, /// pop scope, verify it no longer applies. -#[test] -fn test_scope_local_guardrail_lifecycle() { +#[tokio::test] +async fn test_scope_local_guardrail_lifecycle() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); let handle = setup_isolated_scope("lifecycle_scope"); @@ -1591,7 +1609,7 @@ fn test_scope_local_guardrail_lifecycle() { 1, Arc::new(move |_name, args| { cc.fetch_add(1, Ordering::SeqCst); - args + ready(args) }), ) .unwrap(); @@ -1604,6 +1622,7 @@ fn test_scope_local_guardrail_lifecycle() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!( call_count.load(Ordering::SeqCst), 1, @@ -1700,8 +1719,8 @@ async fn test_scope_local_execution_intercept_cleanup() { /// Register global guardrail at priority 5, scope-local guardrail at priority 3. /// Verify scope-local runs first (lower priority number = higher priority). /// Verify both are applied. -#[test] -fn test_scope_local_and_global_guardrail_merge_priority() { +#[tokio::test] +async fn test_scope_local_and_global_guardrail_merge_priority() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); let handle = setup_isolated_scope("merge_scope"); @@ -1718,7 +1737,7 @@ fn test_scope_local_and_global_guardrail_merge_priority() { args.as_object_mut() .unwrap() .insert("global".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -1734,7 +1753,7 @@ fn test_scope_local_and_global_guardrail_merge_priority() { args.as_object_mut() .unwrap() .insert("local".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -1757,6 +1776,7 @@ fn test_scope_local_and_global_guardrail_merge_priority() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); // Verify order: local (priority 3) runs before global (priority 5) let recorded = order.lock().unwrap(); @@ -1887,7 +1907,7 @@ async fn test_conditional_rejection_prevents_intercepts() { register_tool_conditional_execution_guardrail( "gate", 1, - Arc::new(|_name, _args| Ok(Some("blocked".to_string()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("blocked".to_string())) })), ) .unwrap(); @@ -1899,7 +1919,7 @@ async fn test_conditional_rejection_prevents_intercepts() { false, Arc::new(move |_name, args| { ic.store(true, Ordering::SeqCst); - Ok(args) + ready(args) }), ) .unwrap(); @@ -1938,7 +1958,7 @@ async fn test_conditional_rejection_prevents_execution() { register_tool_conditional_execution_guardrail( "gate2", 1, - Arc::new(|_name, _args| Ok(Some("no execution".to_string()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("no execution".to_string())) })), ) .unwrap(); @@ -1989,8 +2009,8 @@ async fn test_conditional_rejection_prevents_execution() { // ========================================================================= /// Sanitize guardrails pipe data through sequentially. -#[test] -fn test_sanitize_guardrails_pipe_data() { +#[tokio::test] +async fn test_sanitize_guardrails_pipe_data() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -2003,7 +2023,7 @@ fn test_sanitize_guardrails_pipe_data() { args.as_object_mut() .unwrap() .insert("field_a".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -2018,7 +2038,7 @@ fn test_sanitize_guardrails_pipe_data() { args.as_object_mut() .unwrap() .insert("field_b".into(), json!(has_a)); - args + ready(args) }), ) .unwrap(); @@ -2061,8 +2081,8 @@ fn test_sanitize_guardrails_pipe_data() { } /// Response sanitize guardrails also pipe through. -#[test] -fn test_response_sanitize_guardrails_pipe() { +#[tokio::test] +async fn test_response_sanitize_guardrails_pipe() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -2075,7 +2095,7 @@ fn test_response_sanitize_guardrails_pipe() { .as_object_mut() .unwrap() .insert("sanitized".into(), json!(true)); - result + ready(result) }), ) .unwrap(); @@ -2127,8 +2147,8 @@ fn test_response_sanitize_guardrails_pipe() { /// Use multiple threads to register/deregister guardrails concurrently. /// Verify no panics or data races. -#[test] -fn test_concurrent_register_deregister() { +#[tokio::test] +async fn test_concurrent_register_deregister() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); @@ -2145,7 +2165,7 @@ fn test_concurrent_register_deregister() { let res = register_tool_sanitize_request_guardrail( &name, i, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ); assert!(res.is_ok(), "Registration should succeed for {name}"); @@ -2173,8 +2193,8 @@ fn test_concurrent_register_deregister() { } /// Concurrent register/deregister of intercepts across multiple threads. -#[test] -fn test_concurrent_intercept_mutations() { +#[tokio::test] +async fn test_concurrent_intercept_mutations() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); @@ -2191,7 +2211,7 @@ fn test_concurrent_intercept_mutations() { &name, i, false, - Arc::new(|_name, args| Ok(args)), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ); assert!(res.is_ok()); @@ -2217,8 +2237,8 @@ fn test_concurrent_intercept_mutations() { } /// Interleaved register and tool call execution from multiple threads. -#[test] -fn test_concurrent_register_and_read() { +#[tokio::test] +async fn test_concurrent_register_and_read() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); @@ -2227,7 +2247,7 @@ fn test_concurrent_register_and_read() { register_tool_sanitize_request_guardrail( &format!("stable_{i}"), i, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); } @@ -2246,7 +2266,7 @@ fn test_concurrent_register_and_read() { let _ = register_tool_sanitize_request_guardrail( &name, 100 + i, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ); std::thread::yield_now(); let _ = deregister_tool_sanitize_request_guardrail(&name); @@ -2280,8 +2300,8 @@ fn test_concurrent_register_and_read() { // Lock Regression Tests // ========================================================================= -#[test] -fn test_tool_request_intercept_registry_mutations_apply_to_later_calls() { +#[tokio::test] +async fn test_tool_request_intercept_registry_mutations_apply_to_later_calls() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -2308,23 +2328,27 @@ fn test_tool_request_intercept_registry_mutations_apply_to_later_calls() { Arc::new(move |_, args| { record_middleware_callback(&tracked, "tool_request_late"); assert_middleware_callback_locks_are_free(); - Ok(args) + ready(args) }), ) .unwrap(); } - Ok(args) + ready(args) }), ) .unwrap(); - let args = tool_request_intercepts("tool", json!({"round": 1})).unwrap(); + let args = tool_request_intercepts("tool", json!({"round": 1})) + .await + .unwrap(); assert_eq!(args["round"], 1); assert_middleware_callback_labels(&callbacks, &["tool_request_initial"]); callbacks.lock().unwrap().clear(); - let args = tool_request_intercepts("tool", json!({"round": 2})).unwrap(); + let args = tool_request_intercepts("tool", json!({"round": 2})) + .await + .unwrap(); assert_eq!(args["round"], 2); assert_middleware_callback_labels(&callbacks, &["tool_request_initial", "tool_request_late"]); @@ -2332,8 +2356,8 @@ fn test_tool_request_intercept_registry_mutations_apply_to_later_calls() { deregister_tool_request_intercept("snapshot_tool_request_late").unwrap(); } -#[test] -fn test_llm_request_intercept_registry_mutations_apply_to_later_calls() { +#[tokio::test] +async fn test_llm_request_intercept_registry_mutations_apply_to_later_calls() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -2360,7 +2384,7 @@ fn test_llm_request_intercept_registry_mutations_apply_to_later_calls() { Arc::new(move |_, request, annotated| { record_middleware_callback(&tracked, "llm_request_late"); assert_middleware_callback_locks_are_free(); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -2368,7 +2392,7 @@ fn test_llm_request_intercept_registry_mutations_apply_to_later_calls() { .unwrap(); } - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -2381,6 +2405,7 @@ fn test_llm_request_intercept_registry_mutations_apply_to_later_calls() { content: json!({"round": 1}), }, ) + .await .unwrap(); assert_eq!(request.request.content["round"], 1); assert_middleware_callback_labels(&callbacks, &["llm_request_initial"]); @@ -2393,6 +2418,7 @@ fn test_llm_request_intercept_registry_mutations_apply_to_later_calls() { content: json!({"round": 2}), }, ) + .await .unwrap(); assert_eq!(request.request.content["round"], 2); assert_middleware_callback_labels(&callbacks, &["llm_request_initial", "llm_request_late"]); @@ -2415,7 +2441,7 @@ async fn test_tool_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, _| { record_middleware_callback(&tracked, "tool_conditional_global"); assert_middleware_callback_locks_are_free(); - Ok(None) + ready(None) }), ) .unwrap(); @@ -2427,7 +2453,7 @@ async fn test_tool_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, _| { record_middleware_callback(&tracked, "tool_conditional_scope"); assert_middleware_callback_locks_are_free(); - Ok(None) + ready(None) }), ) .unwrap(); @@ -2439,7 +2465,7 @@ async fn test_tool_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, args| { record_middleware_callback(&tracked, "tool_request_global"); assert_middleware_callback_locks_are_free(); - Ok(args) + ready(args) }), ) .unwrap(); @@ -2452,7 +2478,7 @@ async fn test_tool_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, args| { record_middleware_callback(&tracked, "tool_request_scope"); assert_middleware_callback_locks_are_free(); - Ok(args) + ready(args) }), ) .unwrap(); @@ -2463,7 +2489,7 @@ async fn test_tool_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, args| { record_middleware_callback(&tracked, "tool_sanitize_request_global"); assert_middleware_callback_locks_are_free(); - args + ready(args) }), ) .unwrap(); @@ -2475,7 +2501,7 @@ async fn test_tool_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, args| { record_middleware_callback(&tracked, "tool_sanitize_request_scope"); assert_middleware_callback_locks_are_free(); - args + ready(args) }), ) .unwrap(); @@ -2509,7 +2535,7 @@ async fn test_tool_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, result| { record_middleware_callback(&tracked, "tool_sanitize_response_global"); assert_middleware_callback_locks_are_free(); - result + ready(result) }), ) .unwrap(); @@ -2521,7 +2547,7 @@ async fn test_tool_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, result| { record_middleware_callback(&tracked, "tool_sanitize_response_scope"); assert_middleware_callback_locks_are_free(); - result + ready(result) }), ) .unwrap(); @@ -2586,7 +2612,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_| { record_middleware_callback(&tracked, "llm_conditional_global"); assert_middleware_callback_locks_are_free(); - Ok(None) + ready(None) }), ) .unwrap(); @@ -2598,7 +2624,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_| { record_middleware_callback(&tracked, "llm_conditional_scope"); assert_middleware_callback_locks_are_free(); - Ok(None) + ready(None) }), ) .unwrap(); @@ -2610,7 +2636,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, request, annotated| { record_middleware_callback(&tracked, "llm_request_global"); assert_middleware_callback_locks_are_free(); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -2625,7 +2651,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |_, request, annotated| { record_middleware_callback(&tracked, "llm_request_scope"); assert_middleware_callback_locks_are_free(); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -2638,7 +2664,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |request, _context| { record_middleware_callback(&tracked, "llm_sanitize_request_global"); assert_middleware_callback_locks_are_free(); - Some(request) + ready(Some(request)) }), ) .unwrap(); @@ -2650,7 +2676,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |request, _context| { record_middleware_callback(&tracked, "llm_sanitize_request_scope"); assert_middleware_callback_locks_are_free(); - Some(request) + ready(Some(request)) }), ) .unwrap(); @@ -2707,7 +2733,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |response, _context| { record_middleware_callback(&tracked, "llm_sanitize_response_global"); assert_middleware_callback_locks_are_free(); - Some(response) + ready(Some(response)) }), ) .unwrap(); @@ -2719,7 +2745,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { Arc::new(move |response, _context| { record_middleware_callback(&tracked, "llm_sanitize_response_scope"); assert_middleware_callback_locks_are_free(); - Some(response) + ready(Some(response)) }), ) .unwrap(); @@ -2782,6 +2808,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { while let Some(chunk) = stream.next().await { chunk.unwrap(); } + stream.close().await.unwrap(); assert_middleware_callback_labels( &callbacks, &[ @@ -2852,7 +2879,7 @@ async fn test_full_pipeline_integration() { args.as_object_mut() .unwrap() .insert("intercepted".into(), json!(true)); - Ok(args) + ready(args) }), ) .unwrap(); @@ -2864,7 +2891,7 @@ async fn test_full_pipeline_integration() { 1, Arc::new(move |_name, args| { o2.lock().unwrap().push("sanitize_request".into()); - args + ready(args) }), ) .unwrap(); @@ -2876,7 +2903,7 @@ async fn test_full_pipeline_integration() { 1, Arc::new(move |_name, _args| { o3.lock().unwrap().push("conditional".into()); - Ok(None) // Allow + ready(None) // Allow }), ) .unwrap(); @@ -2903,7 +2930,7 @@ async fn test_full_pipeline_integration() { 1, Arc::new(move |_name, result| { o5.lock().unwrap().push("sanitize_response".into()); - result + ready(result) }), ) .unwrap(); @@ -2961,15 +2988,23 @@ async fn test_full_pipeline_integration() { // ========================================================================= /// Attempting to register a guardrail with the same name returns AlreadyExists. -#[test] -fn test_duplicate_guardrail_registration_returns_error() { +#[tokio::test] +async fn test_duplicate_guardrail_registration_returns_error() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); - register_tool_sanitize_request_guardrail("duplicate", 1, Arc::new(|_name, args| args)).unwrap(); + register_tool_sanitize_request_guardrail( + "duplicate", + 1, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ) + .unwrap(); - let err = - register_tool_sanitize_request_guardrail("duplicate", 2, Arc::new(|_name, args| args)); + let err = register_tool_sanitize_request_guardrail( + "duplicate", + 2, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ); assert!(err.is_err()); match err.unwrap_err() { @@ -2984,19 +3019,24 @@ fn test_duplicate_guardrail_registration_returns_error() { } /// Attempting to register an intercept with the same name returns AlreadyExists. -#[test] -fn test_duplicate_intercept_registration_returns_error() { +#[tokio::test] +async fn test_duplicate_intercept_registration_returns_error() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); - register_tool_request_intercept("dup_intercept", 1, false, Arc::new(|_name, args| Ok(args))) - .unwrap(); + register_tool_request_intercept( + "dup_intercept", + 1, + false, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ) + .unwrap(); let err = register_tool_request_intercept( "dup_intercept", 2, false, - Arc::new(|_name, args| Ok(args)), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ); assert!(err.is_err()); @@ -3016,8 +3056,8 @@ fn test_duplicate_intercept_registration_returns_error() { // ========================================================================= /// Deregistering a non-existent guardrail returns false. -#[test] -fn test_deregister_nonexistent_returns_false() { +#[tokio::test] +async fn test_deregister_nonexistent_returns_false() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); @@ -3029,8 +3069,8 @@ fn test_deregister_nonexistent_returns_false() { } /// Deregistering removes the guardrail from the chain. -#[test] -fn test_deregister_removes_from_chain() { +#[tokio::test] +async fn test_deregister_removes_from_chain() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -3043,7 +3083,7 @@ fn test_deregister_removes_from_chain() { 1, Arc::new(move |_name, args| { cc.fetch_add(1, Ordering::SeqCst); - args + ready(args) }), ) .unwrap(); @@ -3056,6 +3096,7 @@ fn test_deregister_removes_from_chain() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!(call_count.load(Ordering::SeqCst), 1); // Deregister @@ -3070,6 +3111,7 @@ fn test_deregister_removes_from_chain() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!( call_count.load(Ordering::SeqCst), 1, @@ -3091,7 +3133,7 @@ async fn test_llm_conditional_guardrail_rejects() { register_llm_conditional_execution_guardrail( "llm_gate", 1, - Arc::new(|_req| Ok(Some("LLM call rejected".to_string()))), + Arc::new(|_req| ready(Some("LLM call rejected".to_string()))), ) .unwrap(); @@ -3142,12 +3184,12 @@ async fn test_llm_conditional_guardrail_emits_guardrail_scope() { ) .unwrap(); - register_llm_conditional_execution_guardrail("llm_scope_allow", 1, Arc::new(|_| Ok(None))) + register_llm_conditional_execution_guardrail("llm_scope_allow", 1, Arc::new(|_| ready(None))) .unwrap(); register_llm_conditional_execution_guardrail( "llm_scope_reject", 2, - Arc::new(|_| Ok(Some("blocked by llm guardrail".to_string()))), + Arc::new(|_| ready(Some("blocked by llm guardrail".to_string()))), ) .unwrap(); @@ -3236,9 +3278,9 @@ async fn test_llm_request_intercept_transforms() { "llm_req_i", 1, false, - Arc::new(|_name: &str, mut req: LlmRequest, annotated| { + Arc::new(|_name: String, mut req: LlmRequest, annotated| { req.headers.insert("x-intercepted".into(), json!(true)); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -3250,15 +3292,15 @@ async fn test_llm_request_intercept_transforms() { content: json!({"prompt": "hello"}), }; - let result = llm_request_intercepts("test_llm", request).unwrap(); + let result = llm_request_intercepts("test_llm", request).await.unwrap(); assert_eq!(result.request.headers["x-intercepted"], true); // Cleanup deregister_llm_request_intercept("llm_req_i").unwrap(); } -#[test] -fn test_llm_request_intercept_pending_marks_preserve_order_and_break_chain() { +#[tokio::test] +async fn test_llm_request_intercept_pending_marks_preserve_order_and_break_chain() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -3273,8 +3315,10 @@ fn test_llm_request_intercept_pending_marks_preserve_order_and_break_chain() { priority, break_chain, Arc::new(move |_name, request, annotated| { - Ok(LlmRequestInterceptOutcome::new(request, annotated) - .with_pending_mark(PendingMarkSpec::builder().name(mark_name).build())) + ready( + LlmRequestInterceptOutcome::new(request, annotated) + .with_pending_mark(PendingMarkSpec::builder().name(mark_name).build()), + ) }), ) .unwrap(); @@ -3287,6 +3331,7 @@ fn test_llm_request_intercept_pending_marks_preserve_order_and_break_chain() { content: json!({"prompt": "hello"}), }, ) + .await .unwrap(); assert_eq!( @@ -3322,7 +3367,7 @@ async fn test_managed_llm_emits_pending_marks_under_started_scope() { 1, Arc::new(|event, mut fields| { fields.metadata = Some(json!({"sanitized_mark": event.name()})); - fields + ready(fields) }), ) .unwrap(); @@ -3331,24 +3376,26 @@ async fn test_managed_llm_emits_pending_marks_under_started_scope() { 1, false, Arc::new(|_name, request, annotated| { - Ok(LlmRequestInterceptOutcome::new(request, annotated) - .with_pending_mark( - PendingMarkSpec::builder() - .name("request.optimized") - .category(EventCategory::custom()) - .category_profile( - CategoryProfile::builder() - .subtype("optimizer.saved_tokens") - .build(), - ) - .data(json!({"saved_tokens": 12})) - .build(), - ) - .with_pending_mark( - PendingMarkSpec::builder() - .name("request.optimized.second") - .build(), - )) + ready( + LlmRequestInterceptOutcome::new(request, annotated) + .with_pending_mark( + PendingMarkSpec::builder() + .name("request.optimized") + .category(EventCategory::custom()) + .category_profile( + CategoryProfile::builder() + .subtype("optimizer.saved_tokens") + .build(), + ) + .data(json!({"saved_tokens": 12})) + .build(), + ) + .with_pending_mark( + PendingMarkSpec::builder() + .name("request.optimized.second") + .build(), + ), + ) }), ) .unwrap(); @@ -3449,7 +3496,7 @@ async fn test_managed_llm_materializes_optimization_mark_and_end_summary() { data.insert("payload".to_string(), json!({"secret": "[redacted]"})); data.remove("future_secret"); } - fields + ready(fields) }), ) .unwrap(); @@ -3466,7 +3513,7 @@ async fn test_managed_llm_materializes_optimization_mark_and_end_summary() { contribution.payload = Some(json!({"secret": "[scope-end-redacted]"})); contribution.extra.remove("future_secret"); } - fields + ready(fields) }), ) .unwrap(); @@ -3492,8 +3539,10 @@ async fn test_managed_llm_materializes_optimization_mark_and_end_summary() { contribution .extra .insert("future_secret".to_string(), json!("classified")); - Ok(LlmRequestInterceptOutcome::new(request, annotated) - .with_optimization_contribution(contribution)) + ready( + LlmRequestInterceptOutcome::new(request, annotated) + .with_optimization_contribution(contribution), + ) }), ) .unwrap(); @@ -3698,7 +3747,7 @@ async fn test_stream_optimization_mark_uses_the_llm_captured_sanitizer_scope() { { data.insert("payload".to_string(), json!({"secret": "[redacted]"})); } - fields + ready(fields) }), ) .unwrap(); @@ -3753,6 +3802,7 @@ async fn test_stream_optimization_mark_uses_the_llm_captured_sanitizer_scope() { while let Some(item) = stream.next().await { item.unwrap(); } + stream.close().await.unwrap(); set_thread_scope_stack(original_stack); let captured = captured_events_snapshot(&events); @@ -3798,8 +3848,10 @@ async fn test_concurrent_managed_llm_calls_keep_optimization_evidence_isolated() saved: Some(LlmOptimizationTokens::saved_prompt(saved_tokens)), ..LlmOptimizationTokenImpact::default() }); - Ok(LlmRequestInterceptOutcome::new(request, annotated) - .with_optimization_contribution(contribution)) + ready( + LlmRequestInterceptOutcome::new(request, annotated) + .with_optimization_contribution(contribution), + ) }), ) .unwrap(); @@ -3890,8 +3942,10 @@ async fn test_failed_request_intercept_does_not_emit_pending_marks_or_start_scop 1, false, Arc::new(|_name, request, annotated| { - Ok(LlmRequestInterceptOutcome::new(request, annotated) - .with_pending_mark(PendingMarkSpec::builder().name("must.not.emit").build())) + ready( + LlmRequestInterceptOutcome::new(request, annotated) + .with_pending_mark(PendingMarkSpec::builder().name("must.not.emit").build()), + ) }), ) .unwrap(); @@ -3900,7 +3954,7 @@ async fn test_failed_request_intercept_does_not_emit_pending_marks_or_start_scop 2, false, Arc::new(|_name, _request, _annotated| { - Err(FlowError::Internal("request intercept failed".into())) + ready_result(Err(FlowError::Internal("request intercept failed".into()))) }), ) .unwrap(); @@ -4015,7 +4069,7 @@ async fn test_llm_start_emits_before_short_circuit_execution_intercept() { .as_object_mut() .unwrap() .insert("phase".into(), json!("request")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -4109,7 +4163,7 @@ async fn test_llm_stream_start_emits_before_short_circuit_execution_intercept() .as_object_mut() .unwrap() .insert("phase".into(), json!("request")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -4162,6 +4216,7 @@ async fn test_llm_stream_start_emits_before_short_circuit_execution_intercept() while let Some(chunk) = stream.next().await { chunk.unwrap(); } + stream.close().await.unwrap(); assert!( !original_called.load(Ordering::SeqCst), @@ -4190,19 +4245,19 @@ async fn test_llm_stream_start_emits_before_short_circuit_execution_intercept() // ========================================================================= /// tool_conditional_execution returns Ok(()) when no guardrails reject. -#[test] -fn test_standalone_conditional_execution_passes() { +#[tokio::test] +async fn test_standalone_conditional_execution_passes() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); - let result = tool_conditional_execution("tool", &json!({})); + let result = tool_conditional_execution("tool", &json!({})).await; assert!(result.is_ok(), "No guardrails means no rejection"); } /// tool_conditional_execution returns GuardrailRejected when a guardrail rejects. -#[test] -fn test_standalone_conditional_execution_rejects() { +#[tokio::test] +async fn test_standalone_conditional_execution_rejects() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); @@ -4210,11 +4265,11 @@ fn test_standalone_conditional_execution_rejects() { register_tool_conditional_execution_guardrail( "standalone_gate", 1, - Arc::new(|_name, _args| Ok(Some("rejected by standalone".to_string()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("rejected by standalone".to_string())) })), ) .unwrap(); - let result = tool_conditional_execution("tool", &json!({})); + let result = tool_conditional_execution("tool", &json!({})).await; assert!(result.is_err()); match result.unwrap_err() { FlowError::GuardrailRejected(reason) => { @@ -4257,12 +4312,14 @@ async fn test_empty_chain_passthrough() { } /// Standalone intercept chain with no registrations returns input unchanged. -#[test] -fn test_empty_request_intercept_chain() { +#[tokio::test] +async fn test_empty_request_intercept_chain() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); - let result = tool_request_intercepts("tool", json!({"key": "val"})).unwrap(); + let result = tool_request_intercepts("tool", json!({"key": "val"})) + .await + .unwrap(); assert_eq!(result["key"], "val"); } diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index ee8a860be..c669f7c46 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -239,6 +239,7 @@ async fn sdk_cdylib_registers_tool_request_intercept() { .expect("outer scope should push"); let outer_uuid = outer.uuid; let rewritten = tool_request_intercepts("demo_tool", json!({ "input": "value" })) + .await .expect("native request intercept should run"); let tool_result = tool_call_execute( ToolCallExecuteParams::builder() @@ -446,6 +447,7 @@ async fn sdk_cdylib_registers_tool_request_intercept() { .expect("thread outer scope should push"); let thread_outer_uuid = thread_outer.uuid; let rewritten = tool_request_intercepts("demo_tool", json!({ "input": "thread" })) + .await .expect("native request intercept should run with thread stack"); assert_eq!(rewritten["native_plugin"], true); pop_scope( @@ -656,6 +658,123 @@ async fn sdk_cdylib_registers_tool_request_intercept() { activation.clear(); } +#[tokio::test] +async fn native_v3_async_registration_supports_all_middleware_kinds() { + let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; + let fixture = build_fixture_plugin(); + let manifest_ref = write_manifest_with_plugin_id_and_symbol( + &fixture, + "fixture_async", + "nemo_relay_fixture_async_entry", + ); + + let activation = load_native_plugins([NativePluginLoadSpec { + plugin_id: "fixture_async".into(), + manifest_ref: manifest_ref.to_string_lossy().into_owned(), + }]) + .expect("v3 async native fixture should load"); + let mut config = PluginConfig::default(); + config.components.push(PluginComponentSpec { + kind: "fixture_async".into(), + enabled: true, + config: Map::new(), + }); + initialize_plugins_exact(config) + .await + .expect("v3 async native fixture should register"); + + let rewritten = tool_request_intercepts("async-tool", json!({"input": true})) + .await + .expect("v3 async request intercept should settle"); + assert_eq!(rewritten["input"], true); + assert_eq!(rewritten["native_async"], true); + + let duplicate = tool_request_intercepts("async-double", json!({"input": true})) + .await + .expect("duplicate v3 async settlement keeps the first result"); + assert_eq!(duplicate["native_async"], true); + + let executed = tool_call_execute( + ToolCallExecuteParams::builder() + .name("async-execution") + .args(json!({"input": true})) + .func(Arc::new(|args| Box::pin(async move { Ok(args) }))) + .build(), + ) + .await + .expect("v3 async execution intercept should continue with next"); + assert_eq!(executed["native_async_execution"], true); + + let llm_response = llm_call_execute( + LlmCallExecuteParams::builder() + .name("async-llm") + .request(LlmRequest { + headers: Map::new(), + content: json!({"prompt": "native async"}), + }) + .func(Arc::new(|_request| { + Box::pin(async move { Ok(json!({"content": "native async response"})) }) + })) + .build(), + ) + .await + .expect("v3 async LLM middleware should settle"); + assert_eq!(llm_response["content"], "native async response"); + flush_subscribers().expect("async native LLM events should flush"); + + let stream_chunks = Arc::new(Mutex::new(Vec::::new())); + let collected_chunks = stream_chunks.clone(); + let finalized_chunks = stream_chunks.clone(); + let mut stream = llm_stream_call_execute( + LlmStreamCallExecuteParams::builder() + .name("async-llm-stream") + .request(LlmRequest { + headers: Map::new(), + content: json!({"prompt": "native async stream"}), + }) + .func(Arc::new(|_request| { + Box::pin(async move { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok(json!({ + "content": "native async stream response" + }))]))) + }) + })) + .collector(Box::new(move |chunk| { + collected_chunks.lock().unwrap().push(chunk); + Ok(()) + })) + .finalizer(Box::new(move || { + Json::Array(finalized_chunks.lock().unwrap().clone()) + })) + .build(), + ) + .await + .expect("v3 async LLM stream middleware should settle"); + assert_eq!( + stream + .next() + .await + .expect("stream should contain a chunk") + .expect("stream chunk should succeed")["content"], + "native async stream response" + ); + assert!(stream.next().await.is_none()); + flush_subscribers().expect("async native LLM stream events should flush"); + + let pending = tokio::spawn(async { + tool_request_intercepts("async-pending", json!({"input": true})).await + }); + tokio::task::yield_now().await; + clear_plugin_configuration().expect("v3 async native fixture should clear while pending"); + let pending = pending + .await + .expect("pending v3 async task should not panic") + .expect("pending v3 async request intercept should settle after clear"); + assert_eq!(pending["native_async"], true); + + drop(activation); +} + #[tokio::test] async fn native_validation_diagnostics_prevent_initialization() { let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; @@ -752,7 +871,7 @@ async fn native_tool_execution_rejects_null_malformed_and_error_outcomes() { } #[tokio::test] -async fn native_event_sanitizer_callback_errors_clear_observability_fields() { +async fn native_event_sanitizer_callback_errors_preserve_observability_fields() { let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; let fixture = build_fixture_plugin(); let manifest_ref = @@ -793,8 +912,8 @@ async fn native_event_sanitizer_callback_errors_clear_observability_fields() { let captured_events = events.lock().unwrap().clone(); let event = find_event(&captured_events, "native-event-sanitize-error", None); - assert_eq!(event.data(), None); - assert_eq!(event.metadata(), None); + assert_eq!(event.data(), Some(&json!({ "secret": true }))); + assert_eq!(event.metadata(), Some(&json!({ "secret": true }))); deregister_subscriber("native_event_sanitizer_error_capture") .expect("test subscriber should deregister"); @@ -1218,6 +1337,7 @@ async fn plugin_host_activation_owns_configuration_until_clear() { .any(|kind| kind == "fixture_native") ); let rewritten = tool_request_intercepts("host-owned-tool", json!({ "input": true })) + .await .expect("host-owned intercept should run"); assert_eq!(rewritten["native_plugin"], true); @@ -1238,6 +1358,7 @@ async fn plugin_host_activation_owns_configuration_until_clear() { .any(|kind| kind == "fixture_native") ); let unchanged = tool_request_intercepts("host-owned-tool", json!({ "input": true })) + .await .expect("cleared intercept chain should be empty"); assert_eq!(unchanged, json!({ "input": true })); } @@ -1393,6 +1514,7 @@ async fn plugin_host_clear_allows_an_in_flight_native_callback_to_finish() { .clear() .expect("host should clear while a callback snapshot remains in flight"); let unchanged = tool_request_intercepts("after-clear", json!({ "input": true })) + .await .expect("new calls should observe the cleared registries"); assert_eq!(unchanged, json!({ "input": true })); diff --git a/crates/core/tests/integration/pipeline_tests.rs b/crates/core/tests/integration/pipeline_tests.rs index 1bc192b4b..9711974de 100644 --- a/crates/core/tests/integration/pipeline_tests.rs +++ b/crates/core/tests/integration/pipeline_tests.rs @@ -9,6 +9,9 @@ use std::sync::{Arc, Mutex}; +mod test_support; +use test_support::ready; + use futures::StreamExt; use serde_json::json; @@ -359,7 +362,7 @@ async fn test_decode_runs_before_intercepts() { false, Arc::new(move |_name, req, annotated| { *cap.lock().unwrap() = Some(annotated.clone()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -414,7 +417,7 @@ async fn test_encode_runs_after_intercepts() { let mut ann = annotated.unwrap(); ann.model = Some("modified".into()); req.headers.insert("x-codec-route".into(), json!("blue")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, Some(ann), )) @@ -517,7 +520,7 @@ async fn anthropic_issue_501_round_trips_and_applies_annotated_edits() { request .headers .insert("x-annotation-seen".into(), json!("yes")); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + ready(LlmRequestInterceptOutcome::new(request, Some(annotated))) }), ) .unwrap(); @@ -571,8 +574,10 @@ async fn test_codec_rejects_raw_content_mutation_before_lifecycle() { false, Arc::new(|_name, mut request, annotated| { request.content["model"] = json!("raw-model-edit"); - Ok(LlmRequestInterceptOutcome::new(request, annotated) - .with_pending_mark(PendingMarkSpec::builder().name("must.not.emit").build())) + ready( + LlmRequestInterceptOutcome::new(request, annotated) + .with_pending_mark(PendingMarkSpec::builder().name("must.not.emit").build()), + ) }), ) .unwrap(); @@ -583,7 +588,7 @@ async fn test_codec_rejects_raw_content_mutation_before_lifecycle() { false, Arc::new(move |_name, request, annotated| { *later_called.lock().unwrap() = true; - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + ready(LlmRequestInterceptOutcome::new(request, annotated)) }), ) .unwrap(); @@ -630,7 +635,9 @@ async fn test_codec_rejects_missing_annotation_before_lifecycle() { "codec_missing_annotation", 1, false, - Arc::new(|_name, request, _annotated| Ok(LlmRequestInterceptOutcome::new(request, None))), + Arc::new(|_name, request, _annotated| { + ready(LlmRequestInterceptOutcome::new(request, None)) + }), ) .unwrap(); @@ -680,7 +687,7 @@ async fn test_stream_codec_rejects_raw_content_mutation_before_lifecycle() { false, Arc::new(|_name, mut request, annotated| { request.content["model"] = json!("raw-stream-edit"); - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + ready(LlmRequestInterceptOutcome::new(request, annotated)) }), ) .unwrap(); @@ -741,7 +748,7 @@ async fn test_annotated_intercept_receives_both() { false, Arc::new(move |_name, req, annotated| { *cp.lock().unwrap() = Some((req.clone(), annotated.clone())); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -799,7 +806,7 @@ async fn test_canonical_intercept_with_and_without_codec() { Arc::new(move |_name, mut req, annotated| { *lc1.lock().unwrap() = true; req.headers.insert("x-legacy".into(), json!("was-here")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -845,7 +852,7 @@ async fn test_canonical_intercept_with_and_without_codec() { Arc::new(move |_name, mut req, annotated| { *lc2.lock().unwrap() = true; req.headers.insert("x-legacy-2".into(), json!("also-here")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -903,7 +910,7 @@ async fn test_stream_path_also_decodes() { false, Arc::new(move |_name, req, annotated| { *ca.lock().unwrap() = Some(annotated.clone()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -930,6 +937,7 @@ async fn test_stream_path_also_decodes() { // Consume the stream to trigger full pipeline while let Some(_chunk) = stream.next().await {} + stream.close().await.unwrap(); // Assert decode was called let dl = decode_log.lock().unwrap(); @@ -969,7 +977,7 @@ async fn test_shared_helper_both_paths() { false, Arc::new(move |_name, req, annotated| { *acc.lock().unwrap() += 1; - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1048,7 +1056,7 @@ async fn test_explicit_codec_param_overrides() { if let Some(ref ann) = annotated { *cm.lock().unwrap() = ann.model.clone(); } - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1097,7 +1105,7 @@ async fn test_encode_merge_not_replace() { Arc::new(|_name, req, annotated| { let mut ann = annotated.unwrap(); ann.model = Some("new_model".into()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, Some(ann), )) @@ -1168,7 +1176,7 @@ async fn test_unified_chain_priority_order() { false, Arc::new(move |_name, req, annotated| { cl1.lock().unwrap().push("legacy_p10".into()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1183,7 +1191,7 @@ async fn test_unified_chain_priority_order() { false, Arc::new(move |_name, req, annotated| { cl2.lock().unwrap().push("annotated_p5".into()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1233,7 +1241,7 @@ async fn test_no_codec_annotated_intercept_receives_none() { false, Arc::new(move |_name, req, annotated| { *ca.lock().unwrap() = Some(annotated.clone()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1448,7 +1456,7 @@ async fn test_response_codec_annotation_uses_sanitized_managed_response() { register_llm_sanitize_response_guardrail( "sanitize_resp_codec_annotation", 1, - Arc::new(|_response, _context| Some(make_openai_chat_response("Sanitized"))), + Arc::new(|_response, _context| ready(Some(make_openai_chat_response("Sanitized")))), ) .unwrap(); @@ -1649,10 +1657,10 @@ async fn test_request_codec_annotation_uses_sanitized_start_payload() { "sanitize_req_codec_annotation", 1, Arc::new(|request, _context| { - Some(LlmRequest { + ready(Some(LlmRequest { headers: request.headers, content: make_openai_chat_request("Sanitized").content, - }) + })) }), ) .unwrap(); @@ -1733,6 +1741,7 @@ async fn test_stream_response_codec_populates_annotated_response() { // Drain the stream to trigger finalization and END event while let Some(_chunk) = stream.next().await {} + stream.close().await.unwrap(); let captured = captured_events_snapshot(&events); let end_event = captured @@ -1818,6 +1827,7 @@ async fn managed_buffered_and_streaming_close_price_the_committed_route_not_resp while let Some(item) = stream.next().await { item.unwrap(); } + stream.close().await.unwrap(); llm_call_execute( LlmCallExecuteParams::builder() @@ -1916,7 +1926,7 @@ async fn test_stream_response_codec_annotation_uses_sanitized_aggregated_respons register_llm_sanitize_response_guardrail( "stream_sanitize_resp_codec_annotation", 1, - Arc::new(|_response, _context| Some(make_openai_chat_response("Sanitized"))), + Arc::new(|_response, _context| ready(Some(make_openai_chat_response("Sanitized")))), ) .unwrap(); @@ -1939,6 +1949,7 @@ async fn test_stream_response_codec_annotation_uses_sanitized_aggregated_respons .unwrap(); while let Some(_chunk) = stream.next().await {} + stream.close().await.unwrap(); let captured = captured_events_snapshot(&events); let end_event = captured diff --git a/crates/core/tests/integration/scope_local_tests.rs b/crates/core/tests/integration/scope_local_tests.rs index 551d39dd7..ceb17e821 100644 --- a/crates/core/tests/integration/scope_local_tests.rs +++ b/crates/core/tests/integration/scope_local_tests.rs @@ -84,7 +84,7 @@ fn test_scope_local_guardrail_registration_and_execution() { args.as_object_mut() .unwrap() .insert("scope_sanitized".into(), json!(true)); - args + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -166,7 +166,7 @@ async fn test_auto_cleanup_on_scope_pop() { args.as_object_mut() .unwrap() .insert("ephemeral".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -234,7 +234,7 @@ async fn test_priority_merge_global_and_scope_local() { args.as_object_mut() .unwrap() .insert("p10".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -250,7 +250,7 @@ async fn test_priority_merge_global_and_scope_local() { args.as_object_mut() .unwrap() .insert("p30".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -267,7 +267,7 @@ async fn test_priority_merge_global_and_scope_local() { args.as_object_mut() .unwrap() .insert("p20".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -326,7 +326,7 @@ fn test_name_coexistence_global_and_scope_local() { 1, Arc::new(move |_name, args| { c1.fetch_add(1, Ordering::SeqCst); - args + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -339,7 +339,7 @@ fn test_name_coexistence_global_and_scope_local() { 2, Arc::new(move |_name, args| { c2.fetch_add(1, Ordering::SeqCst); - args + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -354,6 +354,7 @@ fn test_name_coexistence_global_and_scope_local() { .unwrap(); // Both guardrails with the same name ran. + flush_subscribers().unwrap(); assert_eq!(count.load(Ordering::SeqCst), 2); // Cleanup @@ -400,7 +401,7 @@ async fn test_scope_isolation_between_stacks() { args.as_object_mut() .unwrap() .insert("agent".into(), json!("a")); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -426,7 +427,7 @@ async fn test_scope_isolation_between_stacks() { args.as_object_mut() .unwrap() .insert("agent".into(), json!("b")); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -506,7 +507,7 @@ async fn test_nested_scope_inheritance() { args.as_object_mut() .unwrap() .insert("global".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -530,7 +531,7 @@ async fn test_nested_scope_inheritance() { args.as_object_mut() .unwrap() .insert("scope_a".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -555,7 +556,7 @@ async fn test_nested_scope_inheritance() { args.as_object_mut() .unwrap() .insert("scope_b".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -702,11 +703,13 @@ async fn test_scope_local_conditional_execution_guardrail() { "tool_blocker", 1, Arc::new(|name, _args| { - if name == "banned_tool" { - Ok(Some("banned_tool is not allowed in this scope".to_string())) - } else { - Ok(None) - } + Box::pin(async move { + if name == "banned_tool" { + Ok(Some("banned_tool is not allowed in this scope".to_string())) + } else { + Ok(None) + } + }) }), ) .unwrap(); diff --git a/crates/core/tests/integration/subscriber_dispatcher_tests.rs b/crates/core/tests/integration/subscriber_dispatcher_tests.rs index 2276dee6c..303ee112b 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -7,7 +7,6 @@ use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, mpsc}; use std::time::Duration; -use nemo_relay::api::event::Event; use nemo_relay::api::registry::{ deregister_mark_sanitize_guardrail, register_mark_sanitize_guardrail, }; @@ -16,7 +15,7 @@ use nemo_relay::api::runtime::{ }; use nemo_relay::api::scope::{EmitMarkEventParams, event}; use nemo_relay::api::subscriber::{deregister_subscriber, flush_subscribers, register_subscriber}; -use serde_json::json; +use nemo_relay::error::FlowError; static TEST_MUTEX: Mutex<()> = Mutex::new(()); @@ -103,146 +102,107 @@ fn dispatcher_preserves_event_order() { } #[test] -fn mark_emission_snapshots_sanitizers_and_returns_before_they_finish() { +fn dispatcher_continues_after_subscriber_panic() { let _lock = TEST_MUTEX.lock().unwrap(); flush_subscribers().unwrap(); reset_global(); setup_isolated_thread(); - let (sanitizer_started_tx, sanitizer_started_rx) = mpsc::channel(); - let (release_tx, release_rx) = mpsc::channel(); - let release_rx = Arc::new(Mutex::new(release_rx)); - register_mark_sanitize_guardrail( - "blocking-mark-sanitizer", - 10, - Arc::new(move |_, mut fields| { - sanitizer_started_tx.send(()).unwrap(); - release_rx.lock().unwrap().recv().unwrap(); - fields.data = Some(json!({"sanitized": true})); - fields - }), - ) - .unwrap(); - - let observed = Arc::new(Mutex::new(Vec::::new())); + let observed = Arc::new(Mutex::new(Vec::new())); let observed_events = Arc::clone(&observed); register_subscriber( - "sanitized-mark-subscriber", - Arc::new(move |event| observed_events.lock().unwrap().push(event.clone())), + "panic-isolated-subscriber", + Arc::new(move |event| { + if event.name() == "panic-isolated" { + panic!("subscriber failed"); + } + observed_events + .lock() + .unwrap() + .push(event.name().to_string()); + }), ) .unwrap(); - let (returned_tx, returned_rx) = mpsc::channel(); - let event_thread = std::thread::spawn(move || { - emit_mark("queued-sanitizer"); - returned_tx.send(()).unwrap(); - }); - - sanitizer_started_rx - .recv_timeout(Duration::from_secs(1)) - .expect("sanitizer should start on the dispatcher thread"); - returned_rx - .recv_timeout(Duration::from_secs(1)) - .expect("mark emission should return while its sanitizer is blocked"); - - // Removing the global registration cannot affect the already-snapshotted - // publication chain. - deregister_mark_sanitize_guardrail("blocking-mark-sanitizer").unwrap(); - release_tx.send(()).unwrap(); - event_thread.join().unwrap(); + emit_mark("panic-isolated"); + emit_mark("after-panic"); flush_subscribers().unwrap(); + deregister_subscriber("panic-isolated-subscriber").unwrap(); - let events = observed.lock().unwrap(); - assert_eq!(events.len(), 1); - assert_eq!( - events[0].sanitize_fields().data, - Some(json!({"sanitized": true})) - ); - drop(events); - deregister_subscriber("sanitized-mark-subscriber").unwrap(); + assert_eq!(observed.lock().unwrap().as_slice(), ["after-panic"]); } #[test] -fn mark_emission_skips_sanitizers_without_subscribers() { +fn dispatcher_publishes_the_snapshot_when_an_async_sanitizer_fails() { let _lock = TEST_MUTEX.lock().unwrap(); flush_subscribers().unwrap(); reset_global(); setup_isolated_thread(); - let sanitizer_called = Arc::new(AtomicBool::new(false)); - let called = Arc::clone(&sanitizer_called); + let observed = Arc::new(Mutex::new(Vec::new())); + let observed_events = Arc::clone(&observed); + register_subscriber( + "fail-open-sanitizer-subscriber", + Arc::new(move |event| { + observed_events + .lock() + .unwrap() + .push(event.name().to_string()) + }), + ) + .unwrap(); register_mark_sanitize_guardrail( - "unused-mark-sanitizer", + "fail-open-mark-sanitizer", 10, - Arc::new(move |_, fields| { - called.store(true, Ordering::Release); - fields + Arc::new(|_, _| { + Box::pin(async { + Err(FlowError::Internal( + "intentional event-sanitizer failure".to_string(), + )) + }) }), ) .unwrap(); - emit_mark("no-subscribers"); + emit_mark("unsanitized-fallback"); flush_subscribers().unwrap(); - deregister_mark_sanitize_guardrail("unused-mark-sanitizer").unwrap(); - assert!(!sanitizer_called.load(Ordering::Acquire)); + assert_eq!( + observed.lock().unwrap().as_slice(), + ["unsanitized-fallback"] + ); + deregister_mark_sanitize_guardrail("fail-open-mark-sanitizer").unwrap(); + deregister_subscriber("fail-open-sanitizer-subscriber").unwrap(); } #[test] -fn sanitizer_panic_publishes_the_latest_valid_event() { +fn mark_emission_skips_sanitizers_without_subscribers() { let _lock = TEST_MUTEX.lock().unwrap(); flush_subscribers().unwrap(); reset_global(); setup_isolated_thread(); + let sanitizer_called = Arc::new(AtomicBool::new(false)); + let called = Arc::clone(&sanitizer_called); register_mark_sanitize_guardrail( - "successful-mark-sanitizer", - 0, - Arc::new(move |_, mut fields| { - fields.data = Some(json!({"redacted": true})); - fields - }), - ) - .unwrap(); - register_mark_sanitize_guardrail( - "panicking-mark-sanitizer", + "unused-mark-sanitizer", 10, - Arc::new(move |_, _| panic!("sanitizer failed")), - ) - .unwrap(); - - let observed = Arc::new(Mutex::new(Vec::::new())); - let observed_events = Arc::clone(&observed); - register_subscriber( - "panic-fallback-subscriber", - Arc::new(move |event| observed_events.lock().unwrap().push(event.clone())), + Arc::new(move |_, fields| { + called.store(true, Ordering::Release); + Box::pin(async move { Ok(fields) }) + }), ) .unwrap(); - event( - EmitMarkEventParams::builder() - .name("panic-fallback") - .data(json!({"original": true})) - .build(), - ) - .unwrap(); + emit_mark("no-subscribers"); flush_subscribers().unwrap(); + deregister_mark_sanitize_guardrail("unused-mark-sanitizer").unwrap(); - deregister_mark_sanitize_guardrail("successful-mark-sanitizer").unwrap(); - deregister_mark_sanitize_guardrail("panicking-mark-sanitizer").unwrap(); - deregister_subscriber("panic-fallback-subscriber").unwrap(); - - let events = observed.lock().unwrap(); - assert_eq!(events.len(), 1); - assert_eq!(events[0].name(), "panic-fallback"); - assert_eq!( - events[0].sanitize_fields().data, - Some(json!({"redacted": true})) - ); + assert!(!sanitizer_called.load(Ordering::Acquire)); } #[test] -fn dispatcher_continues_after_subscriber_panic() { +fn dispatcher_publishes_the_snapshot_when_an_async_sanitizer_panics() { let _lock = TEST_MUTEX.lock().unwrap(); flush_subscribers().unwrap(); reset_global(); @@ -251,23 +211,26 @@ fn dispatcher_continues_after_subscriber_panic() { let observed = Arc::new(Mutex::new(Vec::new())); let observed_events = Arc::clone(&observed); register_subscriber( - "panic-isolated-subscriber", + "panic-sanitizer-subscriber", Arc::new(move |event| { - if event.name() == "panic-isolated" { - panic!("subscriber failed"); - } observed_events .lock() .unwrap() - .push(event.name().to_string()); + .push(event.name().to_string()) }), ) .unwrap(); + register_mark_sanitize_guardrail( + "panic-mark-sanitizer", + 10, + Arc::new(|_, _| Box::pin(async { panic!("intentional event-sanitizer panic") })), + ) + .unwrap(); - emit_mark("panic-isolated"); - emit_mark("after-panic"); + emit_mark("panic-fallback"); flush_subscribers().unwrap(); - deregister_subscriber("panic-isolated-subscriber").unwrap(); - assert_eq!(observed.lock().unwrap().as_slice(), ["after-panic"]); + assert_eq!(observed.lock().unwrap().as_slice(), ["panic-fallback"]); + deregister_mark_sanitize_guardrail("panic-mark-sanitizer").unwrap(); + deregister_subscriber("panic-sanitizer-subscriber").unwrap(); } diff --git a/crates/core/tests/integration/test_support.rs b/crates/core/tests/integration/test_support.rs new file mode 100644 index 000000000..bf1507f0f --- /dev/null +++ b/crates/core/tests/integration/test_support.rs @@ -0,0 +1,18 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::future::Future; +use std::pin::Pin; + +pub fn ready( + value: T, +) -> Pin> + Send>> { + Box::pin(async move { Ok(value) }) +} + +#[allow(dead_code)] +pub fn ready_result( + value: nemo_relay::error::Result, +) -> Pin> + Send>> { + Box::pin(async move { value }) +} diff --git a/crates/core/tests/integration/worker_plugin_tests.rs b/crates/core/tests/integration/worker_plugin_tests.rs index 4387b048a..b7eb75611 100644 --- a/crates/core/tests/integration/worker_plugin_tests.rs +++ b/crates/core/tests/integration/worker_plugin_tests.rs @@ -76,6 +76,7 @@ async fn plugin_host_activation_owns_worker_lifecycle() { .any(|kind| kind == "fixture_worker") ); let rewritten = tool_request_intercepts("worker-host-tool", json!({ "input": true })) + .await .expect("worker host intercept should run"); assert_eq!(rewritten["worker_plugin"], true); @@ -86,6 +87,7 @@ async fn plugin_host_activation_owns_worker_lifecycle() { .any(|kind| kind == "fixture_worker") ); let unchanged = tool_request_intercepts("worker-host-tool", json!({ "input": true })) + .await .expect("cleared worker intercept chain should be empty"); assert_eq!(unchanged, json!({ "input": true })); } @@ -109,6 +111,7 @@ async fn plugin_host_clear_surfaces_worker_shutdown_failure_and_releases_safe_ow .expect("worker plugin host should activate"); tool_request_intercepts("terminate-worker", json!({ "input": true })) + .await .expect_err("fixture worker should terminate during callback"); let error = activation .clear() @@ -159,6 +162,7 @@ async fn rust_worker_registers_and_invokes_all_current_surfaces() { .expect("outer scope should push"); let outer_uuid = outer.uuid; let rewritten = tool_request_intercepts("demo_tool", json!({ "input": "value" })) + .await .expect("worker request intercept should run"); let tool_result = tool_call_execute( ToolCallExecuteParams::builder() @@ -473,6 +477,7 @@ async fn worker_request_intercept_callback_error_surfaces_to_host() { .await; let error = tool_request_intercepts("demo_tool", json!({ "input": "value" })) + .await .expect_err("worker callback error should surface"); assert!( error @@ -1120,6 +1125,7 @@ async fn python_worker_host_runtime_mark_and_mutated_request_round_trip() { cleanup.subscriber_name = Some(subscriber_name); let rewritten = tool_request_intercepts("lookup", json!({ "query": "relay" })) + .await .expect("Python callback should emit a mark and return its mutation"); assert_eq!( rewritten["_nemo_relay_plugin"]["tag"], diff --git a/crates/core/tests/unit/context_tests.rs b/crates/core/tests/unit/context_tests.rs index 49e57b2cb..48853939d 100644 --- a/crates/core/tests/unit/context_tests.rs +++ b/crates/core/tests/unit/context_tests.rs @@ -46,7 +46,7 @@ fn scope_stack_tracks_scope_local_registries_and_subscribers() { priority: 10, payload: RequestIntercept { break_chain: false, - callable: Arc::new(|_, value| Ok(value)), + callable: Arc::new(|_, value| Box::pin(async move { Ok(value) })), }, }) .unwrap(); @@ -221,15 +221,17 @@ fn merge_helpers_preserve_global_and_scope_local_priority_order() { assert_eq!(merged_exec, vec![("local", 1), ("global", 15)]); } -#[test] -fn conditional_guardrail_snapshots_keep_names_and_callbacks_after_deregister() { +#[tokio::test] +async fn conditional_guardrail_snapshots_keep_names_and_callbacks_after_deregister() { let mut state = NemoRelayContextState::new(); state .tool_conditional_execution_guardrails .register(Guardrail { name: "snapshot_guardrail".to_string(), priority: 1, - payload: Arc::new(|name, _args| Ok(Some(format!("{name} blocked")))), + payload: Arc::new(|name, _args| { + Box::pin(async move { Ok(Some(format!("{name} blocked"))) }) + }), }) .unwrap(); @@ -257,6 +259,7 @@ fn conditional_guardrail_snapshots_keep_names_and_callbacks_after_deregister() { None, None, ) + .await .unwrap(); assert_eq!(rejection.as_deref(), Some("snapshot_target blocked")); @@ -320,12 +323,14 @@ fn context_state_supports_extensions_events_and_builders() { content: json!({"messages": []}), }; let entries = state.llm_sanitize_request_entries(&[]); - let sanitized = NemoRelayContextState::llm_sanitize_request_snapshot_chain( - request.clone(), - crate::api::runtime::LlmSanitizeRequestContext::default(), - &entries, - ) - .expect("an empty sanitizer chain must retain the request"); + let sanitized = tokio::runtime::Runtime::new() + .unwrap() + .block_on(NemoRelayContextState::llm_sanitize_request_snapshot_chain( + request.clone(), + crate::api::runtime::LlmSanitizeRequestContext::default(), + &entries, + )) + .expect("an empty sanitizer chain must retain the request"); assert!(sanitized.headers.is_empty()); let events = Arc::new(Mutex::new(Vec::::new())); diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index 7754238dc..da3d9da7e 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -461,6 +461,7 @@ async fn callback_helpers_cover_worker_response_edges() { RegistrationSurface::MarkSanitizeGuardrail, &event, ) + .await .expect_err("invalid event sanitizer fields should fail"); assert!( error @@ -474,6 +475,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), LlmSanitizeRequestContext::default(), ) + .await .expect_err("invalid LLM JSON result should fail"); assert!(error.to_string().contains("invalid type")); @@ -484,6 +486,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), None, ) + .await .expect_err("invalid LLM intercept request should fail"); assert!( error @@ -498,6 +501,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), None, ) + .await .expect_err("legacy outcome schema should fail"); assert!( error @@ -512,6 +516,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), None, ) + .await .expect_err("invalid annotated request should fail"); assert!( error @@ -521,6 +526,7 @@ async fn callback_helpers_cover_worker_response_edges() { let error = callback .invoke_llm_request_intercept("llm_intercept_error", "model", valid_llm_request(), None) + .await .expect_err("LLM intercept worker error should surface"); assert!(error.to_string().contains("worker.failed: boom")); @@ -531,6 +537,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), None, ) + .await .expect_err("unexpected LLM intercept result should fail"); assert!( error @@ -586,6 +593,7 @@ async fn llm_worker_sanitizers_forward_codec_context_and_omission() { valid_llm_request(), LlmSanitizeRequestContext::with_identity(identity.clone()), ) + .await .expect("empty worker result must represent request omission") .is_none() ); @@ -596,6 +604,7 @@ async fn llm_worker_sanitizers_forward_codec_context_and_omission() { json!({"secret": "value"}), LlmSanitizeResponseContext::with_identity(identity), ) + .await .expect("empty worker result must represent response omission") .is_none() ); @@ -774,6 +783,7 @@ async fn llm_worker_codec_capabilities_are_active_only_during_sanitizer_invocati request, LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), ) + .await .expect("request sanitizer must succeed") .is_none() ); @@ -793,6 +803,7 @@ async fn llm_worker_codec_capabilities_are_active_only_during_sanitizer_invocati response, LlmSanitizeResponseContext::for_response_codec(Some(codec)), ) + .await .expect_err("worker sanitizer error must surface"); assert!(error.to_string().contains("worker.failed: boom")); @@ -1324,6 +1335,7 @@ async fn install_registrations_covers_registry_error_edges() { } #[tokio::test(flavor = "multi_thread")] +#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. async fn installed_callbacks_apply_surface_specific_fallbacks() { struct RuntimeCleanup { registrations: Option, @@ -1423,9 +1435,9 @@ async fn installed_callbacks_apply_surface_specific_fallbacks() { ] { let entries = NemoRelayContextState::event_sanitize_entries(registry, &[]); let sanitized = - NemoRelayContextState::event_sanitize_snapshot_chain(event.clone(), &entries); - assert_eq!(sanitized.data(), None); - assert_eq!(sanitized.metadata(), None); + NemoRelayContextState::event_sanitize_snapshot_chain(event.clone(), &entries).await; + assert_eq!(sanitized.data(), event.data()); + assert_eq!(sanitized.metadata(), event.metadata()); } let entries = state.tool_sanitize_request_entries(&[]); @@ -1434,7 +1446,8 @@ async fn installed_callbacks_apply_surface_specific_fallbacks() { "tool", tool_request.clone(), &entries, - ), + ) + .await, tool_request ); let entries = state.tool_sanitize_response_entries(&[]); @@ -1443,28 +1456,29 @@ async fn installed_callbacks_apply_surface_specific_fallbacks() { "tool", tool_response.clone(), &entries, - ), + ) + .await, tool_response ); let entries = state.llm_sanitize_request_entries(&[]); - assert!( + assert_eq!( NemoRelayContextState::llm_sanitize_request_snapshot_chain( llm_request.clone(), crate::api::runtime::LlmSanitizeRequestContext::default(), &entries, ) - .is_none(), - "a worker request sanitizer failure must omit the observability payload" + .await, + Some(llm_request), ); let entries = state.llm_sanitize_response_entries(&[]); - assert!( + assert_eq!( NemoRelayContextState::llm_sanitize_response_snapshot_chain( llm_response.clone(), crate::api::runtime::LlmSanitizeResponseContext::default(), &entries, ) - .is_none(), - "a worker response sanitizer failure must omit the observability payload" + .await, + Some(llm_response), ); } crate::api::subscriber::flush_subscribers().expect("subscriber callback should flush"); diff --git a/crates/core/tests/unit/llm_api_tests.rs b/crates/core/tests/unit/llm_api_tests.rs index 65c6e5b17..843fe4bf5 100644 --- a/crates/core/tests/unit/llm_api_tests.rs +++ b/crates/core/tests/unit/llm_api_tests.rs @@ -295,7 +295,7 @@ fn credential_headers_are_removed_before_request_sanitizers_and_event_emission() 1, Arc::new(move |request, _context| { sanitizer_capture.lock().unwrap().push(request.clone()); - Some(request) + Box::pin(async move { Ok(Some(request)) }) }), ) .unwrap(); @@ -431,13 +431,13 @@ fn sanitization_invalidates_manual_annotations_without_a_codec() { register_llm_sanitize_request_guardrail( "manual-annotation-invalidation-request", 1, - Arc::new(|_request, _context| Some(redacted_request())), + Arc::new(|_request, _context| Box::pin(async { Ok(Some(redacted_request())) })), ) .unwrap(); register_llm_sanitize_response_guardrail( "manual-annotation-invalidation-response", 1, - Arc::new(|_response, _context| Some(redacted_response())), + Arc::new(|_response, _context| Box::pin(async { Ok(Some(redacted_response())) })), ) .unwrap(); @@ -500,13 +500,13 @@ fn no_op_sanitizers_keep_manual_annotations() { register_llm_sanitize_request_guardrail( "manual-annotation-noop-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); register_llm_sanitize_response_guardrail( "manual-annotation-noop-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); @@ -558,13 +558,13 @@ fn sanitization_regenerates_annotations_with_active_codecs() { register_llm_sanitize_request_guardrail( "active-codec-annotation-regeneration-request", 1, - Arc::new(|_request, _context| Some(redacted_request())), + Arc::new(|_request, _context| Box::pin(async { Ok(Some(redacted_request())) })), ) .unwrap(); register_llm_sanitize_response_guardrail( "active-codec-annotation-regeneration-response", 1, - Arc::new(|_response, _context| Some(redacted_response())), + Arc::new(|_response, _context| Box::pin(async { Ok(Some(redacted_response())) })), ) .unwrap(); @@ -647,7 +647,7 @@ fn buffered_null_fallback_is_sanitized_before_emission() { 1, Arc::new(move |response, _context| { sanitizer_inputs.lock().unwrap().push(response); - Some(Json::Null) + Box::pin(async { Ok(Some(Json::Null)) }) }), ) .unwrap(); @@ -676,7 +676,7 @@ fn buffered_null_fallback_is_sanitized_before_emission() { register_llm_sanitize_response_guardrail( "buffered-null-fallback-redacted", 1, - Arc::new(|_response, _context| Some(redacted_response())), + Arc::new(|_response, _context| Box::pin(async { Ok(Some(redacted_response())) })), ) .unwrap(); let handle = create_llm_handle( @@ -760,7 +760,7 @@ fn streaming_null_fallback_is_sanitized_before_emission() { 1, Arc::new(move |response, _context| { sanitizer_inputs.lock().unwrap().push(response); - Some(Json::Null) + Box::pin(async { Ok(Some(Json::Null)) }) }), ) .unwrap(); @@ -791,7 +791,7 @@ fn streaming_null_fallback_is_sanitized_before_emission() { register_llm_sanitize_response_guardrail( "streaming-null-fallback-redacted", 1, - Arc::new(|_response, _context| Some(redacted_response())), + Arc::new(|_response, _context| Box::pin(async { Ok(Some(redacted_response())) })), ) .unwrap(); runtime.block_on(async { @@ -1525,11 +1525,12 @@ fn failed_managed_calls_sanitize_fallback_end_data() { "failed-managed-call-sanitization", 1, Arc::new(move |response, context| { - sanitizer_inputs - .lock() - .unwrap() - .push((response, context.codec().clone())); - Some(redacted_response()) + let codec = context.codec().clone(); + let sanitizer_inputs = Arc::clone(&sanitizer_inputs); + Box::pin(async move { + sanitizer_inputs.lock().unwrap().push((response, codec)); + Ok(Some(redacted_response())) + }) }), ) .unwrap(); diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index beec626af..a29955c17 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -257,10 +257,135 @@ fn native_string_and_json_helpers_cover_abi_boundaries() { assert_eq!(host_api.abi_version, NEMO_RELAY_NATIVE_ABI_VERSION); assert_eq!( host_api.struct_size, - std::mem::size_of::() + std::mem::size_of::() ); } +#[test] +fn native_async_next_abi_runs_tool_llm_and_stream_continuations() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let cases: Vec<(NativeAsyncNextInner, Json, Json)> = vec![ + ( + NativeAsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + json!({"tool": true}), + json!({"result": {"tool": true}, "pending_marks": []}), + ), + ( + NativeAsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + serde_json::to_value(LlmRequest { + headers: Map::new(), + content: json!({"llm": true}), + }) + .unwrap(), + json!({"llm": true}), + ), + ( + NativeAsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![ + Ok(json!({"chunk": 1})), + Ok(json!({"chunk": 2})), + ]))) + }) + })), + serde_json::to_value(LlmRequest { + headers: Map::new(), + content: json!({"stream": true}), + }) + .unwrap(), + json!([{"chunk": 1}, {"chunk": 2}]), + ), + ]; + + for (inner, invocation, expected) in cases { + let next = Arc::new(NativeAsyncNext { + inner, + runtime: runtime.handle().clone(), + _callback_user_data: None, + }); + let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invocation = native_string_from_json(&invocation).unwrap(); + assert_eq!( + unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, + NemoRelayStatus::Ok + ); + assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); + unsafe { + native_string_free(invocation); + native_async_next_release(next_ref); + native_async_completion_release(completion_ref); + } + } +} + +#[test] +fn native_async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlement() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invalid = native_string("not-json"); + assert_eq!( + unsafe { native_async_completion_resolve_json(completion_ref, invalid) }, + NemoRelayStatus::InvalidJson + ); + let value = native_string(r#"{"ok":true}"#); + assert_eq!( + unsafe { native_async_completion_resolve_json(completion_ref, value) }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { native_async_completion_resolve_json(completion_ref, value) }, + NemoRelayStatus::InvalidArg + ); + assert_eq!( + runtime.block_on(receiver).unwrap().unwrap(), + json!({"ok": true}) + ); + unsafe { + native_string_free(invalid); + native_string_free(value); + native_async_completion_release(completion_ref); + } + + let (sender, _receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(true), + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + assert!(unsafe { native_async_completion_is_cancelled(completion_ref) }); + assert!(unsafe { native_async_completion_is_cancelled(ptr::null()) }); + assert_eq!( + unsafe { native_async_completion_reject(completion_ref, ptr::null()) }, + NemoRelayStatus::InvalidArg + ); + unsafe { native_async_completion_release(completion_ref) }; +} + #[test] fn native_timestamp_scope_type_and_error_mappings_cover_variants() { assert_eq!(optional_timestamp_from_native(ptr::null()).unwrap(), None); diff --git a/crates/core/tests/unit/plugin_tests.rs b/crates/core/tests/unit/plugin_tests.rs index a0c7503cf..b16b3fb11 100644 --- a/crates/core/tests/unit/plugin_tests.rs +++ b/crates/core/tests/unit/plugin_tests.rs @@ -115,8 +115,10 @@ impl Plugin for TestPlugin { 1, false, Arc::new(|_name, mut request, annotated| { - request.headers.insert("x-plugin".into(), json!(true)); - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + Box::pin(async move { + request.headers.insert("x-plugin".into(), json!(true)); + Ok(LlmRequestInterceptOutcome::new(request, annotated)) + }) }), ) }) @@ -754,27 +756,29 @@ fn test_plugin_registration_context_registers_and_rolls_back() { .block_on(TestPlugin.register(&Map::new(), &mut ctx)) .unwrap(); - let request = llm_request_intercepts( - "model", - LlmRequest { - headers: Map::new(), - content: json!({"messages": []}), - }, - ) - .unwrap(); + let request = runtime + .block_on(llm_request_intercepts( + "model", + LlmRequest { + headers: Map::new(), + content: json!({"messages": []}), + }, + )) + .unwrap(); assert_eq!(request.request.headers.get("x-plugin"), Some(&json!(true))); let mut registrations = ctx.into_registrations(); rollback_registrations(&mut registrations); - let request = llm_request_intercepts( - "model", - LlmRequest { - headers: Map::new(), - content: json!({"messages": []}), - }, - ) - .unwrap(); + let request = runtime + .block_on(llm_request_intercepts( + "model", + LlmRequest { + headers: Map::new(), + content: json!({"messages": []}), + }, + )) + .unwrap(); assert_eq!(request.request.headers.get("x-plugin"), None); reset_global(); } @@ -798,25 +802,27 @@ fn test_initialize_plugins_registers_and_clears_components() { assert!(!report.has_errors()); assert!(active_plugin_report().is_some()); - let request = llm_request_intercepts( - "model", - LlmRequest { - headers: Map::new(), - content: json!({"messages": []}), - }, - ) - .unwrap(); + let request = runtime + .block_on(llm_request_intercepts( + "model", + LlmRequest { + headers: Map::new(), + content: json!({"messages": []}), + }, + )) + .unwrap(); assert_eq!(request.request.headers.get("x-plugin"), Some(&json!(true))); clear_plugin_configuration().unwrap(); - let request = llm_request_intercepts( - "model", - LlmRequest { - headers: Map::new(), - content: json!({"messages": []}), - }, - ) - .unwrap(); + let request = runtime + .block_on(llm_request_intercepts( + "model", + LlmRequest { + headers: Map::new(), + content: json!({"messages": []}), + }, + )) + .unwrap(); assert_eq!(request.request.headers.get("x-plugin"), None); reset_global(); } @@ -1070,8 +1076,13 @@ fn test_plugin_registration_context_covers_all_registration_helpers() { let mut ctx = PluginRegistrationContext::with_namespace("demo::"); ctx.register_subscriber("subscriber", Arc::new(|_event| {})) .unwrap(); - ctx.register_tool_request_intercept("tool-request", 1, false, Arc::new(|_name, args| Ok(args))) - .unwrap(); + ctx.register_tool_request_intercept( + "tool-request", + 1, + false, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ) + .unwrap(); ctx.register_tool_execution_intercept( "tool-exec", 1, @@ -1083,7 +1094,7 @@ fn test_plugin_registration_context_covers_all_registration_helpers() { 1, false, Arc::new(|_name, request, annotated| { - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + Box::pin(async move { Ok(LlmRequestInterceptOutcome::new(request, annotated)) }) }), ) .unwrap(); @@ -1586,69 +1597,78 @@ fn test_plugin_registration_context_supports_guardrail_helpers() { reset_global(); let mut ctx = PluginRegistrationContext::with_namespace("plugin::"); - ctx.register_mark_sanitize_guardrail("mark_sanitize", 1, Arc::new(|_, fields| fields)) - .unwrap(); + ctx.register_mark_sanitize_guardrail( + "mark_sanitize", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); ctx.register_scope_sanitize_start_guardrail( "scope_sanitize_start", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); ctx.register_scope_sanitize_end_guardrail( "scope_sanitize_end", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); ctx.register_tool_sanitize_request_guardrail( "tool_sanitize_request", 1, - Arc::new(|_, args| args), + Arc::new(|_, args| Box::pin(async move { Ok(args) })), ) .unwrap(); ctx.register_tool_sanitize_response_guardrail( "tool_sanitize_response", 1, - Arc::new(|_, response| response), + Arc::new(|_, response| Box::pin(async move { Ok(response) })), ) .unwrap(); ctx.register_tool_conditional_execution_guardrail( "tool_conditional", 1, - Arc::new(|name, _args| Ok((name == "blocked-tool").then(|| "blocked tool".to_string()))), + Arc::new(|name, _args| { + Box::pin( + async move { Ok((name == "blocked-tool").then(|| "blocked tool".to_string())) }, + ) + }), ) .unwrap(); ctx.register_llm_sanitize_request_guardrail( "llm_sanitize_request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); ctx.register_llm_sanitize_response_guardrail( "llm_sanitize_response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); ctx.register_llm_conditional_execution_guardrail( "llm_conditional", 1, Arc::new(|request| { - Ok((request.headers.get("blocked") == Some(&json!(true))) - .then(|| "blocked llm".to_string())) + let blocked = request.headers.get("blocked") == Some(&json!(true)); + Box::pin(async move { Ok(blocked.then(|| "blocked llm".to_string())) }) }), ) .unwrap(); - match tool_conditional_execution("blocked-tool", &json!({})) { + let runtime = tokio::runtime::Runtime::new().unwrap(); + match runtime.block_on(tool_conditional_execution("blocked-tool", &json!({}))) { Err(FlowError::GuardrailRejected(message)) => assert_eq!(message, "blocked tool"), other => panic!("expected tool guardrail rejection, got {other:?}"), } - match llm_conditional_execution(&LlmRequest { + match runtime.block_on(llm_conditional_execution(&LlmRequest { headers: Map::from_iter([(String::from("blocked"), json!(true))]), content: json!({"messages": []}), - }) { + })) { Err(FlowError::GuardrailRejected(message)) => assert_eq!(message, "blocked llm"), other => panic!("expected llm guardrail rejection, got {other:?}"), } @@ -1656,13 +1676,18 @@ fn test_plugin_registration_context_supports_guardrail_helpers() { let mut registrations = ctx.into_registrations(); rollback_registrations(&mut registrations); - assert!(tool_conditional_execution("blocked-tool", &json!({})).is_ok()); assert!( - llm_conditional_execution(&LlmRequest { - headers: Map::from_iter([(String::from("blocked"), json!(true))]), - content: json!({"messages": []}), - }) - .is_ok() + runtime + .block_on(tool_conditional_execution("blocked-tool", &json!({}))) + .is_ok() + ); + assert!( + runtime + .block_on(llm_conditional_execution(&LlmRequest { + headers: Map::from_iter([(String::from("blocked"), json!(true))]), + content: json!({"messages": []}), + })) + .is_ok() ); reset_global(); @@ -1674,22 +1699,46 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { reset_global(); let mut ctx = PluginRegistrationContext::with_namespace("duplicate::"); - ctx.register_mark_sanitize_guardrail("mark", 1, Arc::new(|_, fields| fields)) - .unwrap(); + ctx.register_mark_sanitize_guardrail( + "mark", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); expect_registration_failed( - ctx.register_mark_sanitize_guardrail("mark", 1, Arc::new(|_, fields| fields)), + ctx.register_mark_sanitize_guardrail( + "mark", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ), "mark sanitizer:", ); - ctx.register_scope_sanitize_start_guardrail("scope-start", 1, Arc::new(|_, fields| fields)) - .unwrap(); + ctx.register_scope_sanitize_start_guardrail( + "scope-start", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); expect_registration_failed( - ctx.register_scope_sanitize_start_guardrail("scope-start", 1, Arc::new(|_, fields| fields)), + ctx.register_scope_sanitize_start_guardrail( + "scope-start", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ), "scope-start sanitizer:", ); - ctx.register_scope_sanitize_end_guardrail("scope-end", 1, Arc::new(|_, fields| fields)) - .unwrap(); + ctx.register_scope_sanitize_end_guardrail( + "scope-end", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); expect_registration_failed( - ctx.register_scope_sanitize_end_guardrail("scope-end", 1, Arc::new(|_, fields| fields)), + ctx.register_scope_sanitize_end_guardrail( + "scope-end", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ), "scope-end sanitizer:", ); ctx.register_llm_request_intercept( @@ -1697,7 +1746,7 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { 1, false, Arc::new(|_name, request, annotated| { - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + Box::pin(async move { Ok(LlmRequestInterceptOutcome::new(request, annotated)) }) }), ) .unwrap(); @@ -1707,7 +1756,7 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { 1, false, Arc::new(|_name, request, annotated| { - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + Box::pin(async move { Ok(LlmRequestInterceptOutcome::new(request, annotated)) }) }), ), "llm request intercept:", @@ -1716,14 +1765,14 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { ctx.register_tool_sanitize_request_guardrail( "tool-sanitize-request", 1, - Arc::new(|_, args| args), + Arc::new(|_, args| Box::pin(async move { Ok(args) })), ) .unwrap(); expect_registration_failed( ctx.register_tool_sanitize_request_guardrail( "tool-sanitize-request", 1, - Arc::new(|_, args| args), + Arc::new(|_, args| Box::pin(async move { Ok(args) })), ), "tool sanitize request guardrail:", ); @@ -1731,14 +1780,14 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { ctx.register_tool_sanitize_response_guardrail( "tool-sanitize-response", 1, - Arc::new(|_, response| response), + Arc::new(|_, response| Box::pin(async move { Ok(response) })), ) .unwrap(); expect_registration_failed( ctx.register_tool_sanitize_response_guardrail( "tool-sanitize-response", 1, - Arc::new(|_, response| response), + Arc::new(|_, response| Box::pin(async move { Ok(response) })), ), "tool sanitize response guardrail:", ); @@ -1746,14 +1795,14 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { ctx.register_tool_conditional_execution_guardrail( "tool-conditional", 1, - Arc::new(|_, _| Ok(None)), + Arc::new(|_, _| Box::pin(async { Ok(None) })), ) .unwrap(); expect_registration_failed( ctx.register_tool_conditional_execution_guardrail( "tool-conditional", 1, - Arc::new(|_, _| Ok(None)), + Arc::new(|_, _| Box::pin(async { Ok(None) })), ), "tool conditional execution guardrail:", ); @@ -1761,14 +1810,14 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { ctx.register_llm_sanitize_request_guardrail( "llm-sanitize-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); expect_registration_failed( ctx.register_llm_sanitize_request_guardrail( "llm-sanitize-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ), "llm sanitize request guardrail:", ); @@ -1776,25 +1825,29 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { ctx.register_llm_sanitize_response_guardrail( "llm-sanitize-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); expect_registration_failed( ctx.register_llm_sanitize_response_guardrail( "llm-sanitize-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ), "llm sanitize response guardrail:", ); - ctx.register_llm_conditional_execution_guardrail("llm-conditional", 1, Arc::new(|_| Ok(None))) - .unwrap(); + ctx.register_llm_conditional_execution_guardrail( + "llm-conditional", + 1, + Arc::new(|_| Box::pin(async { Ok(None) })), + ) + .unwrap(); expect_registration_failed( ctx.register_llm_conditional_execution_guardrail( "llm-conditional", 1, - Arc::new(|_| Ok(None)), + Arc::new(|_| Box::pin(async { Ok(None) })), ), "llm conditional execution guardrail:", ); @@ -1841,14 +1894,19 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { "llm stream execution intercept:", ); - ctx.register_tool_request_intercept("tool-request", 1, false, Arc::new(|_name, args| Ok(args))) - .unwrap(); + ctx.register_tool_request_intercept( + "tool-request", + 1, + false, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ) + .unwrap(); expect_registration_failed( ctx.register_tool_request_intercept( "tool-request", 1, false, - Arc::new(|_name, args| Ok(args)), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ), "tool request intercept:", ); @@ -1879,18 +1937,22 @@ fn test_plugin_registration_context_maps_deregistration_errors() { reset_global(); let mut ctx = PluginRegistrationContext::with_namespace("teardown::"); - ctx.register_mark_sanitize_guardrail("mark-sanitize", 1, Arc::new(|_, fields| fields)) - .unwrap(); + ctx.register_mark_sanitize_guardrail( + "mark-sanitize", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); ctx.register_scope_sanitize_start_guardrail( "scope-sanitize-start", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); ctx.register_scope_sanitize_end_guardrail( "scope-sanitize-end", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); ctx.register_subscriber("subscriber", Arc::new(|_event| {})) @@ -1900,42 +1962,46 @@ fn test_plugin_registration_context_maps_deregistration_errors() { 1, false, Arc::new(|_name, request, annotated| { - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + Box::pin(async move { Ok(LlmRequestInterceptOutcome::new(request, annotated)) }) }), ) .unwrap(); ctx.register_tool_sanitize_request_guardrail( "tool-sanitize-request", 1, - Arc::new(|_, args| args), + Arc::new(|_, args| Box::pin(async move { Ok(args) })), ) .unwrap(); ctx.register_tool_sanitize_response_guardrail( "tool-sanitize-response", 1, - Arc::new(|_, response| response), + Arc::new(|_, response| Box::pin(async move { Ok(response) })), ) .unwrap(); ctx.register_tool_conditional_execution_guardrail( "tool-conditional", 1, - Arc::new(|_, _| Ok(None)), + Arc::new(|_, _| Box::pin(async { Ok(None) })), ) .unwrap(); ctx.register_llm_sanitize_request_guardrail( "llm-sanitize-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); ctx.register_llm_sanitize_response_guardrail( "llm-sanitize-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), + ) + .unwrap(); + ctx.register_llm_conditional_execution_guardrail( + "llm-conditional", + 1, + Arc::new(|_| Box::pin(async { Ok(None) })), ) .unwrap(); - ctx.register_llm_conditional_execution_guardrail("llm-conditional", 1, Arc::new(|_| Ok(None))) - .unwrap(); ctx.register_llm_execution_intercept( "llm-exec", 1, @@ -1954,8 +2020,13 @@ fn test_plugin_registration_context_maps_deregistration_errors() { }), ) .unwrap(); - ctx.register_tool_request_intercept("tool-request", 1, false, Arc::new(|_name, args| Ok(args))) - .unwrap(); + ctx.register_tool_request_intercept( + "tool-request", + 1, + false, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ) + .unwrap(); ctx.register_tool_execution_intercept( "tool-exec", 1, diff --git a/crates/core/tests/unit/shared_tests.rs b/crates/core/tests/unit/shared_tests.rs index 9d17ce188..1de4a9ed7 100644 --- a/crates/core/tests/unit/shared_tests.rs +++ b/crates/core/tests/unit/shared_tests.rs @@ -168,8 +168,9 @@ fn stale_process_runtime_owner_is_reclaimed() { reset_global(); } -#[test] -fn test_run_request_intercepts_with_codec_none_and_codec_paths() { +#[tokio::test] +#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. +async fn test_run_request_intercepts_with_codec_none_and_codec_paths() { let _guard = lock_runtime_owner(); reset_global(); @@ -178,11 +179,13 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { 1, false, Arc::new(|_name, mut request, annotated| { - assert!(annotated.is_none()); - request.headers.insert("x-no-codec".into(), json!(true)); - let mut annotated = SharedTestCodec.decode(&request)?; - annotated.model = Some("interceptor-model".into()); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + Box::pin(async move { + assert!(annotated.is_none()); + request.headers.insert("x-no-codec".into(), json!(true)); + let mut annotated = SharedTestCodec.decode(&request)?; + annotated.model = Some("interceptor-model".into()); + Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + }) }), ) .unwrap(); @@ -196,6 +199,7 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { }, None, ) + .await .unwrap(); assert_eq!( request_without_codec.headers.get("x-no-codec"), @@ -215,10 +219,12 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { 1, false, Arc::new(|_name, mut request, annotated| { - let mut annotated = annotated.expect("codec should provide annotated request"); - annotated.model = Some("intercepted-model".into()); - request.headers.insert("x-codec".into(), json!(true)); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + Box::pin(async move { + let mut annotated = annotated.expect("codec should provide annotated request"); + annotated.model = Some("intercepted-model".into()); + request.headers.insert("x-codec".into(), json!(true)); + Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + }) }), ) .unwrap(); @@ -233,6 +239,7 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { }, Some(codec), ) + .await .unwrap(); assert_eq!( @@ -259,8 +266,9 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { reset_global(); } -#[test] -fn managed_request_chain_records_contributions_incrementally_while_standalone_retains_them() { +#[tokio::test] +#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. +async fn managed_request_chain_records_contributions_incrementally_while_standalone_retains_them() { let _guard = lock_runtime_owner(); reset_global(); @@ -269,11 +277,12 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re 1, false, Arc::new(|_name, request, annotated| { - Ok( - LlmRequestInterceptOutcome::new(request, annotated).with_optimization_contribution( - LlmOptimizationContribution::new("accepted", "custom"), - ), - ) + Box::pin(async move { + Ok(LlmRequestInterceptOutcome::new(request, annotated) + .with_optimization_contribution(LlmOptimizationContribution::new( + "accepted", "custom", + ))) + }) }), ) .unwrap(); @@ -282,14 +291,13 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re 2, false, Arc::new(|_name, request, annotated| { - Ok( - LlmRequestInterceptOutcome::new(request, annotated).with_optimization_contribution( - LlmOptimizationContribution::new( + Box::pin(async move { + Ok(LlmRequestInterceptOutcome::new(request, annotated) + .with_optimization_contribution(LlmOptimizationContribution::new( "x".repeat(MAX_LLM_OPTIMIZATION_CONTRIBUTION_BYTES), "custom", - ), - ), - ) + ))) + }) }), ) .unwrap(); @@ -302,6 +310,7 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re }, None, ) + .await .unwrap(); assert_eq!(standalone.3.len(), 2); assert!(standalone.3.iter().all(|item| item.sequence.is_none())); @@ -316,6 +325,7 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re None, &recorder, ) + .await .unwrap(); assert!(managed.3.is_empty()); let recorded = recorder.unemitted(); @@ -326,8 +336,9 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re reset_global(); } -#[test] -fn test_run_request_intercepts_injects_dynamo_agent_lineage() { +#[tokio::test] +#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. +async fn test_run_request_intercepts_injects_dynamo_agent_lineage() { let _guard = lock_runtime_owner(); reset_global(); @@ -369,6 +380,7 @@ fn test_run_request_intercepts_injects_dynamo_agent_lineage() { }, None, ) + .await .unwrap(); assert_eq!( request.headers.get(DYNAMO_SESSION_ID_HEADER_KEY), @@ -396,6 +408,7 @@ fn test_run_request_intercepts_injects_dynamo_agent_lineage() { }, Some(Arc::new(SharedTestCodec)), ) + .await .unwrap(); assert_eq!( request_with_codec.headers.get(DYNAMO_SESSION_ID_HEADER_KEY), @@ -433,6 +446,7 @@ fn test_run_request_intercepts_injects_dynamo_agent_lineage() { }, None, ) + .await .unwrap(); assert_eq!( request.headers.get(DYNAMO_SESSION_ID_HEADER_KEY), @@ -475,6 +489,7 @@ fn test_run_request_intercepts_injects_dynamo_agent_lineage() { }, None, ) + .await .unwrap(); assert!(!request.headers.contains_key(DYNAMO_SESSION_ID_HEADER_KEY)); assert!( diff --git a/crates/ffi/src/api/mod.rs b/crates/ffi/src/api/mod.rs index 03103f58a..f2984f32e 100644 --- a/crates/ffi/src/api/mod.rs +++ b/crates/ffi/src/api/mod.rs @@ -99,6 +99,15 @@ fn tokio_runtime() -> &'static Runtime { }) } +fn block_on_sync_ffi(future: impl Future>) -> FlowResult { + if tokio::runtime::Handle::try_current().is_ok() { + return Err(nemo_relay::error::FlowError::Internal( + "synchronous FFI middleware helpers cannot run on a Tokio runtime thread; use the completion-based async registration API".into(), + )); + } + tokio_runtime().block_on(future) +} + // --------------------------------------------------------------------------- // Standalone middleware chains // --------------------------------------------------------------------------- @@ -141,7 +150,7 @@ pub unsafe extern "C" fn nemo_relay_tool_request_intercepts( Some(a) => a, None => return NemoRelayStatus::InvalidJson, }; - match core_tool_api::tool_request_intercepts(&name, args) { + match block_on_sync_ffi(core_tool_api::tool_request_intercepts(&name, args)) { Ok(result) => { unsafe { *out = json_to_c_string(&result) }; NemoRelayStatus::Ok @@ -179,7 +188,7 @@ pub unsafe extern "C" fn nemo_relay_tool_conditional_execution( Some(a) => a, None => return NemoRelayStatus::InvalidJson, }; - match core_tool_api::tool_conditional_execution(&name, &args) { + match block_on_sync_ffi(core_tool_api::tool_conditional_execution(&name, &args)) { Ok(()) => NemoRelayStatus::Ok, Err(e) => status_from_error(&e), } @@ -233,7 +242,7 @@ pub unsafe extern "C" fn nemo_relay_llm_request_intercepts( return NemoRelayStatus::InvalidJson; } }; - match core_llm_api::llm_request_intercepts(name_str, request) { + match block_on_sync_ffi(core_llm_api::llm_request_intercepts(name_str, request)) { Ok(transformed) => { let result_json = serde_json::to_value(&transformed).unwrap_or(serde_json::Value::Null); unsafe { *out = json_to_c_string(&result_json) }; @@ -396,7 +405,7 @@ pub unsafe extern "C" fn nemo_relay_llm_conditional_execution( return NemoRelayStatus::InvalidJson; } }; - match core_llm_api::llm_conditional_execution(&request) { + match block_on_sync_ffi(core_llm_api::llm_conditional_execution(&request)) { Ok(()) => NemoRelayStatus::Ok, Err(e) => status_from_error(&e), } diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index f327e01d2..f1c952863 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -321,14 +321,17 @@ pub fn wrap_tool_sanitize_fn( free_fn: NemoRelayFreeFn, ) -> ToolSanitizeFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |name: &str, args: Json| { - let c_name = CString::new(name).unwrap_or_default(); - let c_args = json_to_c_string(&args); - let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; - unsafe { nemo_relay_string_free_internal(c_args) }; - let result = ptr_to_json(result_ptr); - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |name: String, args: Json| { + let ud = ud.clone(); + Box::pin(async move { + let c_name = CString::new(name).unwrap_or_default(); + let c_args = json_to_c_string(&args); + let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; + unsafe { nemo_relay_string_free_internal(c_args) }; + let result = ptr_to_json(result_ptr); + unsafe { nemo_relay_string_free_internal(result_ptr) }; + Ok(result) + }) }) } @@ -339,22 +342,25 @@ pub fn wrap_tool_conditional_fn( free_fn: NemoRelayFreeFn, ) -> ToolConditionalFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |name: &str, args: &Json| { - clear_last_error(); - let c_name = CString::new(name).unwrap_or_default(); - let c_args = json_to_c_string(args); - let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; - unsafe { nemo_relay_string_free_internal(c_args) }; - let result = if result_ptr.is_null() { - match last_error_message() { - Some(message) => Err(FlowError::Internal(message)), - None => Ok(None), - } - } else { - Ok(ptr_to_opt_string(result_ptr)) - }; - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |name: String, args: Json| { + let ud = ud.clone(); + Box::pin(async move { + clear_last_error(); + let c_name = CString::new(name).unwrap_or_default(); + let c_args = json_to_c_string(&args); + let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; + unsafe { nemo_relay_string_free_internal(c_args) }; + let result = if result_ptr.is_null() { + match last_error_message() { + Some(message) => Err(FlowError::Internal(message)), + None => Ok(None), + } + } else { + Ok(ptr_to_opt_string(result_ptr)) + }; + unsafe { nemo_relay_string_free_internal(result_ptr) }; + result + }) }) } @@ -365,16 +371,19 @@ pub fn wrap_tool_request_intercept_fn( free_fn: NemoRelayFreeFn, ) -> ToolInterceptFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |name: &str, args: Json| { - clear_last_error(); - let c_name = CString::new(name).unwrap_or_default(); - let c_args = json_to_c_string(&args); - let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; - unsafe { nemo_relay_string_free_internal(c_args) }; - let result = - json_result_from_ptr(result_ptr, "tool request intercept callback returned null"); - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |name: String, args: Json| { + let ud = ud.clone(); + Box::pin(async move { + clear_last_error(); + let c_name = CString::new(name).unwrap_or_default(); + let c_args = json_to_c_string(&args); + let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; + unsafe { nemo_relay_string_free_internal(c_args) }; + let result = + json_result_from_ptr(result_ptr, "tool request intercept callback returned null"); + unsafe { nemo_relay_string_free_internal(result_ptr) }; + result + }) }) } @@ -617,64 +626,67 @@ pub fn wrap_llm_request_intercept_fn( ) -> LlmRequestInterceptFn { let ud = make_user_data(user_data, free_fn); Arc::new( - move |name: &str, request: LlmRequest, annotated: Option| { - clear_last_error(); - let c_name = CString::new(name).unwrap_or_default(); - let ffi_req = Box::into_raw(Box::new(FfiLLMRequest(request))); + move |name: String, request: LlmRequest, annotated: Option| { + let ud = ud.clone(); + Box::pin(async move { + clear_last_error(); + let c_name = CString::new(name).unwrap_or_default(); + let ffi_req = Box::into_raw(Box::new(FfiLLMRequest(request))); + + // Serialize annotated to JSON C string if present, else null + let c_annotated = match &annotated { + Some(a) => { + let s = serde_json::to_string(a).unwrap_or_else(|_| "null".to_string()); + CString::new(s).unwrap_or_default() + } + None => CString::default(), + }; + let annotated_ptr = if annotated.is_some() { + c_annotated.as_ptr() + } else { + std::ptr::null() + }; - // Serialize annotated to JSON C string if present, else null - let c_annotated = match &annotated { - Some(a) => { - let s = serde_json::to_string(a).unwrap_or_else(|_| "null".to_string()); - CString::new(s).unwrap_or_default() - } - None => CString::default(), - }; - let annotated_ptr = if annotated.is_some() { - c_annotated.as_ptr() - } else { - std::ptr::null() - }; + let mut out_outcome: *mut c_char = std::ptr::null_mut(); - let mut out_outcome: *mut c_char = std::ptr::null_mut(); + let status = unsafe { + cb( + ud.ptr, + c_name.as_ptr(), + ffi_req, + annotated_ptr, + &mut out_outcome, + ) + }; - let status = unsafe { - cb( - ud.ptr, - c_name.as_ptr(), - ffi_req, - annotated_ptr, - &mut out_outcome, - ) - }; + // Free the input request + unsafe { drop(Box::from_raw(ffi_req)) }; - // Free the input request - unsafe { drop(Box::from_raw(ffi_req)) }; + if status != NemoRelayStatus::Ok { + unsafe { nemo_relay_string_free_internal(out_outcome) }; + let message = last_error_message() + .unwrap_or_else(|| "request intercept callback failed".to_string()); + return Err(FlowError::Internal(message)); + } - if status != NemoRelayStatus::Ok { + if out_outcome.is_null() { + return Err(FlowError::Internal( + "request intercept returned null out_outcome_json".to_string(), + )); + } + let outcome = unsafe { CStr::from_ptr(out_outcome) } + .to_str() + .map_err(|error| FlowError::Internal(format!("invalid outcome UTF-8: {error}"))) + .and_then(|json| { + serde_json::from_str::(json).map_err(|error| { + FlowError::Internal(format!( + "invalid LLM request intercept outcome JSON: {error}" + )) + }) + }); unsafe { nemo_relay_string_free_internal(out_outcome) }; - let message = last_error_message() - .unwrap_or_else(|| "request intercept callback failed".to_string()); - return Err(FlowError::Internal(message)); - } - - if out_outcome.is_null() { - return Err(FlowError::Internal( - "request intercept returned null out_outcome_json".to_string(), - )); - } - let outcome = unsafe { CStr::from_ptr(out_outcome) } - .to_str() - .map_err(|error| FlowError::Internal(format!("invalid outcome UTF-8: {error}"))) - .and_then(|json| { - serde_json::from_str::(json).map_err(|error| { - FlowError::Internal(format!( - "invalid LLM request intercept outcome JSON: {error}" - )) - }) - }); - unsafe { nemo_relay_string_free_internal(out_outcome) }; - outcome + outcome + }) }, ) } @@ -688,79 +700,85 @@ pub fn wrap_llm_sanitize_request_fn( let ud = make_user_data(user_data, free_fn); Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { + let ud = ud.clone(); + Box::pin(async move { + clear_last_error(); + let (codec_kind, codec_id) = match ffi_codec_identity(context.codec()) { + Ok(identity) => identity, + Err(error) => { + set_last_error(&error.to_string()); + return Ok(None); + } + }; + let codec = context + .resolve_codec() + .map(crate::types::FfiLlmSanitizeRequestCodec); + let ffi_context = NemoRelayLlmSanitizeRequestContext { + codec_kind, + codec_id: codec_id + .as_ref() + .map_or(std::ptr::null(), |name| name.as_ptr()), + codec: codec.as_ref().map_or(std::ptr::null(), std::ptr::from_ref), + }; + let ffi_req = Box::into_raw(Box::new(FfiLLMRequest(request))); + let result_ptr = unsafe { cb(ud.ptr, ffi_req, ffi_context) }; + if result_ptr.is_null() { + unsafe { drop(Box::from_raw(ffi_req)) }; + return Ok(None); + } + if result_ptr == ffi_req { + return Ok(Some(unsafe { Box::from_raw(ffi_req) }.0)); + } + unsafe { drop(Box::from_raw(ffi_req)) }; + Ok(Some(unsafe { Box::from_raw(result_ptr) }.0)) + }) + }, + ) +} + +/// Wrap a C LLM response sanitizer into a Rust closure. +pub fn wrap_llm_sanitize_response_fn( + cb: NemoRelayLlmSanitizeResponseCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmSanitizeResponseFn { + let ud = make_user_data(user_data, free_fn); + Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { + let ud = ud.clone(); + Box::pin(async move { clear_last_error(); let (codec_kind, codec_id) = match ffi_codec_identity(context.codec()) { Ok(identity) => identity, Err(error) => { set_last_error(&error.to_string()); - return None; + return Ok(None); } }; let codec = context .resolve_codec() - .map(crate::types::FfiLlmSanitizeRequestCodec); - let ffi_context = NemoRelayLlmSanitizeRequestContext { + .map(crate::types::FfiLlmSanitizeResponseCodec); + let ffi_context = NemoRelayLlmSanitizeResponseContext { codec_kind, codec_id: codec_id .as_ref() .map_or(std::ptr::null(), |name| name.as_ptr()), codec: codec.as_ref().map_or(std::ptr::null(), std::ptr::from_ref), }; - let ffi_req = Box::into_raw(Box::new(FfiLLMRequest(request))); - let result_ptr = unsafe { cb(ud.ptr, ffi_req, ffi_context) }; + let response_json = json_to_c_string(&response); + let result_ptr = unsafe { cb(ud.ptr, response_json, ffi_context) }; if result_ptr.is_null() { - unsafe { drop(Box::from_raw(ffi_req)) }; - return None; - } - if result_ptr == ffi_req { - return Some(unsafe { Box::from_raw(ffi_req) }.0); + unsafe { nemo_relay_string_free_internal(response_json) }; + return Ok(None); } - unsafe { drop(Box::from_raw(ffi_req)) }; - Some(unsafe { Box::from_raw(result_ptr) }.0) - }, - ) -} - -/// Wrap a C LLM response sanitizer into a Rust closure. -pub fn wrap_llm_sanitize_response_fn( - cb: NemoRelayLlmSanitizeResponseCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, -) -> LlmSanitizeResponseFn { - let ud = make_user_data(user_data, free_fn); - Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { - clear_last_error(); - let (codec_kind, codec_id) = match ffi_codec_identity(context.codec()) { - Ok(identity) => identity, - Err(error) => { - set_last_error(&error.to_string()); - return None; - } - }; - let codec = context - .resolve_codec() - .map(crate::types::FfiLlmSanitizeResponseCodec); - let ffi_context = NemoRelayLlmSanitizeResponseContext { - codec_kind, - codec_id: codec_id - .as_ref() - .map_or(std::ptr::null(), |name| name.as_ptr()), - codec: codec.as_ref().map_or(std::ptr::null(), std::ptr::from_ref), - }; - let response_json = json_to_c_string(&response); - let result_ptr = unsafe { cb(ud.ptr, response_json, ffi_context) }; - if result_ptr.is_null() { - unsafe { nemo_relay_string_free_internal(response_json) }; - return None; - } - let result = c_str_to_json(result_ptr); - unsafe { - nemo_relay_string_free_internal(response_json); - if result_ptr != response_json { - nemo_relay_string_free_internal(result_ptr); + let result = c_str_to_json(result_ptr); + unsafe { + nemo_relay_string_free_internal(response_json); + if result_ptr != response_json { + nemo_relay_string_free_internal(result_ptr); + } } - } - result + Ok(result) + }) }) } @@ -790,20 +808,23 @@ pub fn wrap_llm_conditional_fn( free_fn: NemoRelayFreeFn, ) -> LlmConditionalFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |request: &LlmRequest| { - clear_last_error(); - let ffi_req = FfiLLMRequest(request.clone()); - let result_ptr = unsafe { cb(ud.ptr, &ffi_req) }; - let result = if result_ptr.is_null() { - match last_error_message() { - Some(message) => Err(FlowError::Internal(message)), - None => Ok(None), - } - } else { - Ok(ptr_to_opt_string(result_ptr)) - }; - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |request: LlmRequest| { + let ud = ud.clone(); + Box::pin(async move { + clear_last_error(); + let ffi_req = FfiLLMRequest(request); + let result_ptr = unsafe { cb(ud.ptr, &ffi_req) }; + let result = if result_ptr.is_null() { + match last_error_message() { + Some(message) => Err(FlowError::Internal(message)), + None => Ok(None), + } + } else { + Ok(ptr_to_opt_string(result_ptr)) + }; + unsafe { nemo_relay_string_free_internal(result_ptr) }; + result + }) }) } @@ -918,14 +939,18 @@ pub fn wrap_event_sanitize_fn( free_fn: NemoRelayFreeFn, ) -> EventSanitizeFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |event: &Event, fields: EventSanitizeFields| { - let ffi_event = FfiEvent(event.clone()); - let fields_json = json_to_c_string(&serde_json::to_value(&fields).unwrap_or(Json::Null)); - let result_ptr = unsafe { cb(ud.ptr, &ffi_event, fields_json) }; - unsafe { nemo_relay_string_free_internal(fields_json) }; - let result = serde_json::from_value(ptr_to_json(result_ptr)).unwrap_or_default(); - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |event: Event, fields: EventSanitizeFields| { + let ud = ud.clone(); + Box::pin(async move { + let ffi_event = FfiEvent(event); + let fields_json = + json_to_c_string(&serde_json::to_value(&fields).unwrap_or(Json::Null)); + let result_ptr = unsafe { cb(ud.ptr, &ffi_event, fields_json) }; + unsafe { nemo_relay_string_free_internal(fields_json) }; + let result = serde_json::from_value(ptr_to_json(result_ptr)).unwrap_or_default(); + unsafe { nemo_relay_string_free_internal(result_ptr) }; + Ok(result) + }) }) } diff --git a/crates/ffi/tests/integration/callable_extra_tests.rs b/crates/ffi/tests/integration/callable_extra_tests.rs index b176f4988..69a341b96 100644 --- a/crates/ffi/tests/integration/callable_extra_tests.rs +++ b/crates/ffi/tests/integration/callable_extra_tests.rs @@ -4,10 +4,19 @@ //! Integration tests for callable extra in the NeMo Relay FFI crate. use super::*; +use std::future::Future; use std::ptr; use tokio_stream::StreamExt; +fn resolve(future: impl Future) -> T { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap() + .block_on(future) +} + unsafe extern "C" fn tool_conditional_error_cb( _user_data: *mut libc::c_void, _name: *const c_char, @@ -172,7 +181,7 @@ fn test_callable_extra_trampoline_and_helper_paths() { .unwrap(); let conditional = wrap_tool_conditional_fn(tool_conditional_error_cb, ptr::null_mut(), None); - let conditional_err = conditional("tool", &json!({})).unwrap_err(); + let conditional_err = resolve(conditional("tool".into(), json!({}))).unwrap_err(); assert!( conditional_err .to_string() @@ -243,7 +252,7 @@ fn test_callable_extra_request_intercept_and_codec_paths() { let intercept_error = wrap_llm_request_intercept_fn(llm_request_intercept_status_error_cb, ptr::null_mut(), None); - let err = intercept_error("llm", request.clone(), None).unwrap_err(); + let err = resolve(intercept_error("llm".into(), request.clone(), None)).unwrap_err(); assert!( err.to_string() .contains("request intercept callback failed") @@ -254,7 +263,7 @@ fn test_callable_extra_request_intercept_and_codec_paths() { ptr::null_mut(), None, ); - let err = intercept_null("llm", request.clone(), None).unwrap_err(); + let err = resolve(intercept_null("llm".into(), request.clone(), None)).unwrap_err(); assert!(err.to_string().contains("null out_outcome_json")); let intercept_invalid_annotated = wrap_llm_request_intercept_fn( @@ -262,22 +271,28 @@ fn test_callable_extra_request_intercept_and_codec_paths() { ptr::null_mut(), None, ); - let err = intercept_invalid_annotated("llm", request.clone(), None).unwrap_err(); + let err = resolve(intercept_invalid_annotated( + "llm".into(), + request.clone(), + None, + )) + .unwrap_err(); assert!( err.to_string() .contains("invalid LLM request intercept outcome JSON") ); let sanitize = wrap_llm_sanitize_request_fn(llm_request_passthrough_cb, ptr::null_mut(), None); - let sanitized = sanitize( + let sanitized = resolve(sanitize( request.clone(), nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ) + )) + .unwrap() .expect("non-null sanitizer result"); assert_eq!(sanitized.content, request.content); let conditional = wrap_llm_conditional_fn(llm_conditional_error_cb, ptr::null_mut(), None); - let conditional_err = conditional(&request).unwrap_err(); + let conditional_err = resolve(conditional(request.clone())).unwrap_err(); assert!( conditional_err .to_string() @@ -359,12 +374,16 @@ fn test_sanitizer_context_resolves_directional_ffi_codecs() { "preserve": true }), }; - let sanitized = - wrap_llm_sanitize_request_fn(llm_request_codec_round_trip_cb, ptr::null_mut(), None)( - request.clone(), - LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), - ) - .expect("codec round trip returns a request"); + let sanitized = resolve(wrap_llm_sanitize_request_fn( + llm_request_codec_round_trip_cb, + ptr::null_mut(), + None, + )( + request.clone(), + LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), + )) + .unwrap() + .expect("codec round trip returns a request"); assert_eq!(sanitized.content, request.content); let response = json!({ @@ -376,11 +395,15 @@ fn test_sanitizer_context_resolves_directional_ffi_codecs() { "finish_reason": "stop" }] }); - let sanitized = - wrap_llm_sanitize_response_fn(llm_response_codec_decode_cb, ptr::null_mut(), None)( - response.clone(), - LlmSanitizeResponseContext::for_response_codec(Some(codec)), - ) - .expect("codec decode returns a response"); + let sanitized = resolve(wrap_llm_sanitize_response_fn( + llm_response_codec_decode_cb, + ptr::null_mut(), + None, + )( + response.clone(), + LlmSanitizeResponseContext::for_response_codec(Some(codec)), + )) + .unwrap() + .expect("codec decode returns a response"); assert_eq!(sanitized, response); } diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index c58098b4a..60b260868 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -4,6 +4,7 @@ //! Unit tests for callable in the NeMo Relay FFI crate. use super::*; +use std::future::Future; use std::sync::atomic::{AtomicUsize, Ordering}; use nemo_relay::api::event::{Event, EventSanitizeFields}; @@ -22,6 +23,14 @@ fn user_data_counter() -> (*mut libc::c_void, Arc) { (ptr, counter) } +fn resolve(future: impl Future) -> T { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap() + .block_on(future) +} + unsafe extern "C" fn tool_sanitize_cb( user_data: *mut libc::c_void, name: *const c_char, @@ -328,7 +337,7 @@ fn make_request() -> LlmRequest { fn test_wrap_tool_request_and_conditional_callbacks() { let (user_data, called) = user_data_counter(); let wrapped = wrap_tool_sanitize_fn(tool_sanitize_cb, user_data, Some(free_arc_counter)); - let result = wrapped("tool-name", json!({"value": 1})); + let result = resolve(wrapped("tool-name".into(), json!({"value": 1}))).unwrap(); assert_eq!(result["value"], json!(1)); assert_eq!(result["name"], json!("tool-name")); assert_eq!(called.load(Ordering::SeqCst), 1); @@ -338,11 +347,11 @@ fn test_wrap_tool_request_and_conditional_callbacks() { let wrapped_conditional = wrap_tool_conditional_fn(tool_conditional_cb, std::ptr::null_mut(), None); assert_eq!( - wrapped_conditional("tool", &json!({"block": true})).unwrap(), + resolve(wrapped_conditional("tool".into(), json!({"block": true}))).unwrap(), Some("blocked".into()) ); assert_eq!( - wrapped_conditional("tool", &json!({"block": false})).unwrap(), + resolve(wrapped_conditional("tool".into(), json!({"block": false}))).unwrap(), None ); } @@ -411,56 +420,60 @@ fn test_wrap_tool_exec_and_intercept_callbacks() { fn test_wrap_llm_request_response_and_conditional_callbacks() { let request_intercept = wrap_llm_request_intercept_fn(llm_request_intercept_cb, std::ptr::null_mut(), None); - let outcome = request_intercept("llm", make_request(), None).unwrap(); + let outcome = resolve(request_intercept("llm".into(), make_request(), None)).unwrap(); assert_eq!(outcome.request.content["intercepted"], json!(true)); let sanitize_request = wrap_llm_sanitize_request_fn(llm_request_null_cb, std::ptr::null_mut(), None); assert_eq!( - sanitize_request( + resolve(sanitize_request( make_request(), nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ), + )) + .unwrap(), None ); let alias_request = wrap_llm_sanitize_request_fn(llm_request_alias_cb, std::ptr::null_mut(), None); assert_eq!( - alias_request( + resolve(alias_request( make_request(), nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ), + )) + .unwrap(), Some(make_request()) ); let conditional = wrap_llm_conditional_fn(llm_conditional_cb, std::ptr::null_mut(), None); assert_eq!( - conditional(&LlmRequest { + resolve(conditional(LlmRequest { headers: serde_json::Map::new(), content: json!({"block": true}), - }) + })) .unwrap(), Some("blocked llm".into()) ); - assert_eq!(conditional(&make_request()).unwrap(), None); + assert_eq!(resolve(conditional(make_request())).unwrap(), None); let wrapped_response = wrap_llm_sanitize_response_fn(json_cb, std::ptr::null_mut(), None); assert_eq!( - wrapped_response( + resolve(wrapped_response( json!({"value": 2}), nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ) + )) + .unwrap() .unwrap()["wrapped"], json!(true) ); let alias_response = wrap_llm_sanitize_response_fn(json_alias_cb, std::ptr::null_mut(), None); assert_eq!( - alias_response( + resolve(alias_response( json!({"value": 2}), nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ), + )) + .unwrap(), Some(json!({"value": 2})) ); @@ -468,10 +481,11 @@ fn test_wrap_llm_request_response_and_conditional_callbacks() { let malformed_response = wrap_llm_sanitize_response_fn(callback, std::ptr::null_mut(), None); assert_eq!( - malformed_response( + resolve(malformed_response( json!({"secret": "must be omitted"}), nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ), + )) + .unwrap(), None ); } @@ -484,14 +498,17 @@ fn test_llm_sanitizers_fail_closed_for_runtime_codec_ids_with_embedded_nul() { let request_sanitizer = wrap_llm_sanitize_request_fn(llm_request_alias_cb, std::ptr::null_mut(), None); - assert_eq!( - request_sanitizer( - make_request(), - nemo_relay::api::runtime::LlmSanitizeRequestContext::with_identity( - runtime_identity.clone(), - ), + let request_error = resolve(request_sanitizer( + make_request(), + nemo_relay::api::runtime::LlmSanitizeRequestContext::with_identity( + runtime_identity.clone(), ), - None + )) + .expect_err("an embedded runtime codec ID must fail the callback wrapper"); + assert!( + request_error + .to_string() + .contains("runtime codec ID contains an embedded NUL") ); assert!( last_error_message() @@ -501,12 +518,15 @@ fn test_llm_sanitizers_fail_closed_for_runtime_codec_ids_with_embedded_nul() { let response_sanitizer = wrap_llm_sanitize_response_fn(json_alias_cb, std::ptr::null_mut(), None); - assert_eq!( - response_sanitizer( - json!({"secret": "must be omitted"}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), - ), - None + let response_error = resolve(response_sanitizer( + json!({"secret": "must be omitted"}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), + )) + .expect_err("an embedded runtime codec ID must fail the callback wrapper"); + assert!( + response_error + .to_string() + .contains("runtime codec ID contains an embedded NUL") ); assert!( last_error_message() @@ -542,7 +562,12 @@ fn test_wrap_llm_request_intercept_with_annotated_input() { stream: None, extra: serde_json::Map::from_iter([("annotated".into(), json!(true))]), }; - let outcome = request_intercept("llm", make_request(), Some(annotated)).unwrap(); + let outcome = resolve(request_intercept( + "llm".into(), + make_request(), + Some(annotated), + )) + .unwrap(); assert_eq!(outcome.request.content["intercepted"], json!(true)); let annotated_out = outcome .annotated_request @@ -651,7 +676,7 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { .build(); let (user_data, sanitize_calls) = user_data_counter(); let sanitizer = wrap_event_sanitize_fn(event_sanitize_cb, user_data, Some(free_arc_counter)); - let sanitized = sanitizer(&event, original_fields.clone()); + let sanitized = resolve(sanitizer(event.clone(), original_fields.clone())).unwrap(); assert_eq!(sanitized.data, Some(json!({"safe": true}))); assert_eq!( sanitized @@ -667,12 +692,12 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { let invalid = wrap_event_sanitize_fn(invalid_event_sanitize_cb, std::ptr::null_mut(), None); assert_eq!( - invalid(&event, original_fields.clone()), + resolve(invalid(event.clone(), original_fields.clone())).unwrap(), EventSanitizeFields::default() ); let null = wrap_event_sanitize_fn(null_event_sanitize_cb, std::ptr::null_mut(), None); assert_eq!( - null(&event, original_fields.clone()), + resolve(null(event, original_fields.clone())).unwrap(), EventSanitizeFields::default() ); diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index 24df871fb..2efa65892 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -79,11 +79,9 @@ use crate::convert::{ get_last_callback_error as get_recorded_callback_error, opt_json, parse_timestamp_micros, record_callback_error, to_napi_err, }; +use crate::promise_call::PromiseAwareFn; use crate::stream::LlmStream; -use crate::types::{ - EventSanitizeFields, LlmHandle, ScopeHandle, ScopeStack, ScopeType, ToolHandle, - event_sanitize_fields_from_json, -}; +use crate::types::{LlmHandle, ScopeHandle, ScopeStack, ScopeType, ToolHandle}; #[napi::module_init] fn init() { @@ -756,7 +754,9 @@ fn build_plugin_context( core_registry_api::register_tool_sanitize_request_guardrail( &name, priority, - callable::wrap_js_tool_fn(middleware_tool_callback_tsfn(ctx.env, &callback)?), + callable::wrap_js_tool_sanitize_promise_fn(Arc::new(PromiseAwareFn::new( + ctx.env, &callback, + )?)), ) .map_err(to_napi_err)?; @@ -800,7 +800,9 @@ fn build_plugin_context( core_registry_api::register_tool_sanitize_response_guardrail( &name, priority, - callable::wrap_js_tool_fn(middleware_tool_callback_tsfn(ctx.env, &callback)?), + callable::wrap_js_tool_sanitize_promise_fn(Arc::new(PromiseAwareFn::new( + ctx.env, &callback, + )?)), ) .map_err(to_napi_err)?; @@ -840,9 +842,9 @@ fn build_plugin_context( core_registry_api::register_tool_conditional_execution_guardrail( &name, priority, - callable::wrap_js_tool_conditional_fn(middleware_tool_callback_tsfn( + callable::wrap_js_tool_conditional_promise_fn(Arc::new(PromiseAwareFn::new( ctx.env, &callback, - )?), + )?)), ) .map_err(to_napi_err)?; @@ -888,9 +890,9 @@ fn build_plugin_context( core_registry_api::register_llm_sanitize_request_guardrail( &name, priority, - callable::wrap_js_llm_sanitize_request_fn( - middleware_llm_sanitize_request_callback_tsfn(ctx.env, &callback)?, - ), + callable::wrap_js_llm_sanitize_request_promise_fn(Arc::new(PromiseAwareFn::new( + ctx.env, &callback, + )?)), ) .map_err(to_napi_err)?; @@ -934,9 +936,9 @@ fn build_plugin_context( core_registry_api::register_llm_sanitize_response_guardrail( &name, priority, - callable::wrap_js_llm_sanitize_response_fn( - middleware_llm_sanitize_response_callback_tsfn(ctx.env, &callback)?, - ), + callable::wrap_js_llm_sanitize_response_promise_fn(Arc::new(PromiseAwareFn::new( + ctx.env, &callback, + )?)), ) .map_err(to_napi_err)?; @@ -976,9 +978,9 @@ fn build_plugin_context( core_registry_api::register_llm_conditional_execution_guardrail( &name, priority, - callable::wrap_js_llm_conditional_fn(middleware_json_callback_tsfn( + callable::wrap_js_llm_conditional_promise_fn(Arc::new(PromiseAwareFn::new( ctx.env, &callback, - )?), + )?)), ) .map_err(to_napi_err)?; @@ -1018,12 +1020,13 @@ fn build_plugin_context( let priority = ctx.get::(1)?; let break_chain = ctx.get::(2)?; let callback = ctx.get::(3)?; - let tsfn = middleware_json_callback_tsfn(ctx.env, &callback)?; core_registry_api::register_llm_request_intercept( &name, priority, break_chain, - callable::wrap_js_llm_request_intercept_fn(tsfn), + callable::wrap_js_llm_request_intercept_promise_fn(Arc::new(PromiseAwareFn::new( + ctx.env, &callback, + )?)), ) .map_err(to_napi_err)?; @@ -1147,12 +1150,13 @@ fn build_plugin_context( let priority = ctx.get::(1)?; let break_chain = ctx.get::(2)?; let callback = ctx.get::(3)?; - let callback = middleware_tool_callback_tsfn(ctx.env, &callback)?; core_registry_api::register_tool_request_intercept( &name, priority, break_chain, - callable::wrap_js_tool_request_intercept_fn(callback), + callable::wrap_js_tool_request_intercept_promise_fn(Arc::new(PromiseAwareFn::new( + ctx.env, &callback, + )?)), ) .map_err(to_napi_err)?; @@ -1384,40 +1388,6 @@ impl PersistentJsFunction { unsafe { Option::::from_napi_value(self.env, returned.raw()) }.map(callback_json) } - fn call_event_sanitize(&self, event: Json, fields: EventSanitizeFields) -> napi::Result { - let mut value = ptr::null_mut(); - // SAFETY: `self.reference` is a live N-API reference created in - // `self.env`, and `value` is writable storage for the borrowed - // function value. - let status = - unsafe { napi::sys::napi_get_reference_value(self.env, self.reference, &mut value) }; - if status != napi::sys::Status::napi_ok { - return Err(napi::Error::from_reason( - "failed to borrow event sanitizer function", - )); - } - // SAFETY: `value` was resolved from this struct's function reference, - // so it is a live function value in `self.env` for this call. - let func = unsafe { JsFunction::from_raw_unchecked(self.env, value) }; - // SAFETY: `Json::to_napi_value` created this event value in `self.env`, - // so wrapping it as `JsUnknown` is valid for the immediate callback. - let event = unsafe { - JsUnknown::from_raw_unchecked(self.env, Json::to_napi_value(self.env, event)?) - }; - // SAFETY: `EventSanitizeFields::to_napi_value` created this fields - // value in `self.env`, so wrapping it as `JsUnknown` is valid for the - // immediate callback. - let fields = unsafe { - JsUnknown::from_raw_unchecked( - self.env, - EventSanitizeFields::to_napi_value(self.env, fields)?, - ) - }; - let returned = func.call(None, &[event, fields])?; - // SAFETY: `returned` is the live result of invoking `func` in this environment. - unsafe { Option::::from_napi_value(self.env, returned.raw()) }.map(callback_json) - } - fn call_json(&self, argument: Json) -> napi::Result { let mut value = ptr::null_mut(); // SAFETY: `self.reference` is a live N-API reference created in @@ -1442,89 +1412,9 @@ impl PersistentJsFunction { } } -fn core_event_fields( - fields: EventSanitizeFields, -) -> Option { - Some(nemo_relay::api::event::EventSanitizeFields { - data: fields.data, - category_profile: fields - .category_profile - .map(serde_json::from_value) - .transpose() - .ok()?, - metadata: fields.metadata, - }) -} - -fn js_event_fields(fields: &nemo_relay::api::event::EventSanitizeFields) -> EventSanitizeFields { - EventSanitizeFields { - data: fields.data.clone(), - category_profile: fields - .category_profile - .as_ref() - .and_then(|value| serde_json::to_value(value).ok()), - metadata: fields.metadata.clone(), - } -} - fn node_event_sanitize_fn(env: &Env, func: &JsFunction) -> napi::Result { - let callback = callable::safe_middleware_callback(env, func)?; - let direct = Arc::new(PersistentJsFunction::new(env, &callback)?); - let register_thread = std::thread::current().id(); - let mut tsfn = callback.create_threadsafe_function( - 0, - |ctx: napi::threadsafe_function::ThreadSafeCallContext<(Json, Json)>| { - Ok(vec![ctx.value.0, ctx.value.1]) - }, - )?; - tsfn.unref(env)?; - let background = callable::wrap_js_event_sanitize_fn(tsfn); - Ok(Arc::new(move |event, fields| { - if std::thread::current().id() == register_thread { - let event_json = match event.try_to_json_value() { - Ok(event_json) => event_json, - Err(error) => { - record_callback_error(format!( - "nemo_relay: failed to serialize JS event sanitizer context: {error}" - )); - return nemo_relay::api::event::EventSanitizeFields::default(); - } - }; - let sanitized = (|| -> FlowResult<_> { - let value = direct - .call_event_sanitize(event_json, js_event_fields(&fields)) - .map_err(|error| { - FlowError::Internal(format!( - "nemo_relay: JS event sanitizer callback failed: {error}" - )) - })?; - let value = callable::unwrap_middleware_result( - value, - "nemo_relay: JS event sanitizer callback failed", - )?; - let fields = event_sanitize_fields_from_json(value).map_err(|error| { - FlowError::Internal(format!( - "nemo_relay: JS event sanitizer callback failed: invalid JS event sanitizer result: {error}" - )) - })?; - core_event_fields(fields).ok_or_else(|| { - FlowError::Internal( - "nemo_relay: JS event sanitizer callback failed: invalid JS event sanitizer result" - .to_string(), - ) - }) - })(); - match sanitized { - Ok(sanitized) => sanitized, - Err(error) => { - record_callback_error(error.to_string()); - nemo_relay::api::event::EventSanitizeFields::default() - } - } - } else { - background(event, fields) - } - })) + let callback = Arc::new(crate::promise_call::PromiseAwareFn::new(env, func)?); + Ok(callable::wrap_js_event_sanitize_promise_fn(callback)) } type NodeLlmCodec = ( @@ -1857,15 +1747,22 @@ pub fn clear_last_callback_error() { /// Internal test helper: invoke a closed JS tool callback wrapper and return the fallback value. #[napi(js_name = "__testClosedToolCallback")] -pub fn test_closed_tool_callback( +pub async fn test_closed_tool_callback( callback: ThreadsafeFunction<(String, Json), ErrorStrategy::Fatal>, name: String, args: Json, -) -> Json { +) -> Result { clear_recorded_callback_error(); let _ = callback.clone().abort(); let wrapped = callable::wrap_js_tool_fn(callback); - wrapped(&name, args) + let fallback = args.clone(); + match wrapped(name, args).await { + Ok(value) => Ok(value), + Err(error) => { + record_callback_error(error.to_string()); + Ok(fallback) + } + } } /// Internal test helper: model a closed JS LLM request sanitizer. @@ -2767,8 +2664,10 @@ macro_rules! napi_event_guardrail_api { ($register_name:ident, $deregister_name:ident, $core_register:path, $core_deregister:path) => { /// Register an event sanitize guardrail. /// - /// The callback must be synchronous. Callback, serialization, conversion, or - /// invalid-result failures clear the event fields and record the error for + /// The callback may return fields directly or in a Promise. Scope and mark + /// calls queue the event and return synchronously; publication resumes after + /// the Promise settles. Callback, serialization, conversion, or invalid-result + /// failures preserve the original event fields and record the error for /// `getLastCallbackError()`. #[napi] pub fn $register_name( @@ -2776,7 +2675,7 @@ macro_rules! napi_event_guardrail_api { name: String, priority: i32, #[napi( - ts_arg_type = "(event: Json, fields: EventSanitizeFields) => EventSanitizeFields" + ts_arg_type = "(event: Json, fields: EventSanitizeFields) => EventSanitizeFields | Promise" )] guardrail: JsFunction, ) -> Result<()> { @@ -2822,8 +2721,13 @@ macro_rules! napi_guardrail_tool_api { priority: i32, guardrail: JsFunction, ) -> Result<()> { - let callback = middleware_tool_callback_tsfn(&env, &guardrail)?; - $core_register(&name, priority, $wrapper(callback)).map_err(to_napi_err) + let callback = Arc::new(PromiseAwareFn::new(&env, &guardrail)?); + $core_register( + &name, + priority, + callable::wrap_js_tool_sanitize_promise_fn(callback), + ) + .map_err(to_napi_err) } $(#[doc = $dereg_doc])* @@ -2878,13 +2782,20 @@ pub fn register_tool_conditional_execution_guardrail( env: Env, name: String, priority: i32, + #[napi( + ts_arg_type = "(toolName: string, args: Json) => string | null | Promise" + )] guardrail: JsFunction, ) -> Result<()> { - let callback = middleware_tool_callback_tsfn(&env, &guardrail)?; + let callback = std::sync::Arc::new( + crate::promise_call::PromiseAwareFn::new(&env, &guardrail).map_err(|error| { + napi::Error::from_reason(format!("failed to create PromiseAwareFn: {error}")) + })?, + ); core_registry_api::register_tool_conditional_execution_guardrail( &name, priority, - callable::wrap_js_tool_conditional_fn(callback), + callable::wrap_js_tool_conditional_promise_fn(callback), ) .map_err(to_napi_err) } @@ -2914,8 +2825,18 @@ macro_rules! napi_intercept_tool_api { break_chain: bool, callable: JsFunction, ) -> Result<()> { - let callback = middleware_tool_callback_tsfn(&env, &callable)?; - $core_register(&name, priority, break_chain, $wrapper(callback)).map_err(to_napi_err) + let callback = std::sync::Arc::new( + crate::promise_call::PromiseAwareFn::new(&env, &callable).map_err(|error| { + napi::Error::from_reason(format!("failed to create PromiseAwareFn: {error}")) + })?, + ); + $core_register( + &name, + priority, + break_chain, + callable::wrap_js_tool_request_intercept_promise_fn(callback), + ) + .map_err(to_napi_err) } $(#[doc = $dereg_doc])* @@ -3003,15 +2924,16 @@ pub fn register_llm_sanitize_request_guardrail( name: String, priority: i32, #[napi( - ts_arg_type = "(request: Json, context: import('./plugin').LlmSanitizeRequestContext) => Json | null" + ts_arg_type = "(request: Json, context: import('./plugin').LlmSanitizeRequestContext) => Json | null | Promise" )] guardrail: JsFunction, ) -> Result<()> { - let callback = middleware_llm_sanitize_request_callback_tsfn(&env, &guardrail)?; core_registry_api::register_llm_sanitize_request_guardrail( &name, priority, - callable::wrap_js_llm_sanitize_request_fn(callback), + callable::wrap_js_llm_sanitize_request_promise_fn(Arc::new(PromiseAwareFn::new( + &env, &guardrail, + )?)), ) .map_err(to_napi_err) } @@ -3036,15 +2958,16 @@ pub fn register_llm_sanitize_response_guardrail( name: String, priority: i32, #[napi( - ts_arg_type = "(response: Json, context: import('./plugin').LlmSanitizeResponseContext) => Json | null" + ts_arg_type = "(response: Json, context: import('./plugin').LlmSanitizeResponseContext) => Json | null | Promise" )] guardrail: JsFunction, ) -> Result<()> { - let callback = middleware_llm_sanitize_response_callback_tsfn(&env, &guardrail)?; core_registry_api::register_llm_sanitize_response_guardrail( &name, priority, - callable::wrap_js_llm_sanitize_response_fn(callback), + callable::wrap_js_llm_sanitize_response_promise_fn(Arc::new(PromiseAwareFn::new( + &env, &guardrail, + )?)), ) .map_err(to_napi_err) } @@ -3067,13 +2990,18 @@ pub fn register_llm_conditional_execution_guardrail( env: Env, name: String, priority: i32, + #[napi(ts_arg_type = "(request: Json) => string | null | Promise")] guardrail: JsFunction, ) -> Result<()> { - let callback = middleware_json_callback_tsfn(&env, &guardrail)?; + let callback = std::sync::Arc::new( + crate::promise_call::PromiseAwareFn::new(&env, &guardrail).map_err(|error| { + napi::Error::from_reason(format!("failed to create PromiseAwareFn: {error}")) + })?, + ); core_registry_api::register_llm_conditional_execution_guardrail( &name, priority, - callable::wrap_js_llm_conditional_fn(callback), + callable::wrap_js_llm_conditional_promise_fn(callback), ) .map_err(to_napi_err) } @@ -3103,16 +3031,20 @@ pub fn register_llm_request_intercept( priority: i32, break_chain: bool, #[napi( - ts_arg_type = "(args: { name: string; request: Json; annotated: Json | null }) => { request: Json; annotated?: Json | null; pendingMarks?: Array<{ name: string; category?: string | null; categoryProfile?: Json; data?: Json; metadata?: Json }>; optimizationContributions?: Array<{ id?: string; sequence?: number; producer: string; kind: 'input_compression' | 'model_routing' | (string & {}); applied: boolean; model_transition?: { baseline?: { model: string; provider?: string }; effective?: { model: string; provider?: string } }; token_impact?: { baseline?: { prompt_tokens?: number; completion_tokens?: number; cache_read_tokens?: number; cache_write_tokens?: number; total_tokens?: number }; effective?: { prompt_tokens?: number; completion_tokens?: number; cache_read_tokens?: number; cache_write_tokens?: number; total_tokens?: number }; saved?: { prompt_tokens?: number; completion_tokens?: number; cache_read_tokens?: number; cache_write_tokens?: number; total_tokens?: number }; quality?: 'observed' | 'estimated'; estimation_method?: string }; payload_schema?: { name: string; version: string }; payload?: Json; [key: string]: Json | undefined }> }" + ts_arg_type = "(args: { name: string; request: Json; annotated: Json | null }) => import('./plugin').LlmRequestInterceptOutcome | Promise" )] callable: JsFunction, ) -> Result<()> { - let callback = middleware_json_callback_tsfn(&env, &callable)?; + let callback = std::sync::Arc::new( + crate::promise_call::PromiseAwareFn::new(&env, &callable).map_err(|error| { + napi::Error::from_reason(format!("failed to create PromiseAwareFn: {error}")) + })?, + ); core_registry_api::register_llm_request_intercept( &name, priority, break_chain, - callable::wrap_js_llm_request_intercept_fn(callback), + callable::wrap_js_llm_request_intercept_promise_fn(callback), ) .map_err(to_napi_err) } @@ -3245,7 +3177,7 @@ pub fn deregister_subscriber(name: String) -> Result { /// still run. /// /// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this -/// Promise does not block the Node event loop while event sanitizers settle. +/// Promise does not block the Node event loop while Promise-returning event sanitizers settle. /// /// The Promise rejects if the blocking task fails or the core subscriber flush returns an error. /// Callers should handle errors when awaiting it. @@ -3265,8 +3197,10 @@ macro_rules! napi_scope_event_guardrail_api { ($register_name:ident, $deregister_name:ident, $core_register:path, $core_deregister:path) => { /// Register a scope-local event sanitize guardrail. /// - /// The callback must be synchronous. Callback, serialization, conversion, or - /// invalid-result failures clear the event fields and record the error for + /// The callback may return fields directly or in a Promise. Scope and mark + /// calls queue the event and return synchronously; publication resumes after + /// the Promise settles. Callback, serialization, conversion, or invalid-result + /// failures preserve the original event fields and record the error for /// `getLastCallbackError()`. #[napi] pub fn $register_name( @@ -3275,7 +3209,7 @@ macro_rules! napi_scope_event_guardrail_api { name: String, priority: i32, #[napi( - ts_arg_type = "(event: Json, fields: EventSanitizeFields) => EventSanitizeFields" + ts_arg_type = "(event: Json, fields: EventSanitizeFields) => EventSanitizeFields | Promise" )] guardrail: JsFunction, ) -> Result<()> { @@ -3333,8 +3267,14 @@ macro_rules! napi_scope_guardrail_tool_api { ) -> Result<()> { let uuid = uuid::Uuid::parse_str(&scope_uuid) .map_err(|e| napi::Error::from_reason(format!("invalid UUID: {e}")))?; - let callback = middleware_tool_callback_tsfn(&env, &guardrail)?; - $core_register(&uuid, &name, priority, $wrapper(callback)).map_err(to_napi_err) + let callback = Arc::new(PromiseAwareFn::new(&env, &guardrail)?); + $core_register( + &uuid, + &name, + priority, + callable::wrap_js_tool_sanitize_promise_fn(callback), + ) + .map_err(to_napi_err) } $(#[doc = $dereg_doc])* @@ -3402,7 +3342,11 @@ pub fn scope_register_tool_conditional_execution_guardrail( &uuid, &name, priority, - callable::wrap_js_tool_conditional_fn(middleware_tool_callback_tsfn(&env, &guardrail)?), + callable::wrap_js_tool_conditional_promise_fn(std::sync::Arc::new( + crate::promise_call::PromiseAwareFn::new(&env, &guardrail).map_err(|error| { + napi::Error::from_reason(format!("failed to create PromiseAwareFn: {error}")) + })?, + )), ) .map_err(to_napi_err) } @@ -3441,9 +3385,19 @@ macro_rules! napi_scope_intercept_tool_api { ) -> Result<()> { let uuid = uuid::Uuid::parse_str(&scope_uuid) .map_err(|e| napi::Error::from_reason(format!("invalid UUID: {e}")))?; - let callback = middleware_tool_callback_tsfn(&env, &callable)?; - $core_register(&uuid, &name, priority, break_chain, $wrapper(callback)) - .map_err(to_napi_err) + let callback = std::sync::Arc::new( + crate::promise_call::PromiseAwareFn::new(&env, &callable).map_err(|error| { + napi::Error::from_reason(format!("failed to create PromiseAwareFn: {error}")) + })?, + ); + $core_register( + &uuid, + &name, + priority, + break_chain, + callable::wrap_js_tool_request_intercept_promise_fn(callback), + ) + .map_err(to_napi_err) } $(#[doc = $dereg_doc])* @@ -3546,7 +3500,7 @@ pub fn scope_register_llm_sanitize_request_guardrail( name: String, priority: i32, #[napi( - ts_arg_type = "(request: Json, context: import('./plugin').LlmSanitizeRequestContext) => Json | null" + ts_arg_type = "(request: Json, context: import('./plugin').LlmSanitizeRequestContext) => Json | null | Promise" )] guardrail: JsFunction, ) -> Result<()> { @@ -3556,9 +3510,9 @@ pub fn scope_register_llm_sanitize_request_guardrail( &uuid, &name, priority, - callable::wrap_js_llm_sanitize_request_fn(middleware_llm_sanitize_request_callback_tsfn( + callable::wrap_js_llm_sanitize_request_promise_fn(Arc::new(PromiseAwareFn::new( &env, &guardrail, - )?), + )?)), ) .map_err(to_napi_err) } @@ -3590,7 +3544,7 @@ pub fn scope_register_llm_sanitize_response_guardrail( name: String, priority: i32, #[napi( - ts_arg_type = "(response: Json, context: import('./plugin').LlmSanitizeResponseContext) => Json | null" + ts_arg_type = "(response: Json, context: import('./plugin').LlmSanitizeResponseContext) => Json | null | Promise" )] guardrail: JsFunction, ) -> Result<()> { @@ -3600,9 +3554,9 @@ pub fn scope_register_llm_sanitize_response_guardrail( &uuid, &name, priority, - callable::wrap_js_llm_sanitize_response_fn(middleware_llm_sanitize_response_callback_tsfn( + callable::wrap_js_llm_sanitize_response_promise_fn(Arc::new(PromiseAwareFn::new( &env, &guardrail, - )?), + )?)), ) .map_err(to_napi_err) } @@ -3640,7 +3594,11 @@ pub fn scope_register_llm_conditional_execution_guardrail( &uuid, &name, priority, - callable::wrap_js_llm_conditional_fn(middleware_json_callback_tsfn(&env, &guardrail)?), + callable::wrap_js_llm_conditional_promise_fn(std::sync::Arc::new( + crate::promise_call::PromiseAwareFn::new(&env, &guardrail).map_err(|error| { + napi::Error::from_reason(format!("failed to create PromiseAwareFn: {error}")) + })?, + )), ) .map_err(to_napi_err) } @@ -3677,19 +3635,23 @@ pub fn scope_register_llm_request_intercept( priority: i32, break_chain: bool, #[napi( - ts_arg_type = "(args: { name: string; request: Json; annotated: Json | null }) => { request: Json; annotated?: Json | null; pendingMarks?: Array<{ name: string; category?: string | null; categoryProfile?: Json; data?: Json; metadata?: Json }>; optimizationContributions?: Array<{ id?: string; sequence?: number; producer: string; kind: 'input_compression' | 'model_routing' | (string & {}); applied: boolean; model_transition?: { baseline?: { model: string; provider?: string }; effective?: { model: string; provider?: string } }; token_impact?: { baseline?: { prompt_tokens?: number; completion_tokens?: number; cache_read_tokens?: number; cache_write_tokens?: number; total_tokens?: number }; effective?: { prompt_tokens?: number; completion_tokens?: number; cache_read_tokens?: number; cache_write_tokens?: number; total_tokens?: number }; saved?: { prompt_tokens?: number; completion_tokens?: number; cache_read_tokens?: number; cache_write_tokens?: number; total_tokens?: number }; quality?: 'observed' | 'estimated'; estimation_method?: string }; payload_schema?: { name: string; version: string }; payload?: Json; [key: string]: Json | undefined }> }" + ts_arg_type = "(args: { name: string; request: Json; annotated: Json | null }) => import('./plugin').LlmRequestInterceptOutcome | Promise" )] callable: JsFunction, ) -> Result<()> { let uuid = uuid::Uuid::parse_str(&scope_uuid) .map_err(|e| napi::Error::from_reason(format!("invalid UUID: {e}")))?; - let callback = middleware_json_callback_tsfn(&env, &callable)?; + let callback = std::sync::Arc::new( + crate::promise_call::PromiseAwareFn::new(&env, &callable).map_err(|error| { + napi::Error::from_reason(format!("failed to create PromiseAwareFn: {error}")) + })?, + ); core_registry_api::scope_register_llm_request_intercept( &uuid, &name, priority, break_chain, - callable::wrap_js_llm_request_intercept_fn(callback), + callable::wrap_js_llm_request_intercept_promise_fn(callback), ) .map_err(to_napi_err) } @@ -3865,7 +3827,9 @@ pub fn tool_request_intercepts(env: Env, name: String, args: Json) -> Result Result< async move { TASK_SCOPE_STACK .scope(scope_stack, async move { - core_tool_api::tool_conditional_execution(&name, &args).map_err(to_napi_err) + core_tool_api::tool_conditional_execution(&name, &args) + .await + .map_err(to_napi_err) }) .await }, @@ -3905,6 +3871,7 @@ pub fn llm_request_intercepts(env: Env, name: String, request: Json) -> Result Result { async move { TASK_SCOPE_STACK .scope(scope_stack, async move { - core_llm_api::llm_conditional_execution(&llm_request).map_err(to_napi_err) + core_llm_api::llm_conditional_execution(&llm_request) + .await + .map_err(to_napi_err) }) .await }, diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index ba8abafc2..b40aba073 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -207,6 +207,318 @@ fn recv_middleware_option_string_result( } } +async fn await_middleware_json_result( + rx: tokio::sync::oneshot::Receiver, + error_prefix: &str, +) -> Result { + let value = rx + .await + .map_err(|error| FlowError::Internal(format!("{error_prefix}: {error}")))?; + unwrap_middleware_result(value, error_prefix) +} + +async fn await_middleware_json_or_value( + rx: tokio::sync::oneshot::Receiver, + error_prefix: &str, + fallback: Json, +) -> Json { + match await_middleware_json_result(rx, error_prefix).await { + Ok(value) => value, + Err(error) => { + record_callback_error(error.to_string()); + fallback + } + } +} + +async fn await_middleware_option_string_result( + rx: tokio::sync::oneshot::Receiver, + error_prefix: &str, +) -> Result> { + match await_middleware_json_result(rx, error_prefix).await? { + Json::Null => Ok(None), + Json::String(value) => Ok(Some(value)), + other => Err(FlowError::Internal(format!( + "{error_prefix}: expected string or null, got {other:?}", + ))), + } +} + +/// Wrap a Promise-aware JS `(name, args) => string | null` tool guardrail. +pub fn wrap_js_tool_conditional_promise_fn(func: Arc) -> ToolConditionalFn { + Arc::new(move |name: String, args: Json| { + let func = func.clone(); + Box::pin(async move { + let value = func + .call_spread(vec![Json::String(name), args]) + .await + .inspect_err(|error| record_callback_error(error.to_string()))?; + match value { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => { + let error = FlowError::Internal(format!( + "JS tool conditional callback failed: expected string or null, got {other:?}" + )); + record_callback_error(error.to_string()); + Err(error) + } + } + }) + }) +} + +/// Wrap a Promise-aware JS `(name, args) => Json` tool request intercept. +pub fn wrap_js_tool_request_intercept_promise_fn(func: Arc) -> ToolInterceptFn { + Arc::new(move |name: String, args: Json| { + let func = func.clone(); + Box::pin(async move { + func.call_spread(vec![Json::String(name), args]) + .await + .inspect_err(|error| record_callback_error(error.to_string())) + }) + }) +} + +/// Wrap a Promise-aware JS tool sanitizer. +pub fn wrap_js_tool_sanitize_promise_fn(func: Arc) -> ToolSanitizeFn { + Arc::new(move |name: String, value: Json| { + let func = func.clone(); + Box::pin(async move { + func.call_spread(vec![Json::String(name), value]) + .await + .inspect_err(|error| { + record_callback_error(error.to_string()); + }) + }) + }) +} + +/// Wrap a Promise-aware JS LLM request sanitizer. +pub fn wrap_js_llm_sanitize_request_promise_fn(func: Arc) -> LlmSanitizeRequestFn { + Arc::new( + move |request: LlmRequest, context: LlmSanitizeRequestContext| { + let func = func.clone(); + Box::pin(async move { + let request = serde_json::to_value(request).unwrap_or(Json::Null); + let context = js_llm_sanitize_request_context(&context); + let value = func + .call_spread_with_arg0(Box::new(move |env| { + let mut args = env.create_array_with_length(2)?; + let request = unsafe { + JsUnknown::from_raw_unchecked( + env.raw(), + Json::to_napi_value(env.raw(), request)?, + ) + }; + args.set_element(0, request)?; + args.set_element( + 1, + js_llm_sanitize_request_context_to_napi(env, context)?, + )?; + Ok(js_object_to_unknown(env, args)) + })) + .await + .inspect_err(|error| { + record_callback_error(error.to_string()); + })?; + if value.is_null() { + Ok(None) + } else { + serde_json::from_value(value) + .map(Some) + .map_err(|error| { + let error = FlowError::Internal(format!( + "JS LLM sanitize request callback failed: failed to deserialize LlmRequest: {error}" + )); + record_callback_error(error.to_string()); + error + }) + } + }) + }, + ) +} + +/// Wrap a Promise-aware JS LLM response sanitizer. +pub fn wrap_js_llm_sanitize_response_promise_fn( + func: Arc, +) -> LlmSanitizeResponseFn { + Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { + let func = func.clone(); + Box::pin(async move { + let context = js_llm_sanitize_response_context(&context); + let value = func + .call_spread_with_arg0(Box::new(move |env| { + let mut args = env.create_array_with_length(2)?; + let response = unsafe { + JsUnknown::from_raw_unchecked( + env.raw(), + Json::to_napi_value(env.raw(), response)?, + ) + }; + args.set_element(0, response)?; + args.set_element(1, js_llm_sanitize_response_context_to_napi(env, context)?)?; + Ok(js_object_to_unknown(env, args)) + })) + .await + .inspect_err(|error| { + record_callback_error(error.to_string()); + })?; + Ok((!value.is_null()).then_some(value)) + }) + }) +} + +/// Wrap a Promise-aware JS `(request) => string | null` LLM guardrail. +pub fn wrap_js_llm_conditional_promise_fn(func: Arc) -> LlmConditionalFn { + Arc::new(move |request: LlmRequest| { + let func = func.clone(); + Box::pin(async move { + let value = func + .call(serde_json::to_value(request).unwrap_or(Json::Null)) + .await + .inspect_err(|error| record_callback_error(error.to_string()))?; + match value { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => { + let error = FlowError::Internal(format!( + "JS LLM conditional callback failed: expected string or null, got {other:?}" + )); + record_callback_error(error.to_string()); + Err(error) + } + } + }) + }) +} + +/// Wrap a Promise-aware JS LLM request intercept. +pub fn wrap_js_llm_request_intercept_promise_fn( + func: Arc, +) -> LlmRequestInterceptFn { + Arc::new( + move |name: String, request: LlmRequest, annotated: Option| { + let func = func.clone(); + Box::pin(async move { + let value = func + .call(serde_json::json!({ + "name": name, + "request": request, + "annotated": annotated, + })) + .await + .inspect_err(|error| { + record_callback_error(error.to_string()); + })?; + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct JsOutcome { + request: LlmRequest, + #[serde(default)] + annotated: Option, + #[serde(default)] + pending_marks: Vec, + #[serde(default)] + optimization_contributions: Vec, + } + let outcome: JsOutcome = serde_json::from_value(value).map_err(|error| { + let error = FlowError::Internal(format!( + "invalid JS LLM request intercept outcome: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + Ok(LlmRequestInterceptOutcome { + request: outcome.request, + annotated_request: outcome.annotated, + pending_marks: outcome.pending_marks.into_iter().map(Into::into).collect(), + optimization_contributions: outcome.optimization_contributions, + }) + }) + }, + ) +} + +/// Wrap a Promise-aware JS event sanitizer. +/// +/// Event sanitizers run on Relay's serial publication dispatcher, not on the +/// JavaScript registration thread. Waiting here therefore preserves synchronous +/// scope/mark APIs while allowing the JavaScript callback to settle a Promise +/// on the Node event loop. +pub fn wrap_js_event_sanitize_promise_fn(func: Arc) -> EventSanitizeFn { + Arc::new(move |event: Event, fields: CoreEventSanitizeFields| { + let func = func.clone(); + Box::pin(async move { + let event_json = JsEvent::try_from_event(&event) + .map(JsEvent::into_json) + .map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer context: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let js_fields = EventSanitizeFields { + data: fields.data, + category_profile: fields + .category_profile + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer category profile: {error}" + )); + record_callback_error(error.to_string()); + error + })?, + metadata: fields.metadata, + }; + let value = func + .call_spread(vec![ + event_json, + serde_json::to_value(js_fields).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer fields: {error}" + )); + record_callback_error(error.to_string()); + error + })?, + ]) + .await + .inspect_err(|error| { + // Scope and mark publication happens on the dispatcher + // thread. Preserve the event (the core fails open) while + // making the binding-visible failure available to Node. + record_callback_error(error.to_string()); + })?; + let fields = event_sanitize_fields_from_json(value).map_err(|error| { + let error = + FlowError::Internal(format!("invalid JS event sanitizer result: {error}")); + record_callback_error(error.to_string()); + error + })?; + let category_profile = fields + .category_profile + .map(serde_json::from_value) + .transpose() + .map_err(|error| { + let error = + FlowError::Internal(format!("invalid JS event sanitizer result: {error}")); + record_callback_error(error.to_string()); + error + })?; + Ok(CoreEventSanitizeFields { + data: fields.data, + category_profile, + metadata: fields.metadata, + }) + }) + }) +} + fn recv_json_or_null(rx: std::sync::mpsc::Receiver, error_prefix: &str) -> Json { rx.recv().unwrap_or_else(|e| { record_callback_error(format!("{error_prefix}: {e}")); @@ -249,28 +561,25 @@ pub fn wrap_js_tool_fn( func: ThreadsafeFunction<(String, Json), ErrorStrategy::Fatal>, ) -> ToolSanitizeFn { let func = Arc::new(func); - Arc::new(move |name: &str, args: Json| { + Arc::new(move |name: String, args: Json| { let func = func.clone(); - let name = name.to_string(); - let fallback = args.clone(); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - (name, args), - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - record_callback_error(format!( - "nemo_relay: failed to queue JS tool callback: {status:?}" - )); - return fallback; - } - // TODO: This closure returns Json (not Result), so we cannot propagate - // errors through the type system. Log the error so failures are not silent. - recv_middleware_json_or_value(rx, "nemo_relay: JS tool callback failed", fallback) + Box::pin(async move { + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + (name, args), + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS tool callback: {status:?}" + ))); + } + await_middleware_json_result(rx, "nemo_relay: JS tool callback failed").await + }) }) } @@ -279,25 +588,25 @@ pub fn wrap_js_tool_conditional_fn( func: ThreadsafeFunction<(String, Json), ErrorStrategy::Fatal>, ) -> ToolConditionalFn { let func = Arc::new(func); - Arc::new(move |name: &str, args: &Json| { + Arc::new(move |name: String, args: Json| { let func = func.clone(); - let name = name.to_string(); - let args = args.clone(); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - (name, args), - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - return Err(FlowError::Internal(format!( - "failed to queue JS tool conditional callback: {status:?}", - ))); - } - recv_middleware_option_string_result(rx, "JS tool conditional callback failed") + Box::pin(async move { + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + (name, args), + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS tool conditional callback: {status:?}", + ))); + } + await_middleware_option_string_result(rx, "JS tool conditional callback failed").await + }) }) } @@ -306,24 +615,25 @@ pub fn wrap_js_tool_request_intercept_fn( func: ThreadsafeFunction<(String, Json), ErrorStrategy::Fatal>, ) -> ToolInterceptFn { let func = Arc::new(func); - Arc::new(move |name: &str, args: Json| { + Arc::new(move |name: String, args: Json| { let func = func.clone(); - let name = name.to_string(); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - (name, args), - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - return Err(FlowError::Internal(format!( - "failed to queue JS tool callback: {status:?}", - ))); - } - recv_middleware_json_result(rx, "JS tool callback failed") + Box::pin(async move { + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + (name, args), + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS tool callback: {status:?}", + ))); + } + await_middleware_json_result(rx, "JS tool callback failed").await + }) }) } @@ -367,57 +677,57 @@ pub fn wrap_js_llm_request_intercept_fn( ) -> LlmRequestInterceptFn { let func = Arc::new(func); Arc::new( - move |name: &str, - request: LlmRequest, - annotated: Option| - -> Result { + move |name: String, request: LlmRequest, annotated: Option| { let func = func.clone(); - let req_json = serde_json::to_value(&request).unwrap_or(Json::Null); - let annotated_json = annotated - .as_ref() - .map(|a| serde_json::to_value(a).unwrap_or(Json::Null)) - .unwrap_or(Json::Null); - let arg = serde_json::json!({ - "name": name, - "request": req_json, - "annotated": annotated_json, - }); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - arg, - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - return Err(FlowError::Internal(format!( - "failed to queue JS LLM request intercept callback: {status:?}", - ))); - } - let result = - recv_middleware_json_result(rx, "JS LLM request intercept callback failed")?; + Box::pin(async move { + let req_json = serde_json::to_value(&request).unwrap_or(Json::Null); + let annotated_json = annotated + .as_ref() + .map(|a| serde_json::to_value(a).unwrap_or(Json::Null)) + .unwrap_or(Json::Null); + let arg = serde_json::json!({ + "name": name, + "request": req_json, + "annotated": annotated_json, + }); + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + arg, + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS LLM request intercept callback: {status:?}", + ))); + } + let result = + await_middleware_json_result(rx, "JS LLM request intercept callback failed") + .await?; - #[derive(Deserialize)] - #[serde(rename_all = "camelCase")] - struct JsOutcome { - request: LlmRequest, - #[serde(default)] - annotated: Option, - #[serde(default)] - pending_marks: Vec, - #[serde(default)] - optimization_contributions: Vec, - } - let outcome: JsOutcome = serde_json::from_value(result).map_err(|e| { - FlowError::Internal(format!("invalid JS LLM request intercept outcome: {e}")) - })?; - Ok(LlmRequestInterceptOutcome { - request: outcome.request, - annotated_request: outcome.annotated, - pending_marks: outcome.pending_marks.into_iter().map(Into::into).collect(), - optimization_contributions: outcome.optimization_contributions, + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct JsOutcome { + request: LlmRequest, + #[serde(default)] + annotated: Option, + #[serde(default)] + pending_marks: Vec, + #[serde(default)] + optimization_contributions: Vec, + } + let outcome: JsOutcome = serde_json::from_value(result).map_err(|e| { + FlowError::Internal(format!("invalid JS LLM request intercept outcome: {e}")) + })?; + Ok(LlmRequestInterceptOutcome { + request: outcome.request, + annotated_request: outcome.annotated, + pending_marks: outcome.pending_marks.into_iter().map(Into::into).collect(), + optimization_contributions: outcome.optimization_contributions, + }) }) }, ) @@ -431,11 +741,66 @@ pub fn wrap_js_llm_sanitize_request_fn( let func = Arc::new(func); Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { - let context = js_llm_sanitize_request_context(&context); - let request = serde_json::to_value(request).unwrap_or(Json::Null); - let (tx, rx) = std::sync::mpsc::channel(); + let func = func.clone(); + Box::pin(async move { + let context = js_llm_sanitize_request_context(&context); + let request = serde_json::to_value(request).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM sanitize request: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let (tx, rx) = tokio::sync::oneshot::channel(); + if func.call_with_return_value( + (request.clone(), context), + ThreadsafeFunctionCallMode::Blocking, + move |value: Option| { + let _ = tx.send(callback_json(value)); + Ok(()) + }, + ) != napi::Status::Ok + { + record_callback_error( + "nemo_relay: failed to queue JS LLM sanitize request callback", + ); + return Err(FlowError::Internal( + "failed to queue JS LLM sanitize request callback".into(), + )); + } + let value = await_middleware_json_result( + rx, + "nemo_relay: JS LLM request sanitizer callback failed", + ) + .await + .inspect_err(|error| record_callback_error(error.to_string()))?; + if value.is_null() { + return Ok(None); + } + serde_json::from_value(value) + .map(Some) + .map_err(|error| FlowError::Internal(format!( + "JS LLM sanitize request callback failed: failed to deserialize LlmRequest: {error}" + ))) + .inspect_err(|error| record_callback_error(error.to_string())) + }) + }, + ) +} + +/// Wrap a JS function for LLM response sanitization. The callback receives +/// `(response, context)`; returning `null` omits the event payload. +pub fn wrap_js_llm_sanitize_response_fn( + func: ThreadsafeFunction<(Json, JsLlmSanitizeResponseContext), ErrorStrategy::Fatal>, +) -> LlmSanitizeResponseFn { + let func = Arc::new(func); + Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { + let func = func.clone(); + Box::pin(async move { + let context = js_llm_sanitize_response_context(&context); + let (tx, rx) = tokio::sync::oneshot::channel(); if func.call_with_return_value( - (request.clone(), context), + (response, context), ThreadsafeFunctionCallMode::Blocking, move |value: Option| { let _ = tx.send(callback_json(value)); @@ -444,58 +809,20 @@ pub fn wrap_js_llm_sanitize_request_fn( ) != napi::Status::Ok { record_callback_error( - "nemo_relay: failed to queue JS LLM sanitize request callback", + "nemo_relay: failed to queue JS LLM sanitize response callback", ); - return None; + return Err(FlowError::Internal( + "failed to queue JS LLM sanitize response callback".into(), + )); } - let value = recv_middleware_json_or_value( + let value = await_middleware_json_result( rx, - "nemo_relay: JS LLM request sanitizer callback failed", - Json::Null, - ); - if value.is_null() { - return None; - } - serde_json::from_value(value).map_or_else( - |error| { - record_callback_error(format!( - "nemo_relay: JS LLM sanitize request callback failed: failed to deserialize LlmRequest: {error}" - )); - None - }, - Some, - ) - }, - ) -} - -/// Wrap a JS function for LLM response sanitization. The callback receives -/// `(response, context)`; returning `null` omits the event payload. -pub fn wrap_js_llm_sanitize_response_fn( - func: ThreadsafeFunction<(Json, JsLlmSanitizeResponseContext), ErrorStrategy::Fatal>, -) -> LlmSanitizeResponseFn { - let func = Arc::new(func); - Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { - let context = js_llm_sanitize_response_context(&context); - let (tx, rx) = std::sync::mpsc::channel(); - if func.call_with_return_value( - (response, context), - ThreadsafeFunctionCallMode::Blocking, - move |value: Option| { - let _ = tx.send(callback_json(value)); - Ok(()) - }, - ) != napi::Status::Ok - { - record_callback_error("nemo_relay: failed to queue JS LLM sanitize response callback"); - return None; - } - let value = recv_middleware_json_or_value( - rx, - "nemo_relay: JS LLM response sanitizer callback failed", - Json::Null, - ); - Some(value).and_then(|value| (!value.is_null()).then_some(value)) + "nemo_relay: JS LLM response sanitizer callback failed", + ) + .await + .inspect_err(|error| record_callback_error(error.to_string()))?; + Ok((!value.is_null()).then_some(value)) + }) }) } @@ -649,24 +976,26 @@ pub fn wrap_js_llm_conditional_fn( func: ThreadsafeFunction, ) -> LlmConditionalFn { let func = Arc::new(func); - Arc::new(move |request: &LlmRequest| { + Arc::new(move |request: LlmRequest| { let func = func.clone(); - let req_json = serde_json::to_value(request).unwrap_or(Json::Null); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - req_json, - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - return Err(FlowError::Internal(format!( - "failed to queue JS LLM conditional callback: {status:?}", - ))); - } - recv_middleware_option_string_result(rx, "JS LLM conditional callback failed") + Box::pin(async move { + let req_json = serde_json::to_value(request).unwrap_or(Json::Null); + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + req_json, + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS LLM conditional callback: {status:?}", + ))); + } + await_middleware_option_string_result(rx, "JS LLM conditional callback failed").await + }) }) } @@ -783,72 +1112,93 @@ pub fn wrap_js_event_sanitize_fn( func: ThreadsafeFunction<(Json, Json), ErrorStrategy::Fatal>, ) -> EventSanitizeFn { let func = Arc::new(func); - Arc::new(move |event: &Event, fields: CoreEventSanitizeFields| { - let event_json = match JsEvent::try_from_event(event) { - Ok(event) => event.into_json(), - Err(error) => { + Arc::new(move |event: Event, fields: CoreEventSanitizeFields| { + let func = func.clone(); + Box::pin(async move { + let event_json = match JsEvent::try_from_event(&event) { + Ok(event) => event.into_json(), + Err(error) => { + record_callback_error(format!( + "nemo_relay: failed to serialize JS event sanitizer context: {error}" + )); + return Err(FlowError::Internal(error.to_string())); + } + }; + let js_fields = EventSanitizeFields { + data: fields.data.clone(), + category_profile: fields + .category_profile + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer category profile: {error}" + )); + record_callback_error(error.to_string()); + error + })?, + metadata: fields.metadata.clone(), + }; + let js_fields = serde_json::to_value(js_fields).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer fields: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + (event_json, js_fields), + ThreadsafeFunctionCallMode::Blocking, + move |value: Option| { + let _ = tx.send(callback_json(value)); + Ok(()) + }, + ); + if status != napi::Status::Ok { record_callback_error(format!( - "nemo_relay: failed to serialize JS event sanitizer context: {error}" + "nemo_relay: failed to queue JS event sanitizer callback: {status:?}" )); - return CoreEventSanitizeFields::default(); + return Err(FlowError::Internal(format!( + "failed to queue JS event sanitizer callback: {status:?}" + ))); } - }; - let js_fields = EventSanitizeFields { - data: fields.data.clone(), - category_profile: fields - .category_profile - .as_ref() - .and_then(|value| serde_json::to_value(value).ok()), - metadata: fields.metadata.clone(), - }; - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - ( - event_json, - serde_json::to_value(js_fields).unwrap_or(Json::Null), - ), - ThreadsafeFunctionCallMode::Blocking, - move |value: Option| { - let _ = tx.send(callback_json(value)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - record_callback_error(format!( - "nemo_relay: failed to queue JS event sanitizer callback: {status:?}" - )); - return CoreEventSanitizeFields::default(); - } - let sanitized = (|| -> Result<_> { - let result = - recv_middleware_json_result(rx, "nemo_relay: JS event sanitizer callback failed")?; - let result = event_sanitize_fields_from_json(result).map_err(|error| { - FlowError::Internal(format!( - "nemo_relay: invalid JS event sanitizer result: {error}" - )) - })?; - let category_profile = result - .category_profile - .map(serde_json::from_value) - .transpose() - .map_err(|error| { + let sanitized: Result = async { + let result = await_middleware_json_result( + rx, + "nemo_relay: JS event sanitizer callback failed", + ) + .await?; + let result = event_sanitize_fields_from_json(result).map_err(|error| { FlowError::Internal(format!( "nemo_relay: invalid JS event sanitizer result: {error}" )) })?; - Ok(CoreEventSanitizeFields { - data: result.data, - category_profile, - metadata: result.metadata, - }) - })(); - match sanitized { - Ok(sanitized) => sanitized, - Err(error) => { - record_callback_error(error.to_string()); - CoreEventSanitizeFields::default() + let category_profile = result + .category_profile + .map(serde_json::from_value) + .transpose() + .map_err(|error| { + FlowError::Internal(format!( + "nemo_relay: invalid JS event sanitizer result: {error}" + )) + })?; + Ok(CoreEventSanitizeFields { + data: result.data, + category_profile, + metadata: result.metadata, + }) } - } + .await; + match sanitized { + Ok(sanitized) => Ok(sanitized), + Err(error) => { + record_callback_error(error.to_string()); + Err(error) + } + } + }) }) } diff --git a/crates/node/src/callback_factory.rs b/crates/node/src/callback_factory.rs index 891373dbf..48b301883 100644 --- a/crates/node/src/callback_factory.rs +++ b/crates/node/src/callback_factory.rs @@ -66,14 +66,26 @@ const CALLBACK_FACTORIES_SOURCE: &str = r#"(() => { }, promise(fn) { - return function __nemo_relay_promise_wrapper(error, arg0, next, resolve, reject) { + return function __nemo_relay_promise_wrapper(error, arg0, spread, next, resolve, reject) { if (error != null) { reject(error); return; } Promise.resolve().then(() => ( - next === undefined ? fn(arg0) : fn(arg0, next) - )).then((value) => jsonValue(value === undefined ? null : value)).then(resolve, reject); + next === undefined + ? (spread ? fn(...arg0) : fn(arg0)) + : (spread ? fn(...arg0, next) : fn(arg0, next)) + )).then((value) => jsonValue(value === undefined ? null : value)).then(resolve, (error) => { + let message = 'unknown error'; + try { + if (typeof error === 'string') { + message = error; + } else if (error != null && typeof error.message === 'string') { + message = error.message; + } + } catch {} + reject(message); + }); }; }, }; diff --git a/crates/node/src/promise_call.rs b/crates/node/src/promise_call.rs index bb5435207..cdc15bada 100644 --- a/crates/node/src/promise_call.rs +++ b/crates/node/src/promise_call.rs @@ -53,6 +53,7 @@ enum PrimaryArg { struct CallArgs { arg0: PrimaryArg, + spread: bool, next: Option, completion: CallCompletion, } @@ -76,19 +77,6 @@ impl CallCompletion { } } -fn rejection_message( - string_result: napi::Result, - object_message_result: Option>, -) -> String { - if let Ok(value) = string_result { - value - } else if let Some(message_result) = object_message_result { - message_result.unwrap_or_else(|_| "unknown error".to_string()) - } else { - "unknown error".to_string() - } -} - fn closed_tsfn_error() -> FlowError { FlowError::Internal("PromiseAwareFn threadsafe function closed".into()) } @@ -168,12 +156,12 @@ fn build_completion_unknowns( })?; let reject = env.create_function_from_closure("__nemo_relay_reject", move |ctx| { - let message = rejection_message( - ctx.get::(0), - ctx.get::(0) - .ok() - .map(|value| value.get_named_property::("message")), - ); + // Do not invoke arbitrary `error.message` getters here. A throwing + // getter used to escape this callback and abort the N-API call rather + // than settling the middleware future as a rejection. + let message = ctx + .get::(0) + .unwrap_or_else(|_| "unknown error".to_string()); completion.send(Err(FlowError::Internal(message))); ctx.env.get_undefined() })?; @@ -208,7 +196,13 @@ impl PromiseAwareFn { PrimaryArg::Build(build) => build(&ctx.env)?, }; - let args = vec![arg0, next, resolve, reject]; + let spread = unsafe { + JsUnknown::from_raw_unchecked( + ctx.env.raw(), + ctx.env.get_boolean(ctx.value.spread)?.raw(), + ) + }; + let args = vec![arg0, spread, next, resolve, reject]; Ok(args) })?; @@ -222,7 +216,17 @@ impl PromiseAwareFn { /// Call the JS function with the given args and await the result. pub async fn call(&self, args: Json) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), None).await + self.call_inner(PrimaryArg::Json(args), false, None).await + } + + /// Call a JavaScript callback with several JSON arguments. + /// + /// This retains the normal callback shape for middleware such as tool + /// guardrails, whose public contract is `(name, payload)` rather than a + /// single envelope object. + pub async fn call_spread(&self, args: Vec) -> FlowResult { + self.call_inner(PrimaryArg::Json(Json::Array(args)), true, None) + .await } /// Call the JS function with a builder-constructed first argument and await @@ -232,13 +236,20 @@ impl PromiseAwareFn { /// cannot cross the threadsafe-function boundary as plain JSON, such as a /// `#[napi]` class instance. pub async fn call_with_arg0(&self, build_arg0: Arg0Builder) -> FlowResult { - self.call_inner(PrimaryArg::Build(build_arg0), None).await + self.call_inner(PrimaryArg::Build(build_arg0), false, None) + .await + } + + /// Call a JavaScript callback with builder-constructed spread arguments. + pub async fn call_spread_with_arg0(&self, build_arg0: Arg0Builder) -> FlowResult { + self.call_inner(PrimaryArg::Build(build_arg0), true, None) + .await } /// Call the JS function with a middleware-style `next(arg)` callback that /// resolves to a JSON result. pub async fn call_with_json_next(&self, args: Json, next: JsonNextFn) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), Some(NextFn::Json(next))) + self.call_inner(PrimaryArg::Json(args), false, Some(NextFn::Json(next))) .await } @@ -249,7 +260,7 @@ impl PromiseAwareFn { args: Json, next: JsonStreamNextFn, ) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), Some(NextFn::Stream(next))) + self.call_inner(PrimaryArg::Json(args), false, Some(NextFn::Stream(next))) .await } @@ -260,7 +271,12 @@ impl PromiseAwareFn { } } - async fn call_inner(&self, arg0: PrimaryArg, next: Option) -> FlowResult { + async fn call_inner( + &self, + arg0: PrimaryArg, + spread: bool, + next: Option, + ) -> FlowResult { let (sender, receiver) = tokio::sync::oneshot::channel(); let tsfn = self .tsfn @@ -272,6 +288,7 @@ impl PromiseAwareFn { let status = tsfn.call( Ok(CallArgs { arg0, + spread, next, completion: CallCompletion::new(sender), }), diff --git a/crates/node/tests/callback_error_tests.mjs b/crates/node/tests/callback_error_tests.mjs index 83e684a5f..c05eccef9 100644 --- a/crates/node/tests/callback_error_tests.mjs +++ b/crates/node/tests/callback_error_tests.mjs @@ -67,11 +67,11 @@ describe('callback error helpers', () => { } }); - it('closed tool sanitize callbacks preserve the original payload and record the queue failure', () => { + it('closed tool sanitize callbacks preserve the original payload and record the queue failure', async () => { const args = { value: 1, }; - const result = __testClosedToolCallback( + const result = await __testClosedToolCallback( () => ({ ok: true, }), diff --git a/crates/node/tests/event_sanitizers_tests.mjs b/crates/node/tests/event_sanitizers_tests.mjs index 3679b52cf..c5d6d75ea 100644 --- a/crates/node/tests/event_sanitizers_tests.mjs +++ b/crates/node/tests/event_sanitizers_tests.mjs @@ -25,10 +25,10 @@ async function waitFor(events, count) { assert.ok(events.length >= count, `expected ${count} events, received ${events.length}`); } -function assertSanitizerFieldsCleared(event) { - assert.equal(event.data, null); - assert.equal(event.category_profile, null); - assert.equal(event.metadata, null); +function assertSanitizerFieldsPreserved(event, expectedData, expectedMetadata = expectedData) { + assert.deepEqual(event.data, expectedData); + assert.equal(event.category_profile?.subtype, 'seeded'); + assert.deepEqual(event.metadata, expectedMetadata); } async function initializeWithoutDiscoveredPluginConfig(config) { @@ -107,13 +107,35 @@ describe('event sanitizer registries', () => { assert.ok(lifecycle.every((event) => event.category_profile.subtype === 'sanitized')); }); - it('fails closed and records invalid direct sanitizer results', async () => { + it('awaits Promise-returning mark sanitizers without making event() asynchronous', async () => { + const events = capture('node-event-sanitize-promise-sub'); + let settled = false; + lib.registerMarkSanitizeGuardrail('node-event-promise', 0, async (_event, fields) => { + await new Promise((resolve) => setImmediate(resolve)); + settled = true; + return { ...fields, data: { sanitized: true } }; + }); + try { + const result = lib.event('promise-checkpoint', null, { raw: true }); + assert.equal(result, undefined); + assert.equal(settled, false); + await lib.flushSubscribers(); + await waitFor(events, 1); + } finally { + lib.deregisterMarkSanitizeGuardrail('node-event-promise'); + lib.deregisterSubscriber('node-event-sanitize-promise-sub'); + } + assert.equal(settled, true); + assert.deepEqual(events.at(-1).data, { sanitized: true }); + }); + + it('fails open and records invalid sanitizer results', async () => { const events = capture('node-event-sanitize-invalid-sub'); const invalidResults = { scalar: () => 'invalid', emptyObject: () => ({}), array: () => [], - promise: () => Promise.resolve({ data: { changed: true } }), + promise: () => Promise.resolve([]), }; try { for (const [kind, sanitizer] of Object.entries(invalidResults)) { @@ -135,7 +157,7 @@ describe('event sanitizer registries', () => { lib.deregisterMarkSanitizeGuardrail(seedName); lib.deregisterMarkSanitizeGuardrail(name); } - assertSanitizerFieldsCleared(events.at(-1)); + assertSanitizerFieldsPreserved(events.at(-1), { kept: kind }); assert.match(lib.getLastCallbackError(), /invalid JS event sanitizer result/); } } finally { @@ -163,12 +185,12 @@ describe('event sanitizer registries', () => { assert.equal(start.metadata.background, true); }); - it('fails closed and records invalid thread-safe sanitizer results', async () => { + it('fails open and records invalid queued sanitizer results', async () => { const events = capture('node-event-sanitize-background-invalid-sub'); const invalidResults = { emptyObject: () => ({}), array: () => [], - promise: () => Promise.resolve({ data: { changed: true } }), + promise: () => Promise.resolve([]), }; try { for (const [kind, sanitizer] of Object.entries(invalidResults)) { @@ -193,7 +215,7 @@ describe('event sanitizer registries', () => { const start = events.find( (event) => event.kind === 'scope' && event.name === name && event.scope_category === 'start', ); - assertSanitizerFieldsCleared(start); + assertSanitizerFieldsPreserved(start, { kept: kind }); assert.match(lib.getLastCallbackError(), /invalid JS event sanitizer result/); } } finally { @@ -201,7 +223,7 @@ describe('event sanitizer registries', () => { } }); - it('fails closed when a thread-safe sanitizer throws', async () => { + it('fails open when a queued sanitizer throws', async () => { const events = capture('node-event-sanitize-background-throw-sub'); lib.clearLastCallbackError(); lib.registerScopeSanitizeStartGuardrail('node-background-throw-seed', -1, (_event, fields) => ({ @@ -220,7 +242,7 @@ describe('event sanitizer registries', () => { const start = events.find( (event) => event.kind === 'scope' && event.name === 'background-throw-tool' && event.scope_category === 'start', ); - assertSanitizerFieldsCleared(start); + assertSanitizerFieldsPreserved(start, { kept: true }); assert.match(lib.getLastCallbackError() ?? '', /background sanitizer boom/i); } finally { lib.deregisterScopeSanitizeStartGuardrail('node-background-throw-seed'); @@ -289,7 +311,7 @@ describe('event sanitizer registries', () => { assert.deepEqual(marks.cleared.data, { raw: true }); }); - it('fails closed when a plugin-owned sanitizer throws', async () => { + it('fails open when a plugin-owned sanitizer throws', async () => { const kind = `node.test.event-sanitize-throw.${Date.now()}`; const events = capture('node-event-sanitize-plugin-throw-sub'); plugin.register(kind, { @@ -314,7 +336,7 @@ describe('event sanitizer registries', () => { lib.event('plugin-throw', null, { raw: true }, { raw: true }); await lib.flushSubscribers(); await waitFor(events, 1); - assertSanitizerFieldsCleared(events.at(-1)); + assertSanitizerFieldsPreserved(events.at(-1), { raw: true }); assert.match(lib.getLastCallbackError() ?? '', /plugin sanitizer boom/i); } finally { plugin.clear(); diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index 61eabb89b..c83c78cb9 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -57,6 +57,17 @@ async function flushSubscriberCallbacks() { } } +async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { + flushSubscribers(); + const deadline = Date.now() + timeoutMs; + while (!predicate()) { + if (Date.now() >= deadline) { + throw new Error('timed out waiting for subscriber callbacks'); + } + await new Promise((resolve) => setImmediate(resolve)); + } +} + function makeNative() { return { headers: {}, @@ -291,7 +302,10 @@ describe('LLM execute', () => { /llm status failure/, ); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks( + () => events.some((e) => e.name === 'exec_status_ok_llm' && e.scope_category === 'end') + && events.some((e) => e.name === 'exec_status_error_llm' && e.scope_category === 'end'), + ); const okEnd = events.find( (e) => e.name === 'exec_status_ok_llm' && e.kind === 'scope' && e.category === 'llm' && e.scope_category === 'end', @@ -386,7 +400,11 @@ describe('LLM guardrails', () => { assert.deepEqual(result, { ok: true }); assert.equal(requestContextChecked, true); assert.equal(responseContextChecked, true); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks( + () => + events.some((event) => event.name === 'contextual_sanitize_llm' && event.scope_category === 'start') && + events.some((event) => event.name === 'contextual_sanitize_llm' && event.scope_category === 'end'), + ); const start = events.find( (event) => event.name === 'contextual_sanitize_llm' && event.scope_category === 'start', ); @@ -718,7 +736,7 @@ describe('LLM guardrails', () => { } }); - it('sanitize request guardrail failures omit the payload and remain usable', async () => { + it('sanitize request guardrail failures preserve the payload and remain usable', async () => { const events = []; clearLastCallbackError(); registerSubscriber('node_llm_san_req_throw_sub', (event) => events.push(event)); @@ -728,7 +746,16 @@ describe('LLM guardrails', () => { try { const request = makeNative(); await llmCallExecute('llm_san_req_throw', request, () => ({ ok: true }), null, null, null, null, null); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks( + () => + events.some( + (event) => + event.name === 'llm_san_req_throw' && + event.kind === 'scope' && + event.category === 'llm' && + event.scope_category === 'start', + ), + ); const start = events.find( (event) => event.name === 'llm_san_req_throw' && @@ -736,8 +763,8 @@ describe('LLM guardrails', () => { event.category === 'llm' && event.scope_category === 'start', ); - assert.equal(start.data, null); - assert.match(getLastCallbackError() ?? '', /JavaScript callback threw/i); + assert.deepEqual(start.data, { headers: request.headers, content: request.content }); + assert.match(getLastCallbackError() ?? '', /(unknown error|callback)/i); deregisterLlmSanitizeRequestGuardrail('node_llm_san_req_throw'); const result = await llmCallExecute( @@ -840,7 +867,7 @@ describe('LLM guardrails', () => { } }); - it('sanitize response guardrail failures omit the payload and remain usable', async () => { + it('sanitize response guardrail failures preserve the payload and remain usable', async () => { const events = []; clearLastCallbackError(); registerSubscriber('node_llm_san_resp_throw_sub', (event) => events.push(event)); @@ -850,7 +877,15 @@ describe('LLM guardrails', () => { try { const response = { ok: true }; await llmCallExecute('llm_san_resp_throw', makeNative(), () => response, null, null, null, null, null); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => + events.some( + (event) => + event.name === 'llm_san_resp_throw' && + event.kind === 'scope' && + event.category === 'llm' && + event.scope_category === 'end', + ), + ); const end = events.find( (event) => event.name === 'llm_san_resp_throw' && @@ -858,7 +893,7 @@ describe('LLM guardrails', () => { event.category === 'llm' && event.scope_category === 'end', ); - assert.equal(end.data, null); + assert.deepEqual(end.data, response); assert.match(getLastCallbackError() ?? '', /response sanitizer boom/i); deregisterLlmSanitizeResponseGuardrail('node_llm_san_resp_throw'); @@ -885,6 +920,19 @@ describe('LLM guardrails', () => { deregisterLlmConditionalExecutionGuardrail('node_llm_cond'); }); + it('conditional guardrail awaits a Promise result', async () => { + registerLlmConditionalExecutionGuardrail('node_llm_cond_promise', 10, async () => { + await new Promise((resolve) => setImmediate(resolve)); + return null; + }); + try { + const result = await llmCallExecute('llm_cond_promise', makeNative(), () => ({ ok: true }), null, null, null, null, null); + assert.deepEqual(result, { ok: true }); + } finally { + deregisterLlmConditionalExecutionGuardrail('node_llm_cond_promise'); + } + }); + it('conditional guardrail treats implicit undefined as allow', async () => { registerLlmConditionalExecutionGuardrail('node_llm_cond_undefined', 10, () => undefined); try { @@ -1021,6 +1069,28 @@ describe('LLM intercepts', () => { deregisterLlmRequestIntercept('node_llm_req_mod'); }); + it('request intercept awaits a Promise result', async () => { + registerLlmRequestIntercept('node_llm_req_promise', 10, false, async ({ request, annotated }) => { + await new Promise((resolve) => setImmediate(resolve)); + return { request: { ...request, content: { ...request.content, promised: true } }, annotated }; + }); + try { + const result = await llmCallExecute( + 'llm_req_promise', + makeNative(), + (request) => ({ promised: request.content.promised }), + null, + null, + null, + null, + null, + ); + assert.deepEqual(result, { promised: true }); + } finally { + deregisterLlmRequestIntercept('node_llm_req_promise'); + } + }); + it('request intercept throws a catchable error without terminating Node', async () => { registerLlmRequestIntercept('node_llm_req_throw', 10, false, () => { throw new Error('llm request intercept boom'); @@ -1340,11 +1410,21 @@ describe('LLM intercepts', () => { deregisterLlmRequestIntercept('node_llm_req_helper'); }); - it('generated request-intercept declarations preserve the open optimization kind', () => { + it('generated request-intercept declarations reference the canonical open optimization type', () => { const declarations = readFileSync(new URL('../index.d.ts', import.meta.url), 'utf8'); + const pluginDeclarations = readFileSync(new URL('../plugin.d.ts', import.meta.url), 'utf8'); const openKind = "kind: 'input_compression' | 'model_routing' | (string & {})"; - assert.equal(declarations.split(openKind).length - 1, 3); + assert.equal(declarations.split(openKind).length - 1, 1); + assert.equal(pluginDeclarations.split(openKind).length - 1, 1); + assert.match( + declarations, + /registerLlmRequestIntercept\([^\n]*import\('\.\/plugin'\)\.LlmRequestInterceptOutcome/, + ); + assert.match( + declarations, + /scopeRegisterLlmRequestIntercept\([^\n]*import\('\.\/plugin'\)\.LlmRequestInterceptOutcome/, + ); }); it('generated LLM sanitizer declarations expose directional codec contexts', () => { diff --git a/crates/node/tests/scope_tests.mjs b/crates/node/tests/scope_tests.mjs index cb50f1ade..9ecd3ebb2 100644 --- a/crates/node/tests/scope_tests.mjs +++ b/crates/node/tests/scope_tests.mjs @@ -29,9 +29,13 @@ function rejectWithPrimitive(value) { return Promise.reject(value); } -async function flushSubscriberCallbacks() { +async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { await flushSubscribers(); - for (let i = 0; i < 10; i += 1) { + const deadline = Date.now() + timeoutMs; + while (!predicate()) { + if (Date.now() >= deadline) { + throw new Error('timed out waiting for subscriber callbacks'); + } await new Promise((resolve) => setImmediate(resolve)); } } @@ -104,7 +108,7 @@ describe('Scope operations', () => { try { const scope = pushScope('pop_metadata_scope', ScopeType.Agent, null, null, null, { a: 1, b: 2, c: 3 }); popScope(scope, null, null, { c: 3.5, d: 4 }); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => events.some((e) => e.name === 'pop_metadata_scope' && e.scope_category === 'end')); const end = events.find( (e) => e.name === 'pop_metadata_scope' && e.kind === 'scope' && e.scope_category === 'end', @@ -208,7 +212,7 @@ describe('withScope', () => { await withScope('with_scope_ok_status', ScopeType.Function, () => ({ ok: true }), null, null, null, { caller: 'node', }); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => events.some((e) => e.name === 'with_scope_ok_status' && e.scope_category === 'end')); const end = events.find( (e) => e.name === 'with_scope_ok_status' && e.kind === 'scope' && e.scope_category === 'end', @@ -260,7 +264,7 @@ describe('withScope', () => { }), /node status failure/, ); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => events.some((e) => e.name === 'with_scope_error_status' && e.scope_category === 'end')); const end = events.find( (e) => e.name === 'with_scope_error_status' && e.kind === 'scope' && e.scope_category === 'end', @@ -355,8 +359,7 @@ describe('Subscribers', () => { try { const scope = pushScope('sub_test', ScopeType.Agent, null, null); popScope(scope); - await flushSubscriberCallbacks(); - assert.ok(events.length > 0, 'Expected at least one event'); + await waitForSubscriberCallbacks(() => events.length > 0); } finally { deregisterSubscriber('node_event_collector'); } @@ -367,8 +370,7 @@ describe('Subscribers', () => { registerSubscriber('node_flush_collector', (e) => events.push(e)); try { event('node_flush_mark', null, null, null); - await flushSubscribers(); - await new Promise((resolve) => setImmediate(resolve)); + await waitForSubscriberCallbacks(() => events.some((e) => e.kind === 'mark' && e.name === 'node_flush_mark')); assert.ok(events.some((e) => e.kind === 'mark' && e.name === 'node_flush_mark')); } finally { deregisterSubscriber('node_flush_collector'); @@ -383,7 +385,7 @@ describe('Subscribers', () => { try { const scope = pushScope('prop_test', ScopeType.Function, null, null); popScope(scope); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => captured !== null); assert.ok(captured, 'Expected an event'); assert.ok(typeof captured.uuid === 'string'); assert.ok(typeof captured.timestamp === 'string'); @@ -406,7 +408,7 @@ describe('Subscribers', () => { }, null, ); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => events.some((e) => e.kind === 'mark')); const found = events.some((e) => e.kind === 'mark'); assert.ok(found, 'Expected a Mark event'); } finally { diff --git a/crates/node/tests/tools_tests.mjs b/crates/node/tests/tools_tests.mjs index 3d20628b6..b60642aef 100644 --- a/crates/node/tests/tools_tests.mjs +++ b/crates/node/tests/tools_tests.mjs @@ -608,6 +608,33 @@ describe('Tool guardrails', () => { deregisterToolConditionalExecutionGuardrail('node_tool_cond'); }); + it('conditional guardrail awaits a Promise result', async () => { + registerToolConditionalExecutionGuardrail('node_tool_cond_promise', 10, async () => { + await new Promise((resolve) => setImmediate(resolve)); + return null; + }); + try { + const result = await toolCallExecute('tool_cond_promise', { ok: true }, (args) => args, null, null, null, null); + assert.deepEqual(result, { ok: true }); + } finally { + deregisterToolConditionalExecutionGuardrail('node_tool_cond_promise'); + } + }); + + it('conditional guardrail propagates a rejected Promise', async () => { + registerToolConditionalExecutionGuardrail('node_tool_cond_reject', 10, async () => { + throw new Error('guardrail rejected promise'); + }); + try { + await assert.rejects( + () => toolCallExecute('tool_cond_reject', {}, () => ({ should_not: 'run' }), null, null, null, null), + /guardrail rejected promise/i, + ); + } finally { + deregisterToolConditionalExecutionGuardrail('node_tool_cond_reject'); + } + }); + it('conditional guardrail treats implicit undefined as allow', async () => { registerToolConditionalExecutionGuardrail('node_tool_cond_undefined', 10, () => undefined); try { @@ -782,6 +809,41 @@ describe('Tool intercepts', () => { deregisterToolRequestIntercept('node_tool_req_mod'); }); + it('request intercept awaits a Promise result', async () => { + registerToolRequestIntercept('node_tool_req_promise', 10, false, async (_name, args) => { + await new Promise((resolve) => setImmediate(resolve)); + return { ...args, promised: true }; + }); + try { + const result = await toolCallExecute( + 'tool_req_promise', + { original: true }, + (args) => args, + null, + null, + null, + null, + ); + assert.deepEqual(result, { original: true, promised: true }); + } finally { + deregisterToolRequestIntercept('node_tool_req_promise'); + } + }); + + it('request intercept propagates a rejected Promise', async () => { + registerToolRequestIntercept('node_tool_req_reject', 10, false, async () => { + throw new Error('request intercept rejected promise'); + }); + try { + await assert.rejects( + () => toolCallExecute('tool_req_reject', {}, () => ({ should_not: 'run' }), null, null, null, null), + /request intercept rejected promise/i, + ); + } finally { + deregisterToolRequestIntercept('node_tool_req_reject'); + } + }); + it('request intercept throws a catchable error without terminating Node', async () => { registerToolRequestIntercept('node_tool_req_throw', 10, false, () => { throw new Error('tool request intercept boom'); diff --git a/crates/pii-redaction/src/builtin.rs b/crates/pii-redaction/src/builtin.rs index 414fb13f6..71d421588 100644 --- a/crates/pii-redaction/src/builtin.rs +++ b/crates/pii-redaction/src/builtin.rs @@ -461,12 +461,15 @@ impl CompiledBuiltinBackend { } pub(super) fn tool_sanitize_callback(backend: CompiledBuiltinBackend) -> ToolSanitizeFn { - Arc::new( - move |_name: &str, payload: Json| match backend.trajectory.as_ref() { - Some(trajectory) => trajectory.sanitize_tool_payload(payload), - None => backend.sanitize_json_preorder_dfs(payload), - }, - ) + Arc::new(move |_name: String, payload: Json| { + let backend = backend.clone(); + Box::pin(async move { + Ok(match backend.trajectory.as_ref() { + Some(trajectory) => trajectory.sanitize_tool_payload(payload), + None => backend.sanitize_json_preorder_dfs(payload), + }) + }) + }) } pub(super) fn event_sanitize_callback(backend: CompiledBuiltinBackend) -> EventSanitizeFn { @@ -486,40 +489,43 @@ fn event_sanitize_callback_with_scope_categories( scope_categories: Option<(bool, bool)>, ) -> EventSanitizeFn { Arc::new(move |event, mut fields| { - if scope_categories.is_some_and(|(sanitize_llm, sanitize_tool)| { - matches!(event, Event::Scope(_)) + let backend = backend.clone(); + Box::pin(async move { + if scope_categories.is_some_and(|(sanitize_llm, sanitize_tool)| { + matches!(event, Event::Scope(_)) + && event + .category() + .is_some_and(|category| match category.as_str() { + "llm" => !sanitize_llm, + "tool" => !sanitize_tool, + _ => false, + }) + }) { + return Ok(fields); + } + + if let Some(trajectory) = backend.trajectory.as_ref() { + return Ok(trajectory.sanitize_event_fields(&event, fields)); + } + let specialized_scope = matches!(event, Event::Scope(_)) && event .category() - .is_some_and(|category| match category.as_str() { - "llm" => !sanitize_llm, - "tool" => !sanitize_tool, - _ => false, - }) - }) { - return fields; - } - - if let Some(trajectory) = backend.trajectory.as_ref() { - return trajectory.sanitize_event_fields(event, fields); - } - let specialized_scope = matches!(event, Event::Scope(_)) - && event - .category() - .is_some_and(|category| matches!(category.as_str(), "tool" | "llm")); - - if !specialized_scope { - fields.data = fields - .data - .map(|data| backend.sanitize_json_preorder_dfs(data)); - fields.category_profile = fields.category_profile.and_then(|profile| { - sanitize_serializable_with_backend::(&backend, profile).ok() - }); - } + .is_some_and(|category| matches!(category.as_str(), "tool" | "llm")); + + if !specialized_scope { + fields.data = fields + .data + .map(|data| backend.sanitize_json_preorder_dfs(data)); + fields.category_profile = fields.category_profile.and_then(|profile| { + sanitize_serializable_with_backend::(&backend, profile).ok() + }); + } - fields.metadata = fields - .metadata - .map(|metadata| backend.sanitize_json_preorder_dfs(metadata)); - fields + fields.metadata = fields + .metadata + .map(|metadata| backend.sanitize_json_preorder_dfs(metadata)); + Ok(fields) + }) }) } @@ -527,41 +533,44 @@ pub(super) fn llm_sanitize_request_callback( backend: CompiledBuiltinBackend, ) -> LlmSanitizeRequestFn { Arc::new(move |mut request: LlmRequest, context| { - if let Some(trajectory) = backend.trajectory.as_ref() { - request.headers = trajectory - .sanitize_tool_payload(Json::Object(request.headers)) - .as_object() - .cloned() - .unwrap_or_default(); - request.content = trajectory.sanitize_provider_payload(request.content); - return Some(request); - } - request.headers = backend.sanitize_request_headers(request.headers); - if backend.target_paths.is_empty() { - request.content = backend.sanitize_json_preorder_dfs(request.content); - return Some(request); - } - let resolved = context.resolve_codec(); - let fallback = if resolved.is_none() { - backend - .selected_surface(context.codec()) - .map(build_request_codec) - } else { - None - }; - let Some(codec) = resolved.as_deref().or(fallback.as_deref()) else { - log_llm_payload_omitted("request", context.codec(), "no usable request codec"); - return None; - }; - let sanitized = backend.sanitize_request_with_codec(codec, &request); - if sanitized.is_none() { - log_llm_payload_omitted( - "request", - context.codec(), - "codec decode, sanitize, or encode failure", - ); - } - sanitized + let backend = backend.clone(); + Box::pin(async move { + if let Some(trajectory) = backend.trajectory.as_ref() { + request.headers = trajectory + .sanitize_tool_payload(Json::Object(request.headers)) + .as_object() + .cloned() + .unwrap_or_default(); + request.content = trajectory.sanitize_provider_payload(request.content); + return Ok(Some(request)); + } + request.headers = backend.sanitize_request_headers(request.headers); + if backend.target_paths.is_empty() { + request.content = backend.sanitize_json_preorder_dfs(request.content); + return Ok(Some(request)); + } + let resolved = context.resolve_codec(); + let fallback = if resolved.is_none() { + backend + .selected_surface(context.codec()) + .map(build_request_codec) + } else { + None + }; + let Some(codec) = resolved.as_deref().or(fallback.as_deref()) else { + log_llm_payload_omitted("request", context.codec(), "no usable request codec"); + return Ok(None); + }; + let sanitized = backend.sanitize_request_with_codec(codec, &request); + if sanitized.is_none() { + log_llm_payload_omitted( + "request", + context.codec(), + "codec decode, sanitize, or encode failure", + ); + } + Ok(sanitized) + }) }) } @@ -569,49 +578,52 @@ pub(super) fn llm_sanitize_response_callback( backend: CompiledBuiltinBackend, ) -> LlmSanitizeResponseFn { Arc::new(move |payload: Json, context| { - if let Some(trajectory) = backend.trajectory.as_ref() { - return Some(trajectory.sanitize_provider_payload(payload)); - } - if backend.target_paths.is_empty() { - return Some(backend.sanitize_json_preorder_dfs(payload)); - } - if matches!(context.codec(), LlmCodecIdentity::None) - && !backend.uses_compatible_legacy_response_codec(&payload) - { - log_llm_payload_omitted( - "response", - context.codec(), - "no active response codec or compatible legacy codec", - ); - return None; - } - let Some(surface) = backend.selected_surface(context.codec()) else { - log_llm_payload_omitted( - "response", - context.codec(), - "no recognized response codec surface", - ); - return None; - }; - let resolved = context.resolve_codec(); - let fallback = if resolved.is_none() { - Some(build_response_codec(surface)) - } else { - None - }; - let Some(codec) = resolved.as_deref().or(fallback.as_deref()) else { - log_llm_payload_omitted("response", context.codec(), "no usable response codec"); - return None; - }; - let sanitized = backend.sanitize_response_with_codec(codec, surface, payload); - if sanitized.is_none() { - log_llm_payload_omitted( - "response", - context.codec(), - "codec decode, sanitize, or encode failure", - ); - } - sanitized + let backend = backend.clone(); + Box::pin(async move { + if let Some(trajectory) = backend.trajectory.as_ref() { + return Ok(Some(trajectory.sanitize_provider_payload(payload))); + } + if backend.target_paths.is_empty() { + return Ok(Some(backend.sanitize_json_preorder_dfs(payload))); + } + if matches!(context.codec(), LlmCodecIdentity::None) + && !backend.uses_compatible_legacy_response_codec(&payload) + { + log_llm_payload_omitted( + "response", + context.codec(), + "no active response codec or compatible legacy codec", + ); + return Ok(None); + } + let Some(surface) = backend.selected_surface(context.codec()) else { + log_llm_payload_omitted( + "response", + context.codec(), + "no recognized response codec surface", + ); + return Ok(None); + }; + let resolved = context.resolve_codec(); + let fallback = if resolved.is_none() { + Some(build_response_codec(surface)) + } else { + None + }; + let Some(codec) = resolved.as_deref().or(fallback.as_deref()) else { + log_llm_payload_omitted("response", context.codec(), "no usable response codec"); + return Ok(None); + }; + let sanitized = backend.sanitize_response_with_codec(codec, surface, payload); + if sanitized.is_none() { + log_llm_payload_omitted( + "response", + context.codec(), + "codec decode, sanitize, or encode failure", + ); + } + Ok(sanitized) + }) }) } diff --git a/crates/pii-redaction/tests/unit/component_tests.rs b/crates/pii-redaction/tests/unit/component_tests.rs index d4fe4e8be..2152c5d6d 100644 --- a/crates/pii-redaction/tests/unit/component_tests.rs +++ b/crates/pii-redaction/tests/unit/component_tests.rs @@ -296,8 +296,8 @@ impl LlmCodec for IdentifiedRequestCodec { } } -#[test] -fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs() { +#[tokio::test] +async fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".to_string(), @@ -324,7 +324,9 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( }, LlmSanitizeRequestContext::for_request_codec(Some(Arc::new(OpenAIResponsesCodec))), ) - .expect("the active OpenAI Responses codec must override the legacy fallback"); + .await + .expect("the active OpenAI Responses codec must override the legacy fallback") + .expect("the active OpenAI Responses codec must retain the payload"); assert_eq!( active_request.content["input"][0]["content"][0]["text"], json!("[REDACTED]") @@ -350,7 +352,9 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( inner: OpenAIResponsesCodec, }))), ) - .expect("an active runtime or opaque request codec must remain usable"); + .await + .expect("an active runtime or opaque request codec must remain usable") + .expect("an active runtime or opaque request codec must retain the payload"); assert_eq!( active_request.content["input"][0]["content"][0]["text"], json!("[REDACTED]") @@ -373,14 +377,19 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( BuiltinLlmCodec::OpenAiResponses, )), ) - .expect("the active OpenAI Responses codec must override the legacy fallback"); + .await + .expect("the active OpenAI Responses codec must override the legacy fallback") + .expect("the active OpenAI Responses codec must retain the payload"); assert_eq!( active_responses["output"][0]["content"][0]["text"], json!("[REDACTED]") ); assert!( - sanitize_response(responses_payload.clone(), no_codec_context()).is_none(), + sanitize_response(responses_payload.clone(), no_codec_context()) + .await + .expect("sanitizer callback must succeed") + .is_none(), "an incompatible configured fallback codec must omit a normalized payload" ); @@ -389,6 +398,8 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( responses_payload, LlmSanitizeResponseContext::with_identity(LlmCodecIdentity::Opaque), ) + .await + .expect("sanitizer callback must succeed") .is_none(), "a normalized-path policy must omit an unknown active provider payload" ); @@ -403,13 +414,15 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( "com.example.chat.v1".to_owned(), )), ) + .await + .expect("sanitizer callback must succeed") .is_none(), "a normalized-path policy must omit a runtime codec until it has a compatible projection" ); } -#[test] -fn normalized_llm_paths_omit_payloads_when_legacy_codec_decode_fails() { +#[tokio::test] +async fn normalized_llm_paths_omit_payloads_when_legacy_codec_decode_fails() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".to_string(), @@ -432,17 +445,22 @@ fn normalized_llm_paths_omit_payloads_when_legacy_codec_decode_fails() { }, no_codec_request_context(), ) + .await + .expect("sanitizer callback must succeed") .is_none(), "a shallow legacy surface match must not enable a raw-payload fallback" ); assert!( - sanitize_response(json!({"choices": "sk-response-secret"}), no_codec_context()).is_none(), + sanitize_response(json!({"choices": "sk-response-secret"}), no_codec_context()) + .await + .expect("sanitizer callback must succeed") + .is_none(), "a legacy response codec failure must omit the payload instead of emitting raw content" ); } -#[test] -fn normalized_openai_chat_api_specific_policy_omits_multiple_choices() { +#[tokio::test] +async fn normalized_openai_chat_api_specific_policy_omits_multiple_choices() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "remove".to_string(), @@ -475,12 +493,14 @@ fn normalized_openai_chat_api_specific_policy_omits_multiple_choices() { BuiltinLlmCodec::OpenAiChat, )), ) + .await + .expect("sanitizer callback must succeed") .is_none() ); } -#[test] -fn normalized_llm_paths_use_configured_anthropic_codec_without_a_system_message() { +#[tokio::test] +async fn normalized_llm_paths_use_configured_anthropic_codec_without_a_system_message() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".to_string(), @@ -504,7 +524,9 @@ fn normalized_llm_paths_use_configured_anthropic_codec_without_a_system_message( }, no_codec_request_context(), ) - .expect("the configured Anthropic codec must sanitize a valid message-only request"); + .await + .expect("the configured Anthropic codec must sanitize a valid message-only request") + .expect("the configured Anthropic codec must retain the payload"); assert_eq!( sanitized.content["messages"][0]["content"], @@ -512,8 +534,8 @@ fn normalized_llm_paths_use_configured_anthropic_codec_without_a_system_message( ); } -#[test] -fn trajectory_preset_redacts_chat_content_without_erasing_request_structure() { +#[tokio::test] +async fn trajectory_preset_redacts_chat_content_without_erasing_request_structure() { let callback = crate::builtin::llm_sanitize_request_callback(trajectory_backend( Some("openai_chat"), "preserve", @@ -548,6 +570,8 @@ fn trajectory_preset_redacts_chat_content_without_erasing_request_structure() { "person_name": "Alice Example" }), }, no_codec_request_context()) + .await + .unwrap() .unwrap(); assert_eq!(request.content["model"], "claude-sonnet-4-6"); @@ -595,8 +619,8 @@ fn trajectory_preset_redacts_chat_content_without_erasing_request_structure() { ); } -#[test] -fn trajectory_preset_preserves_response_analytics_and_redacts_response_content() { +#[tokio::test] +async fn trajectory_preset_preserves_response_analytics_and_redacts_response_content() { let callback = crate::builtin::llm_sanitize_response_callback(trajectory_backend( Some("openai_chat"), "preserve", @@ -617,6 +641,8 @@ fn trajectory_preset_preserves_response_analytics_and_redacts_response_content() }), no_codec_context(), ) + .await + .unwrap() .unwrap(); assert_eq!(sanitized["id"], "chatcmpl_1"); @@ -643,8 +669,8 @@ fn trajectory_preset_preserves_response_analytics_and_redacts_response_content() assert_eq!(sanitized["cost"]["total"], 1.25); } -#[test] -fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { +#[tokio::test] +async fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { let responses_request = crate::builtin::llm_sanitize_request_callback(trajectory_backend( Some("openai_responses"), "preserve", @@ -657,6 +683,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { "max_output_tokens": 100 }), }, no_codec_request_context()) + .await + .unwrap() .unwrap(); assert_eq!(responses_request.content["model"], "gpt-5"); assert_eq!(responses_request.content["input"][0]["role"], "user"); @@ -681,6 +709,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { }), no_codec_context(), ) + .await + .unwrap() .unwrap(); assert_eq!(responses_response["id"], "resp_1"); assert_eq!(responses_response["status"], "completed"); @@ -706,6 +736,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { "max_tokens": 128 }), }, no_codec_request_context()) + .await + .unwrap() .unwrap(); assert_eq!(anthropic_request.content["model"], "claude-sonnet-4-6"); assert_eq!(anthropic_request.content["system"], "[REDACTED]"); @@ -735,6 +767,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { "stop_reason": "end_turn", "usage": {"input_tokens": 12, "output_tokens": 6, "cache_read_input_tokens": 8} }), no_codec_context()) + .await + .unwrap() .unwrap(); assert_eq!(anthropic_response["id"], "msg_1"); assert_eq!(anthropic_response["role"], "assistant"); @@ -745,8 +779,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { assert_eq!(anthropic_response["usage"]["cache_read_input_tokens"], 8); } -#[test] -fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { +#[tokio::test] +async fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { let callback = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); let chunk = Event::Mark(MarkEvent::new( BaseEvent::builder().name("llm.chunk").build(), @@ -754,7 +788,7 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { Some(CategoryProfile::builder().subtype("llm.chunk").build()), )); let sanitized = callback( - &chunk, + chunk.clone(), EventSanitizeFields { data: Some(json!({ "chunk_index": 2, @@ -764,7 +798,9 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { category_profile: chunk.category_profile().cloned(), metadata: None, }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.data.as_ref().unwrap()["chunk_index"], 2); assert_eq!( sanitized.data.as_ref().unwrap()["event_type"], @@ -787,7 +823,7 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { ), )); let sanitized = callback( - &optimization, + optimization.clone(), EventSanitizeFields { data: Some(json!({ "producer": "neutral.router", @@ -803,7 +839,9 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { category_profile: optimization.category_profile().cloned(), metadata: None, }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.data.as_ref().unwrap()["producer"], "neutral.router" @@ -829,7 +867,7 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { None, )); let sanitized = callback( - &nested_agent, + nested_agent, EventSanitizeFields { data: Some(json!({ "request_id": "request-1", @@ -839,7 +877,9 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { category_profile: None, metadata: Some(json!({"parent_scope_id": "scope-1", "note": "private note"})), }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.data.as_ref().unwrap()["request_id"], "request-1"); assert_eq!( sanitized.data.as_ref().unwrap()["instruction"], @@ -860,8 +900,8 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { assert_eq!(sanitized.metadata.as_ref().unwrap()["note"], "[REDACTED]"); } -#[test] -fn trajectory_preset_preserves_trusted_scope_metadata_only() { +#[tokio::test] +async fn trajectory_preset_preserves_trusted_scope_metadata_only() { let callback = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); let metadata = json!({ "nemo_relay_scope_role": "turn", @@ -922,13 +962,15 @@ fn trajectory_preset_preserves_trusted_scope_metadata_only() { None, )); let sanitized = callback( - &event, + event, EventSanitizeFields { data: None, category_profile: None, metadata: Some(metadata.clone()), }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.metadata, Some(expected_metadata.clone())); } @@ -940,7 +982,7 @@ fn trajectory_preset_preserves_trusted_scope_metadata_only() { None, )); let sanitized = callback( - &malformed, + malformed, EventSanitizeFields { data: None, category_profile: None, @@ -951,7 +993,9 @@ fn trajectory_preset_preserves_trusted_scope_metadata_only() { "provider_payload_exact": "private context" })), }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.metadata, Some(json!({ @@ -968,21 +1012,23 @@ fn trajectory_preset_preserves_trusted_scope_metadata_only() { Some(CategoryProfile::builder().subtype("llm.chunk").build()), )); let sanitized = callback( - &mark, + mark.clone(), EventSanitizeFields { data: None, category_profile: mark.category_profile().cloned(), metadata: Some(json!({"harness": "codex", "source": "hook"})), }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.metadata, Some(json!({"harness": "[REDACTED]", "source": "[REDACTED]"})) ); } -#[test] -fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { +#[tokio::test] +async fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { let event = Event::Mark(MarkEvent::new( BaseEvent::builder().name("neutral.plugin.evidence").build(), Some(EventCategory::custom()), @@ -999,11 +1045,14 @@ fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { }; let preserve = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); - assert_eq!(preserve(&event, fields.clone()), fields); + assert_eq!( + preserve(event.clone(), fields.clone()).await.unwrap(), + fields + ); let redact = crate::builtin::event_sanitize_callback(trajectory_backend(None, "redact_all_leaves")); - let sanitized = redact(&event, fields); + let sanitized = redact(event, fields).await.unwrap(); assert_eq!( sanitized.data.unwrap(), json!({ @@ -1016,8 +1065,8 @@ fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { assert_eq!(profile.extra["opaque"]["label"], "[REDACTED]"); } -#[test] -fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations() { +#[tokio::test] +async fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations() { let callback = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); let annotated_response: nemo_relay::codec::response::AnnotatedLlmResponse = serde_json::from_value(json!({ @@ -1068,7 +1117,7 @@ fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations None, )); let sanitized = callback( - &event, + event, EventSanitizeFields { data: Some(json!({"already": "sanitized by the response callback"})), category_profile: Some( @@ -1079,7 +1128,9 @@ fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations ), metadata: None, }, - ); + ) + .await + .unwrap(); let profile = sanitized.category_profile.unwrap(); assert_eq!(profile.model_name.as_deref(), Some("claude-sonnet-4-6")); @@ -1116,8 +1167,8 @@ fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations ); } -#[test] -fn preserved_custom_marks_remain_eligible_for_a_later_email_profile() { +#[tokio::test] +async fn preserved_custom_marks_remain_eligible_for_a_later_email_profile() { let event = Event::Mark(MarkEvent::new( BaseEvent::builder().name("neutral.plugin.evidence").build(), Some(EventCategory::custom()), @@ -1141,7 +1192,8 @@ fn preserved_custom_marks_remain_eligible_for_a_later_email_profile() { .unwrap(), ); - let sanitized = email(&event, trajectory(&event, fields)); + let fields = trajectory(event.clone(), fields).await.unwrap(); + let sanitized = email(event, fields).await.unwrap(); assert_eq!(sanitized.data.as_ref().unwrap()["owner"], "[REDACTED]"); assert_eq!(sanitized.data.as_ref().unwrap()["score"], 0.9); assert_eq!( @@ -1365,7 +1417,11 @@ fn local_profile_registrations_receive_generated_namespaces() { reset_runtime(); register_local_backend_provider(Arc::new(|_, ctx| { - ctx.register_mark_sanitize_guardrail("shared", 100, Arc::new(|_, fields| fields)) + ctx.register_mark_sanitize_guardrail( + "shared", + 100, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) })) .unwrap(); @@ -1430,8 +1486,8 @@ fn failed_later_profile_rolls_back_earlier_profile_registrations() { deregister_subscriber("pii-profile-rollback").unwrap(); } -#[test] -fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() { +#[tokio::test] +async fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".into(), @@ -1449,7 +1505,7 @@ fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() None, )); let sanitized = callback( - &event, + event, EventSanitizeFields { data: Some(json!({"email": "person@example.com"})), category_profile: Some( @@ -1459,7 +1515,9 @@ fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() ), metadata: Some(json!({"owner": "person@example.com"})), }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.data.unwrap()["email"], "[REDACTED]"); assert_eq!( sanitized.category_profile.unwrap().subtype.as_deref(), @@ -1468,8 +1526,8 @@ fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() assert_eq!(sanitized.metadata.unwrap()["owner"], "[REDACTED]"); } -#[test] -fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() { +#[tokio::test] +async fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "redact".into(), @@ -1493,13 +1551,15 @@ fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() .subtype("person@example.com") .build(); let sanitized = callback( - &event, + event, EventSanitizeFields { data: Some(json!({"content": "person@example.com"})), category_profile: Some(original_profile.clone()), metadata: Some(json!({"owner": "person@example.com"})), }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.data.unwrap()["content"], @@ -1511,8 +1571,8 @@ fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() } } -#[test] -fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { +#[tokio::test] +async fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "redact".into(), @@ -1547,13 +1607,15 @@ fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { .subtype("person@example.com") .build(); let sanitized = callback( - &event, + event, EventSanitizeFields { data: Some(json!({"content": "person@example.com"})), category_profile: Some(original_profile.clone()), metadata: Some(json!({"owner": "person@example.com"})), }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.data.unwrap()["content"], "person@example.com"); assert_eq!(sanitized.category_profile.unwrap(), original_profile); @@ -1562,8 +1624,8 @@ fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { } } -#[test] -fn event_sanitizer_discards_category_profile_when_sanitization_fails() { +#[tokio::test] +async fn event_sanitizer_discards_category_profile_when_sanitization_fails() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".into(), @@ -1581,7 +1643,7 @@ fn event_sanitizer_discards_category_profile_when_sanitization_fails() { None, )); let sanitized = callback( - &event, + event, EventSanitizeFields { data: None, category_profile: Some(CategoryProfile { @@ -1593,7 +1655,9 @@ fn event_sanitizer_discards_category_profile_when_sanitization_fails() { }), metadata: None, }, - ); + ) + .await + .unwrap(); assert!(sanitized.category_profile.is_none()); } diff --git a/crates/plugin/README.md b/crates/plugin/README.md index 7f3dd6c42..0b68a888e 100644 --- a/crates/plugin/README.md +++ b/crates/plugin/README.md @@ -42,8 +42,12 @@ the dynamic-library boundary on the stable C-compatible ABI. - **`PluginContext`**: Component-scoped registration APIs for middleware and subscribers. - **`PluginRuntime`**: Typed helpers for Relay-owned scopes and marks. -- **Stable native ABI v1**: C-compatible host and plugin tables behind the - safe Rust authoring interface. +- **Stable native ABI v3**: C-compatible host and plugin tables behind the + safe Rust authoring interface, with a v2-compatible prefix for existing + plugins. +- **Raw async middleware**: Completion-based raw registrations for plugins + that need asynchronous guardrails, intercepts, or event sanitizers. Typed + Rust callbacks remain synchronous convenience APIs. ## Installation diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index 727e09aa2..5f5500072 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -36,7 +36,15 @@ use serde::{Serialize, de::DeserializeOwned}; use serde_json::Map; /// Native plugin ABI version supported by this crate. -pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 2; +/// +/// Version 3 reserves the native async middleware extension. Hosts retain a +/// version-2 table for already-built plugins during entry-point negotiation. +pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 3; +/// ABI version that introduced completion-based asynchronous middleware. +pub const NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE: u32 = 3; + +/// Legacy native plugin ABI accepted by Relay hosts for compatibility. +pub const NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY: u32 = 2; /// Built-in LLM codec identities available to native plugins. #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, serde::Deserialize)] @@ -754,6 +762,133 @@ pub struct NemoRelayNativeHostApiV1 { ) -> NemoRelayStatus, } +/// Middleware surface selected by the native async registration hook. +/// +/// The host only exposes this through the ABI-v3 extension table. It keeps +/// every asynchronous callback shape uniform while allowing the host to +/// deserialize the surface-specific invocation and result payloads. +#[repr(u32)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum NemoRelayNativeAsyncMiddlewareKind { + /// Tool start-event request sanitizer. + ToolSanitizeRequest = 0, + /// Tool end-event response sanitizer. + ToolSanitizeResponse = 1, + /// Tool execution admission guardrail. + ToolConditionalExecution = 2, + /// Tool request rewrite intercept. + ToolRequestIntercept = 3, + /// Tool execution intercept with a continuation. + ToolExecutionIntercept = 4, + /// LLM start-event request sanitizer. + LlmSanitizeRequest = 5, + /// LLM end-event response sanitizer. + LlmSanitizeResponse = 6, + /// LLM execution admission guardrail. + LlmConditionalExecution = 7, + /// LLM request rewrite intercept. + LlmRequestIntercept = 8, + /// LLM execution intercept with a continuation. + LlmExecutionIntercept = 9, + /// Streaming LLM execution intercept with a continuation. + LlmStreamExecutionIntercept = 10, + /// Mark event sanitizer. + MarkSanitize = 11, + /// Scope-start event sanitizer. + ScopeSanitizeStart = 12, + /// Scope-end event sanitizer. + ScopeSanitizeEnd = 13, +} + +/// Indicates whether an asynchronous native callback settled before returning. +#[repr(u32)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum NemoRelayNativeAsyncCallbackState { + /// The callback settled its completion before returning. + Complete = 0, + /// The callback retained its completion for later settlement. + Pending = 1, +} + +/// Opaque one-shot completion retained by a pending native callback. +#[repr(C)] +pub struct NemoRelayNativeAsyncCompletion { + _private: [u8; 0], + _marker: PhantomData<(*mut u8, PhantomPinned)>, +} + +/// Opaque native execution continuation supplied only to execution intercepts. +#[repr(C)] +pub struct NemoRelayNativeAsyncNext { + _private: [u8; 0], + _marker: PhantomData<(*mut u8, PhantomPinned)>, +} + +/// Completion-based native middleware callback. +/// +/// `invocation_json` is borrowed for the call. A callback that returns +/// [`NemoRelayNativeAsyncCallbackState::Pending`] owns one completion +/// reference and must settle it then call the v3 `async_completion_release` +/// hook. When `next` is non-null, the callback owns that handle for the +/// invocation and must call `async_next_release` after its final use. `next` +/// is null for non-execution middleware. +pub type NemoRelayNativeAsyncMiddlewareCb = + unsafe extern "C" fn( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, + ) -> NemoRelayNativeAsyncCallbackState; + +/// ABI-v3 host extension appended to [`NemoRelayNativeHostApiV1`]. +/// +/// Its first field is the complete v1/v2 table, so legacy plugins can keep +/// treating the pointer as a [`NemoRelayNativeHostApiV1`]. +#[repr(C)] +#[derive(Clone, Copy)] +pub struct NemoRelayNativeHostApiV3 { + /// Compatibility prefix for ABI-v1/v2 plugins. + pub v1: NemoRelayNativeHostApiV1, + /// Resolves an async callback completion with a JSON value. + pub async_completion_resolve_json: unsafe extern "C" fn( + completion: *const NemoRelayNativeAsyncCompletion, + value_json: *const NemoRelayNativeString, + ) -> NemoRelayStatus, + /// Rejects an async callback completion with a UTF-8 message. + pub async_completion_reject: unsafe extern "C" fn( + completion: *const NemoRelayNativeAsyncCompletion, + message: *const NemoRelayNativeString, + ) -> NemoRelayStatus, + /// Returns true after the awaiting runtime has cancelled the invocation. + pub async_completion_is_cancelled: + unsafe extern "C" fn(completion: *const NemoRelayNativeAsyncCompletion) -> bool, + /// Releases the callback-owned reference after a pending completion settles. + pub async_completion_release: + unsafe extern "C" fn(completion: *const NemoRelayNativeAsyncCompletion), + /// Invokes an execution continuation and settles a supplied completion. + pub async_next_invoke: unsafe extern "C" fn( + next: *const NemoRelayNativeAsyncNext, + invocation_json: *const NemoRelayNativeString, + completion: *const NemoRelayNativeAsyncCompletion, + ) -> NemoRelayStatus, + /// Releases the callback-owned continuation reference for a pending callback. + pub async_next_release: unsafe extern "C" fn(next: *const NemoRelayNativeAsyncNext), + /// Registers any completion-based asynchronous middleware surface. + pub plugin_context_register_async_middleware: unsafe extern "C" fn( + ctx: *mut NemoRelayNativePluginContext, + kind: NemoRelayNativeAsyncMiddlewareKind, + name: *const NemoRelayNativeString, + priority: i32, + break_chain: bool, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, + ) -> NemoRelayStatus, +} + +unsafe impl Send for NemoRelayNativeHostApiV3 {} +unsafe impl Sync for NemoRelayNativeHostApiV3 {} + // The host API table is immutable after construction. Function pointers and // the null-terminated version string pointer are safe to share across threads. unsafe impl Send for NemoRelayNativeHostApiV1 {} @@ -2231,6 +2366,47 @@ impl<'a> PluginContext<'a> { }) } + /// Registers completion-based asynchronous middleware through the ABI-v3 + /// extension table. + /// + /// Plugins built against older hosts receive [`NemoRelayStatus::InvalidArg`] + /// instead of attempting to read beyond the legacy host table. + /// + /// # Safety + /// `cb`, `user_data`, and `free_fn` must remain valid until the host + /// deregisters the callback or invokes `free_fn`. A callback returning + /// `Pending` must settle and release its completion/next references. + #[allow(clippy::too_many_arguments)] // Mirrors the native C ABI registration callback. + pub unsafe fn register_async_middleware_raw( + &mut self, + kind: NemoRelayNativeAsyncMiddlewareKind, + name: &str, + priority: i32, + break_chain: bool, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, + ) -> NemoRelayStatus { + if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE + || self.host.struct_size < std::mem::size_of::() + { + return NemoRelayStatus::InvalidArg; + } + let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV3) }; + self.with_name(name, |_, name| unsafe { + (host.plugin_context_register_async_middleware)( + self.raw, + kind, + name, + priority, + break_chain, + cb, + user_data, + free_fn, + ) + }) + } + fn with_name( &self, name: &str, diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index fb3562bd3..865f05e16 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -300,8 +300,8 @@ static LLM_REQUEST_INTERCEPT_REGISTRATION: Mutex(), test_host().struct_size diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index f5c2ccd51..80b9d672e 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1313,12 +1313,42 @@ fn deregister_llm_stream_execution_intercept(name: &str) -> PyResult { #[pyfunction] fn tool_request_intercepts<'py>( py: Python<'py>, - name: &str, + name: String, args: &Bound<'py, PyAny>, -) -> PyResult> { +) -> PyResult> { let args_json = py_to_json(args)?; - let result = core_tool_api::tool_request_intercepts(name, args_json).map_err(to_py_err)?; - json_to_py(py, &result) + // Preserve the established synchronous helper behavior when no Python + // event loop is active. Awaitable middleware is supported from async + // callers below; a synchronous caller can continue using direct + // callbacks without manufacturing an asyncio loop. + if py + .import("asyncio")? + .call_method0("get_running_loop") + .is_err() + { + let scope_stack = current_scope_stack_handle(); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; + let result = runtime + .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { + core_tool_api::tool_request_intercepts(&name, args_json).await + })) + .map_err(to_py_err)?; + return json_to_py(py, &result).map(|value| value.into_bound(py)); + } + let scope_stack = current_scope_stack_handle(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + TASK_SCOPE_STACK + .scope(scope_stack, async move { + let result = core_tool_api::tool_request_intercepts(&name, args_json) + .await + .map_err(to_py_err)?; + Python::attach(|py| json_to_py(py, &result)) + }) + .await + }) } /// Run the registered tool conditional execution guardrail chain. @@ -1329,9 +1359,39 @@ fn tool_request_intercepts<'py>( /// name: Tool name. /// args: Tool arguments (any JSON-serializable object). #[pyfunction] -fn tool_conditional_execution(name: &str, args: &Bound<'_, PyAny>) -> PyResult<()> { +fn tool_conditional_execution<'py>( + py: Python<'py>, + name: String, + args: &Bound<'py, PyAny>, +) -> PyResult> { let args_json = py_to_json(args)?; - core_tool_api::tool_conditional_execution(name, &args_json).map_err(to_py_err) + if py + .import("asyncio")? + .call_method0("get_running_loop") + .is_err() + { + let scope_stack = current_scope_stack_handle(); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; + runtime + .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { + core_tool_api::tool_conditional_execution(&name, &args_json).await + })) + .map_err(to_py_err)?; + return Ok(py.None().into_bound(py)); + } + let scope_stack = current_scope_stack_handle(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + TASK_SCOPE_STACK + .scope(scope_stack, async move { + core_tool_api::tool_conditional_execution(&name, &args_json) + .await + .map_err(to_py_err) + }) + .await + }) } /// Run the registered LLM request intercept chain on the given request. @@ -1344,12 +1404,43 @@ fn tool_conditional_execution(name: &str, args: &Bound<'_, PyAny>) -> PyResult<( /// Returns: /// The (possibly transformed) ``LlmRequest``. #[pyfunction] -fn llm_request_intercepts( - name: &str, +fn llm_request_intercepts<'py>( + py: Python<'py>, + name: String, request: PyLLMRequest, -) -> PyResult { - let result = core_llm_api::llm_request_intercepts(name, request.inner).map_err(to_py_err)?; - Ok(crate::py_types::PyLLMRequestInterceptOutcome { inner: result }) +) -> PyResult> { + if py + .import("asyncio")? + .call_method0("get_running_loop") + .is_err() + { + let scope_stack = current_scope_stack_handle(); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; + let result = runtime + .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { + core_llm_api::llm_request_intercepts(&name, request.inner).await + })) + .map_err(to_py_err)?; + return Py::new( + py, + crate::py_types::PyLLMRequestInterceptOutcome { inner: result }, + ) + .map(|value| value.into_bound(py).into_any()); + } + let scope_stack = current_scope_stack_handle(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + TASK_SCOPE_STACK + .scope(scope_stack, async move { + let result = core_llm_api::llm_request_intercepts(&name, request.inner) + .await + .map_err(to_py_err)?; + Ok(crate::py_types::PyLLMRequestInterceptOutcome { inner: result }) + }) + .await + }) } /// Run the registered LLM conditional execution guardrail chain. @@ -1359,8 +1450,37 @@ fn llm_request_intercepts( /// Args: /// request: An ``LlmRequest`` object. #[pyfunction] -fn llm_conditional_execution(request: PyLLMRequest) -> PyResult<()> { - core_llm_api::llm_conditional_execution(&request.inner).map_err(to_py_err) +fn llm_conditional_execution<'py>( + py: Python<'py>, + request: PyLLMRequest, +) -> PyResult> { + if py + .import("asyncio")? + .call_method0("get_running_loop") + .is_err() + { + let scope_stack = current_scope_stack_handle(); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; + runtime + .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { + core_llm_api::llm_conditional_execution(&request.inner).await + })) + .map_err(to_py_err)?; + return Ok(py.None().into_bound(py)); + } + let scope_stack = current_scope_stack_handle(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + TASK_SCOPE_STACK + .scope(scope_stack, async move { + core_llm_api::llm_conditional_execution(&request.inner) + .await + .map_err(to_py_err) + }) + .await + }) } // --------------------------------------------------------------------------- diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 57fffd5d5..f5cb02e55 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -39,7 +39,7 @@ use tokio_stream::Stream; use tokio_stream::wrappers::ReceiverStream; use nemo_relay::api::event::{Event, EventSanitizeFields}; -use nemo_relay::api::llm::{LlmRequest, LlmRequestInterceptOutcome}; +use nemo_relay::api::llm::LlmRequest; use nemo_relay::api::tool::ToolExecutionInterceptOutcome; use nemo_relay::codec::request::AnnotatedLlmRequest as AnnotatedLLMRequest; use nemo_relay::codec::response::AnnotatedLlmResponse as AnnotatedLLMResponse; @@ -408,71 +408,69 @@ fn stream_from_async_iter(async_iter: Py) -> FlowResult { /// Wrap a Python callable `(str, Json) -> Json` for tool sanitize/intercept fns. pub fn wrap_py_tool_fn(py_fn: Py) -> ToolSanitizeFn { - Arc::new(move |name: &str, args: Json| { - Python::attach(|py| { - let py_args = match json_to_py(py, &args) { - Ok(v) => v, - Err(e) => { - eprintln!("nemo_relay: json_to_py failed in tool fn for '{name}': {e}"); - return args.clone(); - } - }; - let result = match py_fn.call1(py, (name, py_args)) { - Ok(v) => v, - Err(e) => { - eprintln!("nemo_relay: Python tool callable failed for '{name}': {e}"); - return args.clone(); - } - }; - py_to_json(result.bind(py)).unwrap_or_else(|e| { - eprintln!("nemo_relay: py_to_json failed in tool fn for '{name}': {e}"); - args.clone() - }) + let py_fn = Arc::new(py_fn); + Arc::new(move |name: String, args: Json| { + let py_fn = py_fn.clone(); + Box::pin(async move { + resolve_json_or_future(Python::attach(|py| { + let py_args = json_to_py(py, &args) + .map_err(|e| FlowError::Internal(format!("tool json_to_py failed: {e}")))?; + let result = py_fn.call1(py, (name, py_args)).map_err(|e| { + FlowError::Internal(format!("Python tool callback failed: {e}")) + })?; + split_json_or_future(py, result) + })) + .await }) }) } /// Wrap a Python callable `(str, Json) -> Optional[str]` for tool conditional guardrails. pub fn wrap_py_tool_conditional_fn(py_fn: Py) -> ToolConditionalFn { - Arc::new(move |name: &str, args: &Json| { - Python::attach(|py| { - let py_args = json_to_py(py, args).map_err(|e| { - FlowError::Internal(format!( - "tool conditional json_to_py failed for '{name}': {e}" - )) - })?; - let result = py_fn.call1(py, (name, py_args)).map_err(|e| { - FlowError::Internal(format!( - "Python tool conditional callable failed for '{name}': {e}" - )) - })?; - let bound = result.bind(py); - if bound.is_none() { - Ok(None) - } else { - bound.extract::().map(Some).map_err(|e| { - FlowError::Internal(format!( - "tool conditional guardrail for '{name}' returned unexpected type (expected str or None): {e}" - )) - }) - } + let py_fn = Arc::new(py_fn); + Arc::new(move |name: String, args: Json| { + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let py_args = + json_to_py(py, &args).map_err(|e| FlowError::Internal(e.to_string()))?; + let result = py_fn + .call1(py, (name, py_args)) + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + let bound = result.bind(py); + if bound.is_none() { + Ok(None) + } else { + bound.extract::().map(Some).map_err(|e| { + FlowError::Internal(format!( + "tool conditional guardrail returned unexpected type: {e}" + )) + }) + } + }) }) }) } /// Wrap a Python callable `(str, Json) -> Json` for tool request intercepts. pub fn wrap_py_tool_request_intercept_fn(py_fn: Py) -> ToolInterceptFn { - Arc::new(move |name: &str, args: Json| { - Python::attach(|py| { - let py_args = json_to_py(py, &args).map_err(|e| { - FlowError::Internal(format!("tool callback json_to_py failed for '{name}': {e}")) - })?; - let result = py_fn.call1(py, (name, py_args)).map_err(|e| { - FlowError::Internal(format!("Python tool callable failed for '{name}': {e}")) - })?; - py_to_json(result.bind(py)).map_err(|e| { - FlowError::Internal(format!("tool callback py_to_json failed for '{name}': {e}")) - }) + let py_fn = Arc::new(py_fn); + Arc::new(move |name: String, args: Json| { + let py_fn = py_fn.clone(); + Box::pin(async move { + resolve_json_or_future(Python::attach(|py| { + let py_args = + json_to_py(py, &args).map_err(|e| FlowError::Internal(e.to_string()))?; + let result = py_fn + .call1(py, (name, py_args)) + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_json_or_future(py, result) + })) + .await }) }) } @@ -815,30 +813,38 @@ pub fn wrap_py_llm_stream_exec_intercept_fn( /// Wrap a Python callable `(LlmRequest, LlmSanitizeRequestContext) -> Optional`. fn wrap_py_llm_sanitize_request_callback(py_fn: Py) -> LlmSanitizeRequestFn { + let py_fn = Arc::new(py_fn); Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { - Python::attach(|py| { - let py_context = PyLlmSanitizeRequestContext { inner: context }; - let py_request = PyLLMRequest { inner: request }; - let result = match py_fn.call1(py, (py_request, py_context)) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: LLM sanitize request callable failed: {error}"); - return None; - } - }; - if result.is_none(py) { - return None; - } - match result.extract::(py) { - Ok(request) => Some(request.inner), - Err(error) => { - eprintln!( - "nemo_relay: LLM sanitize request callable returned unexpected type: {error}" - ); - None + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let result = py_fn + .call1( + py, + ( + PyLLMRequest { inner: request }, + PyLlmSanitizeRequestContext { inner: context }, + ), + ) + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + if result.is_none(py) { + Ok(None) + } else { + result + .extract::(py) + .map(|request| Some(request.inner)) + .map_err(|error| { + FlowError::Internal(format!( + "LLM sanitize request returned unexpected type: {error}" + )) + }) } - } + }) }) }, ) @@ -846,24 +852,29 @@ fn wrap_py_llm_sanitize_request_callback(py_fn: Py) -> LlmSanitizeRequest /// Wrap a Python callable `(LlmRequest) -> Optional[str]` for LLM conditional guardrails. pub fn wrap_py_llm_conditional_fn(py_fn: Py) -> LlmConditionalFn { - Arc::new(move |request: &LlmRequest| { - Python::attach(|py| { - let py_req = PyLLMRequest { - inner: request.clone(), - }; - let result = py_fn.call1(py, (py_req,)).map_err(|e| { - FlowError::Internal(format!("LLM conditional guardrail callable failed: {e}")) - })?; - let bound = result.bind(py); - if bound.is_none() { - Ok(None) - } else { - bound.extract::().map(Some).map_err(|e| { - FlowError::Internal(format!( - "LLM conditional guardrail returned unexpected type (expected str or None): {e}" - )) - }) - } + let py_fn = Arc::new(py_fn); + Arc::new(move |request: LlmRequest| { + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let result = py_fn + .call1(py, (PyLLMRequest { inner: request },)) + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + let bound = result.bind(py); + if bound.is_none() { + Ok(None) + } else { + bound.extract::().map(Some).map_err(|e| { + FlowError::Internal(format!( + "LLM conditional guardrail returned unexpected type: {e}" + )) + }) + } + }) }) }) } @@ -875,42 +886,45 @@ pub fn wrap_py_llm_conditional_fn(py_fn: Py) -> LlmConditionalFn { /// When ``annotated`` is present, request content is read-only and provider-body /// edits must be made through the returned annotation; headers remain writable. pub fn wrap_py_llm_request_intercept_fn(py_fn: Py) -> LlmRequestInterceptFn { + let py_fn = Arc::new(py_fn); Arc::new( - move |name: &str, - request: LlmRequest, - annotated: Option| - -> FlowResult { - Python::attach(|py| { - let py_req = PyLLMRequest { - inner: request.clone(), - }; - let py_ann: Py = match annotated { - Some(ann) => { - let wrapper = PyAnnotatedLLMRequest { inner: ann }; - wrapper - .into_pyobject(py) - .map_err(|e| { - FlowError::Internal(format!( - "Failed to convert AnnotatedLLMRequest to Python: {e}" - )) - })? - .into_any() - .unbind() - } - None => py.None(), - }; - let result = py_fn.call1(py, (name, py_req, py_ann)).map_err(|e| { - FlowError::Internal(format!("LLM request intercept callable failed: {e}")) - })?; + move |name: String, request: LlmRequest, annotated: Option| { + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let py_req = PyLLMRequest { inner: request }; + let py_ann: Py = match annotated { + Some(ann) => { + let wrapper = PyAnnotatedLLMRequest { inner: ann }; + wrapper + .into_pyobject(py) + .map_err(|e| { + FlowError::Internal(format!( + "Failed to convert AnnotatedLLMRequest to Python: {e}" + )) + })? + .into_any() + .unbind() + } + None => py.None(), + }; + let result = py_fn.call1(py, (name, py_req, py_ann)).map_err(|e| { + FlowError::Internal(format!("LLM request intercept callable failed: {e}")) + })?; - result - .extract::(py) - .map(|value| value.inner) - .map_err(|e| { - FlowError::Internal(format!( - "LLM request intercept must return LLMRequestInterceptOutcome: {e}" - )) - }) + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + result + .extract::(py) + .map(|value| value.inner) + .map_err(|e| { + FlowError::Internal(format!( + "LLM request intercept must return LLMRequestInterceptOutcome: {e}" + )) + }) + }) }) }, ) @@ -1014,33 +1028,29 @@ pub fn wrap_py_finalizer_fn(py_fn: Py) -> Box Json + Send /// Wrap a Python callable `(Json, LlmSanitizeResponseContext) -> Optional[Json]`. fn wrap_py_llm_sanitize_response_callback(py_fn: Py) -> LlmSanitizeResponseFn { + let py_fn = Arc::new(py_fn); Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { - Python::attach(|py| { - let py_context = PyLlmSanitizeResponseContext { inner: context }; - let py_response = match json_to_py(py, &response) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: json_to_py failed in LLM sanitize response: {error}"); - return None; - } - }; - let result = match py_fn.call1(py, (py_response, py_context)) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: LLM sanitize response callable failed: {error}"); - return None; - } - }; - if result.is_none(py) { - return None; - } - match py_to_json(result.bind(py)) { - Ok(response) => Some(response), - Err(error) => { - eprintln!("nemo_relay: py_to_json failed in LLM sanitize response: {error}"); - None + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let py_context = PyLlmSanitizeResponseContext { inner: context }; + let py_response = json_to_py(py, &response) + .map_err(|error| FlowError::Internal(error.to_string()))?; + let result = py_fn + .call1(py, (py_response, py_context)) + .map_err(|error| FlowError::Internal(error.to_string()))?; + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + if result.is_none(py) { + Ok(None) + } else { + py_to_json(result.bind(py)) + .map(Some) + .map_err(|error| FlowError::Internal(error.to_string())) } - } + }) }) }) } @@ -1084,61 +1094,71 @@ pub fn wrap_py_event_subscriber(py_fn: Py) -> EventSubscriberFn { /// Wrap a Python callable ``(Event, EventSanitizeFields) -> EventSanitizeFields``. pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { - Arc::new(move |event: &Event, fields: EventSanitizeFields| { - Python::attach(|py| { - let py_event = match event { - Event::Scope(inner) => Py::new( - py, - crate::py_types::PyScopeEvent { - inner: inner.clone(), - }, - ) - .map(|value| value.into_any()), - Event::Mark(inner) => Py::new( - py, - crate::py_types::PyMarkEvent { - inner: inner.clone(), - }, - ) - .map(|value| value.into_any()), - }; - let py_event = match py_event { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: failed to convert event sanitizer context: {error}"); - return EventSanitizeFields::default(); - } - }; - let fields_json = match serde_json::to_value(&fields) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: failed to serialize event sanitizer fields: {error}"); - return EventSanitizeFields::default(); - } - }; - let py_fields = match json_to_py(py, &fields_json) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: failed to convert event sanitizer fields: {error}"); - return EventSanitizeFields::default(); - } - }; - let result = match py_fn.call1(py, (py_event, py_fields)) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: Python event sanitizer callable failed: {error}"); - return EventSanitizeFields::default(); - } - }; - py_to_json(result.bind(py)) - .ok() - .and_then(|value| serde_json::from_value(value).ok()) - .unwrap_or_else(|| { - eprintln!( - "nemo_relay: event sanitizer must return data, category_profile, and metadata fields" - ); - EventSanitizeFields::default() - }) + let py_fn = Arc::new(py_fn); + Arc::new(move |event: Event, fields: EventSanitizeFields| { + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = Python::attach( + |py| -> FlowResult, PyValueFuture>> { + let py_event = match &event { + Event::Scope(inner) => Py::new( + py, + crate::py_types::PyScopeEvent { + inner: inner.clone(), + }, + ) + .map(|value| value.into_any()), + Event::Mark(inner) => Py::new( + py, + crate::py_types::PyMarkEvent { + inner: inner.clone(), + }, + ) + .map(|value| value.into_any()), + }; + let py_event = match py_event { + Ok(value) => value, + Err(error) => { + eprintln!( + "nemo_relay: failed to convert event sanitizer context: {error}" + ); + return Err(FlowError::Internal(error.to_string())); + } + }; + let fields_json = match serde_json::to_value(&fields) { + Ok(value) => value, + Err(error) => { + eprintln!( + "nemo_relay: failed to serialize event sanitizer fields: {error}" + ); + return Err(FlowError::Internal(error.to_string())); + } + }; + let py_fields = match json_to_py(py, &fields_json) { + Ok(value) => value, + Err(error) => { + eprintln!( + "nemo_relay: failed to convert event sanitizer fields: {error}" + ); + return Err(FlowError::Internal(error.to_string())); + } + }; + let result = py_fn + .call1(py, (py_event, py_fields)) + .map_err(|error| FlowError::Internal(error.to_string()))?; + split_py_object_or_future(py, result) + }, + ); + let result = resolve_py_object_or_future(result).await?; + Python::attach(|py| { + py_to_json(result.bind(py)) + .map_err(|error| FlowError::Internal(error.to_string())) + .and_then(|value| { + serde_json::from_value(value).map_err(|error| { + FlowError::Internal(format!("invalid event sanitizer result: {error}")) + }) + }) + }) }) }) } diff --git a/crates/python/tests/coverage/coverage_tests.rs b/crates/python/tests/coverage/coverage_tests.rs index 67cebd3ba..2fa805abf 100644 --- a/crates/python/tests/coverage/coverage_tests.rs +++ b/crates/python/tests/coverage/coverage_tests.rs @@ -661,51 +661,68 @@ def event_fail(event): "#, ); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); let tool_ok = wrap_py_tool_fn(module.getattr("tool_ok").unwrap().unbind()); assert_eq!( - tool_ok("demo", json!({"x": 1})), + runtime + .block_on(tool_ok("demo".to_string(), json!({"x": 1}))) + .unwrap(), json!({"seen": 1, "name": "demo"}) ); let tool_fail = wrap_py_tool_fn(module.getattr("tool_fail").unwrap().unbind()); - assert_eq!(tool_fail("demo", json!({"x": 1})), json!({"x": 1})); + assert!( + runtime + .block_on(tool_fail("demo".to_string(), json!({"x": 1}))) + .is_err() + ); let tool_cond = wrap_py_tool_conditional_fn(module.getattr("tool_cond_bad").unwrap().unbind()); + let error = runtime + .block_on(tool_cond("demo".to_string(), json!({"x": 1}))) + .unwrap_err(); assert!( - tool_cond("demo", &json!({"x": 1})) - .unwrap_err() - .to_string() - .contains("expected str or None") + error.to_string().contains("unexpected type"), + "unexpected tool conditional error: {error}" ); let request = make_request(); let llm_sanitize = wrap_py_llm_sanitize_request_fn(module.getattr("llm_sanitize_bad").unwrap().unbind()) .unwrap(); - assert_eq!( - llm_sanitize( - request.clone(), - nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ), - None + assert!( + runtime + .block_on(llm_sanitize( + request.clone(), + nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), + )) + .is_err() ); let llm_cond = wrap_py_llm_conditional_fn(module.getattr("llm_cond_bad").unwrap().unbind()); assert!( - llm_cond(&request) + runtime + .block_on(llm_cond(request.clone())) .unwrap_err() .to_string() - .contains("expected str or None") + .contains("unexpected type") ); let llm_cond_none = wrap_py_llm_conditional_fn(module.getattr("llm_cond_none").unwrap().unbind()); - assert_eq!(llm_cond_none(&request).unwrap(), None); + assert_eq!( + runtime.block_on(llm_cond_none(request.clone())).unwrap(), + None + ); let llm_req = wrap_py_llm_request_intercept_fn(module.getattr("llm_req_bad").unwrap().unbind()); assert!( - llm_req("demo", request.clone(), None) + runtime + .block_on(llm_req("demo".to_string(), request.clone(), None)) .unwrap_err() .to_string() .contains("intercept callable failed") @@ -714,21 +731,21 @@ def event_fail(event): let tool_req = wrap_py_tool_request_intercept_fn(module.getattr("tool_fail").unwrap().unbind()); assert!( - tool_req("demo", json!({"x": 1})) - .unwrap_err() - .to_string() - .contains("Python tool callable failed") + runtime + .block_on(tool_req("demo".to_string(), json!({"x": 1}))) + .is_err() ); let llm_resp = wrap_py_llm_sanitize_response_fn(module.getattr("llm_resp_fail").unwrap().unbind()) .unwrap(); - assert_eq!( - llm_resp( - json!({"ok": true}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ), - None + assert!( + runtime + .block_on(llm_resp( + json!({"ok": true}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), + )) + .is_err() ); let mut collector = diff --git a/crates/python/tests/coverage/py_api_coverage_tests.rs b/crates/python/tests/coverage/py_api_coverage_tests.rs index 0bfe38573..c1cc50293 100644 --- a/crates/python/tests/coverage/py_api_coverage_tests.rs +++ b/crates/python/tests/coverage/py_api_coverage_tests.rs @@ -472,18 +472,31 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute ) .unwrap(); - let tool_intercepted = - tool_request_intercepts(py, "demo-tool", &py_dict(py, json!({"value": 1}))).unwrap(); + let tool_intercepted = tool_request_intercepts( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": 1})), + ) + .unwrap(); assert_eq!( - crate::convert::py_to_json(tool_intercepted.bind(py)).unwrap(), + crate::convert::py_to_json(&tool_intercepted).unwrap(), json!({"value": 3}) ); - tool_conditional_execution("demo-tool", &py_dict(py, json!({"value": 1}))).unwrap(); + tool_conditional_execution( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": 1})), + ) + .unwrap(); assert!( - tool_conditional_execution("demo-tool", &py_dict(py, json!({"value": -1}))) - .unwrap_err() - .to_string() - .contains("blocked") + tool_conditional_execution( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": -1})) + ) + .unwrap_err() + .to_string() + .contains("blocked") ); let llm_request = PyLLMRequest { @@ -492,7 +505,10 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute content: json!({"messages": [{"role": "user", "content": "hello"}], "model": "demo-model"}), }, }; - let intercepted_request = llm_request_intercepts("demo-llm", llm_request.clone()).unwrap(); + let intercepted_request = + llm_request_intercepts(py, "demo-llm".to_string(), llm_request.clone()).unwrap(); + let intercepted_request: PyRef<'_, crate::py_types::PyLLMRequestInterceptOutcome> = + intercepted_request.extract().unwrap(); assert_eq!( intercepted_request .inner @@ -501,14 +517,17 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute .get("x-intercepted"), Some(&json!("1")) ); - llm_conditional_execution(llm_request.clone()).unwrap(); + llm_conditional_execution(py, llm_request.clone()).unwrap(); assert!( - llm_conditional_execution(PyLLMRequest { - inner: nemo_relay::api::llm::LlmRequest { - headers: serde_json::Map::new(), - content: json!({"messages": [], "model": "blocked"}), - }, - }) + llm_conditional_execution( + py, + PyLLMRequest { + inner: nemo_relay::api::llm::LlmRequest { + headers: serde_json::Map::new(), + content: json!({"messages": [], "model": "blocked"}), + }, + } + ) .unwrap_err() .to_string() .contains("blocked") diff --git a/crates/python/tests/coverage/py_callable_coverage_tests.rs b/crates/python/tests/coverage/py_callable_coverage_tests.rs index fe815cc1b..89a382f88 100644 --- a/crates/python/tests/coverage/py_callable_coverage_tests.rs +++ b/crates/python/tests/coverage/py_callable_coverage_tests.rs @@ -153,7 +153,13 @@ class RaisingResponseCodec: "model": "codec-model" })) .unwrap(); - let outcome = request_intercept("llm", make_request(), Some(annotated.clone())).unwrap(); + let outcome = runtime + .block_on(request_intercept( + "llm".to_string(), + make_request(), + Some(annotated.clone()), + )) + .unwrap(); assert_eq!( outcome.annotated_request.unwrap().last_user_message(), Some("annotated") @@ -163,7 +169,12 @@ class RaisingResponseCodec: module.getattr("request_bad_annotated").unwrap().unbind(), ); assert!( - bad_request_intercept("llm", make_request(), Some(annotated)) + runtime + .block_on(bad_request_intercept( + "llm".to_string(), + make_request(), + Some(annotated), + )) .unwrap_err() .to_string() .contains("must return LLMRequestInterceptOutcome") @@ -173,7 +184,12 @@ class RaisingResponseCodec: module.getattr("request_short_tuple").unwrap().unbind(), ); assert!( - short_request_intercept("llm", make_request(), None) + runtime + .block_on(short_request_intercept( + "llm".to_string(), + make_request(), + None + )) .unwrap_err() .to_string() .contains("must return LLMRequestInterceptOutcome") @@ -189,12 +205,13 @@ class RaisingResponseCodec: let llm_response = wrap_py_llm_sanitize_response_fn(module.getattr("llm_resp_bad_json").unwrap().unbind()) .unwrap(); - assert_eq!( - llm_response( - json!({"ok": true}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::default() - ), - None + assert!( + runtime + .block_on(llm_response( + json!({"ok": true}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::default() + )) + .is_err() ); let bad_codec = PyLlmCodecWrapper { @@ -683,23 +700,31 @@ def invalid(event, fields): metadata: Some(json!({"secret": true})), }; - let sanitized = wrap_py_event_sanitize_fn(module.getattr("sanitize").unwrap().unbind())( - &event, - fields.clone(), - ); + let runtime = tokio::runtime::Runtime::new().unwrap(); + let sanitized = runtime + .block_on(wrap_py_event_sanitize_fn( + module.getattr("sanitize").unwrap().unbind(), + )(event.clone(), fields.clone())) + .unwrap(); assert_eq!(sanitized.data, Some(json!({"safe": "checkpoint"}))); assert_eq!(sanitized.metadata, None); - let raised = wrap_py_event_sanitize_fn(module.getattr("raises").unwrap().unbind())( - &event, - fields.clone(), - ); - assert_eq!(raised, EventSanitizeFields::default()); - - let invalid = wrap_py_event_sanitize_fn(module.getattr("invalid").unwrap().unbind())( - &event, - fields.clone(), + let raised = runtime + .block_on(wrap_py_event_sanitize_fn( + module.getattr("raises").unwrap().unbind(), + )(event.clone(), fields.clone())) + .unwrap_err(); + assert!(raised.to_string().contains("sanitize boom")); + + let invalid = runtime + .block_on(wrap_py_event_sanitize_fn( + module.getattr("invalid").unwrap().unbind(), + )(event, fields.clone())) + .unwrap_err(); + assert!( + invalid + .to_string() + .contains("invalid event sanitizer result") ); - assert_eq!(invalid, EventSanitizeFields::default()); }); } diff --git a/docs/about-nemo-relay/concepts/middleware.mdx b/docs/about-nemo-relay/concepts/middleware.mdx index 0afcd2d56..afd1ae0cf 100644 --- a/docs/about-nemo-relay/concepts/middleware.mdx +++ b/docs/about-nemo-relay/concepts/middleware.mdx @@ -19,6 +19,19 @@ events. NeMo Relay applies each surface at a specific lifecycle point. Middleware is organized by lifecycle meaning rather than as one undifferentiated hook system. +## Asynchronous Callbacks + +All middleware families accept asynchronous callbacks. Rust callbacks return a +future; Python callbacks may return a value or an awaitable; and Node callbacks +may return a value or a Promise. Relay awaits entries sequentially in priority +order, so later callbacks observe earlier middleware output. + +Managed execution and standalone conditional/request-intercept helpers are +asynchronous because their result depends on middleware completion. Manual +lifecycle APIs (`tool_call`, `tool_call_end`, `llm_call`, and `llm_call_end`) +remain synchronous: they create or close their handle immediately and queue +observability work rather than awaiting it. + ## Registration Levels Middleware and subscribers can be registered at different levels depending on their @@ -125,6 +138,19 @@ context. For the callback contract and binding APIs, refer to Sanitize guardrails are observability-oriented. They do not rewrite the real arguments passed to the callback or the real value returned to the caller. +## Queued Event Publication + +Scope operations, marks, and manual tool/LLM lifecycle calls never become +awaitable because an event sanitizer is asynchronous. At emission time Relay +snapshots the event, visible sanitizer chain, and subscribers, then places the +work on a serial dispatcher. The dispatcher awaits sanitizers and publishes the +event later in FIFO order. + +Subscriber and exporter delivery is therefore delayed, while start/end/mark +order is preserved. Closing a scope or deregistering middleware after emission +does not affect queued snapshots. Sanitizer failures fail open: Relay records +the callback failure and publishes the last valid event snapshot. + ## Managed Execution Order For managed execution, NeMo Relay applies middleware and emits lifecycle events diff --git a/docs/build-plugins/dynamic-plugins/native-dynamic/about.mdx b/docs/build-plugins/dynamic-plugins/native-dynamic/about.mdx index a4c162368..8bb88a328 100644 --- a/docs/build-plugins/dynamic-plugins/native-dynamic/about.mdx +++ b/docs/build-plugins/dynamic-plugins/native-dynamic/about.mdx @@ -1,6 +1,6 @@ --- title: "Native Dynamic Plugins (Rust)" -description: "Build in-process Rust shared-library plugins against the NeMo Relay Native ABI v2." +description: "Build in-process Rust shared-library plugins against the NeMo Relay Native ABI v3." position: 10 --- {/* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. @@ -116,10 +116,12 @@ path, then replace `` with that library's SHA-256 digest. Use Native Plugin](/build-plugins/dynamic-plugins/native-dynamic/rust-native-plugin-example) for a complete example with validation, middleware, scopes, and configuration schema support. -## Native ABI v2 +## Native ABI v3 -The host passes a `NemoRelayNativeHostApiV1` table to the entry symbol. The -plugin returns a `NemoRelayNativePluginV1` descriptor: +The entry symbol receives a `*const NemoRelayNativeHostApiV1` pointer. It +points at the v1 prefix of a v3 `NemoRelayNativeHostApiV3` table; check +`abi_version` and `struct_size` before casting. The plugin returns a +`NemoRelayNativePluginV1` descriptor: ```rust extern "C" fn nemo_relay_register_plugin( @@ -128,6 +130,35 @@ extern "C" fn nemo_relay_register_plugin( ) -> NemoRelayStatus ``` +The v3 host table retains the frozen legacy prefix and appends a +completion-based asynchronous middleware extension. An entry that rejects the +v3 table with `InvalidArg` is retried with the legacy table. Rust plugins using the +typed `NativePlugin` APIs continue to work unchanged. Raw ABI plugins can use +`PluginContext::register_async_middleware_raw` when a callback must complete +later. The callback receives a JSON invocation, an optional continuation for +execution intercepts, and a one-shot completion handle. + +Return `Complete` after resolving or rejecting the completion before the +callback returns. Return `Pending` only when retaining the completion; settle +it exactly once, call `async_completion_release`, and release an async `next` +handle after use. The host marks a completion cancelled when the awaiting +runtime work is dropped; late and duplicate settlement is rejected safely. + +Event sanitizers registered through this extension still run on Relay's serial +publication dispatcher. Scope and mark emission remain synchronous and their +sanitized events are delivered later in emission order. + +The v3 completion and continuation ABI settles one JSON value. Consequently, +an async native LLM stream execution intercept currently receives and returns +the complete JSON array of chunks: Relay buffers the provider stream before +replaying it to the caller. It is not an incremental streaming transport and +does not provide per-chunk backpressure. Use a synchronous native stream +intercept or a worker plugin when first-token latency is required. + +Legacy v1/v2 middleware callbacks are synchronous and run on the runtime's +execution path. They must not block on I/O; use the v3 completion-based API for +long-running work. + Text and JSON data cross this boundary as host-owned `NemoRelayNativeString` handles. ABI structs also carry scalars, opaque handles, callback pointers, and plugin-owned `user_data`. Do not pass Rust diff --git a/docs/reference/event-sanitizers.mdx b/docs/reference/event-sanitizers.mdx index 8a4689662..6ff89e7e7 100644 --- a/docs/reference/event-sanitizers.mdx +++ b/docs/reference/event-sanitizers.mdx @@ -49,22 +49,25 @@ semantic category, attributes, semantic input and output meaning, or schemas. Registries run in priority order. Lower priorities run first, and each callback receives the fields returned by the callback before it. Invalid -binding callback results fail open and preserve the current fields. In Node.js, +binding callback results fail open and preserve the current fields. The same +rule applies to tool and LLM request/response sanitizer errors: Relay preserves +the last valid observability payload without changing provider execution. In Node.js, a synchronous sanitizer callback that throws also fails open; Relay records the error for `getLastCallbackError()`. -## Publication Semantics +## Async Delivery and Ordering -Scope and mark emission APIs remain synchronous. They snapshot the event, -visible sanitizer chain, and subscribers, then enqueue that snapshot for -sanitization and publication on a serial background dispatcher. Subscribers -and exporters therefore receive the sanitized event after the emission call -returns. +Event sanitizer callbacks may be asynchronous: use an `async def` callback in +Python, return a Promise in Node.js, or return a future in Rust. Scope and mark +emission remains synchronous. Relay snapshots the event, sanitizers, and +subscribers and queues them on one serial publication dispatcher; that +dispatcher awaits sanitizers before delivering the event to subscribers and +exporters. -The dispatcher processes snapshots in FIFO order, preserving scope start/end -and mark ordering. Closing a scope or deregistering middleware after emission -does not alter an already-snapshotted publication chain. Use the binding's -subscriber flush API when a test or shutdown path must wait for queued delivery. +This preserves FIFO start/end/mark delivery without making `push_scope`, +`pop_scope`, or `event` awaitable. A scope-local sanitizer removed after an +event is emitted still applies to its queued snapshot. An asynchronous +sanitizer rejection fails open and preserves the last valid event fields. ## Registration Lifetimes @@ -187,17 +190,23 @@ activation fails. ## Experimental C and Go Bindings -The source-first C API uses `NemoRelayEventSanitizeCb`. It provides global, -scope-local, and plugin-context registration functions for all three surfaces. -Global names start with `nemo_relay_register_`, and scope-local names start -with `nemo_relay_scope_register_`. - -The Go binding provides `EventSanitizeFields`, `EventSanitizeFunc`, global -`Register*SanitizeGuardrail` helpers, scope-local -`ScopeRegister*SanitizeGuardrail` helpers, and the same methods on -`PluginContext`. The `guardrails` package provides shorter aliases. Because a -returned `EventSanitizeFields` replaces all three fields, copy the supplied -value and modify only the fields that should change. +The source-first C API retains `NemoRelayEventSanitizeCb` and adds parallel +completion-based async registration APIs. An async callback returns `Complete` +or `Pending` and settles its one-shot completion handle with resolve or reject; +there is no implicit timeout. A callback that returns `Pending` must settle the +handle exactly once, or serial event publication remains blocked. Relay cancels +the handle when the invocation is abandoned; late or duplicate settlement after +cancellation is rejected safely. After resolving or rejecting a retained +completion, call `nemo_relay_async_completion_release` to release the +callback-owned reference. Global names start with `nemo_relay_register_`, and +scope-local names start with `nemo_relay_scope_register_`. + +The Go binding provides `EventSanitizeFields`, `EventSanitizeFunc`, and +`AsyncMiddlewareFunc` variants for global and scope-local event sanitizers. +Async Go callbacks receive a `context.Context`; Relay cancels it when the +invocation is abandoned. Because a returned `EventSanitizeFields` replaces all +three fields, copy the supplied value and modify only the fields that should +change. ## Related Topics diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index 5f5b6aca3..dac53db81 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -1,6 +1,6 @@ --- title: "Migration Guides" -description: "Upgrade NeMo Relay integrations and migrate LLM sanitizer callbacks, plugins, workers, and PII policy." +description: "Upgrade NeMo Relay integrations and migrate async middleware, plugins, workers, LLM sanitizers, and PII policy." position: 6 --- {/* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. @@ -13,17 +13,64 @@ intervening release in sequence. ## Upgrade to NeMo Relay 0.7 -NeMo Relay 0.7 changes the LLM observability sanitizer contract across -in-process bindings, native plugins, raw C FFI consumers, and worker plugins. -Complete the following migrations before you run an existing sanitizer with a -0.7 host. +NeMo Relay 0.7 makes the Rust middleware callback contract asynchronous and +adds awaitable middleware support across the in-process bindings, native +plugins, raw C FFI consumers, and worker plugins. It also changes the LLM +observability sanitizer contract. Complete the following migrations before you +run existing middleware or a sanitizer with a 0.7 host. -Do not deploy a 0.6 sanitizer plugin or worker against a 0.7 host. The LLM -callback signature, native ABI layout, and worker invocation schema changed. -NeMo Relay does not adapt one-argument LLM sanitizer callbacks. +Do not deploy a 0.6 plugin or worker against a 0.7 host. The middleware +callback contract, LLM callback signature, native ABI layout, and worker +invocation schema changed. NeMo Relay does not adapt synchronous Rust +middleware callbacks or one-argument LLM sanitizer callbacks. +### Migrate Middleware Callbacks + +The following callback families are now asynchronous: conditional execution +guardrails, request intercepts, execution intercepts, tool and LLM sanitizers, +and event sanitizers. Relay awaits each registered callback sequentially in +priority order. A callback that rejects or returns an error preserves the +existing error behavior for its middleware family. + +| Surface | 0.6 Callback | 0.7 Callback | +| --- | --- | --- | +| Rust | `Fn(...) -> Result` | `Fn(...) -> Pin> + Send>>` | +| Python | Direct return value | Direct return value or awaitable | +| Node.js | Direct return value | Direct return value or `Promise` | +| Go / raw C FFI | Synchronous callback | Existing synchronous callback, or the new `Async` / completion-based registration API | + +For Rust, wrap the existing result in a ready async future, or use an async +block when the callback needs to await work: + +```rust +use std::sync::Arc; + +use nemo_relay::api::registry::register_tool_conditional_execution_guardrail; + +register_tool_conditional_execution_guardrail( + "policy", + 10, + Arc::new(|_name, _args| { + Box::pin(async move { + // Await policy I/O here when needed. + Ok(None) // Return Some(reason) to block execution. + }) + }), +)?; +``` + +Python and Node.js registration names are unchanged. Mark a Python callback +`async def`, or return a Promise from Node.js, only when it needs asynchronous +work; existing direct-value callbacks remain supported. + +Scope lifecycle and mark emission remain synchronous. `push_scope`, +`pop_scope`, and mark APIs snapshot the event and visible sanitizer/subscriber +chain, then enqueue sanitization and publication on a serial dispatcher. Event +subscribers and exporters therefore receive sanitized events later, in emission +order. Do not add `await` to scope or mark emission calls. + ### Update LLM Sanitizer Callbacks The registration names remain unchanged for global, plugin-context, and @@ -49,9 +96,11 @@ register_llm_sanitize_request_guardrail( "redact-request", 10, Arc::new(|request, context| { - let _active_codec = context.resolve_codec(); - // Apply policy, using _active_codec when normalized access is required. - Some(request) + Box::pin(async move { + let _active_codec = context.resolve_codec(); + // Apply policy, using _active_codec when normalized access is required. + Ok(Some(request)) + }) }), )?; ``` @@ -150,15 +199,15 @@ For complete in-process examples, refer to ### Migrate Worker Sanitizers -All Rust worker sanitizer registrations now require callbacks that return -futures. This change applies to mark, scope-start, scope-end, tool-request, -tool-response, LLM-request, and LLM-response sanitizers. Conditional guardrails -and request intercepts keep their existing synchronous contracts. +All Rust worker middleware registrations now require callbacks that return +futures. This includes conditional guardrails, request and execution intercepts, +mark and scope event sanitizers, tool request/response sanitizers, and LLM +request/response sanitizers. Python worker middleware can return either an +immediate value or an awaitable. Python LLM sanitizers must still accept both +the payload and directional context. -Update Rust worker callbacks to use `async move` and return `Result` from the -future. Python worker sanitizers can return either an immediate value or an -awaitable, but Python LLM sanitizers must still accept both the payload and -directional context. +Update Rust worker callbacks to return `Box::pin(async move { ... })` and +resolve to `Result` from the future. @@ -167,13 +216,15 @@ directional context. ctx.register_llm_sanitize_request_guardrail( "redact-request", 10, - |request, context| async move { - if let Some(codec) = context.resolve_codec() { - let annotated = codec.decode(&request).await?; - let request = codec.encode(&annotated, &request).await?; - return Ok(Some(request)); - } - Ok(Some(request)) + |request, context| { + Box::pin(async move { + if let Some(codec) = context.resolve_codec() { + let annotated = codec.decode(&request).await?; + let request = codec.encode(&annotated, &request).await?; + return Ok(Some(request)); + } + Ok(Some(request)) + }) }, ); ``` @@ -220,13 +271,14 @@ wrong-direction capability IDs. ### Rebuild Native and Raw FFI Plugins -NeMo Relay 0.7 uses native ABI v2. Recompile native plugins against the 0.7 +NeMo Relay 0.7 uses native ABI v3. Recompile native plugins against the 0.7 `nemo-relay-plugin` crate and rebuild raw FFI consumers against the generated 0.7 header. -If you already built a plugin against an earlier 0.7 ABI v2 prerelease, rebuild -it again. The ABI version remains 2, but the prerelease LLM sanitizer callback -slots and context layouts changed before release. +The v3 table preserves the v2 prefix, and Relay retries a legacy v2 table when +loading a plugin that rejects v3. Rebuild anyway if a plugin uses raw ABI +callbacks: v3 adds completion-based async middleware registration, async +execution continuations, and explicit cancellation/late-settlement behavior. The plugin manifest value remains `compat.native_api = "1"`. This manifest contract version is separate from the host ABI version; do not change it to @@ -243,7 +295,7 @@ after the callback returns. Release host-owned output strings with the standard host string release operation. For the complete ABI contract, refer to -[Native ABI v2](/build-plugins/dynamic-plugins/native-dynamic/about#native-abi-v2). +[Native ABI v3](/build-plugins/dynamic-plugins/native-dynamic/about#native-abi-v3). ### Update PII Redaction Configuration diff --git a/integrations/openclaw/test/live-smoke.test.ts b/integrations/openclaw/test/live-smoke.test.ts index c8b3de692..a7a788509 100644 --- a/integrations/openclaw/test/live-smoke.test.ts +++ b/integrations/openclaw/test/live-smoke.test.ts @@ -17,6 +17,17 @@ import { callGatewayStatus, type TestGatewayMethodHandler } from './gateway-stat const liveSmokeEnabled = process.env.NEMO_RELAY_OPENCLAW_LIVE_SMOKE === '1'; +async function waitForExportFile(outputDir: string, prefix: string, timeoutMs = 2_000): Promise { + const deadline = Date.now() + timeoutMs; + while (Date.now() < deadline) { + const files = await fs.readdir(outputDir); + const exportedPath = files.find((file) => file.startsWith(prefix) && file.endsWith('.json')); + if (exportedPath) return exportedPath; + await new Promise((resolve) => setTimeout(resolve, 10)); + } + return undefined; +} + it( 'runs a live NeMo Relay binding smoke for session ATIF export and hook replay', { skip: !liveSmokeEnabled }, @@ -133,8 +144,7 @@ it( { sessionId: '../live-session:1' }, ); - const files = await fs.readdir(outputDir); - const exportedPath = files.find((file) => file.startsWith('live-') && file.endsWith('.json')); + const exportedPath = await waitForExportFile(outputDir, 'live-'); assert.ok(exportedPath, 'expected generic observability ATIF export'); const exported = JSON.parse(await fs.readFile(path.join(outputDir, exportedPath), 'utf8')) as unknown; assert.equal(typeof exported, 'object'); diff --git a/python/nemo_relay/__init__.py b/python/nemo_relay/__init__.py index 9a7690b96..1b6a3a9ee 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -167,26 +167,33 @@ class EventSanitizeFields(TypedDict): #: Arguments are the tool name and JSON payload. The return value is the JSON #: payload recorded on the emitted event. Exceptions propagate through the #: lifecycle call that invoked the guardrail. -ToolSanitizeGuardrail: TypeAlias = Callable[[str, Json], Json] -EventSanitizeGuardrail: TypeAlias = Callable[["Event", EventSanitizeFields], EventSanitizeFields] +ToolSanitizeGuardrail: TypeAlias = Callable[[str, Json], Json | Awaitable[Json]] +EventSanitizeGuardrail: TypeAlias = Callable[ + ["Event", EventSanitizeFields], EventSanitizeFields | Awaitable[EventSanitizeFields] +] #: Guardrail callback that can block tool execution by returning a rejection #: message. Returning ``None`` allows execution to continue. -ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, Json], Optional[str]] +ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, Json], Optional[str] | Awaitable[Optional[str]]] #: Guardrail callback that sanitizes an ``LLMRequest`` used for emitted events. #: Callbacks receive ``(request, context)``. Returning ``None`` omits the LLM observability #: payload and annotation without changing the caller-visible request. -LlmSanitizeRequestGuardrail: TypeAlias = Callable[[LLMRequest, "LlmSanitizeRequestContext"], Optional[LLMRequest]] +LlmSanitizeRequestGuardrail: TypeAlias = Callable[ + [LLMRequest, "LlmSanitizeRequestContext"], + Optional[LLMRequest] | Awaitable[Optional[LLMRequest]], +] #: Guardrail callback that sanitizes an emitted JSON LLM response payload. #: Callbacks receive ``(response, context)`` and can return ``None`` to omit #: observability payload and annotation without changing the caller response. -LlmSanitizeResponseGuardrail: TypeAlias = Callable[[Json, "LlmSanitizeResponseContext"], Optional[Json]] +LlmSanitizeResponseGuardrail: TypeAlias = Callable[ + [Json, "LlmSanitizeResponseContext"], Optional[Json] | Awaitable[Optional[Json]] +] #: Guardrail callback that can block an LLM call by returning a rejection #: message. Returning ``None`` allows execution to continue. -LlmConditionalExecutionGuardrail: TypeAlias = Callable[[LLMRequest], Optional[str]] +LlmConditionalExecutionGuardrail: TypeAlias = Callable[[LLMRequest], Optional[str] | Awaitable[Optional[str]]] #: Request intercept callback that rewrites tool arguments before execution. #: Arguments are the tool name and current JSON payload. The return value #: becomes the payload seen by later request intercepts and tool execution. -ToolRequestIntercept: TypeAlias = AbcCallable[[str, Json], Json] +ToolRequestIntercept: TypeAlias = AbcCallable[[str, Json], Json | Awaitable[Json]] #: Execution intercept callback that wraps tool execution with middleware #: behavior. The callback receives the tool name, current arguments, and the #: next callable. It may await and return ``next(args)`` or short-circuit. @@ -198,7 +205,7 @@ class EventSanitizeFields(TypedDict): #: and pending-mark outcome passed to later intercepts and managed execution. LlmRequestIntercept: TypeAlias = Callable[ [str, LLMRequest, AnnotatedLLMRequest | None], - LLMRequestInterceptOutcome, + LLMRequestInterceptOutcome | Awaitable[LLMRequestInterceptOutcome], ] #: Execution intercept callback that wraps non-streaming LLM execution. The #: callback receives the logical LLM name, request, and next callable. It may diff --git a/python/nemo_relay/__init__.pyi b/python/nemo_relay/__init__.pyi index 2e8f5b0fc..f640a2f72 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -161,8 +161,11 @@ class EventSanitizeFields(TypedDict): category_profile: JsonObject | None metadata: Json | None -ToolSanitizeGuardrail: TypeAlias = Callable[[str, Json], Json] -EventSanitizeGuardrail: TypeAlias = Callable[[Event, EventSanitizeFields], EventSanitizeFields] +ToolSanitizeGuardrail: TypeAlias = Callable[[str, Json], Json | Awaitable[Json]] +EventSanitizeGuardrail: TypeAlias = Callable[ + [Event, EventSanitizeFields], + EventSanitizeFields | Awaitable[EventSanitizeFields], +] """Guardrail callback that sanitizes emitted tool request or response payloads. Arguments: @@ -175,7 +178,7 @@ Exceptional flow: Exceptions raised by the callback propagate through the lifecycle operation that invoked the guardrail. """ -ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, Json], Optional[str]] +ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, Json], Optional[str] | Awaitable[Optional[str]]] """Guardrail callback that can block tool execution. Arguments: @@ -184,7 +187,10 @@ Arguments: Return: ``None`` to allow execution, or a rejection message to block it. """ -LlmSanitizeRequestGuardrail: TypeAlias = Callable[[LLMRequest, "LlmSanitizeRequestContext"], Optional[LLMRequest]] +LlmSanitizeRequestGuardrail: TypeAlias = Callable[ + [LLMRequest, "LlmSanitizeRequestContext"], + Optional[LLMRequest] | Awaitable[Optional[LLMRequest]], +] """Guardrail callback that sanitizes an ``LLMRequest`` used for emitted events. Arguments: @@ -197,7 +203,10 @@ Return: Request object recorded on the emitted lifecycle event, or ``None`` to omit the LLM observability payload and annotation. """ -LlmSanitizeResponseGuardrail: TypeAlias = Callable[[Json, "LlmSanitizeResponseContext"], Optional[Json]] +LlmSanitizeResponseGuardrail: TypeAlias = Callable[ + [Json, "LlmSanitizeResponseContext"], + Optional[Json] | Awaitable[Optional[Json]], +] """Guardrail callback that sanitizes an emitted JSON LLM response payload. Arguments: @@ -210,7 +219,7 @@ Return: Response object recorded on the emitted lifecycle event, or ``None`` to omit the LLM observability payload and annotation. """ -LlmConditionalExecutionGuardrail: TypeAlias = Callable[[LLMRequest], Optional[str]] +LlmConditionalExecutionGuardrail: TypeAlias = Callable[[LLMRequest], Optional[str] | Awaitable[Optional[str]]] """Guardrail callback that can block an LLM call. Arguments: @@ -219,7 +228,7 @@ Arguments: Return: ``None`` to allow execution, or a rejection message to block it. """ -ToolRequestIntercept: TypeAlias = Callable[[str, Json], Json] +ToolRequestIntercept: TypeAlias = Callable[[str, Json], Json | Awaitable[Json]] """Request intercept callback that rewrites tool arguments before execution. Arguments: @@ -246,7 +255,7 @@ Exceptional flow: """ LlmRequestIntercept: TypeAlias = Callable[ [str, LLMRequest, AnnotatedLLMRequest | None], - LLMRequestInterceptOutcome, + LLMRequestInterceptOutcome | Awaitable[LLMRequestInterceptOutcome], ] """Request intercept callback that rewrites raw and annotated LLM requests. diff --git a/python/nemo_relay/_native.pyi b/python/nemo_relay/_native.pyi index 1b4e7354a..d4f409405 100644 --- a/python/nemo_relay/_native.pyi +++ b/python/nemo_relay/_native.pyi @@ -37,20 +37,29 @@ class _EventSanitizeFields(TypedDict): category_profile: _JsonObject | None metadata: _Json | None -_ToolSanitizeGuardrail: TypeAlias = Callable[[str, _Json], _Json] -_ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, _Json], Optional[str]] -_LlmSanitizeRequestGuardrail: TypeAlias = Callable[["LLMRequest", "LlmSanitizeRequestContext"], Optional["LLMRequest"]] -_LlmSanitizeResponseGuardrail: TypeAlias = Callable[[_Json, "LlmSanitizeResponseContext"], Optional[_Json]] -_EventSanitizeGuardrail: TypeAlias = Callable[[ScopeEvent | MarkEvent, _EventSanitizeFields], _EventSanitizeFields] -_LlmConditionalExecutionGuardrail: TypeAlias = Callable[["LLMRequest"], Optional[str]] -_ToolRequestIntercept: TypeAlias = Callable[[str, _Json], _Json] +_ToolSanitizeGuardrail: TypeAlias = Callable[[str, _Json], _Json | Awaitable[_Json]] +_ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, _Json], Optional[str] | Awaitable[Optional[str]]] +_LlmSanitizeRequestGuardrail: TypeAlias = Callable[ + ["LLMRequest", "LlmSanitizeRequestContext"], + Optional["LLMRequest"] | Awaitable[Optional["LLMRequest"]], +] +_LlmSanitizeResponseGuardrail: TypeAlias = Callable[ + [_Json, "LlmSanitizeResponseContext"], + Optional[_Json] | Awaitable[Optional[_Json]], +] +_EventSanitizeGuardrail: TypeAlias = Callable[ + [ScopeEvent | MarkEvent, _EventSanitizeFields], + _EventSanitizeFields | Awaitable[_EventSanitizeFields], +] +_LlmConditionalExecutionGuardrail: TypeAlias = Callable[["LLMRequest"], Optional[str] | Awaitable[Optional[str]]] +_ToolRequestIntercept: TypeAlias = Callable[[str, _Json], _Json | Awaitable[_Json]] _ToolExecutionIntercept: TypeAlias = Callable[ [str, _Json, Callable[[_Json], Awaitable[_Json]]], "ToolExecutionInterceptOutcome | Awaitable[ToolExecutionInterceptOutcome]", ] _LlmRequestIntercept: TypeAlias = Callable[ [str, "LLMRequest", "AnnotatedLLMRequest | None"], - "LLMRequestInterceptOutcome", + "LLMRequestInterceptOutcome | Awaitable[LLMRequestInterceptOutcome]", ] _LlmExecutionIntercept: TypeAlias = Callable[ [str, "LLMRequest", Callable[["LLMRequest"], Awaitable[_Json]]], @@ -1615,7 +1624,7 @@ def llm_stream_call_execute( """ ... -def tool_request_intercepts(name: str, args: _Json) -> _Json: +def tool_request_intercepts(name: str, args: _Json) -> _Json | Awaitable[_Json]: """Run the registered tool request-intercept chain. Args: @@ -1623,14 +1632,15 @@ def tool_request_intercepts(name: str, args: _Json) -> _Json: args: Current JSON-compatible tool arguments. Returns: - Transformed tool arguments after all applicable request intercepts. + Transformed tool arguments directly outside an event loop, or an + awaitable resolving to them from an async caller. Exceptional flow: Callback exceptions and native middleware errors propagate unchanged. """ ... -def tool_conditional_execution(name: str, args: _Json) -> None: +def tool_conditional_execution(name: str, args: _Json) -> None | Awaitable[None]: """Run tool conditional-execution guardrails. Args: @@ -1638,7 +1648,8 @@ def tool_conditional_execution(name: str, args: _Json) -> None: args: Current JSON-compatible tool arguments. Returns: - ``None`` when all guardrails allow execution. + ``None`` when all guardrails allow execution, directly outside an event + loop or through an awaitable from an async caller. Exceptional flow: Raises a native rejection error when a guardrail returns a rejection @@ -1646,7 +1657,9 @@ def tool_conditional_execution(name: str, args: _Json) -> None: """ ... -def llm_request_intercepts(name: str, request: LLMRequest) -> LLMRequestInterceptOutcome: +def llm_request_intercepts( + name: str, request: LLMRequest +) -> LLMRequestInterceptOutcome | Awaitable[LLMRequestInterceptOutcome]: """Run the registered LLM request-intercept chain. Args: @@ -1654,21 +1667,23 @@ def llm_request_intercepts(name: str, request: LLMRequest) -> LLMRequestIntercep request: Current LLM request. Returns: - Transformed request after all applicable request intercepts. + Transformed request directly outside an event loop, or an awaitable + resolving to it from an async caller. Exceptional flow: Callback exceptions and native middleware errors propagate unchanged. """ ... -def llm_conditional_execution(request: LLMRequest) -> None: +def llm_conditional_execution(request: LLMRequest) -> None | Awaitable[None]: """Run LLM conditional-execution guardrails. Args: request: LLM request to evaluate. Returns: - ``None`` when all guardrails allow execution. + ``None`` when all guardrails allow execution, directly outside an event + loop or through an awaitable from an async caller. Exceptional flow: Raises a native rejection error when a guardrail returns a rejection diff --git a/python/tests/test_adaptive.py b/python/tests/test_adaptive.py index c7e956e7b..899c3b2e5 100644 --- a/python/tests/test_adaptive.py +++ b/python/tests/test_adaptive.py @@ -222,7 +222,7 @@ async def test_adaptive_runtime_bind_scope_passes_through_without_state(self): ) with scope.scope("adaptive-runtime-translate", ScopeType.Agent) as handle: runtime.bind_scope(handle) - translated = llm.request_intercepts("anthropic", request) + translated = await llm.request_intercepts("anthropic", request) assert translated.request.content == { "messages": [{"role": "user", "content": "Hello"}], "system": "You are helpful.", diff --git a/python/tests/test_builtin_codecs.py b/python/tests/test_builtin_codecs.py index 156122583..c1fcef898 100644 --- a/python/tests/test_builtin_codecs.py +++ b/python/tests/test_builtin_codecs.py @@ -12,8 +12,6 @@ from typing import cast -import pytest - import nemo_relay from nemo_relay import ( AnnotatedLLMRequest, @@ -434,8 +432,8 @@ def sanitize_response(response, context): subscribers.deregister("test-manual-call-end-sanitized-response-codec") guardrails.deregister_llm_sanitize_response("test-call-end-codec-sanitizer") - def test_manual_call_end_response_codec_failure_raises_after_end_event(self): - """manual llm.call_end() surfaces response codec failures instead of dropping them.""" + def test_manual_call_end_response_codec_failure_defers_without_raising(self): + """manual llm.call_end() records deferred response codec failures without blocking.""" captured_events = [] def capture(event): @@ -448,8 +446,7 @@ def capture(event): "manual-codec-error-llm", LLMRequest({}, {"model": "gpt-4", "messages": []}), ) - with pytest.raises(RuntimeError, match="OpenAI Chat response decode"): - llm.call_end(handle, "malformed response", response_codec=OpenAIChatCodec()) + llm.call_end(handle, "malformed response", response_codec=OpenAIChatCodec()) subscribers.flush() end_events = [ diff --git a/python/tests/test_context_isolation.py b/python/tests/test_context_isolation.py index 49314e312..7ee5a7531 100644 --- a/python/tests/test_context_isolation.py +++ b/python/tests/test_context_isolation.py @@ -199,8 +199,8 @@ async def run_tool(owner): ) await asyncio.sleep(0) - args = nemo_relay.tools.request_intercepts("task-tool", {"owner": owner}) - nemo_relay.tools.conditional_execution("task-tool", args) + args = await nemo_relay.tools.request_intercepts("task-tool", {"owner": owner}) + await nemo_relay.tools.conditional_execution("task-tool", args) manual_handle = nemo_relay.tools.call(f"manual-tool-{owner}", args) await asyncio.sleep(0) @@ -258,9 +258,9 @@ def intercept(name, request, annotated): request = nemo_relay.LLMRequest({}, {"messages": [], "owner": owner}) await asyncio.sleep(0) - intercepted = nemo_relay.llm.request_intercepts("task-llm", request) + intercepted = await nemo_relay.llm.request_intercepts("task-llm", request) assert intercepted.request.content["intercepted_by"] == owner - nemo_relay.llm.conditional_execution(request) + await nemo_relay.llm.conditional_execution(request) manual_handle = nemo_relay.llm.call(f"manual-llm-{owner}", request) await asyncio.sleep(0) diff --git a/python/tests/test_event_sanitizers.py b/python/tests/test_event_sanitizers.py index 89d97d9d6..08234b70e 100644 --- a/python/tests/test_event_sanitizers.py +++ b/python/tests/test_event_sanitizers.py @@ -56,7 +56,7 @@ def second(event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitiz assert calls == [("checkpoint", {"secret": "raw"}), ("mark", {"stage": "first"})] -def test_mark_sanitizer_exception_clears_observability_fields(capture_events): +def test_mark_sanitizer_exception_preserves_observability_fields(capture_events): _capture_name, events = capture_events def raises(_event: nemo_relay.Event, _fields: EventSanitizeFields) -> EventSanitizeFields: @@ -69,7 +69,7 @@ def raises(_event: nemo_relay.Event, _fields: EventSanitizeFields) -> EventSanit finally: guardrails.deregister_mark_sanitize("python-mark-raises") - assert events[-1].data is None + assert events[-1].data == {"kept": True} assert events[-1].metadata is None diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index 9a3662be4..411629401 100644 --- a/python/tests/test_llm.py +++ b/python/tests/test_llm.py @@ -167,6 +167,7 @@ def sanitize_response(response, context): try: handle = llm.call("py_llm_structured_context", make_request()) llm.call_end(handle, {"response": "ok"}) + subscribers.flush() finally: guardrails.deregister_llm_sanitize_request("py_llm_structured_context_request") guardrails.deregister_llm_sanitize_response("py_llm_structured_context_response") @@ -275,7 +276,7 @@ def test_duplicate_raises(self): guardrails.register_llm_sanitize_request("py_llm_dup", 1, lambda r, context: r) guardrails.deregister_llm_sanitize_request("py_llm_dup") - def test_sanitize_request_callable_error_omits_observability_input(self): + def test_sanitize_request_callable_error_preserves_observability_input(self): events = [] subscribers.register("py_llm_sanitize_req_sub", lambda event: events.append(event)) guardrails.register_llm_sanitize_request( @@ -284,7 +285,10 @@ def test_sanitize_request_callable_error_omits_observability_input(self): lambda request, context: raise_runtime_error("boom"), ) try: - request = make_request() + request = LLMRequest( + {"authorization": "secret", "x-request-id": "safe"}, + make_request().content, + ) handle = llm.call("llm_sanitize_req_fail", request) llm.call_end(handle, {"ok": True}) finally: @@ -295,10 +299,10 @@ def test_sanitize_request_callable_error_omits_observability_input(self): subscribers.deregister("py_llm_sanitize_req_sub") start = _llm_event(events, "llm_sanitize_req_fail", "start") - assert start.data is None + assert start.data == {"headers": {"x-request-id": "safe"}, "content": request.content} assert start.annotated_request is None - def test_sanitize_request_invalid_return_omits_observability_input(self): + def test_sanitize_request_invalid_return_preserves_observability_input(self): events = [] subscribers.register("py_llm_sanitize_req_bad_sub", lambda event: events.append(event)) guardrails.register_llm_sanitize_request( @@ -307,7 +311,10 @@ def test_sanitize_request_invalid_return_omits_observability_input(self): cast(guardrails.LlmSanitizeRequestGuardrail, lambda request, context: object()), ) try: - request = make_request() + request = LLMRequest( + {"authorization": "secret", "x-request-id": "safe"}, + make_request().content, + ) handle = llm.call("llm_sanitize_req_bad", request) llm.call_end(handle, {"ok": True}) finally: @@ -318,10 +325,10 @@ def test_sanitize_request_invalid_return_omits_observability_input(self): subscribers.deregister("py_llm_sanitize_req_bad_sub") start = _llm_event(events, "llm_sanitize_req_bad", "start") - assert start.data is None + assert start.data == {"headers": {"x-request-id": "safe"}, "content": request.content} assert start.annotated_request is None - def test_sanitize_response_callable_error_omits_observability_output(self): + def test_sanitize_response_callable_error_preserves_observability_output(self): events = [] subscribers.register("py_llm_sanitize_resp_sub", lambda event: events.append(event)) guardrails.register_llm_sanitize_response( @@ -340,10 +347,10 @@ def test_sanitize_response_callable_error_omits_observability_output(self): subscribers.deregister("py_llm_sanitize_resp_sub") end = _llm_event(events, "llm_sanitize_resp_fail", "end") - assert end.data is None + assert end.data == {"ok": True} assert end.annotated_response is None - def test_sanitize_response_invalid_return_omits_observability_output(self): + def test_sanitize_response_invalid_return_preserves_observability_output(self): events = [] subscribers.register("py_llm_sanitize_resp_bad_sub", lambda event: events.append(event)) guardrails.register_llm_sanitize_response( @@ -362,7 +369,7 @@ def test_sanitize_response_invalid_return_omits_observability_output(self): subscribers.deregister("py_llm_sanitize_resp_bad_sub") end = _llm_event(events, "llm_sanitize_resp_bad", "end") - assert end.data is None + assert end.data == {"ok": True} assert end.annotated_response is None def test_sanitize_response_guardrail_accepts_scalar_json_payloads(self): @@ -398,7 +405,7 @@ def test_conditional_execution_invalid_return_type_raises(self): cast(guardrails.LlmConditionalExecutionGuardrail, lambda request: 123), ) try: - with pytest.raises(RuntimeError, match="expected str or None"): + with pytest.raises(RuntimeError, match="unexpected type"): llm.conditional_execution(make_request()) finally: guardrails.deregister_llm_conditional_execution("py_llm_cond_bad_type") @@ -410,7 +417,7 @@ def test_conditional_execution_callable_error_raises(self): lambda request: raise_runtime_error("boom"), ) try: - with pytest.raises(RuntimeError, match="callable failed"): + with pytest.raises(RuntimeError, match="RuntimeError: boom"): llm.conditional_execution(make_request()) finally: guardrails.deregister_llm_conditional_execution("py_llm_cond_error") diff --git a/python/tests/test_tools.py b/python/tests/test_tools.py index 3619902f3..a401eb73b 100644 --- a/python/tests/test_tools.py +++ b/python/tests/test_tools.py @@ -361,7 +361,7 @@ def test_duplicate_intercept_raises(self): def test_request_intercept_raises_on_exception(self): intercepts.register_tool_request("py_req_raise", 1, False, lambda n, a: raise_runtime_error("boom")) try: - with pytest.raises(RuntimeError, match="callable failed"): + with pytest.raises(RuntimeError, match="RuntimeError: boom"): tools.request_intercepts("raise_tool", {"value": 1}) finally: intercepts.deregister_tool_request("py_req_raise") @@ -374,7 +374,7 @@ def test_request_intercept_raises_on_unserializable_return(self): cast(intercepts.ToolRequestIntercept, lambda n, a: object()), ) try: - with pytest.raises(RuntimeError, match="py_to_json failed"): + with pytest.raises(RuntimeError, match="unsupported type object"): tools.request_intercepts("bad_return_tool", {"value": 1}) finally: intercepts.deregister_tool_request("py_req_bad_return") @@ -485,7 +485,7 @@ def test_conditional_execution_invalid_return_type_raises(self): cast(guardrails.ToolConditionalExecutionGuardrail, lambda name, args: 123), ) try: - with pytest.raises(RuntimeError, match="expected str or None"): + with pytest.raises(RuntimeError, match="unexpected type"): tools.conditional_execution("bad_type_tool", {}) finally: guardrails.deregister_tool_conditional_execution("py_cond_bad_type") @@ -497,7 +497,7 @@ def test_conditional_execution_callable_error_raises(self): lambda name, args: raise_runtime_error("boom"), ) try: - with pytest.raises(RuntimeError, match="callable failed"): + with pytest.raises(RuntimeError, match="RuntimeError: boom"): tools.conditional_execution("error_tool", {}) finally: guardrails.deregister_tool_conditional_execution("py_cond_error") From 941ca3768d0cbcef1c52ed62f30e4a0417ee7d31 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 15:18:14 -0400 Subject: [PATCH 10/52] test: preserve legacy FFI sanitizer errors Signed-off-by: Will Killian --- crates/ffi/tests/unit/callable_tests.rs | 20 ++++++-------------- 1 file changed, 6 insertions(+), 14 deletions(-) diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 60b260868..4b5a2d251 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -498,18 +498,14 @@ fn test_llm_sanitizers_fail_closed_for_runtime_codec_ids_with_embedded_nul() { let request_sanitizer = wrap_llm_sanitize_request_fn(llm_request_alias_cb, std::ptr::null_mut(), None); - let request_error = resolve(request_sanitizer( + let request_result = resolve(request_sanitizer( make_request(), nemo_relay::api::runtime::LlmSanitizeRequestContext::with_identity( runtime_identity.clone(), ), )) - .expect_err("an embedded runtime codec ID must fail the callback wrapper"); - assert!( - request_error - .to_string() - .contains("runtime codec ID contains an embedded NUL") - ); + .expect("legacy sanitizer wrappers report callback errors out of band"); + assert_eq!(request_result, None); assert!( last_error_message() .unwrap() @@ -518,16 +514,12 @@ fn test_llm_sanitizers_fail_closed_for_runtime_codec_ids_with_embedded_nul() { let response_sanitizer = wrap_llm_sanitize_response_fn(json_alias_cb, std::ptr::null_mut(), None); - let response_error = resolve(response_sanitizer( + let response_result = resolve(response_sanitizer( json!({"secret": "must be omitted"}), nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), )) - .expect_err("an embedded runtime codec ID must fail the callback wrapper"); - assert!( - response_error - .to_string() - .contains("runtime codec ID contains an embedded NUL") - ); + .expect("legacy sanitizer wrappers report callback errors out of band"); + assert_eq!(response_result, None); assert!( last_error_message() .unwrap() From 1f64cee7280d6821df973d757177d1557fc4c5cf Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 15:49:53 -0400 Subject: [PATCH 11/52] fix: address async middleware review feedback Signed-off-by: Will Killian --- crates/core/src/api/llm.rs | 238 +++++++++--------- crates/core/src/api/runtime/callbacks.rs | 2 +- crates/core/src/api/runtime/state.rs | 3 +- .../src/api/runtime/subscriber_dispatcher.rs | 3 + crates/core/src/api/scope.rs | 8 +- crates/core/src/api/tool.rs | 86 ++++--- crates/core/src/plugin/dynamic/native.rs | 107 +++++--- crates/core/src/plugin/dynamic/worker.rs | 7 +- crates/core/src/stream.rs | 19 +- .../tests/fixtures/native_plugin/src/lib.rs | 6 + crates/core/tests/unit/native_plugin_tests.rs | 63 +++-- crates/ffi/src/api/mod.rs | 2 + crates/ffi/src/callable.rs | 8 +- .../tests/integration/callable_extra_tests.rs | 9 +- crates/ffi/tests/integration/main.rs | 2 + crates/ffi/tests/support/mod.rs | 14 ++ crates/ffi/tests/unit/callable_tests.rs | 17 +- crates/node/plugin.d.ts | 4 +- crates/node/src/callable.rs | 59 +++-- crates/node/src/callback_factory.rs | 8 +- crates/node/tests/llm_tests.mjs | 13 +- crates/node/tests/scope_tests.mjs | 7 +- crates/node/tests/tools_tests.mjs | 6 +- crates/pii-redaction/src/builtin.rs | 16 +- .../tests/unit/component_tests.rs | 32 +-- crates/python/src/py_api/mod.rs | 60 +++-- crates/python/src/py_callable.rs | 32 ++- .../tests/coverage/py_api_coverage_tests.rs | 48 ++++ .../coverage/py_callable_coverage_tests.rs | 61 ++++- docs/reference/migration-guides.mdx | 6 + 30 files changed, 596 insertions(+), 350 deletions(-) create mode 100644 crates/ffi/tests/support/mod.rs diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index 06ac5119b..d0bcc79a3 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -745,48 +745,47 @@ pub fn llm_call(params: LlmCallParams<'_>) -> Result { state.build_llm_start_event(&handle, None, None) }; let queued_handle = handle.clone(); - if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { - dispatch_transformed_event( - event, - Box::new(move |event| { - Box::pin(async move { - let mut sanitized_request = - NemoRelayContextState::llm_sanitize_request_snapshot_chain( - request.clone(), - LlmSanitizeRequestContext::default(), - &entries, - ) - .await; - let request_changed = sanitized_request - .as_ref() - .is_some_and(|sanitized| sanitized != &request); - let mut annotation = if sanitized_request.is_none() || request_changed { - None - } else { - annotated_request - }; - if !agent_is_fresh && let Some(sanitized_request) = sanitized_request.as_mut() { - project_llm_request_to_current_user_turn( - sanitized_request, - &mut annotation, - None, - ); - } - let input = sanitized_request - .as_ref() - .and_then(|request| serde_json::to_value(request).ok()); - let context = global_context(); - match context.read() { - Ok(state) => state.build_llm_start_event(&queued_handle, input, annotation), - Err(_) => event, - } - }) - }), - event_sanitizers, - &subscribers, - scope_stack, - ); - } + let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default(); + dispatch_transformed_event( + event, + Box::new(move |event| { + Box::pin(async move { + let mut sanitized_request = + NemoRelayContextState::llm_sanitize_request_snapshot_chain( + request.clone(), + LlmSanitizeRequestContext::default(), + &entries, + ) + .await; + let request_changed = sanitized_request + .as_ref() + .is_some_and(|sanitized| sanitized != &request); + let mut annotation = if sanitized_request.is_none() || request_changed { + None + } else { + annotated_request + }; + if !agent_is_fresh && let Some(sanitized_request) = sanitized_request.as_mut() { + project_llm_request_to_current_user_turn( + sanitized_request, + &mut annotation, + None, + ); + } + let input = sanitized_request + .as_ref() + .and_then(|request| serde_json::to_value(request).ok()); + let context = global_context(); + match context.read() { + Ok(state) => state.build_llm_start_event(&queued_handle, input, annotation), + Err(_) => event, + } + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); Ok(handle) } @@ -873,87 +872,84 @@ pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> { .build(), ) }; - if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { - dispatch_transformed_event( - event, - Box::new(move |event| { - Box::pin(async move { - let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( - response.clone(), - LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), - &entries, - ) - .await; - let changed = sanitized - .as_ref() - .is_some_and(|sanitized| sanitized != &response); - let data = match sanitized { - Some(response) - if response_was_null_without_fallback && response.is_null() => - { - None - } - response => response, - }; - let annotation_omitted = data.as_ref().is_none_or(Json::is_null); - let (mut annotation, decode_error) = if annotation_omitted { - (None, None) - } else { - resolve_llm_end_annotation( - (!changed).then_some(annotated_response).flatten(), - response_codec, - data.as_ref(), - &LlmCallEndBehavior { - response_codec_errors_fatal: false, - attach_estimated_cost: false, - }, - &handle.name, - ) - }; - if let Some(error) = decode_error { - log::error!( - target: "nemo_relay.runtime", - event = "manual_llm_response_codec_failed"; - "Manual LLM response annotation failed during queued publication: {error}" - ); - } - let pricing = crate::codec::response::active_pricing_resolver(); - let summary = finalize_optimization_summary( - &handle.optimization_recorder, - annotation.as_mut(), - handle.model_name.as_deref(), - &pricing, - ); - if !annotation_omitted - && annotation.is_none() - && let Some(summary) = summary - { - annotation = Some(AnnotatedLlmResponse { - optimization_summary: Some(summary), - ..AnnotatedLlmResponse::default() - }); + let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default(); + dispatch_transformed_event( + event, + Box::new(move |event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( + response.clone(), + LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), + &entries, + ) + .await; + let changed = sanitized + .as_ref() + .is_some_and(|sanitized| sanitized != &response); + let data = match sanitized { + Some(response) if response_was_null_without_fallback && response.is_null() => { + None } - let context = global_context(); - let Ok(state) = context.read() else { - return event; - }; - let end_metadata = metadata_with_otel_status(metadata, "OK", None); - state.build_llm_end_event( - EndLlmHandleParams::builder() - .handle(&handle) - .data_opt(data) - .metadata_opt(end_metadata) - .annotated_response_opt(annotation.map(Arc::new)) - .timestamp_opt(timestamp) - .build(), + response => response, + }; + let annotation_omitted = data.as_ref().is_none_or(Json::is_null); + let (mut annotation, decode_error) = if annotation_omitted { + (None, None) + } else { + resolve_llm_end_annotation( + (!changed).then_some(annotated_response).flatten(), + response_codec, + data.as_ref(), + &LlmCallEndBehavior { + response_codec_errors_fatal: false, + attach_estimated_cost: false, + }, + &handle.name, ) - }) - }), - event_sanitizers, - &subscribers, - scope_stack, - ); - } + }; + if let Some(error) = decode_error { + log::error!( + target: "nemo_relay.runtime", + event = "manual_llm_response_codec_failed"; + "Manual LLM response annotation failed during queued publication: {error}" + ); + } + let pricing = crate::codec::response::active_pricing_resolver(); + let summary = finalize_optimization_summary( + &handle.optimization_recorder, + annotation.as_mut(), + handle.model_name.as_deref(), + &pricing, + ); + if !annotation_omitted + && annotation.is_none() + && let Some(summary) = summary + { + annotation = Some(AnnotatedLlmResponse { + optimization_summary: Some(summary), + ..AnnotatedLlmResponse::default() + }); + } + let context = global_context(); + let Ok(state) = context.read() else { + return event; + }; + let end_metadata = metadata_with_otel_status(metadata, "OK", None); + state.build_llm_end_event( + EndLlmHandleParams::builder() + .handle(&handle) + .data_opt(data) + .metadata_opt(end_metadata) + .annotated_response_opt(annotation.map(Arc::new)) + .timestamp_opt(timestamp) + .build(), + ) + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); Ok(()) } diff --git a/crates/core/src/api/runtime/callbacks.rs b/crates/core/src/api/runtime/callbacks.rs index b07a82fde..d52577f07 100644 --- a/crates/core/src/api/runtime/callbacks.rs +++ b/crates/core/src/api/runtime/callbacks.rs @@ -29,7 +29,7 @@ use crate::json::Json; /// it may replace. Later callbacks observe fields returned by earlier entries. pub type EventSanitizeFn = Arc< dyn Fn( - Event, + Arc, EventSanitizeFields, ) -> Pin> + Send>> + Send diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index e2098b9a7..a70d3a996 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -637,9 +637,10 @@ impl NemoRelayContextState { mut event: Event, entries: &[Guardrail], ) -> Event { + let event_context = Arc::new(event.clone()); for entry in entries { let fields = event.sanitize_fields(); - match (entry.payload)(event.clone(), fields).await { + match (entry.payload)(Arc::clone(&event_context), fields).await { Ok(fields) => event.apply_sanitize_fields(fields), Err(error) => log::error!( target: "nemo_relay.runtime", diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index bdd48d20f..475a0434c 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -121,6 +121,9 @@ mod native { subscribers: &[EventSubscriberFn], scope_stack: ScopeStackHandle, ) -> bool { + if subscribers.is_empty() { + return true; + } let message = DispatcherMessage::Deliver { event: Box::new(event), transform: Some(transform), diff --git a/crates/core/src/api/scope.rs b/crates/core/src/api/scope.rs index 24e5d8cc3..635c6c76c 100644 --- a/crates/core/src/api/scope.rs +++ b/crates/core/src/api/scope.rs @@ -217,8 +217,8 @@ pub fn get_handle() -> Result { /// cannot be read safely. /// /// # Notes -/// Scope-local subscribers attached to ancestor scopes observe the emitted -/// start event before the function returns. +/// The start event is queued with subscriber and sanitizer snapshots captured +/// while the new scope is active. pub fn push_scope(params: PushScopeParams<'_>) -> Result { ensure_runtime_owner()?; let parent_uuid = resolve_parent_uuid(params.parent); @@ -353,8 +353,8 @@ pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { /// cannot be read safely. /// /// # Notes -/// Scope-local subscribers attached to ancestor scopes observe the emitted -/// mark event just like scope, tool, and LLM lifecycle events. +/// The mark event is queued with subscriber and sanitizer snapshots captured +/// from the active scope stack. pub fn event(params: EmitMarkEventParams<'_>) -> Result<()> { ensure_runtime_owner()?; let parent_uuid = resolve_parent_uuid(params.parent); diff --git a/crates/core/src/api/tool.rs b/crates/core/src/api/tool.rs index 7d6d10f71..0c1128d39 100644 --- a/crates/core/src/api/tool.rs +++ b/crates/core/src/api/tool.rs @@ -281,26 +281,25 @@ pub fn tool_call(params: ToolCallParams<'_>) -> Result { (handle, event, marks) }; let tool_name = handle.name.clone(); - if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { - dispatch_transformed_event( - event, - Box::new(move |mut event| { - Box::pin(async move { - let sanitized = NemoRelayContextState::tool_sanitize_request_snapshot_chain( - &tool_name, raw_args, &entries, - ) - .await; - let mut fields = event.sanitize_fields(); - fields.data = Some(sanitized); - event.apply_sanitize_fields(fields); - event - }) - }), - event_sanitizers, - &subscribers, - scope_stack.clone(), - ); - } + let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default(); + dispatch_transformed_event( + event, + Box::new(move |mut event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::tool_sanitize_request_snapshot_chain( + &tool_name, raw_args, &entries, + ) + .await; + let mut fields = event.sanitize_fields(); + fields.data = Some(sanitized); + event.apply_sanitize_fields(fields); + event + }) + }), + event_sanitizers, + &subscribers, + scope_stack.clone(), + ); for mark in marks { if let Some(sanitizers) = snapshot_event_sanitizers(&mark, &scope_stack) { dispatch_sanitized_event(mark, sanitizers, &subscribers, scope_stack.clone()); @@ -463,30 +462,29 @@ pub fn tool_call_end(params: ToolCallEndParams<'_>) -> Result<()> { ) }; let tool_name = params.handle.name.clone(); - if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { - dispatch_transformed_event( - event, - Box::new(move |mut event| { - Box::pin(async move { - let sanitized = NemoRelayContextState::tool_sanitize_response_snapshot_chain( - &tool_name, result, &entries, - ) - .await; - let mut fields = event.sanitize_fields(); - fields.data = if sanitized.is_null() { - fallback - } else { - Some(sanitized) - }; - event.apply_sanitize_fields(fields); - event - }) - }), - event_sanitizers, - &subscribers, - scope_stack, - ); - } + let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default(); + dispatch_transformed_event( + event, + Box::new(move |mut event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::tool_sanitize_response_snapshot_chain( + &tool_name, result, &entries, + ) + .await; + let mut fields = event.sanitize_fields(); + fields.data = if sanitized.is_null() { + fallback + } else { + Some(sanitized) + }; + event.apply_sanitize_fields(fields); + event + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); Ok(()) } diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 7481d86ee..084a8900b 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -1341,8 +1341,39 @@ fn make_user_data( } /// One-shot state retained by a v3 native async callback. +enum NativeAsyncResult { + Json(Json), + LlmStream(LlmJsonStream), +} + +impl std::fmt::Debug for NativeAsyncResult { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Json(value) => formatter.debug_tuple("Json").field(value).finish(), + Self::LlmStream(_) => formatter.write_str("LlmStream(..)"), + } + } +} + +impl PartialEq for NativeAsyncResult { + fn eq(&self, other: &Json) -> bool { + matches!(self, Self::Json(value) if value == other) + } +} + +impl NativeAsyncResult { + fn into_json(self) -> FlowResult { + match self { + Self::Json(value) => Ok(value), + Self::LlmStream(_) => Err(FlowError::Internal( + "native async callback returned a stream for a non-stream invocation".into(), + )), + } + } +} + struct NativeAsyncCompletion { - sender: Mutex>>>, + sender: Mutex>>>, cancelled: AtomicBool, // A pending native callback can continue running after its completion // wakes the awaiting task. Keep the callback's dynamic-library instance @@ -1352,7 +1383,7 @@ struct NativeAsyncCompletion { struct NativeAsyncWait { completion: Arc, - receiver: tokio::sync::oneshot::Receiver>, + receiver: tokio::sync::oneshot::Receiver>, } impl Drop for NativeAsyncWait { @@ -1380,7 +1411,7 @@ async fn invoke_native_async_callback( user_data: Arc, invocation: Json, next: Option, -) -> FlowResult { +) -> FlowResult { let runtime = if next.is_some() { Some(tokio::runtime::Handle::try_current().map_err(|error| { FlowError::Internal(format!( @@ -1486,7 +1517,7 @@ unsafe extern "C" fn native_async_completion_resolve_json( else { return NemoRelayStatus::InvalidArg; }; - let _ = sender.send(Ok(value)); + let _ = sender.send(Ok(NativeAsyncResult::Json(value))); NemoRelayStatus::Ok } @@ -1563,11 +1594,14 @@ unsafe extern "C" fn native_async_next_invoke( }; unsafe { Arc::increment_strong_count(completion as *const NativeAsyncCompletion) }; let completion = unsafe { Arc::from_raw(completion as *const NativeAsyncCompletion) }; - let future: Pin> + Send>> = match &next.inner { + let future: Pin> + Send>> = match &next + .inner + { NativeAsyncNextInner::Tool(next) => { let next = next.clone(); Box::pin(async move { serde_json::to_value(ToolExecutionInterceptOutcome::new(next(invocation).await?)) + .map(NativeAsyncResult::Json) .map_err(|error| { FlowError::Internal(format!( "failed to serialize native async tool outcome: {error}" @@ -1589,7 +1623,7 @@ unsafe extern "C" fn native_async_next_invoke( } }; let next = next.clone(); - Box::pin(async move { next(request).await }) + Box::pin(async move { next(request).await.map(NativeAsyncResult::Json) }) } NativeAsyncNextInner::LlmStream(next) => { let request = match serde_json::from_value(invocation) { @@ -1605,14 +1639,7 @@ unsafe extern "C" fn native_async_next_invoke( } }; let next = next.clone(); - Box::pin(async move { - let mut stream = next(request).await?; - let mut chunks = Vec::new(); - while let Some(chunk) = stream.next().await { - chunks.push(chunk?); - } - Ok(Json::Array(chunks)) - }) + Box::pin(async move { next(request).await.map(NativeAsyncResult::LlmStream) }) } }; next.runtime.spawn(async move { @@ -1645,7 +1672,8 @@ fn wrap_native_async_tool_json( serde_json::json!({"name": name, "value": value}), None, ) - .await?; + .await? + .into_json()?; Ok(value) }) }) @@ -1668,6 +1696,7 @@ fn wrap_native_async_tool_conditional( None, ) .await? + .into_json()? { Json::Null => Ok(None), Json::String(reason) => Ok(Some(reason)), @@ -1696,6 +1725,7 @@ fn wrap_native_async_llm_conditional( None, ) .await? + .into_json()? { Json::Null => Ok(None), Json::String(reason) => Ok(Some(reason)), @@ -1724,7 +1754,8 @@ fn wrap_native_async_llm_sanitize_request( serde_json::json!({"request": request, "context": {"codec": codec}}), None, ) - .await?; + .await? + .into_json()?; if value.is_null() { Ok(None) } else { @@ -1753,7 +1784,8 @@ fn wrap_native_async_llm_sanitize_response( serde_json::json!({"response": response, "context": {"codec": codec}}), None, ) - .await?; + .await? + .into_json()?; Ok((!value.is_null()).then_some(value)) }) }) @@ -1780,7 +1812,8 @@ fn wrap_native_async_llm_request_intercept( }), None, ) - .await?, + .await? + .into_json()?, ) .map_err(|error| { FlowError::Internal(format!( @@ -1808,7 +1841,8 @@ fn wrap_native_async_event_sanitize( serde_json::json!({"event": event, "fields": fields}), None, ) - .await?, + .await? + .into_json()?, ) .map_err(|error| { FlowError::Internal(format!("invalid native async event fields: {error}")) @@ -1835,7 +1869,8 @@ fn wrap_native_async_tool_execution( invocation, Some(NativeAsyncNextInner::Tool(next)), ) - .await?, + .await? + .into_json()?, ) .map_err(|error| { FlowError::Internal(format!("invalid native async tool outcome: {error}")) @@ -1853,12 +1888,17 @@ fn wrap_native_async_llm_execution( let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |name, request, next| { let user_data = user_data.clone(); - Box::pin(invoke_native_async_callback( - cb, - user_data, - serde_json::json!({"name": name, "request": request}), - Some(NativeAsyncNextInner::Llm(next)), - )) + let name = name.to_owned(); + Box::pin(async move { + invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "request": request}), + Some(NativeAsyncNextInner::Llm(next)), + ) + .await? + .into_json() + }) }) } @@ -1880,14 +1920,15 @@ fn wrap_native_async_llm_stream_execution( Some(NativeAsyncNextInner::LlmStream(next)), ) .await?; - let chunks = value.as_array().cloned().ok_or_else(|| { - FlowError::Internal( + match value { + NativeAsyncResult::LlmStream(stream) => Ok(stream), + NativeAsyncResult::Json(Json::Array(chunks)) => Ok(LlmJsonStream::new( + tokio_stream::iter(chunks.into_iter().map(Ok)), + )), + NativeAsyncResult::Json(_) => Err(FlowError::Internal( "native async LLM stream intercept must resolve to an array".into(), - ) - })?; - Ok(LlmJsonStream::new(tokio_stream::iter( - chunks.into_iter().map(Ok), - ))) + )), + } }) }) } diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index 5262dd784..1d0bb2ecf 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -1120,7 +1120,7 @@ impl WorkerPluginInstance { let instance = Arc::new(self.clone_for_callback()); let callback_name = name.to_owned(); let callback: EventSanitizeFn = - Arc::new(move |event: Event, _fields: EventSanitizeFields| { + Arc::new(move |event: Arc, _fields: EventSanitizeFields| { let instance = instance.clone(); let callback_name = callback_name.clone(); Box::pin(async move { @@ -1875,12 +1875,15 @@ impl WorkerPluginCallback { .invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) .await; if let Err(error) = &result { + let surface_name = RegistrationSurface::try_from(surface) + .map(|surface| surface.as_str_name()) + .unwrap_or("UNKNOWN"); log::warn!( target: "nemo_relay.worker", event = "worker_callback_failed", plugin_id = self.plugin_kind.as_str(), callback = callback_name.as_str(), - surface; + surface = surface_name; "Worker plugin callback failed: {error}" ); } diff --git a/crates/core/src/stream.rs b/crates/core/src/stream.rs index 756f066a7..2e742fea3 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -344,12 +344,14 @@ impl LlmStreamWrapper { match tokio::runtime::Handle::try_current() { Ok(handle) => Some(handle.spawn(finalize)), Err(_) => { - if let Ok(runtime) = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - { - runtime.block_on(finalize); - } + std::thread::spawn(move || { + if let Ok(runtime) = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + { + runtime.block_on(finalize); + } + }); None } } @@ -420,7 +422,10 @@ impl Stream for LlmStreamWrapper { } if this.ended { - return Poll::Ready(None); + return match this.terminal_result.take() { + Some(result) => Poll::Ready(Some(result)), + None => Poll::Ready(None), + }; } // Poll the inner stream diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index acade0d0c..1367f4d74 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -797,6 +797,9 @@ unsafe extern "C" fn raw_async_tool_execution_callback( }; if next.is_null() || completion.is_null() { unsafe { reject_async_completion(host, completion, "async tool execution requires next and completion") }; + if !next.is_null() { + unsafe { (host.async_next_release)(next) }; + } return NemoRelayNativeAsyncCallbackState::Complete; } let value = unsafe { raw_host_string_value(&host.v1, invocation_json) } @@ -812,11 +815,13 @@ unsafe extern "C" fn raw_async_tool_execution_callback( .and_then(|value| serde_json::to_string(&value).ok()); let Some(value) = value else { unsafe { reject_async_completion(host, completion, "invalid async tool execution invocation") }; + unsafe { (host.async_next_release)(next) }; return NemoRelayNativeAsyncCallbackState::Complete; }; let value = unsafe { raw_host_string(&host.v1, &value) }; if value.is_null() { unsafe { reject_async_completion(host, completion, "failed to allocate async tool execution invocation") }; + unsafe { (host.async_next_release)(next) }; return NemoRelayNativeAsyncCallbackState::Complete; } let status = unsafe { (host.async_next_invoke)(next, value, completion) }; @@ -831,6 +836,7 @@ unsafe extern "C" fn raw_async_tool_execution_callback( NemoRelayNativeAsyncCallbackState::Pending } else { unsafe { reject_async_completion(host, completion, "failed to invoke async tool execution next") }; + unsafe { (host.async_next_release)(next) }; NemoRelayNativeAsyncCallbackState::Complete } } diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index a29955c17..934e0c3ee 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -284,22 +284,6 @@ fn native_async_next_abi_runs_tool_llm_and_stream_continuations() { .unwrap(), json!({"llm": true}), ), - ( - NativeAsyncNextInner::LlmStream(Arc::new(|_request| { - Box::pin(async { - Ok(LlmJsonStream::new(tokio_stream::iter(vec![ - Ok(json!({"chunk": 1})), - Ok(json!({"chunk": 2})), - ]))) - }) - })), - serde_json::to_value(LlmRequest { - headers: Map::new(), - content: json!({"stream": true}), - }) - .unwrap(), - json!([{"chunk": 1}, {"chunk": 2}]), - ), ]; for (inner, invocation, expected) in cases { @@ -329,6 +313,53 @@ fn native_async_next_abi_runs_tool_llm_and_stream_continuations() { native_async_completion_release(completion_ref); } } + + let next = Arc::new(NativeAsyncNext { + inner: NativeAsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![ + Ok(json!({"chunk": 1})), + Ok(json!({"chunk": 2})), + ]))) + }) + })), + runtime: runtime.handle().clone(), + _callback_user_data: None, + }); + let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invocation = native_string_from_json( + &serde_json::to_value(LlmRequest { + headers: Map::new(), + content: json!({"stream": true}), + }) + .unwrap(), + ) + .unwrap(); + assert_eq!( + unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, + NemoRelayStatus::Ok + ); + let NativeAsyncResult::LlmStream(mut stream) = runtime.block_on(receiver).unwrap().unwrap() + else { + panic!("stream continuation should preserve the downstream stream"); + }; + assert_eq!( + runtime.block_on(stream.next()).unwrap().unwrap(), + json!({"chunk": 1}) + ); + unsafe { + native_string_free(invocation); + native_async_next_release(next_ref); + native_async_completion_release(completion_ref); + } } #[test] diff --git a/crates/ffi/src/api/mod.rs b/crates/ffi/src/api/mod.rs index f2984f32e..1fc8977fd 100644 --- a/crates/ffi/src/api/mod.rs +++ b/crates/ffi/src/api/mod.rs @@ -100,6 +100,8 @@ fn tokio_runtime() -> &'static Runtime { } fn block_on_sync_ffi(future: impl Future>) -> FlowResult { + // Embedded hosts must not call synchronous middleware helpers from a Tokio + // runtime thread. Use the completion-based async registration API there. if tokio::runtime::Handle::try_current().is_ok() { return Err(nemo_relay::error::FlowError::Internal( "synchronous FFI middleware helpers cannot run on a Tokio runtime thread; use the completion-based async registration API".into(), diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index f1c952863..e876cbe34 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -939,10 +939,10 @@ pub fn wrap_event_sanitize_fn( free_fn: NemoRelayFreeFn, ) -> EventSanitizeFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |event: Event, fields: EventSanitizeFields| { + Arc::new(move |event: Arc, fields: EventSanitizeFields| { let ud = ud.clone(); Box::pin(async move { - let ffi_event = FfiEvent(event); + let ffi_event = FfiEvent((*event).clone()); let fields_json = json_to_c_string(&serde_json::to_value(&fields).unwrap_or(Json::Null)); let result_ptr = unsafe { cb(ud.ptr, &ffi_event, fields_json) }; @@ -1071,6 +1071,10 @@ unsafe fn nemo_relay_string_free_internal(ptr: *mut c_char) { } } +#[cfg(test)] +#[path = "../tests/support/mod.rs"] +mod test_support; + #[cfg(test)] #[path = "../tests/unit/callable_tests.rs"] mod tests; diff --git a/crates/ffi/tests/integration/callable_extra_tests.rs b/crates/ffi/tests/integration/callable_extra_tests.rs index 69a341b96..2a284f07e 100644 --- a/crates/ffi/tests/integration/callable_extra_tests.rs +++ b/crates/ffi/tests/integration/callable_extra_tests.rs @@ -4,18 +4,11 @@ //! Integration tests for callable extra in the NeMo Relay FFI crate. use super::*; -use std::future::Future; use std::ptr; use tokio_stream::StreamExt; -fn resolve(future: impl Future) -> T { - tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .unwrap() - .block_on(future) -} +use super::test_support::resolve; unsafe extern "C" fn tool_conditional_error_cb( _user_data: *mut libc::c_void, diff --git a/crates/ffi/tests/integration/main.rs b/crates/ffi/tests/integration/main.rs index baed8d570..e8f3035f4 100644 --- a/crates/ffi/tests/integration/main.rs +++ b/crates/ffi/tests/integration/main.rs @@ -36,5 +36,7 @@ mod convert_coverage_tests; #[path = "../coverage/error_tests.rs"] mod error_coverage_tests; mod plugin_activation_tests; +#[path = "../support/mod.rs"] +mod test_support; #[path = "../unit/types_tests.rs"] mod types_tests; diff --git a/crates/ffi/tests/support/mod.rs b/crates/ffi/tests/support/mod.rs new file mode 100644 index 000000000..8e557b330 --- /dev/null +++ b/crates/ffi/tests/support/mod.rs @@ -0,0 +1,14 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Shared helpers for FFI tests. + +use std::future::Future; + +pub(crate) fn resolve(future: impl Future) -> T { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap() + .block_on(future) +} diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 4b5a2d251..1ae03a985 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -4,7 +4,6 @@ //! Unit tests for callable in the NeMo Relay FFI crate. use super::*; -use std::future::Future; use std::sync::atomic::{AtomicUsize, Ordering}; use nemo_relay::api::event::{Event, EventSanitizeFields}; @@ -12,6 +11,8 @@ use nemo_relay::api::llm::{LlmAttributes, LlmHandle}; use serde_json::json; use tokio_stream::StreamExt; +use super::test_support::resolve; + extern "C" fn free_arc_counter(user_data: *mut libc::c_void) { let counter = unsafe { Box::from_raw(user_data as *mut Arc) }; counter.fetch_add(1, Ordering::SeqCst); @@ -23,14 +24,6 @@ fn user_data_counter() -> (*mut libc::c_void, Arc) { (ptr, counter) } -fn resolve(future: impl Future) -> T { - tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .unwrap() - .block_on(future) -} - unsafe extern "C" fn tool_sanitize_cb( user_data: *mut libc::c_void, name: *const c_char, @@ -668,7 +661,7 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { .build(); let (user_data, sanitize_calls) = user_data_counter(); let sanitizer = wrap_event_sanitize_fn(event_sanitize_cb, user_data, Some(free_arc_counter)); - let sanitized = resolve(sanitizer(event.clone(), original_fields.clone())).unwrap(); + let sanitized = resolve(sanitizer(Arc::new(event.clone()), original_fields.clone())).unwrap(); assert_eq!(sanitized.data, Some(json!({"safe": true}))); assert_eq!( sanitized @@ -684,12 +677,12 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { let invalid = wrap_event_sanitize_fn(invalid_event_sanitize_cb, std::ptr::null_mut(), None); assert_eq!( - resolve(invalid(event.clone(), original_fields.clone())).unwrap(), + resolve(invalid(Arc::new(event.clone()), original_fields.clone())).unwrap(), EventSanitizeFields::default() ); let null = wrap_event_sanitize_fn(null_event_sanitize_cb, std::ptr::null_mut(), None); assert_eq!( - resolve(null(event, original_fields.clone())).unwrap(), + resolve(null(Arc::new(event), original_fields.clone())).unwrap(), EventSanitizeFields::default() ); diff --git a/crates/node/plugin.d.ts b/crates/node/plugin.d.ts index 158e25b2f..cc0a130b3 100644 --- a/crates/node/plugin.d.ts +++ b/crates/node/plugin.d.ts @@ -225,7 +225,7 @@ export interface PluginContext { registerToolConditionalExecutionGuardrail( name: string, priority: number, - callback: (name: string, args: Json) => string | null, + callback: (name: string, args: Json) => string | null | Promise, ): void; /** Register an LLM sanitize-request guardrail. The callback receives `(request, context)`. */ registerLlmSanitizeRequestGuardrail( @@ -272,7 +272,7 @@ export interface PluginContext { name: string, priority: number, breakChain: boolean, - callback: (name: string, args: Json) => Json, + callback: (name: string, args: Json) => Json | Promise, ): void; /** * Register tool execution middleware that returns a canonical outcome. diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index b40aba073..c2ef47490 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -300,7 +300,13 @@ pub fn wrap_js_llm_sanitize_request_promise_fn(func: Arc) -> Llm move |request: LlmRequest, context: LlmSanitizeRequestContext| { let func = func.clone(); Box::pin(async move { - let request = serde_json::to_value(request).unwrap_or(Json::Null); + let request = serde_json::to_value(request).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM sanitize request: {error}" + )); + record_callback_error(error.to_string()); + error + })?; let context = js_llm_sanitize_request_context(&context); let value = func .call_spread_with_arg0(Box::new(move |env| { @@ -375,8 +381,15 @@ pub fn wrap_js_llm_conditional_promise_fn(func: Arc) -> LlmCondi Arc::new(move |request: LlmRequest| { let func = func.clone(); Box::pin(async move { + let request = serde_json::to_value(request).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM conditional request: {error}" + )); + record_callback_error(error.to_string()); + error + })?; let value = func - .call(serde_json::to_value(request).unwrap_or(Json::Null)) + .call(request) .await .inspect_err(|error| record_callback_error(error.to_string()))?; match value { @@ -402,16 +415,28 @@ pub fn wrap_js_llm_request_intercept_promise_fn( move |name: String, request: LlmRequest, annotated: Option| { let func = func.clone(); Box::pin(async move { - let value = func - .call(serde_json::json!({ - "name": name, - "request": request, - "annotated": annotated, - })) - .await - .inspect_err(|error| { - record_callback_error(error.to_string()); - })?; + let request = serde_json::to_value(request).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM request intercept request: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let annotated = serde_json::to_value(annotated).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM request intercept annotation: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let value = serde_json::json!({ + "name": name, + "request": request, + "annotated": annotated, + }); + let value = func.call(value).await.inspect_err(|error| { + record_callback_error(error.to_string()); + })?; #[derive(Deserialize)] #[serde(rename_all = "camelCase")] struct JsOutcome { @@ -448,7 +473,7 @@ pub fn wrap_js_llm_request_intercept_promise_fn( /// scope/mark APIs while allowing the JavaScript callback to settle a Promise /// on the Node event loop. pub fn wrap_js_event_sanitize_promise_fn(func: Arc) -> EventSanitizeFn { - Arc::new(move |event: Event, fields: CoreEventSanitizeFields| { + Arc::new(move |event: Arc, fields: CoreEventSanitizeFields| { let func = func.clone(); Box::pin(async move { let event_json = JsEvent::try_from_event(&event) @@ -753,7 +778,7 @@ pub fn wrap_js_llm_sanitize_request_fn( })?; let (tx, rx) = tokio::sync::oneshot::channel(); if func.call_with_return_value( - (request.clone(), context), + (request, context), ThreadsafeFunctionCallMode::Blocking, move |value: Option| { let _ = tx.send(callback_json(value)); @@ -1112,7 +1137,7 @@ pub fn wrap_js_event_sanitize_fn( func: ThreadsafeFunction<(Json, Json), ErrorStrategy::Fatal>, ) -> EventSanitizeFn { let func = Arc::new(func); - Arc::new(move |event: Event, fields: CoreEventSanitizeFields| { + Arc::new(move |event: Arc, fields: CoreEventSanitizeFields| { let func = func.clone(); Box::pin(async move { let event_json = match JsEvent::try_from_event(&event) { @@ -1125,7 +1150,7 @@ pub fn wrap_js_event_sanitize_fn( } }; let js_fields = EventSanitizeFields { - data: fields.data.clone(), + data: fields.data, category_profile: fields .category_profile .as_ref() @@ -1138,7 +1163,7 @@ pub fn wrap_js_event_sanitize_fn( record_callback_error(error.to_string()); error })?, - metadata: fields.metadata.clone(), + metadata: fields.metadata, }; let js_fields = serde_json::to_value(js_fields).map_err(|error| { let error = FlowError::Internal(format!( diff --git a/crates/node/src/callback_factory.rs b/crates/node/src/callback_factory.rs index 48b301883..5fcd45ed8 100644 --- a/crates/node/src/callback_factory.rs +++ b/crates/node/src/callback_factory.rs @@ -68,7 +68,11 @@ const CALLBACK_FACTORIES_SOURCE: &str = r#"(() => { promise(fn) { return function __nemo_relay_promise_wrapper(error, arg0, spread, next, resolve, reject) { if (error != null) { - reject(error); + let message = 'unknown error'; + try { + message = String(error?.message ?? error); + } catch {} + reject(message); return; } Promise.resolve().then(() => ( @@ -80,6 +84,8 @@ const CALLBACK_FACTORIES_SOURCE: &str = r#"(() => { try { if (typeof error === 'string') { message = error; + } else if (error === null || (typeof error !== 'object' && typeof error !== 'function')) { + message = String(error); } else if (error != null && typeof error.message === 'string') { message = error.message; } diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index c83c78cb9..59f63ba33 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -58,14 +58,15 @@ async function flushSubscriberCallbacks() { } async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { - flushSubscribers(); const deadline = Date.now() + timeoutMs; while (!predicate()) { + await flushSubscribers(); if (Date.now() >= deadline) { throw new Error('timed out waiting for subscriber callbacks'); } await new Promise((resolve) => setImmediate(resolve)); } + await flushSubscribers(); } function makeNative() { @@ -764,7 +765,7 @@ describe('LLM guardrails', () => { event.scope_category === 'start', ); assert.deepEqual(start.data, { headers: request.headers, content: request.content }); - assert.match(getLastCallbackError() ?? '', /(unknown error|callback)/i); + assert.equal(getLastCallbackError(), 'internal error: unknown error'); deregisterLlmSanitizeRequestGuardrail('node_llm_san_req_throw'); const result = await llmCallExecute( @@ -1197,7 +1198,7 @@ describe('LLM intercepts', () => { deregisterLlmExecutionIntercept('node_llm_exec_invalid_next'); }); - it('execution intercept propagates primitive rejection values as unknown error', async () => { + it('execution intercept preserves primitive rejection values', async () => { registerLlmExecutionIntercept('node_llm_exec_unknown_err', 10, async () => { return rejectWith(42); }); @@ -1216,14 +1217,14 @@ describe('LLM intercepts', () => { null, null, ), - /unknown error/i, + /internal error: 42/i, ); } finally { deregisterLlmExecutionIntercept('node_llm_exec_unknown_err'); } }); - it('async execute falls back to unknown error for primitive rejections', async () => { + it('async execute preserves primitive rejection values', async () => { await assert.rejects( () => llmCallExecuteAsync( @@ -1236,7 +1237,7 @@ describe('LLM intercepts', () => { null, null, ), - /unknown error/i, + /internal error: 42/i, ); }); diff --git a/crates/node/tests/scope_tests.mjs b/crates/node/tests/scope_tests.mjs index 9ecd3ebb2..0332790ac 100644 --- a/crates/node/tests/scope_tests.mjs +++ b/crates/node/tests/scope_tests.mjs @@ -30,14 +30,15 @@ function rejectWithPrimitive(value) { } async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { - await flushSubscribers(); const deadline = Date.now() + timeoutMs; while (!predicate()) { + await flushSubscribers(); if (Date.now() >= deadline) { throw new Error('timed out waiting for subscriber callbacks'); } await new Promise((resolve) => setImmediate(resolve)); } + await flushSubscribers(); } // =========================================================================== @@ -277,14 +278,14 @@ describe('withScope', () => { } }); - it('surfaces primitive rejection values as unknown error and still pops the scope', async () => { + it('surfaces primitive rejection values and still pops the scope', async () => { const before = getHandle(); await assert.rejects( () => withScope('primitive_reject_test', ScopeType.Tool, async () => { return rejectWithPrimitive(123); }), - /unknown error/i, + /internal error: 123/i, ); const after = getHandle(); assert.equal(after.uuid, before.uuid, 'scope should be popped after primitive rejection'); diff --git a/crates/node/tests/tools_tests.mjs b/crates/node/tests/tools_tests.mjs index b60642aef..e104d0611 100644 --- a/crates/node/tests/tools_tests.mjs +++ b/crates/node/tests/tools_tests.mjs @@ -55,11 +55,13 @@ async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { // callback state is ready, with a timeout to avoid hanging the test forever. const deadline = Date.now() + timeoutMs; while (!predicate()) { + await flushSubscribers(); if (Date.now() >= deadline) { throw new Error('timed out waiting for subscriber callbacks'); } await new Promise((resolve) => setImmediate(resolve)); } + await flushSubscribers(); } // =========================================================================== @@ -996,7 +998,7 @@ describe('Tool intercepts', () => { } }); - it('async execute falls back to unknown error for primitive rejections', async () => { + it('async execute preserves primitive rejection values', async () => { await assert.rejects( () => toolCallExecuteAsync( @@ -1010,7 +1012,7 @@ describe('Tool intercepts', () => { null, null, ), - /unknown error/i, + /internal error: 42/i, ); }); diff --git a/crates/pii-redaction/src/builtin.rs b/crates/pii-redaction/src/builtin.rs index 71d421588..8852f8543 100644 --- a/crates/pii-redaction/src/builtin.rs +++ b/crates/pii-redaction/src/builtin.rs @@ -461,8 +461,9 @@ impl CompiledBuiltinBackend { } pub(super) fn tool_sanitize_callback(backend: CompiledBuiltinBackend) -> ToolSanitizeFn { + let backend = Arc::new(backend); Arc::new(move |_name: String, payload: Json| { - let backend = backend.clone(); + let backend = Arc::clone(&backend); Box::pin(async move { Ok(match backend.trajectory.as_ref() { Some(trajectory) => trajectory.sanitize_tool_payload(payload), @@ -488,11 +489,12 @@ fn event_sanitize_callback_with_scope_categories( backend: CompiledBuiltinBackend, scope_categories: Option<(bool, bool)>, ) -> EventSanitizeFn { + let backend = Arc::new(backend); Arc::new(move |event, mut fields| { - let backend = backend.clone(); + let backend = Arc::clone(&backend); Box::pin(async move { if scope_categories.is_some_and(|(sanitize_llm, sanitize_tool)| { - matches!(event, Event::Scope(_)) + matches!(event.as_ref(), Event::Scope(_)) && event .category() .is_some_and(|category| match category.as_str() { @@ -507,7 +509,7 @@ fn event_sanitize_callback_with_scope_categories( if let Some(trajectory) = backend.trajectory.as_ref() { return Ok(trajectory.sanitize_event_fields(&event, fields)); } - let specialized_scope = matches!(event, Event::Scope(_)) + let specialized_scope = matches!(event.as_ref(), Event::Scope(_)) && event .category() .is_some_and(|category| matches!(category.as_str(), "tool" | "llm")); @@ -532,8 +534,9 @@ fn event_sanitize_callback_with_scope_categories( pub(super) fn llm_sanitize_request_callback( backend: CompiledBuiltinBackend, ) -> LlmSanitizeRequestFn { + let backend = Arc::new(backend); Arc::new(move |mut request: LlmRequest, context| { - let backend = backend.clone(); + let backend = Arc::clone(&backend); Box::pin(async move { if let Some(trajectory) = backend.trajectory.as_ref() { request.headers = trajectory @@ -577,8 +580,9 @@ pub(super) fn llm_sanitize_request_callback( pub(super) fn llm_sanitize_response_callback( backend: CompiledBuiltinBackend, ) -> LlmSanitizeResponseFn { + let backend = Arc::new(backend); Arc::new(move |payload: Json, context| { - let backend = backend.clone(); + let backend = Arc::clone(&backend); Box::pin(async move { if let Some(trajectory) = backend.trajectory.as_ref() { return Ok(Some(trajectory.sanitize_provider_payload(payload))); diff --git a/crates/pii-redaction/tests/unit/component_tests.rs b/crates/pii-redaction/tests/unit/component_tests.rs index 2152c5d6d..2e80f1172 100644 --- a/crates/pii-redaction/tests/unit/component_tests.rs +++ b/crates/pii-redaction/tests/unit/component_tests.rs @@ -788,7 +788,7 @@ async fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { Some(CategoryProfile::builder().subtype("llm.chunk").build()), )); let sanitized = callback( - chunk.clone(), + Arc::new(chunk.clone()), EventSanitizeFields { data: Some(json!({ "chunk_index": 2, @@ -823,7 +823,7 @@ async fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { ), )); let sanitized = callback( - optimization.clone(), + Arc::new(optimization.clone()), EventSanitizeFields { data: Some(json!({ "producer": "neutral.router", @@ -867,7 +867,7 @@ async fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { None, )); let sanitized = callback( - nested_agent, + Arc::new(nested_agent), EventSanitizeFields { data: Some(json!({ "request_id": "request-1", @@ -962,7 +962,7 @@ async fn trajectory_preset_preserves_trusted_scope_metadata_only() { None, )); let sanitized = callback( - event, + Arc::new(event), EventSanitizeFields { data: None, category_profile: None, @@ -982,7 +982,7 @@ async fn trajectory_preset_preserves_trusted_scope_metadata_only() { None, )); let sanitized = callback( - malformed, + Arc::new(malformed), EventSanitizeFields { data: None, category_profile: None, @@ -1012,7 +1012,7 @@ async fn trajectory_preset_preserves_trusted_scope_metadata_only() { Some(CategoryProfile::builder().subtype("llm.chunk").build()), )); let sanitized = callback( - mark.clone(), + Arc::new(mark.clone()), EventSanitizeFields { data: None, category_profile: mark.category_profile().cloned(), @@ -1046,13 +1046,15 @@ async fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { let preserve = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); assert_eq!( - preserve(event.clone(), fields.clone()).await.unwrap(), + preserve(Arc::new(event.clone()), fields.clone()) + .await + .unwrap(), fields ); let redact = crate::builtin::event_sanitize_callback(trajectory_backend(None, "redact_all_leaves")); - let sanitized = redact(event, fields).await.unwrap(); + let sanitized = redact(Arc::new(event), fields).await.unwrap(); assert_eq!( sanitized.data.unwrap(), json!({ @@ -1117,7 +1119,7 @@ async fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annot None, )); let sanitized = callback( - event, + Arc::new(event), EventSanitizeFields { data: Some(json!({"already": "sanitized by the response callback"})), category_profile: Some( @@ -1192,8 +1194,8 @@ async fn preserved_custom_marks_remain_eligible_for_a_later_email_profile() { .unwrap(), ); - let fields = trajectory(event.clone(), fields).await.unwrap(); - let sanitized = email(event, fields).await.unwrap(); + let fields = trajectory(Arc::new(event.clone()), fields).await.unwrap(); + let sanitized = email(Arc::new(event), fields).await.unwrap(); assert_eq!(sanitized.data.as_ref().unwrap()["owner"], "[REDACTED]"); assert_eq!(sanitized.data.as_ref().unwrap()["score"], 0.9); assert_eq!( @@ -1505,7 +1507,7 @@ async fn event_sanitizer_transforms_data_category_profile_and_metadata_independe None, )); let sanitized = callback( - event, + Arc::new(event), EventSanitizeFields { data: Some(json!({"email": "person@example.com"})), category_profile: Some( @@ -1551,7 +1553,7 @@ async fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fie .subtype("person@example.com") .build(); let sanitized = callback( - event, + Arc::new(event), EventSanitizeFields { data: Some(json!({"content": "person@example.com"})), category_profile: Some(original_profile.clone()), @@ -1607,7 +1609,7 @@ async fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { .subtype("person@example.com") .build(); let sanitized = callback( - event, + Arc::new(event), EventSanitizeFields { data: Some(json!({"content": "person@example.com"})), category_profile: Some(original_profile.clone()), @@ -1643,7 +1645,7 @@ async fn event_sanitizer_discards_category_profile_when_sanitization_fails() { None, )); let sanitized = callback( - event, + Arc::new(event), EventSanitizeFields { data: None, category_profile: Some(CategoryProfile { diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index 80b9d672e..bf6280c04 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1327,14 +1327,13 @@ fn tool_request_intercepts<'py>( .is_err() { let scope_stack = current_scope_stack_handle(); - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; - let result = runtime - .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { - core_tool_api::tool_request_intercepts(&name, args_json).await - })) + let result = pyo3_async_runtimes::tokio::get_runtime() + .block_on(py_callable::PY_AWAITABLES_ALLOWED.scope( + false, + TASK_SCOPE_STACK.scope(scope_stack, async move { + core_tool_api::tool_request_intercepts(&name, args_json).await + }), + )) .map_err(to_py_err)?; return json_to_py(py, &result).map(|value| value.into_bound(py)); } @@ -1371,14 +1370,13 @@ fn tool_conditional_execution<'py>( .is_err() { let scope_stack = current_scope_stack_handle(); - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; - runtime - .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { - core_tool_api::tool_conditional_execution(&name, &args_json).await - })) + pyo3_async_runtimes::tokio::get_runtime() + .block_on(py_callable::PY_AWAITABLES_ALLOWED.scope( + false, + TASK_SCOPE_STACK.scope(scope_stack, async move { + core_tool_api::tool_conditional_execution(&name, &args_json).await + }), + )) .map_err(to_py_err)?; return Ok(py.None().into_bound(py)); } @@ -1415,14 +1413,13 @@ fn llm_request_intercepts<'py>( .is_err() { let scope_stack = current_scope_stack_handle(); - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; - let result = runtime - .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { - core_llm_api::llm_request_intercepts(&name, request.inner).await - })) + let result = pyo3_async_runtimes::tokio::get_runtime() + .block_on(py_callable::PY_AWAITABLES_ALLOWED.scope( + false, + TASK_SCOPE_STACK.scope(scope_stack, async move { + core_llm_api::llm_request_intercepts(&name, request.inner).await + }), + )) .map_err(to_py_err)?; return Py::new( py, @@ -1460,14 +1457,13 @@ fn llm_conditional_execution<'py>( .is_err() { let scope_stack = current_scope_stack_handle(); - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; - runtime - .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { - core_llm_api::llm_conditional_execution(&request.inner).await - })) + pyo3_async_runtimes::tokio::get_runtime() + .block_on(py_callable::PY_AWAITABLES_ALLOWED.scope( + false, + TASK_SCOPE_STACK.scope(scope_stack, async move { + core_llm_api::llm_conditional_execution(&request.inner).await + }), + )) .map_err(to_py_err)?; return Ok(py.None().into_bound(py)); } diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index f5cb02e55..71d423e49 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -53,6 +53,23 @@ use crate::py_types::{ type PyValueFuture = Pin>> + Send>>; +tokio::task_local! { + pub(crate) static PY_AWAITABLES_ALLOWED: bool; +} + +fn reject_awaitable_from_sync_caller(result: &Bound<'_, PyAny>) -> FlowResult<()> { + if PY_AWAITABLES_ALLOWED + .try_with(|allowed| *allowed) + .unwrap_or(true) + { + return Ok(()); + } + let _ = result.call_method0("close"); + Err(FlowError::Internal( + "awaitable Python middleware requires an async caller".into(), + )) +} + fn validate_python_llm_sanitizer_signature(py_fn: &Py) -> PyResult<()> { Python::attach(|py| { let inspect = py.import("inspect")?; @@ -79,6 +96,7 @@ fn split_json_or_future( ) -> FlowResult> { let bound = result.bind(py); if bound.getattr("__await__").is_ok() { + reject_awaitable_from_sync_caller(bound)?; let future = pyo3_async_runtimes::tokio::into_future(result.into_bound(py)) .map_err(|e| FlowError::Internal(e.to_string()))?; Ok(Err(Box::pin(future) as PyValueFuture)) @@ -111,6 +129,7 @@ fn split_py_object_or_future( ) -> FlowResult, PyValueFuture>> { let bound = result.bind(py); if bound.getattr("__await__").is_ok() { + reject_awaitable_from_sync_caller(bound)?; let future = pyo3_async_runtimes::tokio::into_future(result.into_bound(py)) .map_err(|e| FlowError::Internal(e.to_string()))?; Ok(Err(Box::pin(future) as PyValueFuture)) @@ -1095,12 +1114,12 @@ pub fn wrap_py_event_subscriber(py_fn: Py) -> EventSubscriberFn { /// Wrap a Python callable ``(Event, EventSanitizeFields) -> EventSanitizeFields``. pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let py_fn = Arc::new(py_fn); - Arc::new(move |event: Event, fields: EventSanitizeFields| { + Arc::new(move |event: Arc, fields: EventSanitizeFields| { let py_fn = py_fn.clone(); Box::pin(async move { let result = Python::attach( |py| -> FlowResult, PyValueFuture>> { - let py_event = match &event { + let py_event = match event.as_ref() { Event::Scope(inner) => Py::new( py, crate::py_types::PyScopeEvent { @@ -1119,27 +1138,18 @@ pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let py_event = match py_event { Ok(value) => value, Err(error) => { - eprintln!( - "nemo_relay: failed to convert event sanitizer context: {error}" - ); return Err(FlowError::Internal(error.to_string())); } }; let fields_json = match serde_json::to_value(&fields) { Ok(value) => value, Err(error) => { - eprintln!( - "nemo_relay: failed to serialize event sanitizer fields: {error}" - ); return Err(FlowError::Internal(error.to_string())); } }; let py_fields = match json_to_py(py, &fields_json) { Ok(value) => value, Err(error) => { - eprintln!( - "nemo_relay: failed to convert event sanitizer fields: {error}" - ); return Err(FlowError::Internal(error.to_string())); } }; diff --git a/crates/python/tests/coverage/py_api_coverage_tests.rs b/crates/python/tests/coverage/py_api_coverage_tests.rs index c1cc50293..5c5d35def 100644 --- a/crates/python/tests/coverage/py_api_coverage_tests.rs +++ b/crates/python/tests/coverage/py_api_coverage_tests.rs @@ -185,6 +185,9 @@ def tool_sanitize_response(name, result): def tool_conditional(name, args): return None if args["value"] >= 0 else "blocked" +async def async_tool_conditional(name, args): + return None + def tool_request_intercept(name, args): updated = dict(args) updated["value"] = updated["value"] + 2 @@ -323,6 +326,16 @@ async def run_llm(api, request, func, handle, attributes, codec, response_codec) response_codec=response_codec, ) +async def run_standalone(api, request): + tool_args = await api.tool_request_intercepts("demo-tool", {"value": 1}) + await api.tool_conditional_execution("demo-tool", tool_args) + llm_outcome = await api.llm_request_intercepts("demo-llm", request) + await api.llm_conditional_execution(llm_outcome.request) + return { + "tool_value": tool_args["value"], + "llm_header": llm_outcome.request.headers["x-intercepted"], + } + async def run_stream(api, request, func, collector, finalizer, handle, attributes, codec, response_codec): stream = await api.llm_stream_call_execute( "demo-stream", @@ -498,6 +511,26 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute .to_string() .contains("blocked") ); + let async_sync_rejection_name = format!("async-sync-{}", Uuid::now_v7()); + register_tool_conditional_execution_guardrail( + &async_sync_rejection_name, + 20, + helpers.getattr("async_tool_conditional").unwrap().unbind(), + ) + .unwrap(); + assert!( + tool_conditional_execution( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": 1})), + ) + .unwrap_err() + .to_string() + .contains("requires an async caller") + ); + assert!( + deregister_tool_conditional_execution_guardrail(&async_sync_rejection_name).unwrap() + ); let llm_request = PyLLMRequest { inner: nemo_relay::api::llm::LlmRequest { @@ -534,6 +567,21 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute ); with_event_loop(py, |event_loop| { + let standalone = event_loop + .call_method1( + "run_until_complete", + (runner + .getattr("run_standalone") + .unwrap() + .call1((api_module.clone(), llm_request.clone())) + .unwrap(),), + ) + .unwrap(); + assert_eq!( + crate::convert::py_to_json(&standalone).unwrap(), + json!({"tool_value": 3, "llm_header": "1"}) + ); + let tool_result = event_loop .call_method1( "run_until_complete", diff --git a/crates/python/tests/coverage/py_callable_coverage_tests.rs b/crates/python/tests/coverage/py_callable_coverage_tests.rs index 89a382f88..53245133c 100644 --- a/crates/python/tests/coverage/py_callable_coverage_tests.rs +++ b/crates/python/tests/coverage/py_callable_coverage_tests.rs @@ -668,7 +668,7 @@ async def collect_stream(awaitable): } #[test] -fn event_sanitize_wrapper_covers_conversion_success_and_fail_closed_paths() { +fn event_sanitize_wrapper_covers_conversion_success_and_error_propagation() { use nemo_relay::api::event::{BaseEvent, MarkEvent}; let _python = crate::test_support::init_python_test(); @@ -704,7 +704,7 @@ def invalid(event, fields): let sanitized = runtime .block_on(wrap_py_event_sanitize_fn( module.getattr("sanitize").unwrap().unbind(), - )(event.clone(), fields.clone())) + )(Arc::new(event.clone()), fields.clone())) .unwrap(); assert_eq!(sanitized.data, Some(json!({"safe": "checkpoint"}))); assert_eq!(sanitized.metadata, None); @@ -712,14 +712,14 @@ def invalid(event, fields): let raised = runtime .block_on(wrap_py_event_sanitize_fn( module.getattr("raises").unwrap().unbind(), - )(event.clone(), fields.clone())) + )(Arc::new(event.clone()), fields.clone())) .unwrap_err(); assert!(raised.to_string().contains("sanitize boom")); let invalid = runtime .block_on(wrap_py_event_sanitize_fn( module.getattr("invalid").unwrap().unbind(), - )(event, fields.clone())) + )(Arc::new(event), fields.clone())) .unwrap_err(); assert!( invalid @@ -728,3 +728,56 @@ def invalid(event, fields): ); }); } + +#[test] +fn awaitable_middleware_wrappers_cover_success_and_failure() { + let _python = crate::test_support::init_python_test(); + Python::attach(|py| { + let module = load_module( + py, + r#" +async def tool_ok(name, args): + return {"name": name, "value": args["value"] + 1} + +async def tool_fail(name, args): + raise RuntimeError("async tool boom") + +async def llm_ok(request): + return None + +async def llm_fail(request): + raise RuntimeError("async llm boom") +"#, + ); + let tool_ok = wrap_py_tool_fn(module.getattr("tool_ok").unwrap().unbind()); + let tool_fail = wrap_py_tool_fn(module.getattr("tool_fail").unwrap().unbind()); + let llm_ok = wrap_py_llm_conditional_fn(module.getattr("llm_ok").unwrap().unbind()); + let llm_fail = wrap_py_llm_conditional_fn(module.getattr("llm_fail").unwrap().unbind()); + + with_event_loop(py, |event_loop| { + pyo3_async_runtimes::tokio::run_until_complete(event_loop, async move { + assert_eq!( + tool_ok("demo".into(), json!({"value": 1})).await.unwrap(), + json!({"name": "demo", "value": 2}) + ); + assert!( + tool_fail("demo".into(), json!({"value": 1})) + .await + .unwrap_err() + .to_string() + .contains("async tool boom") + ); + assert_eq!(llm_ok(make_request()).await.unwrap(), None); + assert!( + llm_fail(make_request()) + .await + .unwrap_err() + .to_string() + .contains("async llm boom") + ); + Ok(()) + }) + .unwrap(); + }); + }); +} diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index dac53db81..9621cf4a1 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -41,6 +41,12 @@ existing error behavior for its middleware family. | Node.js | Direct return value | Direct return value or `Promise` | | Go / raw C FFI | Synchronous callback | Existing synchronous callback, or the new `Async` / completion-based registration API | +Python's standalone middleware helpers preserve their direct synchronous +return when called without a running `asyncio` loop. In that mode, registered +callbacks must also return direct values; an awaitable callback raises a clear +runtime error. Call the helper from async Python and await its result when any +entry may return an awaitable. + For Rust, wrap the existing result in a ready async future, or use an async block when the callback needs to await work: From 02303bc491d003dc930ec7aa2f5aff31ea259a3d Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 16:06:03 -0400 Subject: [PATCH 12/52] docs: clarify Python awaitable middleware callers Signed-off-by: Will Killian --- docs/about-nemo-relay/concepts/middleware.mdx | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/docs/about-nemo-relay/concepts/middleware.mdx b/docs/about-nemo-relay/concepts/middleware.mdx index afd1ae0cf..1af705901 100644 --- a/docs/about-nemo-relay/concepts/middleware.mdx +++ b/docs/about-nemo-relay/concepts/middleware.mdx @@ -22,9 +22,13 @@ hook system. ## Asynchronous Callbacks All middleware families accept asynchronous callbacks. Rust callbacks return a -future; Python callbacks may return a value or an awaitable; and Node callbacks -may return a value or a Promise. Relay awaits entries sequentially in priority -order, so later callbacks observe earlier middleware output. +future, and Node callbacks may return a value or a Promise. Python registrations +accept callbacks that return a value or an awaitable when invoked through an +asynchronous Relay API or queued event publication. Synchronous standalone +Python helpers cannot drive an awaitable callback and raise an error directing +the caller to the corresponding asynchronous helper. Relay awaits entries +sequentially in priority order, so later callbacks observe earlier middleware +output. Managed execution and standalone conditional/request-intercept helpers are asynchronous because their result depends on middleware completion. Manual From 86162fffce51c81b206fcdd7bbb961aeb3c42802 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 16:20:25 -0400 Subject: [PATCH 13/52] fix: isolate async event sanitizer panics Signed-off-by: Will Killian --- crates/core/Cargo.toml | 3 +- crates/core/src/api/runtime/state.rs | 21 +++++++++++-- .../src/api/runtime/subscriber_dispatcher.rs | 21 +++---------- .../subscriber_dispatcher_tests.rs | 30 ++++++++++++++----- 4 files changed, 46 insertions(+), 29 deletions(-) diff --git a/crates/core/Cargo.toml b/crates/core/Cargo.toml index 5578e9393..922539111 100644 --- a/crates/core/Cargo.toml +++ b/crates/core/Cargo.toml @@ -19,7 +19,6 @@ default = [ "object-store", ] atof-streaming = [ - "dep:futures-util", "dep:tokio-tungstenite", "tokio/io-util", "tokio/net", @@ -63,7 +62,7 @@ strum = { version = "0.27", features = ["derive"] } tokio = { version = "1", default-features = false, features = ["rt", "rt-multi-thread", "macros", "sync", "time"] } tokio-stream = { version = "0.1", default-features = false, features = ["sync"] } typed-builder = "0.23.2" -futures-util = { version = "0.3", optional = true } +futures-util = "0.3" opentelemetry = { workspace = true, features = ["trace"] } opentelemetry-semantic-conventions.workspace = true opentelemetry_sdk = { workspace = true, features = ["trace"] } diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index a70d3a996..44eb7a916 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -10,9 +10,12 @@ use std::any::Any; use std::collections::HashMap; +use std::panic::AssertUnwindSafe; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; +use futures_util::FutureExt; + use crate::api::event::{ BaseEvent, CategoryProfile, Event, EventCategory, MarkEvent, ScopeCategory, ScopeEvent, llm_attributes_to_strings, scope_attributes_to_strings, tool_attributes_to_strings, @@ -640,15 +643,27 @@ impl NemoRelayContextState { let event_context = Arc::new(event.clone()); for entry in entries { let fields = event.sanitize_fields(); - match (entry.payload)(Arc::clone(&event_context), fields).await { - Ok(fields) => event.apply_sanitize_fields(fields), - Err(error) => log::error!( + let callback = Arc::clone(&entry.payload); + let context = Arc::clone(&event_context); + match AssertUnwindSafe(async move { callback(context, fields).await }) + .catch_unwind() + .await + { + Ok(Ok(fields)) => event.apply_sanitize_fields(fields), + Ok(Err(error)) => log::error!( target: "nemo_relay.runtime", event = "event_sanitizer_failed", sanitizer = entry.name.as_str(), event_name = event.name(); "Event sanitizer failed; preserving the last valid event snapshot: {error}" ), + Err(_) => log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_panicked", + sanitizer = entry.name.as_str(), + event_name = event.name(); + "Event sanitizer panicked; publishing the latest valid event snapshot" + ), } } event diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 475a0434c..341017a28 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -339,24 +339,11 @@ mod native { if sanitizers.is_empty() { return Some(transformed); } - let original = transformed.clone(); Some( - match catch_unwind(AssertUnwindSafe(|| { - runtime.block_on(NemoRelayContextState::event_sanitize_snapshot_chain( - transformed, - &sanitizers, - )) - })) { - Ok(event) => event, - Err(_) => { - log::error!( - target: "nemo_relay.runtime", - event = "event_sanitizer_panicked"; - "Event sanitizer panicked; publishing the transformed event snapshot" - ); - original - } - }, + runtime.block_on(NemoRelayContextState::event_sanitize_snapshot_chain( + transformed, + &sanitizers, + )), ) } } diff --git a/crates/core/tests/integration/subscriber_dispatcher_tests.rs b/crates/core/tests/integration/subscriber_dispatcher_tests.rs index 303ee112b..07b6429de 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -7,6 +7,7 @@ use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, mpsc}; use std::time::Duration; +use nemo_relay::api::event::Event; use nemo_relay::api::registry::{ deregister_mark_sanitize_guardrail, register_mark_sanitize_guardrail, }; @@ -16,6 +17,7 @@ use nemo_relay::api::runtime::{ use nemo_relay::api::scope::{EmitMarkEventParams, event}; use nemo_relay::api::subscriber::{deregister_subscriber, flush_subscribers, register_subscriber}; use nemo_relay::error::FlowError; +use serde_json::json; static TEST_MUTEX: Mutex<()> = Mutex::new(()); @@ -208,15 +210,21 @@ fn dispatcher_publishes_the_snapshot_when_an_async_sanitizer_panics() { reset_global(); setup_isolated_thread(); - let observed = Arc::new(Mutex::new(Vec::new())); + let observed = Arc::new(Mutex::new(Vec::::new())); let observed_events = Arc::clone(&observed); register_subscriber( "panic-sanitizer-subscriber", - Arc::new(move |event| { - observed_events - .lock() - .unwrap() - .push(event.name().to_string()) + Arc::new(move |event| observed_events.lock().unwrap().push(event.clone())), + ) + .unwrap(); + register_mark_sanitize_guardrail( + "successful-mark-sanitizer", + 0, + Arc::new(|_, mut fields| { + Box::pin(async move { + fields.data = Some(json!({"redacted": true})); + Ok(fields) + }) }), ) .unwrap(); @@ -230,7 +238,15 @@ fn dispatcher_publishes_the_snapshot_when_an_async_sanitizer_panics() { emit_mark("panic-fallback"); flush_subscribers().unwrap(); - assert_eq!(observed.lock().unwrap().as_slice(), ["panic-fallback"]); + let events = observed.lock().unwrap(); + assert_eq!(events.len(), 1); + assert_eq!(events[0].name(), "panic-fallback"); + assert_eq!( + events[0].sanitize_fields().data, + Some(json!({"redacted": true})) + ); + drop(events); + deregister_mark_sanitize_guardrail("successful-mark-sanitizer").unwrap(); deregister_mark_sanitize_guardrail("panic-mark-sanitizer").unwrap(); deregister_subscriber("panic-sanitizer-subscriber").unwrap(); } From 71386c2ec1996fc643d041804df5f2b9fea8b8b2 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 16:50:18 -0400 Subject: [PATCH 14/52] fix: preserve queued sanitizer execution semantics Signed-off-by: Will Killian --- crates/core/src/api/runtime/state.rs | 3 +-- crates/core/src/api/runtime/subscriber_dispatcher.rs | 3 --- 2 files changed, 1 insertion(+), 5 deletions(-) diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index 44eb7a916..0e5a2766c 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -640,11 +640,10 @@ impl NemoRelayContextState { mut event: Event, entries: &[Guardrail], ) -> Event { - let event_context = Arc::new(event.clone()); for entry in entries { let fields = event.sanitize_fields(); let callback = Arc::clone(&entry.payload); - let context = Arc::clone(&event_context); + let context = Arc::new(event.clone()); match AssertUnwindSafe(async move { callback(context, fields).await }) .catch_unwind() .await diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 341017a28..102aa0adb 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -121,9 +121,6 @@ mod native { subscribers: &[EventSubscriberFn], scope_stack: ScopeStackHandle, ) -> bool { - if subscribers.is_empty() { - return true; - } let message = DispatcherMessage::Deliver { event: Box::new(event), transform: Some(transform), From c97baa34668fed54f6d681fe932c465b9753da07 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 17:09:28 -0400 Subject: [PATCH 15/52] fix: address async middleware review follow-ups Signed-off-by: Will Killian --- crates/core/src/api/llm.rs | 207 +++++++------- crates/core/src/api/runtime/state.rs | 254 ++++++++++++++++-- .../src/api/runtime/subscriber_dispatcher.rs | 33 ++- crates/core/src/api/scope.rs | 40 +-- crates/core/src/api/shared.rs | 4 +- crates/core/src/api/tool.rs | 54 ++-- crates/core/src/stream.rs | 42 +-- .../tests/fixtures/native_plugin/src/lib.rs | 6 +- .../tests/integration/middleware_tests.rs | 1 + .../tests/integration/native_plugin_tests.rs | 4 + .../subscriber_dispatcher_tests.rs | 28 +- .../core/tests/unit/dynamic_worker_tests.rs | 127 +++++---- crates/core/tests/unit/llm_api_tests.rs | 19 +- crates/ffi/src/callable.rs | 6 +- crates/ffi/tests/unit/callable_tests.rs | 16 +- crates/node/src/callable.rs | 27 +- crates/node/tests/llm_tests.mjs | 13 +- crates/node/tests/scope_tests.mjs | 13 +- crates/node/tests/test_support.mjs | 24 ++ crates/node/tests/tools_tests.mjs | 18 +- docs/reference/event-sanitizers.mdx | 13 +- docs/reference/migration-guides.mdx | 15 +- 22 files changed, 635 insertions(+), 329 deletions(-) create mode 100644 crates/node/tests/test_support.mjs diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index d0bcc79a3..0301aeeda 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -494,15 +494,16 @@ fn emit_llm_start( let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())? }; - tokio::runtime::Runtime::new() - .map_err(|error| FlowError::Internal(error.to_string()))? - .block_on(emit_llm_start_with_subscribers( + crate::api::runtime::subscriber_dispatcher::block_on_sanitizer_future( + emit_llm_start_with_subscribers( handle, request, annotated_request, request_codec, &subscribers, - )) + ), + ) + .map_err(FlowError::Internal)? } async fn emit_pending_request_marks( @@ -795,6 +796,75 @@ struct LlmCallEndBehavior { attach_estimated_cost: bool, } +struct LlmEndPayload { + data: Option, + annotated_response: Option>, + decode_error: Option, +} + +async fn build_llm_end_payload( + handle: &LlmHandle, + response: Json, + fallback_data: Option, + annotated_response: Option>, + response_codec: Option>, + entries: &[crate::api::registry::Guardrail], + behavior: LlmCallEndBehavior, +) -> LlmEndPayload { + let response_was_null_without_fallback = response.is_null() && fallback_data.is_none(); + let response = if response.is_null() { + fallback_data.unwrap_or(response) + } else { + response + }; + let sanitized_response = NemoRelayContextState::llm_sanitize_response_snapshot_chain( + response.clone(), + LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), + entries, + ) + .await; + let response_changed = sanitized_response + .as_ref() + .is_some_and(|sanitized_response| sanitized_response != &response); + let data = match sanitized_response { + Some(response) if response_was_null_without_fallback && response.is_null() => None, + response => response, + }; + let annotation_omitted = data.as_ref().is_none_or(Json::is_null); + let (mut annotated_response, decode_error) = if annotation_omitted { + (None, None) + } else { + resolve_llm_end_annotation( + (!response_changed).then_some(annotated_response).flatten(), + response_codec, + data.as_ref(), + &behavior, + &handle.name, + ) + }; + let pricing = crate::codec::response::active_pricing_resolver(); + let summary = finalize_optimization_summary( + &handle.optimization_recorder, + annotated_response.as_mut(), + handle.model_name.as_deref(), + &pricing, + ); + if !annotation_omitted + && annotated_response.is_none() + && let Some(summary) = summary + { + annotated_response = Some(AnnotatedLlmResponse { + optimization_summary: Some(summary), + ..AnnotatedLlmResponse::default() + }); + } + LlmEndPayload { + data, + annotated_response: annotated_response.map(Arc::new), + decode_error, + } +} + /// Finish a manual LLM lifecycle span. /// /// This emits an LLM-end event for a handle previously returned by @@ -844,12 +914,8 @@ pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> { subscribers, ) }; - let response = if params.response.is_null() { - params.data.unwrap_or(params.response) - } else { - params.response - }; - let response_was_null_without_fallback = response.is_null(); + let response = params.response; + let fallback_data = params.data; let handle = params.handle.clone(); let metadata = params.metadata; let timestamp = params.timestamp; @@ -877,59 +943,26 @@ pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> { event, Box::new(move |event| { Box::pin(async move { - let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( - response.clone(), - LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), + let payload = build_llm_end_payload( + &handle, + response, + fallback_data, + annotated_response, + response_codec, &entries, + LlmCallEndBehavior { + response_codec_errors_fatal: false, + attach_estimated_cost: false, + }, ) .await; - let changed = sanitized - .as_ref() - .is_some_and(|sanitized| sanitized != &response); - let data = match sanitized { - Some(response) if response_was_null_without_fallback && response.is_null() => { - None - } - response => response, - }; - let annotation_omitted = data.as_ref().is_none_or(Json::is_null); - let (mut annotation, decode_error) = if annotation_omitted { - (None, None) - } else { - resolve_llm_end_annotation( - (!changed).then_some(annotated_response).flatten(), - response_codec, - data.as_ref(), - &LlmCallEndBehavior { - response_codec_errors_fatal: false, - attach_estimated_cost: false, - }, - &handle.name, - ) - }; - if let Some(error) = decode_error { + if let Some(error) = payload.decode_error { log::error!( target: "nemo_relay.runtime", event = "manual_llm_response_codec_failed"; "Manual LLM response annotation failed during queued publication: {error}" ); } - let pricing = crate::codec::response::active_pricing_resolver(); - let summary = finalize_optimization_summary( - &handle.optimization_recorder, - annotation.as_mut(), - handle.model_name.as_deref(), - &pricing, - ); - if !annotation_omitted - && annotation.is_none() - && let Some(summary) = summary - { - annotation = Some(AnnotatedLlmResponse { - optimization_summary: Some(summary), - ..AnnotatedLlmResponse::default() - }); - } let context = global_context(); let Ok(state) = context.read() else { return event; @@ -938,9 +971,9 @@ pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> { state.build_llm_end_event( EndLlmHandleParams::builder() .handle(&handle) - .data_opt(data) + .data_opt(payload.data) .metadata_opt(end_metadata) - .annotated_response_opt(annotation.map(Arc::new)) + .annotated_response_opt(payload.annotated_response) .timestamp_opt(timestamp) .build(), ) @@ -986,56 +1019,18 @@ async fn llm_call_end_with_behavior( let entries = state.llm_sanitize_response_entries(&scope_locals); (entries, subscribers) }; - let response_was_null_without_fallback = response.is_null() && data.is_none(); - let response = if response.is_null() { - data.unwrap_or(response) - } else { - response - }; - let sanitized_response = NemoRelayContextState::llm_sanitize_response_snapshot_chain( - response.clone(), - LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), + handle.optimization_recorder.close_for_finalization(None); + emit_optimization_marks(handle, &subscribers).await; + let payload = build_llm_end_payload( + handle, + response, + data, + annotated_response, + response_codec, &entries, + behavior, ) .await; - let response_changed = sanitized_response - .as_ref() - .is_some_and(|sanitized_response| sanitized_response != &response); - let data = match sanitized_response { - Some(response) if response_was_null_without_fallback && response.is_null() => None, - response => response, - }; - let annotation_omitted = data.as_ref().is_none_or(Json::is_null); - let (mut annotated_response, decode_error) = if annotation_omitted { - (None, None) - } else { - resolve_llm_end_annotation( - (!response_changed).then_some(annotated_response).flatten(), - response_codec, - data.as_ref(), - &behavior, - &handle.name, - ) - }; - handle.optimization_recorder.close_for_finalization(None); - emit_optimization_marks(handle, &subscribers).await; - let pricing = crate::codec::response::active_pricing_resolver(); - let summary = finalize_optimization_summary( - &handle.optimization_recorder, - annotated_response.as_mut(), - handle.model_name.as_deref(), - &pricing, - ); - if !annotation_omitted - && annotated_response.is_none() - && let Some(summary) = summary - { - annotated_response = Some(AnnotatedLlmResponse { - optimization_summary: Some(summary), - ..AnnotatedLlmResponse::default() - }); - } - let annotated_response = annotated_response.map(Arc::new); let event = { let context = global_context(); let state = context @@ -1045,9 +1040,9 @@ async fn llm_call_end_with_behavior( state.build_llm_end_event( EndLlmHandleParams::builder() .handle(handle) - .data_opt(data) + .data_opt(payload.data) .metadata_opt(end_metadata) - .annotated_response_opt(annotated_response) + .annotated_response_opt(payload.annotated_response) .timestamp_opt(timestamp) .build(), ) @@ -1056,7 +1051,7 @@ async fn llm_call_end_with_behavior( { NemoRelayContextState::emit_event(&event, &subscribers); } - if let Some(error) = decode_error + if let Some(error) = payload.decode_error && behavior.response_codec_errors_fatal { Err(error) diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index 0e5a2766c..b7c6182d4 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -42,6 +42,7 @@ use crate::codec::response::AnnotatedLlmResponse; use crate::context::registries::{ merge_execution_intercept_callables, merge_guardrail_entries, merge_intercept_entries, }; +use crate::error::FlowError; use crate::json::{Json, merge_json}; use crate::registry::SortedRegistry; use chrono::{Duration, Utc}; @@ -703,15 +704,28 @@ impl NemoRelayContextState { ) -> Json { let mut value = args; for entry in entries { - match (entry.payload)(name.to_string(), value.clone()).await { - Ok(next) => value = next, - Err(error) => log::error!( + let callback = Arc::clone(&entry.payload); + let callback_name = name.to_string(); + let current = value.clone(); + match AssertUnwindSafe(async move { callback(callback_name, current).await }) + .catch_unwind() + .await + { + Ok(Ok(next)) => value = next, + Ok(Err(error)) => log::error!( target: "nemo_relay.runtime", event = "tool_request_sanitizer_failed", sanitizer = entry.name.as_str(), tool_name = name; "Tool request sanitizer failed; preserving the last valid payload: {error}" ), + Err(_) => log::error!( + target: "nemo_relay.runtime", + event = "tool_request_sanitizer_panicked", + sanitizer = entry.name.as_str(), + tool_name = name; + "Tool request sanitizer panicked; preserving the last valid payload" + ), } } value @@ -752,15 +766,28 @@ impl NemoRelayContextState { ) -> Json { let mut value = result; for entry in entries { - match (entry.payload)(name.to_string(), value.clone()).await { - Ok(next) => value = next, - Err(error) => log::error!( + let callback = Arc::clone(&entry.payload); + let callback_name = name.to_string(); + let current = value.clone(); + match AssertUnwindSafe(async move { callback(callback_name, current).await }) + .catch_unwind() + .await + { + Ok(Ok(next)) => value = next, + Ok(Err(error)) => log::error!( target: "nemo_relay.runtime", event = "tool_response_sanitizer_failed", sanitizer = entry.name.as_str(), tool_name = name; "Tool response sanitizer failed; preserving the last valid payload: {error}" ), + Err(_) => log::error!( + target: "nemo_relay.runtime", + event = "tool_response_sanitizer_panicked", + sanitizer = entry.name.as_str(), + tool_name = name; + "Tool response sanitizer panicked; preserving the last valid payload" + ), } } value @@ -831,7 +858,20 @@ impl NemoRelayContextState { subscribers, ) .await; - let result = (entry.payload)(name.to_string(), args.clone()).await; + let callback = Arc::clone(&entry.payload); + let callback_name = name.to_string(); + let callback_args = args.clone(); + let result = + match AssertUnwindSafe(async move { callback(callback_name, callback_args).await }) + .catch_unwind() + .await + { + Ok(result) => result, + Err(_) => Err(FlowError::Internal(format!( + "tool conditional guardrail '{}' panicked", + entry.name + ))), + }; let output = match &result { Ok(Some(reason)) => json!({ "allowed": false, @@ -897,7 +937,20 @@ impl NemoRelayContextState { ) -> crate::error::Result { let mut value = args; for entry in entries { - value = (entry.payload.callable)(name.to_string(), value).await?; + let callback = Arc::clone(&entry.payload.callable); + let callback_name = name.to_string(); + value = match AssertUnwindSafe(async move { callback(callback_name, value).await }) + .catch_unwind() + .await + { + Ok(result) => result?, + Err(_) => { + return Err(FlowError::Internal(format!( + "tool request intercept '{}' panicked", + entry.name + ))); + } + }; if entry.payload.break_chain { break; } @@ -1015,9 +1068,17 @@ impl NemoRelayContextState { let mut value = Some(request); for entry in entries { if let Some(current) = value.take() { - match (entry.payload)(current.clone(), context.clone()).await { - Ok(next) => value = next, - Err(error) => { + let callback = Arc::clone(&entry.payload); + let callback_value = current.clone(); + let callback_context = context.clone(); + match AssertUnwindSafe( + async move { callback(callback_value, callback_context).await }, + ) + .catch_unwind() + .await + { + Ok(Ok(next)) => value = next, + Ok(Err(error)) => { log::error!( target: "nemo_relay.runtime", event = "llm_request_sanitizer_failed", @@ -1027,6 +1088,15 @@ impl NemoRelayContextState { ); value = Some(current); } + Err(_) => { + log::error!( + target: "nemo_relay.runtime", + event = "llm_request_sanitizer_panicked", + sanitizer = entry.name.as_str(); + "LLM request sanitizer panicked; preserving the last valid request" + ); + value = Some(current); + } } } } @@ -1068,9 +1138,17 @@ impl NemoRelayContextState { let mut value = Some(response); for entry in entries { if let Some(current) = value.take() { - match (entry.payload)(current.clone(), context.clone()).await { - Ok(next) => value = next, - Err(error) => { + let callback = Arc::clone(&entry.payload); + let callback_value = current.clone(); + let callback_context = context.clone(); + match AssertUnwindSafe( + async move { callback(callback_value, callback_context).await }, + ) + .catch_unwind() + .await + { + Ok(Ok(next)) => value = next, + Ok(Err(error)) => { log::error!( target: "nemo_relay.runtime", event = "llm_response_sanitizer_failed", @@ -1080,6 +1158,15 @@ impl NemoRelayContextState { ); value = Some(current); } + Err(_) => { + log::error!( + target: "nemo_relay.runtime", + event = "llm_response_sanitizer_panicked", + sanitizer = entry.name.as_str(); + "LLM response sanitizer panicked; preserving the last valid response" + ); + value = Some(current); + } } } } @@ -1148,7 +1235,18 @@ impl NemoRelayContextState { subscribers, ) .await; - let result = (entry.payload)(request.clone()).await; + let callback = Arc::clone(&entry.payload); + let callback_request = request.clone(); + let result = match AssertUnwindSafe(async move { callback(callback_request).await }) + .catch_unwind() + .await + { + Ok(result) => result, + Err(_) => Err(FlowError::Internal(format!( + "LLM conditional guardrail '{}' panicked", + entry.name + ))), + }; let output = match &result { Ok(Some(reason)) => json!({ "allowed": false, @@ -1245,8 +1343,22 @@ impl NemoRelayContextState { let mut optimization_contributions = Vec::new(); for entry in entries { let input_content = request_value.content.clone(); - let outcome = - (entry.payload.callable)(name.to_string(), request_value, annotated_value).await?; + let callback = Arc::clone(&entry.payload.callable); + let callback_name = name.to_string(); + let outcome = match AssertUnwindSafe(async move { + callback(callback_name, request_value, annotated_value).await + }) + .catch_unwind() + .await + { + Ok(result) => result?, + Err(_) => { + return Err(FlowError::Internal(format!( + "LLM request intercept '{}' panicked", + entry.name + ))); + } + }; if codec_active && outcome.request.content != input_content { return Err(crate::error::FlowError::InvalidArgument(format!( "LLM request intercept '{}' changed request.content while a request codec is active; modify annotated_request instead", @@ -1354,3 +1466,111 @@ impl Default for NemoRelayContextState { Self::new() } } + +#[cfg(test)] +mod panic_tests { + use super::*; + use crate::api::registry::{RegistryRecord, RequestIntercept}; + use serde_json::{Map, json}; + + #[tokio::test] + async fn middleware_snapshot_chains_contain_callback_panics() { + let tool_payload = json!({"tool": "preserved"}); + let tool_sanitizer: ToolSanitizeFn = + Arc::new(|_, _| Box::pin(async { panic!("tool sanitizer panic") })); + let tool_entries = vec![RegistryRecord::new("tool-panic", 0, tool_sanitizer)]; + assert_eq!( + NemoRelayContextState::tool_sanitize_request_snapshot_chain( + "tool", + tool_payload.clone(), + &tool_entries, + ) + .await, + tool_payload + ); + + let request = LlmRequest { + headers: Map::new(), + content: json!({"llm": "preserved"}), + }; + let llm_sanitizer: LlmSanitizeRequestFn = + Arc::new(|_, _| Box::pin(async { panic!("LLM sanitizer panic") })); + let llm_entries = vec![RegistryRecord::new("llm-panic", 0, llm_sanitizer)]; + assert_eq!( + NemoRelayContextState::llm_sanitize_request_snapshot_chain( + request.clone(), + LlmSanitizeRequestContext::default(), + &llm_entries, + ) + .await, + Some(request.clone()) + ); + + let tool_conditional: ToolConditionalFn = + Arc::new(|_, _| Box::pin(async { panic!("tool conditional panic") })); + let error = NemoRelayContextState::tool_conditional_execution_snapshot_chain( + "tool", + &tool_payload, + &[RegistryRecord::new( + "tool-conditional-panic", + 0, + tool_conditional, + )], + &[], + None, + None, + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("tool-conditional-panic")); + + let llm_conditional: LlmConditionalFn = + Arc::new(|_| Box::pin(async { panic!("LLM conditional panic") })); + let error = NemoRelayContextState::llm_conditional_execution_snapshot_chain( + &request, + &[RegistryRecord::new( + "llm-conditional-panic", + 0, + llm_conditional, + )], + &[], + None, + None, + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("llm-conditional-panic")); + + let tool_intercept: ToolInterceptFn = + Arc::new(|_, _| Box::pin(async { panic!("tool intercept panic") })); + let error = NemoRelayContextState::tool_request_intercepts_snapshot_chain( + "tool", + tool_payload, + &[RegistryRecord::new( + "tool-intercept-panic", + 0, + RequestIntercept::new(false, tool_intercept), + )], + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("tool-intercept-panic")); + + let llm_intercept: LlmRequestInterceptFn = + Arc::new(|_, _, _| Box::pin(async { panic!("LLM intercept panic") })); + let error = NemoRelayContextState::llm_request_intercepts_snapshot_chain( + "llm", + request, + None, + &[RegistryRecord::new( + "llm-intercept-panic", + 0, + RequestIntercept::new(false, llm_intercept), + )], + false, + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("llm-intercept-panic")); + } +} diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 102aa0adb..5421bdca9 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -57,6 +57,25 @@ mod native { static IN_DISPATCHER: Cell = const { Cell::new(false) }; } + fn sanitizer_runtime() -> std::result::Result<&'static tokio::runtime::Runtime, String> { + SANITIZER_RUNTIME + .get_or_init(|| { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| error.to_string()) + }) + .as_ref() + .map_err(Clone::clone) + } + + #[cfg(test)] + pub(super) fn block_on_sanitizer_future( + future: F, + ) -> std::result::Result { + sanitizer_runtime().map(|runtime| runtime.block_on(future)) + } + pub(super) fn dispatch_event(event: &Event, subscribers: &[EventSubscriberFn]) -> bool { if subscribers.is_empty() { return true; @@ -297,12 +316,7 @@ mod native { transform: Option, sanitizers: Vec>, ) -> Option { - let runtime = match SANITIZER_RUNTIME.get_or_init(|| { - tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .map_err(|error| error.to_string()) - }) { + let runtime = match sanitizer_runtime() { Ok(runtime) => runtime, Err(error) => { if !SANITIZER_RUNTIME_FAILURE_LOGGED.swap(true, Ordering::AcqRel) { @@ -345,6 +359,13 @@ mod native { } } +#[cfg(test)] +pub(crate) fn block_on_sanitizer_future( + future: F, +) -> std::result::Result { + native::block_on_sanitizer_future(future) +} + /// Queue an event for subscriber delivery. pub(crate) fn dispatch_event(event: &Event, subscribers: &[EventSubscriberFn]) -> bool { native::dispatch_event(event, subscribers) diff --git a/crates/core/src/api/scope.rs b/crates/core/src/api/scope.rs index 635c6c76c..98afdfbb7 100644 --- a/crates/core/src/api/scope.rs +++ b/crates/core/src/api/scope.rs @@ -52,6 +52,16 @@ pub struct ScopeHandle { pub parent_uuid: Option, } +fn scope_stack_lock_error(error: impl std::fmt::Display, operation: &'static str) -> FlowError { + log::error!( + target: "nemo_relay.runtime", + event = "scope_stack_unavailable", + operation = operation; + "Scope operation failed because the scope stack lock is poisoned: {error}" + ); + FlowError::Internal(error.to_string()) +} + /// Builder parameters for [`push_scope`]. #[derive(TypedBuilder)] #[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))] @@ -224,7 +234,9 @@ pub fn push_scope(params: PushScopeParams<'_>) -> Result { let parent_uuid = resolve_parent_uuid(params.parent); let (handle, event, subscribers, emission_scope_stack) = { let scope_stack = current_scope_stack(); - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_guard = scope_stack + .read() + .map_err(|error| scope_stack_lock_error(error, "push"))?; let scope_subscribers = scope_guard.collect_scope_local_subscribers(); let subscribers = snapshot_event_subscribers(scope_subscribers)?; let context = global_context(); @@ -287,7 +299,9 @@ pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { ensure_runtime_owner()?; let scope_stack = current_scope_stack(); let (scope, event, subscribers, emission_scope_stack) = { - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_guard = scope_stack + .read() + .map_err(|error| scope_stack_lock_error(error, "pop"))?; let top = scope_guard.top(); if top.uuid != *params.handle_uuid { if scope_guard.find(params.handle_uuid).is_some() { @@ -361,27 +375,17 @@ pub fn event(params: EmitMarkEventParams<'_>) -> Result<()> { let scope_stack = current_scope_stack(); let (event, subscribers, emission_scope_stack) = { let subscribers = if params.name == COMPACTION_EVENT_NAME { - let mut scope_guard = scope_stack.write().map_err(|error| { - log::error!( - target: "nemo_relay.runtime", - event = "mark_event_scope_stack_unavailable"; - "Mark event was dropped because the scope stack lock is poisoned: {error}" - ); - FlowError::Internal(error.to_string()) - })?; + let mut scope_guard = scope_stack + .write() + .map_err(|error| scope_stack_lock_error(error, "mark"))?; let subscribers = snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; scope_guard.mark_agent_fresh(parent_uuid); subscribers } else { - let scope_guard = scope_stack.read().map_err(|error| { - log::error!( - target: "nemo_relay.runtime", - event = "mark_event_scope_stack_unavailable"; - "Mark event was dropped because the scope stack lock is poisoned: {error}" - ); - FlowError::Internal(error.to_string()) - })?; + let scope_guard = scope_stack + .read() + .map_err(|error| scope_stack_lock_error(error, "mark"))?; snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())? }; let context = global_context(); diff --git a/crates/core/src/api/shared.rs b/crates/core/src/api/shared.rs index b97c147bb..5491b3182 100644 --- a/crates/core/src/api/shared.rs +++ b/crates/core/src/api/shared.rs @@ -260,7 +260,9 @@ async fn run_request_intercepts_with_codec_inner( let entries = { let scope_stack = current_scope_stack(); - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_guard = scope_stack + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; let scope_locals = scope_guard .collect_scope_local_registries(|registries| ®istries.llm_request_intercepts); diff --git a/crates/core/src/api/tool.rs b/crates/core/src/api/tool.rs index 0c1128d39..768d31670 100644 --- a/crates/core/src/api/tool.rs +++ b/crates/core/src/api/tool.rs @@ -84,6 +84,25 @@ pub struct CreateToolHandleParams<'a> { pub timestamp: Option>, } +fn resolve_skill_loads( + name: &str, + args: &Json, + metadata: Option<&Json>, +) -> Vec { + let already_handled = metadata + .and_then(Json::as_object) + .and_then(|metadata| metadata.get(skill_load::HANDLED_METADATA_KEY)) + .and_then(Json::as_bool) + .unwrap_or(false); + if already_handled { + Vec::new() + } else if let Some(skill_loads) = skill_load::precomputed(metadata) { + skill_loads + } else { + skill_load::detect(name, args) + } +} + /// Builder parameters for [`NemoRelayContextState::build_tool_end_event`]. #[derive(Debug, Clone, TypedBuilder)] #[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))] @@ -227,20 +246,7 @@ pub fn tool_call(params: ToolCallParams<'_>) -> Result { subscribers, ) }; - let handled_skill_loads = params - .metadata - .as_ref() - .and_then(Json::as_object) - .and_then(|metadata| metadata.get(skill_load::HANDLED_METADATA_KEY)) - .and_then(Json::as_bool) - .is_some_and(|handled| handled); - let skill_loads = if handled_skill_loads { - Vec::new() - } else if let Some(skill_loads) = skill_load::precomputed(params.metadata.as_ref()) { - skill_loads - } else { - skill_load::detect(params.name, ¶ms.args) - }; + let skill_loads = resolve_skill_loads(params.name, ¶ms.args, params.metadata.as_ref()); let raw_args = params.args; let (handle, event, marks) = { let context = global_context(); @@ -301,9 +307,8 @@ pub fn tool_call(params: ToolCallParams<'_>) -> Result { scope_stack.clone(), ); for mark in marks { - if let Some(sanitizers) = snapshot_event_sanitizers(&mark, &scope_stack) { - dispatch_sanitized_event(mark, sanitizers, &subscribers, scope_stack.clone()); - } + let sanitizers = snapshot_event_sanitizers(&mark, &scope_stack).unwrap_or_default(); + dispatch_sanitized_event(mark, sanitizers, &subscribers, scope_stack.clone()); } Ok(handle) } @@ -328,20 +333,7 @@ async fn tool_call_with_subscriber_snapshot( let entries = state.tool_sanitize_request_entries(&scope_locals); (entries, subscribers) }; - let handled_skill_loads = params - .metadata - .as_ref() - .and_then(Json::as_object) - .and_then(|metadata| metadata.get(skill_load::HANDLED_METADATA_KEY)) - .and_then(Json::as_bool) - .is_some_and(|handled| handled); - let skill_loads = if handled_skill_loads { - Vec::new() - } else if let Some(skill_loads) = skill_load::precomputed(params.metadata.as_ref()) { - skill_loads - } else { - skill_load::detect(params.name, ¶ms.args) - }; + let skill_loads = resolve_skill_loads(params.name, ¶ms.args, params.metadata.as_ref()); let sanitized_args = NemoRelayContextState::tool_sanitize_request_snapshot_chain( params.name, params.args, diff --git a/crates/core/src/stream.rs b/crates/core/src/stream.rs index 2e742fea3..2bae33344 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -178,7 +178,7 @@ impl LlmStreamWrapper { &self.scope_stack } - fn finish(&mut self) { + fn finish(&mut self, background_thread: bool) { if self.ended { return; } @@ -193,7 +193,7 @@ impl LlmStreamWrapper { self.handle .optimization_recorder .close_for_finalization(Some("stream_interrupted")); - self.finalization = self.emit_end_event(metadata, true, true); + self.finalization = self.emit_end_event(metadata, true, background_thread); } fn finish_with_status( @@ -236,21 +236,31 @@ impl LlmStreamWrapper { aggregated }; - let snapshot = { - let ss_guard = self.scope_stack.read().expect("scope stack lock poisoned"); - let sl = - ss_guard.collect_scope_local_registries(|r| &r.llm_sanitize_response_guardrails); - let ctx = global_context(); - let state = ctx.read(); - match state { - Ok(state) => { - let entries = state.llm_sanitize_response_entries(&sl); - Some(entries) + let entries = match self.scope_stack.read() { + Ok(scope_guard) => { + let scope_locals = scope_guard + .collect_scope_local_registries(|r| &r.llm_sanitize_response_guardrails); + match global_context().read() { + Ok(state) => state.llm_sanitize_response_entries(&scope_locals), + Err(error) => { + log::error!( + target: "nemo_relay.runtime", + event = "stream_end_sanitizer_snapshot_failed"; + "LLM stream END sanitizer snapshot failed open: {error}" + ); + Vec::new() + } } - Err(_) => None, + } + Err(error) => { + log::error!( + target: "nemo_relay.runtime", + event = "stream_end_sanitizer_snapshot_failed"; + "LLM stream END sanitizer snapshot failed open: {error}" + ); + Vec::new() } }; - let entries = snapshot?; let handle = self.handle.clone(); let scope_stack = self.scope_stack.clone(); let subscribers = self.subscribers.clone(); @@ -468,7 +478,7 @@ impl LlmStreamInner for LlmStreamWrapper { return result.clone(); } let result = this.inner.close().await; - this.finish(); + this.finish(false); if let Some(finalization) = this.finalization.take() { finalization.await.map_err(|error| { FlowError::Internal(format!("stream finalization task failed: {error}")) @@ -741,7 +751,7 @@ fn non_empty_object(object: Map) -> Option { impl Drop for LlmStreamWrapper { fn drop(&mut self) { - self.finish(); + self.finish(true); } } diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index 1367f4d74..d1928528e 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -296,13 +296,13 @@ pub unsafe extern "C" fn nemo_relay_fixture_async_entry( { return NemoRelayStatus::InvalidArg; } - let host_v2 = unsafe { &*(host as *const NemoRelayNativeHostApiV3) }; + let host_v3 = unsafe { &*(host as *const NemoRelayNativeHostApiV3) }; let mut plugin = NemoRelayNativePluginV1::default(); - plugin.plugin_kind = unsafe { raw_host_string(&host_v2.v1, "fixture_async") }; + plugin.plugin_kind = unsafe { raw_host_string(&host_v3.v1, "fixture_async") }; if plugin.plugin_kind.is_null() { return NemoRelayStatus::Internal; } - plugin.user_data = Box::into_raw(Box::new(*host_v2)).cast(); + plugin.user_data = Box::into_raw(Box::new(*host_v3)).cast(); plugin.register = Some(raw_register_async_tool_request); plugin.drop = Some(raw_drop_async_host); unsafe { *out = plugin }; diff --git a/crates/core/tests/integration/middleware_tests.rs b/crates/core/tests/integration/middleware_tests.rs index d448e7867..7e8763850 100644 --- a/crates/core/tests/integration/middleware_tests.rs +++ b/crates/core/tests/integration/middleware_tests.rs @@ -1645,6 +1645,7 @@ async fn test_scope_local_guardrail_lifecycle() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!( call_count.load(Ordering::SeqCst), 1, diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index c669f7c46..893c5ca15 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -673,6 +673,7 @@ async fn native_v3_async_registration_supports_all_middleware_kinds() { manifest_ref: manifest_ref.to_string_lossy().into_owned(), }]) .expect("v3 async native fixture should load"); + let mut cleanup = NativePluginTestCleanup::new(); let mut config = PluginConfig::default(); config.components.push(PluginComponentSpec { kind: "fixture_async".into(), @@ -682,6 +683,7 @@ async fn native_v3_async_registration_supports_all_middleware_kinds() { initialize_plugins_exact(config) .await .expect("v3 async native fixture should register"); + cleanup.mark_plugin_configuration_active(); let rewritten = tool_request_intercepts("async-tool", json!({"input": true})) .await @@ -766,12 +768,14 @@ async fn native_v3_async_registration_supports_all_middleware_kinds() { }); tokio::task::yield_now().await; clear_plugin_configuration().expect("v3 async native fixture should clear while pending"); + cleanup.plugin_configuration_active = false; let pending = pending .await .expect("pending v3 async task should not panic") .expect("pending v3 async request intercept should settle after clear"); assert_eq!(pending["native_async"], true); + drop(cleanup); drop(activation); } diff --git a/crates/core/tests/integration/subscriber_dispatcher_tests.rs b/crates/core/tests/integration/subscriber_dispatcher_tests.rs index 07b6429de..5bdfd7184 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -145,12 +145,7 @@ fn dispatcher_publishes_the_snapshot_when_an_async_sanitizer_fails() { let observed_events = Arc::clone(&observed); register_subscriber( "fail-open-sanitizer-subscriber", - Arc::new(move |event| { - observed_events - .lock() - .unwrap() - .push(event.name().to_string()) - }), + Arc::new(move |event| observed_events.lock().unwrap().push(event.clone())), ) .unwrap(); register_mark_sanitize_guardrail( @@ -166,13 +161,28 @@ fn dispatcher_publishes_the_snapshot_when_an_async_sanitizer_fails() { ) .unwrap(); - emit_mark("unsanitized-fallback"); + event( + EmitMarkEventParams::builder() + .name("unsanitized-fallback") + .data(json!({"original_data": true})) + .metadata(json!({"original_metadata": true})) + .build(), + ) + .unwrap(); flush_subscribers().unwrap(); + let observed = observed.lock().unwrap(); + assert_eq!(observed.len(), 1); + assert_eq!(observed[0].name(), "unsanitized-fallback"); + assert_eq!( + observed[0].sanitize_fields().data, + Some(json!({"original_data": true})) + ); assert_eq!( - observed.lock().unwrap().as_slice(), - ["unsanitized-fallback"] + observed[0].sanitize_fields().metadata, + Some(json!({"original_metadata": true})) ); + drop(observed); deregister_mark_sanitize_guardrail("fail-open-mark-sanitizer").unwrap(); deregister_subscriber("fail-open-sanitizer-subscriber").unwrap(); } diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index da3d9da7e..63320c7a8 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -1335,7 +1335,7 @@ async fn install_registrations_covers_registry_error_edges() { } #[tokio::test(flavor = "multi_thread")] -#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. +#[allow(clippy::await_holding_lock)] // The process-wide test mutex intentionally serializes runtime state. async fn installed_callbacks_apply_surface_specific_fallbacks() { struct RuntimeCleanup { registrations: Option, @@ -1423,64 +1423,79 @@ async fn installed_callbacks_apply_surface_specific_fallbacks() { let llm_request = valid_llm_request(); let llm_response = json!({"response": "preserved"}); - { + let ( + subscribers, + mark_entries, + scope_start_entries, + scope_end_entries, + tool_request_entries, + tool_response_entries, + llm_request_entries, + llm_response_entries, + ) = { let state = context.read().unwrap(); - let subscribers = state.collect_event_subscribers(&[]); - NemoRelayContextState::emit_event(&event, &subscribers); - - for registry in [ - &state.mark_sanitize_guardrails, - &state.scope_sanitize_start_guardrails, - &state.scope_sanitize_end_guardrails, - ] { - let entries = NemoRelayContextState::event_sanitize_entries(registry, &[]); - let sanitized = - NemoRelayContextState::event_sanitize_snapshot_chain(event.clone(), &entries).await; - assert_eq!(sanitized.data(), event.data()); - assert_eq!(sanitized.metadata(), event.metadata()); - } + ( + state.collect_event_subscribers(&[]), + NemoRelayContextState::event_sanitize_entries(&state.mark_sanitize_guardrails, &[]), + NemoRelayContextState::event_sanitize_entries( + &state.scope_sanitize_start_guardrails, + &[], + ), + NemoRelayContextState::event_sanitize_entries( + &state.scope_sanitize_end_guardrails, + &[], + ), + state.tool_sanitize_request_entries(&[]), + state.tool_sanitize_response_entries(&[]), + state.llm_sanitize_request_entries(&[]), + state.llm_sanitize_response_entries(&[]), + ) + }; + NemoRelayContextState::emit_event(&event, &subscribers); - let entries = state.tool_sanitize_request_entries(&[]); - assert_eq!( - NemoRelayContextState::tool_sanitize_request_snapshot_chain( - "tool", - tool_request.clone(), - &entries, - ) - .await, - tool_request - ); - let entries = state.tool_sanitize_response_entries(&[]); - assert_eq!( - NemoRelayContextState::tool_sanitize_response_snapshot_chain( - "tool", - tool_response.clone(), - &entries, - ) - .await, - tool_response - ); - let entries = state.llm_sanitize_request_entries(&[]); - assert_eq!( - NemoRelayContextState::llm_sanitize_request_snapshot_chain( - llm_request.clone(), - crate::api::runtime::LlmSanitizeRequestContext::default(), - &entries, - ) - .await, - Some(llm_request), - ); - let entries = state.llm_sanitize_response_entries(&[]); - assert_eq!( - NemoRelayContextState::llm_sanitize_response_snapshot_chain( - llm_response.clone(), - crate::api::runtime::LlmSanitizeResponseContext::default(), - &entries, - ) - .await, - Some(llm_response), - ); + for entries in [mark_entries, scope_start_entries, scope_end_entries] { + let sanitized = + NemoRelayContextState::event_sanitize_snapshot_chain(event.clone(), &entries).await; + assert_eq!(sanitized.data(), event.data()); + assert_eq!(sanitized.metadata(), event.metadata()); } + + assert_eq!( + NemoRelayContextState::tool_sanitize_request_snapshot_chain( + "tool", + tool_request.clone(), + &tool_request_entries, + ) + .await, + tool_request + ); + assert_eq!( + NemoRelayContextState::tool_sanitize_response_snapshot_chain( + "tool", + tool_response.clone(), + &tool_response_entries, + ) + .await, + tool_response + ); + assert_eq!( + NemoRelayContextState::llm_sanitize_request_snapshot_chain( + llm_request.clone(), + crate::api::runtime::LlmSanitizeRequestContext::default(), + &llm_request_entries, + ) + .await, + Some(llm_request), + ); + assert_eq!( + NemoRelayContextState::llm_sanitize_response_snapshot_chain( + llm_response.clone(), + crate::api::runtime::LlmSanitizeResponseContext::default(), + &llm_response_entries, + ) + .await, + Some(llm_response), + ); crate::api::subscriber::flush_subscribers().expect("subscriber callback should flush"); } diff --git a/crates/core/tests/unit/llm_api_tests.rs b/crates/core/tests/unit/llm_api_tests.rs index 843fe4bf5..8e90747da 100644 --- a/crates/core/tests/unit/llm_api_tests.rs +++ b/crates/core/tests/unit/llm_api_tests.rs @@ -713,10 +713,25 @@ fn buffered_null_fallback_is_sanitized_before_emission() { ) .unwrap(); + let handle = create_llm_handle( + CreateLlmHandleParams::builder() + .name("buffered-explicit-null-fallback") + .build(), + ) + .unwrap(); + llm_call_end( + LlmCallEndParams::builder() + .handle(&handle) + .response(Json::Null) + .data(Json::Null) + .build(), + ) + .unwrap(); + flush_subscribers().unwrap(); let captured = events.lock().unwrap(); assert_eq!(*seen.lock().unwrap(), vec![fallback]); - assert_eq!(captured.len(), 3); + assert_eq!(captured.len(), 4); assert_eq!(captured[0].output(), Some(&Json::Null)); assert!(captured[0].annotated_response().is_none()); assert_eq!(captured[1].output(), Some(&redacted_response())); @@ -726,6 +741,8 @@ fn buffered_null_fallback_is_sanitized_before_emission() { ); assert!(captured[2].output().is_none()); assert!(captured[2].annotated_response().is_none()); + assert_eq!(captured[3].output(), Some(&Json::Null)); + assert!(captured[3].annotated_response().is_none()); assert!( captured .iter() diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index e876cbe34..0fc04fa9f 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -947,9 +947,11 @@ pub fn wrap_event_sanitize_fn( json_to_c_string(&serde_json::to_value(&fields).unwrap_or(Json::Null)); let result_ptr = unsafe { cb(ud.ptr, &ffi_event, fields_json) }; unsafe { nemo_relay_string_free_internal(fields_json) }; - let result = serde_json::from_value(ptr_to_json(result_ptr)).unwrap_or_default(); + let result = serde_json::from_value(ptr_to_json(result_ptr)).map_err(|error| { + FlowError::Internal(format!("invalid event sanitizer result: {error}")) + }); unsafe { nemo_relay_string_free_internal(result_ptr) }; - Ok(result) + result }) }) } diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 1ae03a985..70469b4f4 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -676,14 +676,18 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { assert_eq!(sanitize_calls.load(Ordering::SeqCst), 2); let invalid = wrap_event_sanitize_fn(invalid_event_sanitize_cb, std::ptr::null_mut(), None); - assert_eq!( - resolve(invalid(Arc::new(event.clone()), original_fields.clone())).unwrap(), - EventSanitizeFields::default() + assert!( + resolve(invalid(Arc::new(event.clone()), original_fields.clone())) + .unwrap_err() + .to_string() + .contains("invalid event sanitizer result") ); let null = wrap_event_sanitize_fn(null_event_sanitize_cb, std::ptr::null_mut(), None); - assert_eq!( - resolve(null(Arc::new(event), original_fields.clone())).unwrap(), - EventSanitizeFields::default() + assert!( + resolve(null(Arc::new(event), original_fields.clone())) + .unwrap_err() + .to_string() + .contains("invalid event sanitizer result") ); let handle = LlmHandle::builder() diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index c2ef47490..4a7eb2a58 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -705,11 +705,20 @@ pub fn wrap_js_llm_request_intercept_fn( move |name: String, request: LlmRequest, annotated: Option| { let func = func.clone(); Box::pin(async move { - let req_json = serde_json::to_value(&request).unwrap_or(Json::Null); - let annotated_json = annotated - .as_ref() - .map(|a| serde_json::to_value(a).unwrap_or(Json::Null)) - .unwrap_or(Json::Null); + let req_json = serde_json::to_value(&request).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM request intercept request: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let annotated_json = serde_json::to_value(annotated).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM request intercept annotation: {error}" + )); + record_callback_error(error.to_string()); + error + })?; let arg = serde_json::json!({ "name": name, "request": req_json, @@ -1004,7 +1013,13 @@ pub fn wrap_js_llm_conditional_fn( Arc::new(move |request: LlmRequest| { let func = func.clone(); Box::pin(async move { - let req_json = serde_json::to_value(request).unwrap_or(Json::Null); + let req_json = serde_json::to_value(request).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM conditional request: {error}" + )); + record_callback_error(error.to_string()); + error + })?; let (tx, rx) = tokio::sync::oneshot::channel(); let status = func.call_with_return_value( req_json, diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index 59f63ba33..5408f6fa9 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -8,6 +8,7 @@ import { createRequire } from 'node:module'; import { readFileSync } from 'node:fs'; import path from 'node:path'; import { fileURLToPath } from 'node:url'; +import { waitForSubscriberCallbacks } from './test_support.mjs'; const require = createRequire(import.meta.url); const lib = require('../index.js'); @@ -57,18 +58,6 @@ async function flushSubscriberCallbacks() { } } -async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { - const deadline = Date.now() + timeoutMs; - while (!predicate()) { - await flushSubscribers(); - if (Date.now() >= deadline) { - throw new Error('timed out waiting for subscriber callbacks'); - } - await new Promise((resolve) => setImmediate(resolve)); - } - await flushSubscribers(); -} - function makeNative() { return { headers: {}, diff --git a/crates/node/tests/scope_tests.mjs b/crates/node/tests/scope_tests.mjs index 0332790ac..3a184bfa7 100644 --- a/crates/node/tests/scope_tests.mjs +++ b/crates/node/tests/scope_tests.mjs @@ -4,6 +4,7 @@ import { describe, it } from 'node:test'; import assert from 'node:assert/strict'; import { createRequire } from 'node:module'; +import { waitForSubscriberCallbacks } from './test_support.mjs'; const require = createRequire(import.meta.url); const lib = require('../index.js'); @@ -29,18 +30,6 @@ function rejectWithPrimitive(value) { return Promise.reject(value); } -async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { - const deadline = Date.now() + timeoutMs; - while (!predicate()) { - await flushSubscribers(); - if (Date.now() >= deadline) { - throw new Error('timed out waiting for subscriber callbacks'); - } - await new Promise((resolve) => setImmediate(resolve)); - } - await flushSubscribers(); -} - // =========================================================================== // Scope operations // =========================================================================== diff --git a/crates/node/tests/test_support.mjs b/crates/node/tests/test_support.mjs new file mode 100644 index 000000000..17274dcf0 --- /dev/null +++ b/crates/node/tests/test_support.mjs @@ -0,0 +1,24 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import { createRequire } from 'node:module'; + +const require = createRequire(import.meta.url); +const { flushSubscribers } = require('../index.js'); + +export async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { + await flushSubscribers(); + // flushSubscribers() waits for Relay's Rust subscriber dispatcher, but JS + // subscriber callbacks are queued onto Node's event loop through N-API + // ThreadsafeFunction. Yield event-loop turns until the observed JS-side + // callback state is ready, with a timeout to avoid hanging the test forever. + const deadline = Date.now() + timeoutMs; + while (!predicate()) { + await flushSubscribers(); + if (Date.now() >= deadline) { + throw new Error('timed out waiting for subscriber callbacks'); + } + await new Promise((resolve) => setImmediate(resolve)); + } + await flushSubscribers(); +} diff --git a/crates/node/tests/tools_tests.mjs b/crates/node/tests/tools_tests.mjs index e104d0611..c9884bc05 100644 --- a/crates/node/tests/tools_tests.mjs +++ b/crates/node/tests/tools_tests.mjs @@ -4,6 +4,7 @@ import { describe, it } from 'node:test'; import assert from 'node:assert/strict'; import { createRequire } from 'node:module'; +import { waitForSubscriberCallbacks } from './test_support.mjs'; const require = createRequire(import.meta.url); const lib = require('../index.js'); @@ -47,23 +48,6 @@ function sparseArray() { return values; } -async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { - await flushSubscribers(); - // flushSubscribers() waits for Relay's Rust subscriber dispatcher, but JS - // subscriber callbacks are queued onto Node's event loop through N-API - // ThreadsafeFunction. Yield event-loop turns until the observed JS-side - // callback state is ready, with a timeout to avoid hanging the test forever. - const deadline = Date.now() + timeoutMs; - while (!predicate()) { - await flushSubscribers(); - if (Date.now() >= deadline) { - throw new Error('timed out waiting for subscriber callbacks'); - } - await new Promise((resolve) => setImmediate(resolve)); - } - await flushSubscribers(); -} - // =========================================================================== // Tool lifecycle // =========================================================================== diff --git a/docs/reference/event-sanitizers.mdx b/docs/reference/event-sanitizers.mdx index 6ff89e7e7..6122394b6 100644 --- a/docs/reference/event-sanitizers.mdx +++ b/docs/reference/event-sanitizers.mdx @@ -192,10 +192,15 @@ activation fails. The source-first C API retains `NemoRelayEventSanitizeCb` and adds parallel completion-based async registration APIs. An async callback returns `Complete` -or `Pending` and settles its one-shot completion handle with resolve or reject; -there is no implicit timeout. A callback that returns `Pending` must settle the -handle exactly once, or serial event publication remains blocked. Relay cancels -the handle when the invocation is abandoned; late or duplicate settlement after +or `Pending` and settles its one-shot completion handle with resolve or reject. +The absence of an implicit timeout is intentional: Relay preserves strict FIFO +publication, so one unsettled `Pending` completion blocks every later event in +that publication queue. Plugin authors should arrange their own operation +deadline and settle each retained completion exactly once on every success, +failure, and cancellation path. + +Relay cancels the handle when its invocation is abandoned, which is the +host-supported recovery mechanism; late or duplicate settlement after cancellation is rejected safely. After resolving or rejecting a retained completion, call `nemo_relay_async_completion_release` to release the callback-owned reference. Global names start with `nemo_relay_register_`, and diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index 9621cf4a1..e344cf738 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -20,9 +20,10 @@ observability sanitizer contract. Complete the following migrations before you run existing middleware or a sanitizer with a 0.7 host. -Do not deploy a 0.6 plugin or worker against a 0.7 host. The middleware -callback contract, LLM callback signature, native ABI layout, and worker -invocation schema changed. NeMo Relay does not adapt synchronous Rust +A 0.6 native plugin can still load through the legacy v2 table fallback, but it +is not compatible with the changed middleware and LLM sanitizer callback +contracts or other changed ABI and schema behavior. Rebuild plugins and workers +for 0.7 before using those surfaces. NeMo Relay does not adapt synchronous Rust middleware callbacks or one-argument LLM sanitizer callbacks. @@ -282,9 +283,11 @@ NeMo Relay 0.7 uses native ABI v3. Recompile native plugins against the 0.7 0.7 header. The v3 table preserves the v2 prefix, and Relay retries a legacy v2 table when -loading a plugin that rejects v3. Rebuild anyway if a plugin uses raw ABI -callbacks: v3 adds completion-based async middleware registration, async -execution continuations, and explicit cancellation/late-settlement behavior. +loading a plugin that rejects v3. That fallback supports loading, not +compatibility with changed middleware, LLM sanitizer, ABI, or schema contracts. +Rebuild plugins that use raw ABI callbacks: v3 adds completion-based async +middleware registration, async execution continuations, and explicit +cancellation/late-settlement behavior. The plugin manifest value remains `compat.native_api = "1"`. This manifest contract version is separate from the host ABI version; do not change it to From 436d46bfebaaafd2386c7985f18b0dc6d26d0399 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 17:36:09 -0400 Subject: [PATCH 16/52] test: align sanitizer failure expectations Signed-off-by: Will Killian --- crates/ffi/tests/unit/api/registry_tests.rs | 4 ++-- go/nemo_relay/event_sanitizers_test.go | 8 +++++--- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/crates/ffi/tests/unit/api/registry_tests.rs b/crates/ffi/tests/unit/api/registry_tests.rs index 97452897c..4ae23ce05 100644 --- a/crates/ffi/tests/unit/api/registry_tests.rs +++ b/crates/ffi/tests/unit/api/registry_tests.rs @@ -413,8 +413,8 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { .iter() .find(|event| event["name"] == "ffi-invalid-callback-mark") .expect("invalid callback mark should be delivered"); - assert_eq!(invalid_callback_event["data"], Json::Null); - assert_eq!(invalid_callback_event["metadata"], Json::Null); + assert_eq!(invalid_callback_event["data"], json!({"secret": true})); + assert_eq!(invalid_callback_event["metadata"], json!({"secret": true})); for name in ["ffi-local-child", "ffi-local-mark"] { for event in events.iter().filter(|event| event["name"] == name) { assert_eq!(event["data"], json!({"sanitized_by": name})); diff --git a/go/nemo_relay/event_sanitizers_test.go b/go/nemo_relay/event_sanitizers_test.go index 5c461fd9a..0a0b91426 100644 --- a/go/nemo_relay/event_sanitizers_test.go +++ b/go/nemo_relay/event_sanitizers_test.go @@ -25,7 +25,7 @@ func TestEventSanitizerRegistries(t *testing.T) { runTestWithScopeStack(t, testEventSanitizerRegistries) } -func TestEventSanitizerMarshalFailureClearsObservabilityFields(t *testing.T) { +func TestEventSanitizerMarshalFailurePreservesObservabilityFields(t *testing.T) { runTestWithScopeStack(t, func(t *testing.T) { var mu sync.Mutex var events []Event @@ -50,8 +50,10 @@ func TestEventSanitizerMarshalFailureClearsObservabilityFields(t *testing.T) { if len(events) != 1 { t.Fatalf("expected one event, got %d", len(events)) } - if len(events[0].Data()) != 0 || len(events[0].CategoryProfile()) != 0 || len(events[0].Metadata()) != 0 { - t.Fatalf("expected cleared observability fields, got data=%s category_profile=%s metadata=%s", events[0].Data(), events[0].CategoryProfile(), events[0].Metadata()) + if string(events[0].Data()) != `{"secret":true}` || + len(events[0].CategoryProfile()) != 0 || + string(events[0].Metadata()) != `{"secret":true}` { + t.Fatalf("expected the last valid observability fields, got data=%s category_profile=%s metadata=%s", events[0].Data(), events[0].CategoryProfile(), events[0].Metadata()) } }) } From 00ccbd7b2dcbe3db70d947a4e88e4017cfccaa14 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 17:39:24 -0400 Subject: [PATCH 17/52] test: move runtime panic coverage out of source Signed-off-by: Will Killian --- crates/core/src/api/runtime/state.rs | 108 +---------------- crates/core/tests/unit/runtime_state_tests.rs | 110 ++++++++++++++++++ 2 files changed, 112 insertions(+), 106 deletions(-) create mode 100644 crates/core/tests/unit/runtime_state_tests.rs diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index b7c6182d4..1af1314d2 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -1468,109 +1468,5 @@ impl Default for NemoRelayContextState { } #[cfg(test)] -mod panic_tests { - use super::*; - use crate::api::registry::{RegistryRecord, RequestIntercept}; - use serde_json::{Map, json}; - - #[tokio::test] - async fn middleware_snapshot_chains_contain_callback_panics() { - let tool_payload = json!({"tool": "preserved"}); - let tool_sanitizer: ToolSanitizeFn = - Arc::new(|_, _| Box::pin(async { panic!("tool sanitizer panic") })); - let tool_entries = vec![RegistryRecord::new("tool-panic", 0, tool_sanitizer)]; - assert_eq!( - NemoRelayContextState::tool_sanitize_request_snapshot_chain( - "tool", - tool_payload.clone(), - &tool_entries, - ) - .await, - tool_payload - ); - - let request = LlmRequest { - headers: Map::new(), - content: json!({"llm": "preserved"}), - }; - let llm_sanitizer: LlmSanitizeRequestFn = - Arc::new(|_, _| Box::pin(async { panic!("LLM sanitizer panic") })); - let llm_entries = vec![RegistryRecord::new("llm-panic", 0, llm_sanitizer)]; - assert_eq!( - NemoRelayContextState::llm_sanitize_request_snapshot_chain( - request.clone(), - LlmSanitizeRequestContext::default(), - &llm_entries, - ) - .await, - Some(request.clone()) - ); - - let tool_conditional: ToolConditionalFn = - Arc::new(|_, _| Box::pin(async { panic!("tool conditional panic") })); - let error = NemoRelayContextState::tool_conditional_execution_snapshot_chain( - "tool", - &tool_payload, - &[RegistryRecord::new( - "tool-conditional-panic", - 0, - tool_conditional, - )], - &[], - None, - None, - ) - .await - .unwrap_err(); - assert!(error.to_string().contains("tool-conditional-panic")); - - let llm_conditional: LlmConditionalFn = - Arc::new(|_| Box::pin(async { panic!("LLM conditional panic") })); - let error = NemoRelayContextState::llm_conditional_execution_snapshot_chain( - &request, - &[RegistryRecord::new( - "llm-conditional-panic", - 0, - llm_conditional, - )], - &[], - None, - None, - ) - .await - .unwrap_err(); - assert!(error.to_string().contains("llm-conditional-panic")); - - let tool_intercept: ToolInterceptFn = - Arc::new(|_, _| Box::pin(async { panic!("tool intercept panic") })); - let error = NemoRelayContextState::tool_request_intercepts_snapshot_chain( - "tool", - tool_payload, - &[RegistryRecord::new( - "tool-intercept-panic", - 0, - RequestIntercept::new(false, tool_intercept), - )], - ) - .await - .unwrap_err(); - assert!(error.to_string().contains("tool-intercept-panic")); - - let llm_intercept: LlmRequestInterceptFn = - Arc::new(|_, _, _| Box::pin(async { panic!("LLM intercept panic") })); - let error = NemoRelayContextState::llm_request_intercepts_snapshot_chain( - "llm", - request, - None, - &[RegistryRecord::new( - "llm-intercept-panic", - 0, - RequestIntercept::new(false, llm_intercept), - )], - false, - ) - .await - .unwrap_err(); - assert!(error.to_string().contains("llm-intercept-panic")); - } -} +#[path = "../../../tests/unit/runtime_state_tests.rs"] +mod tests; diff --git a/crates/core/tests/unit/runtime_state_tests.rs b/crates/core/tests/unit/runtime_state_tests.rs new file mode 100644 index 000000000..fd6a04f7f --- /dev/null +++ b/crates/core/tests/unit/runtime_state_tests.rs @@ -0,0 +1,110 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Unit tests for runtime middleware snapshot chains. + +use serde_json::{Map, json}; + +use super::*; +use crate::api::registry::{RegistryRecord, RequestIntercept}; + +#[tokio::test] +async fn middleware_snapshot_chains_contain_callback_panics() { + let tool_payload = json!({"tool": "preserved"}); + let tool_sanitizer: ToolSanitizeFn = + Arc::new(|_, _| Box::pin(async { panic!("tool sanitizer panic") })); + let tool_entries = vec![RegistryRecord::new("tool-panic", 0, tool_sanitizer)]; + assert_eq!( + NemoRelayContextState::tool_sanitize_request_snapshot_chain( + "tool", + tool_payload.clone(), + &tool_entries, + ) + .await, + tool_payload + ); + + let request = LlmRequest { + headers: Map::new(), + content: json!({"llm": "preserved"}), + }; + let llm_sanitizer: LlmSanitizeRequestFn = + Arc::new(|_, _| Box::pin(async { panic!("LLM sanitizer panic") })); + let llm_entries = vec![RegistryRecord::new("llm-panic", 0, llm_sanitizer)]; + assert_eq!( + NemoRelayContextState::llm_sanitize_request_snapshot_chain( + request.clone(), + LlmSanitizeRequestContext::default(), + &llm_entries, + ) + .await, + Some(request.clone()) + ); + + let tool_conditional: ToolConditionalFn = + Arc::new(|_, _| Box::pin(async { panic!("tool conditional panic") })); + let error = NemoRelayContextState::tool_conditional_execution_snapshot_chain( + "tool", + &tool_payload, + &[RegistryRecord::new( + "tool-conditional-panic", + 0, + tool_conditional, + )], + &[], + None, + None, + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("tool-conditional-panic")); + + let llm_conditional: LlmConditionalFn = + Arc::new(|_| Box::pin(async { panic!("LLM conditional panic") })); + let error = NemoRelayContextState::llm_conditional_execution_snapshot_chain( + &request, + &[RegistryRecord::new( + "llm-conditional-panic", + 0, + llm_conditional, + )], + &[], + None, + None, + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("llm-conditional-panic")); + + let tool_intercept: ToolInterceptFn = + Arc::new(|_, _| Box::pin(async { panic!("tool intercept panic") })); + let error = NemoRelayContextState::tool_request_intercepts_snapshot_chain( + "tool", + tool_payload, + &[RegistryRecord::new( + "tool-intercept-panic", + 0, + RequestIntercept::new(false, tool_intercept), + )], + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("tool-intercept-panic")); + + let llm_intercept: LlmRequestInterceptFn = + Arc::new(|_, _, _| Box::pin(async { panic!("LLM intercept panic") })); + let error = NemoRelayContextState::llm_request_intercepts_snapshot_chain( + "llm", + request, + None, + &[RegistryRecord::new( + "llm-intercept-panic", + 0, + RequestIntercept::new(false, llm_intercept), + )], + false, + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("llm-intercept-panic")); +} From fa55b029e3e4f2fdc2c13b60bb93e65055c498cb Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 17:48:26 -0400 Subject: [PATCH 18/52] test: cover all sanitizer panic paths Signed-off-by: Will Killian --- crates/core/tests/unit/runtime_state_tests.rs | 51 +++++++++++++++++++ 1 file changed, 51 insertions(+) diff --git a/crates/core/tests/unit/runtime_state_tests.rs b/crates/core/tests/unit/runtime_state_tests.rs index fd6a04f7f..8e6709b91 100644 --- a/crates/core/tests/unit/runtime_state_tests.rs +++ b/crates/core/tests/unit/runtime_state_tests.rs @@ -10,6 +10,25 @@ use crate::api::registry::{RegistryRecord, RequestIntercept}; #[tokio::test] async fn middleware_snapshot_chains_contain_callback_panics() { + let event = Event::Mark(MarkEvent::new( + BaseEvent::builder() + .name("preserved-event") + .data(json!({"event": "preserved"})) + .metadata(json!({"metadata": "preserved"})) + .build(), + None, + None, + )); + let event_sanitizer: EventSanitizeFn = + Arc::new(|_, _| Box::pin(async { panic!("event sanitizer panic") })); + let sanitized_event = NemoRelayContextState::event_sanitize_snapshot_chain( + event.clone(), + &[RegistryRecord::new("event-panic", 0, event_sanitizer)], + ) + .await; + assert_eq!(sanitized_event.data(), event.data()); + assert_eq!(sanitized_event.metadata(), event.metadata()); + let tool_payload = json!({"tool": "preserved"}); let tool_sanitizer: ToolSanitizeFn = Arc::new(|_, _| Box::pin(async { panic!("tool sanitizer panic") })); @@ -23,6 +42,22 @@ async fn middleware_snapshot_chains_contain_callback_panics() { .await, tool_payload ); + let tool_response = json!({"tool_response": "preserved"}); + let tool_response_sanitizer: ToolSanitizeFn = + Arc::new(|_, _| Box::pin(async { panic!("tool response sanitizer panic") })); + assert_eq!( + NemoRelayContextState::tool_sanitize_response_snapshot_chain( + "tool", + tool_response.clone(), + &[RegistryRecord::new( + "tool-response-panic", + 0, + tool_response_sanitizer, + )], + ) + .await, + tool_response + ); let request = LlmRequest { headers: Map::new(), @@ -40,6 +75,22 @@ async fn middleware_snapshot_chains_contain_callback_panics() { .await, Some(request.clone()) ); + let llm_response = json!({"llm_response": "preserved"}); + let llm_response_sanitizer: LlmSanitizeResponseFn = + Arc::new(|_, _| Box::pin(async { panic!("LLM response sanitizer panic") })); + assert_eq!( + NemoRelayContextState::llm_sanitize_response_snapshot_chain( + llm_response.clone(), + LlmSanitizeResponseContext::default(), + &[RegistryRecord::new( + "llm-response-panic", + 0, + llm_response_sanitizer, + )], + ) + .await, + Some(llm_response) + ); let tool_conditional: ToolConditionalFn = Arc::new(|_, _| Box::pin(async { panic!("tool conditional panic") })); From 165b080239bccf192062a95873fc1d2835b93f77 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 17:56:52 -0400 Subject: [PATCH 19/52] test: assert middleware panic error variants Signed-off-by: Will Killian --- crates/core/tests/unit/runtime_state_tests.rs | 20 +++++++++++++++---- 1 file changed, 16 insertions(+), 4 deletions(-) diff --git a/crates/core/tests/unit/runtime_state_tests.rs b/crates/core/tests/unit/runtime_state_tests.rs index 8e6709b91..09f01f279 100644 --- a/crates/core/tests/unit/runtime_state_tests.rs +++ b/crates/core/tests/unit/runtime_state_tests.rs @@ -108,7 +108,10 @@ async fn middleware_snapshot_chains_contain_callback_panics() { ) .await .unwrap_err(); - assert!(error.to_string().contains("tool-conditional-panic")); + assert!(matches!( + error, + FlowError::Internal(ref message) if message.contains("tool-conditional-panic") + )); let llm_conditional: LlmConditionalFn = Arc::new(|_| Box::pin(async { panic!("LLM conditional panic") })); @@ -125,7 +128,10 @@ async fn middleware_snapshot_chains_contain_callback_panics() { ) .await .unwrap_err(); - assert!(error.to_string().contains("llm-conditional-panic")); + assert!(matches!( + error, + FlowError::Internal(ref message) if message.contains("llm-conditional-panic") + )); let tool_intercept: ToolInterceptFn = Arc::new(|_, _| Box::pin(async { panic!("tool intercept panic") })); @@ -140,7 +146,10 @@ async fn middleware_snapshot_chains_contain_callback_panics() { ) .await .unwrap_err(); - assert!(error.to_string().contains("tool-intercept-panic")); + assert!(matches!( + error, + FlowError::Internal(ref message) if message.contains("tool-intercept-panic") + )); let llm_intercept: LlmRequestInterceptFn = Arc::new(|_, _, _| Box::pin(async { panic!("LLM intercept panic") })); @@ -157,5 +166,8 @@ async fn middleware_snapshot_chains_contain_callback_panics() { ) .await .unwrap_err(); - assert!(error.to_string().contains("llm-intercept-panic")); + assert!(matches!( + error, + FlowError::Internal(ref message) if message.contains("llm-intercept-panic") + )); } From 0c20cc1e9edc8010198fd9833df76ed5eaa7a993 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:10:47 -0400 Subject: [PATCH 20/52] fix: address hidden middleware review findings Signed-off-by: Will Killian --- crates/core/src/api/runtime/state.rs | 11 +++--- crates/plugin/tests/typed_callbacks.rs | 50 ++++++++++++++++++++------ 2 files changed, 45 insertions(+), 16 deletions(-) diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index 1af1314d2..6b36739fa 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -641,10 +641,11 @@ impl NemoRelayContextState { mut event: Event, entries: &[Guardrail], ) -> Event { + let context = Arc::new(event.clone()); for entry in entries { let fields = event.sanitize_fields(); let callback = Arc::clone(&entry.payload); - let context = Arc::new(event.clone()); + let context = Arc::clone(&context); match AssertUnwindSafe(async move { callback(context, fields).await }) .catch_unwind() .await @@ -1083,8 +1084,8 @@ impl NemoRelayContextState { target: "nemo_relay.runtime", event = "llm_request_sanitizer_failed", sanitizer = entry.name.as_str(), - preserved_value = "unsanitized_request"; - "LLM request sanitizer failed; preserving the last valid unsanitized request: {error}" + preserved_value = "last_valid_request"; + "LLM request sanitizer failed; preserving the last valid request: {error}" ); value = Some(current); } @@ -1153,8 +1154,8 @@ impl NemoRelayContextState { target: "nemo_relay.runtime", event = "llm_response_sanitizer_failed", sanitizer = entry.name.as_str(), - preserved_value = "unsanitized_response"; - "LLM response sanitizer failed; preserving the last valid unsanitized response: {error}" + preserved_value = "last_valid_response"; + "LLM response sanitizer failed; preserving the last valid response: {error}" ); value = Some(current); } diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index 865f05e16..821b8015f 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -18,17 +18,17 @@ use nemo_relay_plugin::{ LlmRequest, LlmRequestInterceptOutcome, LlmStream, LlmStreamNext, NEMO_RELAY_NATIVE_ABI_VERSION, NativePlugin, NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, - NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, - NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, - NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, - NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, - NemoRelayNativeLlmStreamV1, NemoRelayNativePluginContext, NemoRelayNativePluginV1, - NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, - NemoRelayNativeScopeType, NemoRelayNativeString, NemoRelayNativeToolConditionalCb, - NemoRelayNativeToolExecutionCb, NemoRelayNativeToolJsonCb, NemoRelayNativeWithScopeStackCb, - NemoRelayStatus, PendingMarkSpec, PluginContext, PluginRuntime, ScopeType, - ToolExecutionInterceptOutcome, ToolNext, + NemoRelayNativeHostApiV3, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, + NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmRequestCodec, + NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, + NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, + NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, + NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativePluginContext, + NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, + NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, NemoRelayNativeString, + NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, NemoRelayNativeToolJsonCb, + NemoRelayNativeWithScopeStackCb, NemoRelayStatus, PendingMarkSpec, PluginContext, + PluginRuntime, ScopeType, ToolExecutionInterceptOutcome, ToolNext, }; use serde_json::{Map, json}; @@ -328,6 +328,12 @@ fn native_abi_v3_struct_sizes_are_self_describing() { 280, 288, 296, 304, 312, ] ); + assert_eq!(align_of::(), 8); + assert_eq!(size_of::(), 376); + assert_eq!( + host_api_v3_offsets(), + [0, 320, 328, 336, 344, 352, 360, 368] + ); assert_eq!(align_of::(), 8); assert_eq!(size_of::(), 56); assert_eq!(plugin_offsets(), [0, 8, 16, 24, 32, 40, 48]); @@ -348,6 +354,12 @@ fn native_abi_v3_struct_sizes_are_self_describing() { 152, 156, ] ); + assert_eq!(align_of::(), 4); + assert_eq!(size_of::(), 188); + assert_eq!( + host_api_v3_offsets(), + [0, 160, 164, 168, 172, 176, 180, 184] + ); assert_eq!(align_of::(), 4); assert_eq!(size_of::(), 28); assert_eq!(plugin_offsets(), [0, 4, 8, 12, 16, 20, 24]); @@ -357,6 +369,22 @@ fn native_abi_v3_struct_sizes_are_self_describing() { } } +fn host_api_v3_offsets() -> [usize; 8] { + [ + offset_of!(NemoRelayNativeHostApiV3, v1), + offset_of!(NemoRelayNativeHostApiV3, async_completion_resolve_json), + offset_of!(NemoRelayNativeHostApiV3, async_completion_reject), + offset_of!(NemoRelayNativeHostApiV3, async_completion_is_cancelled), + offset_of!(NemoRelayNativeHostApiV3, async_completion_release), + offset_of!(NemoRelayNativeHostApiV3, async_next_invoke), + offset_of!(NemoRelayNativeHostApiV3, async_next_release), + offset_of!( + NemoRelayNativeHostApiV3, + plugin_context_register_async_middleware + ), + ] +} + fn host_api_offsets() -> [usize; 40] { [ offset_of!(NemoRelayNativeHostApiV1, abi_version), From b3df9db5412fb18fd7522c02f9b3bb54fce2e13e Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:23:03 -0400 Subject: [PATCH 21/52] fix: retain Python loop context for sanitizers Signed-off-by: Will Killian --- crates/python/src/py_callable.rs | 35 +++++++++++++++++++++++---- python/tests/test_event_sanitizers.py | 24 ++++++++++++++++++ python/tests/test_llm.py | 33 +++++++++++++++++++++++++ 3 files changed, 87 insertions(+), 5 deletions(-) diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 71d423e49..bea95c16d 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -34,6 +34,7 @@ use nemo_relay::api::runtime::{ use nemo_relay::error::{FlowError, Result as FlowResult}; use pyo3::prelude::*; use pyo3::types::PyDict; +use pyo3_async_runtimes::TaskLocals; use serde_json::Value as Json; use tokio_stream::Stream; use tokio_stream::wrappers::ReceiverStream; @@ -126,18 +127,38 @@ async fn resolve_json_or_future( fn split_py_object_or_future( py: Python<'_>, result: Py, +) -> FlowResult, PyValueFuture>> { + split_py_object_or_future_with_locals(py, result, None) +} + +fn split_py_object_or_future_with_locals( + py: Python<'_>, + result: Py, + task_locals: Option<&TaskLocals>, ) -> FlowResult, PyValueFuture>> { let bound = result.bind(py); if bound.getattr("__await__").is_ok() { reject_awaitable_from_sync_caller(bound)?; - let future = pyo3_async_runtimes::tokio::into_future(result.into_bound(py)) - .map_err(|e| FlowError::Internal(e.to_string()))?; - Ok(Err(Box::pin(future) as PyValueFuture)) + let future: PyValueFuture = match task_locals { + Some(locals) => Box::pin( + pyo3_async_runtimes::into_future_with_locals(locals, result.into_bound(py)) + .map_err(|e| FlowError::Internal(e.to_string()))?, + ), + None => Box::pin( + pyo3_async_runtimes::tokio::into_future(result.into_bound(py)) + .map_err(|e| FlowError::Internal(e.to_string()))?, + ), + }; + Ok(Err(future)) } else { Ok(Ok(result)) } } +fn capture_python_task_locals() -> Option { + Python::attach(|py| pyo3_async_runtimes::tokio::get_current_locals(py).ok()) +} + async fn resolve_py_object_or_future( outcome: FlowResult, PyValueFuture>>, ) -> FlowResult> { @@ -1048,8 +1069,10 @@ pub fn wrap_py_finalizer_fn(py_fn: Py) -> Box Json + Send /// Wrap a Python callable `(Json, LlmSanitizeResponseContext) -> Optional[Json]`. fn wrap_py_llm_sanitize_response_callback(py_fn: Py) -> LlmSanitizeResponseFn { let py_fn = Arc::new(py_fn); + let task_locals = capture_python_task_locals(); Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { let py_fn = py_fn.clone(); + let task_locals = task_locals.clone(); Box::pin(async move { let result = resolve_py_object_or_future(Python::attach(|py| { let py_context = PyLlmSanitizeResponseContext { inner: context }; @@ -1058,7 +1081,7 @@ fn wrap_py_llm_sanitize_response_callback(py_fn: Py) -> LlmSanitizeRespon let result = py_fn .call1(py, (py_response, py_context)) .map_err(|error| FlowError::Internal(error.to_string()))?; - split_py_object_or_future(py, result) + split_py_object_or_future_with_locals(py, result, task_locals.as_ref()) })) .await?; Python::attach(|py| { @@ -1114,8 +1137,10 @@ pub fn wrap_py_event_subscriber(py_fn: Py) -> EventSubscriberFn { /// Wrap a Python callable ``(Event, EventSanitizeFields) -> EventSanitizeFields``. pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let py_fn = Arc::new(py_fn); + let task_locals = capture_python_task_locals(); Arc::new(move |event: Arc, fields: EventSanitizeFields| { let py_fn = py_fn.clone(); + let task_locals = task_locals.clone(); Box::pin(async move { let result = Python::attach( |py| -> FlowResult, PyValueFuture>> { @@ -1156,7 +1181,7 @@ pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let result = py_fn .call1(py, (py_event, py_fields)) .map_err(|error| FlowError::Internal(error.to_string()))?; - split_py_object_or_future(py, result) + split_py_object_or_future_with_locals(py, result, task_locals.as_ref()) }, ); let result = resolve_py_object_or_future(result).await?; diff --git a/python/tests/test_event_sanitizers.py b/python/tests/test_event_sanitizers.py index 08234b70e..4a91bf658 100644 --- a/python/tests/test_event_sanitizers.py +++ b/python/tests/test_event_sanitizers.py @@ -3,6 +3,7 @@ from __future__ import annotations +import asyncio from collections.abc import Iterator from typing import cast @@ -73,6 +74,29 @@ def raises(_event: nemo_relay.Event, _fields: EventSanitizeFields) -> EventSanit assert events[-1].metadata is None +async def test_async_mark_sanitizer_runs_on_originating_loop(capture_events): + _capture_name, events = capture_events + originating_loop = asyncio.get_running_loop() + + async def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + await asyncio.sleep(0) + assert asyncio.get_running_loop() is originating_loop + return { + "data": {"async": True}, + "category_profile": fields["category_profile"], + "metadata": fields["metadata"], + } + + guardrails.register_mark_sanitize("python-async-mark", 0, sanitize) + try: + scope.event("async-checkpoint", data={"raw": True}) + await asyncio.to_thread(subscribers.flush) + finally: + guardrails.deregister_mark_sanitize("python-async-mark") + + assert events[-1].data == {"async": True} + + def test_scope_start_and_end_sanitizers_cover_category_profile(capture_events): _capture_name, events = capture_events diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index 411629401..3da6c3d68 100644 --- a/python/tests/test_llm.py +++ b/python/tests/test_llm.py @@ -665,6 +665,39 @@ def finalizer(): # Collector should have received all chunks assert len(collected) == len(chunks) + async def test_async_response_sanitizer_runs_during_stream_finalization(self): + events = [] + originating_loop = asyncio.get_running_loop() + subscribers.register("py_llm_async_stream_sanitizer_sub", events.append) + + async def sanitize_response(response, context): + del context + await asyncio.sleep(0) + assert asyncio.get_running_loop() is originating_loop + return {"sanitized": response["raw"]} + + async def stream_func(request): + del request + yield {"token": "hello"} + + guardrails.register_llm_sanitize_response("py_llm_async_stream_sanitizer", 1, sanitize_response) + try: + stream = await llm.stream_execute( + "stream_async_response_sanitizer", + make_request(), + stream_func, + lambda chunk: None, + lambda: {"raw": True}, + ) + assert [chunk async for chunk in stream] == [{"token": "hello"}] + await asyncio.to_thread(subscribers.flush) + finally: + guardrails.deregister_llm_sanitize_response("py_llm_async_stream_sanitizer") + subscribers.deregister("py_llm_async_stream_sanitizer_sub") + + end = _llm_event(events, "stream_async_response_sanitizer", "end") + assert end.data == {"sanitized": True} + async def test_stream_execute_aclose_stops_partially_consumed_stream(self): producer_closed = asyncio.Event() wait_for_more_chunks = asyncio.Event() From b22988009fcfd1786b1ca3fb412a41df689db484 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:41:07 -0400 Subject: [PATCH 22/52] fix: preserve progressive event sanitizer context Signed-off-by: Will Killian --- crates/core/src/api/runtime/state.rs | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index 6b36739fa..b83e3d170 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -641,15 +641,16 @@ impl NemoRelayContextState { mut event: Event, entries: &[Guardrail], ) -> Event { - let context = Arc::new(event.clone()); for entry in entries { let fields = event.sanitize_fields(); let callback = Arc::clone(&entry.payload); - let context = Arc::clone(&context); - match AssertUnwindSafe(async move { callback(context, fields).await }) + let context = Arc::new(event); + let callback_context = Arc::clone(&context); + let outcome = AssertUnwindSafe(async move { callback(callback_context, fields).await }) .catch_unwind() - .await - { + .await; + event = Arc::try_unwrap(context).unwrap_or_else(|context| (*context).clone()); + match outcome { Ok(Ok(fields)) => event.apply_sanitize_fields(fields), Ok(Err(error)) => log::error!( target: "nemo_relay.runtime", From 74289a91b28c4677d1dffa1f56b3058ce1e13948 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 19:12:55 -0400 Subject: [PATCH 23/52] fix: address async middleware review findings Signed-off-by: Will Killian --- .../src/api/runtime/subscriber_dispatcher.rs | 17 ++++- crates/core/src/api/tool.rs | 24 ++---- crates/core/src/plugin/dynamic/native.rs | 25 ++++++- crates/core/src/stream.rs | 6 +- .../tests/fixtures/native_plugin/src/lib.rs | 49 +++++++----- .../tests/integration/native_plugin_tests.rs | 17 ++++- crates/plugin/src/lib.rs | 67 +++++++++++++---- crates/plugin/tests/typed_callbacks.rs | 17 ++++- crates/python/src/py_api/mod.rs | 74 +++++++++++-------- python/tests/test_llm.py | 7 +- 10 files changed, 211 insertions(+), 92 deletions(-) diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 5421bdca9..1d46f3fe6 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -350,12 +350,23 @@ mod native { if sanitizers.is_empty() { return Some(transformed); } - Some( + let fallback = transformed.clone(); + match catch_unwind(AssertUnwindSafe(|| { runtime.block_on(NemoRelayContextState::event_sanitize_snapshot_chain( transformed, &sanitizers, - )), - ) + )) + })) { + Ok(event) => Some(event), + Err(_) => { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_panicked"; + "Event sanitizer panicked; preserving the last valid event snapshot" + ); + Some(fallback) + } + } } } diff --git a/crates/core/src/api/tool.rs b/crates/core/src/api/tool.rs index 768d31670..671471748 100644 --- a/crates/core/src/api/tool.rs +++ b/crates/core/src/api/tool.rs @@ -377,17 +377,13 @@ async fn tool_call_with_subscriber_snapshot( .collect::>(); (handle, event, marks) }; - let mut sanitized_marks = Vec::with_capacity(marks.len()); - for mark in marks { - if let Some(mark) = sanitize_event(mark).await { - sanitized_marks.push(mark); - } - } if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, &subscribers); } - for mark in sanitized_marks { - NemoRelayContextState::emit_event(&mark, &subscribers); + for mark in marks { + if let Some(mark) = sanitize_event(mark).await { + NemoRelayContextState::emit_event(&mark, &subscribers); + } } Ok((handle, subscribers)) } @@ -549,17 +545,13 @@ async fn tool_call_end_with_pending_marks( )) }) .collect::>(); - let mut sanitized_marks = Vec::with_capacity(marks.len()); - for mark in marks { - if let Some(mark) = sanitize_event(mark).await { - sanitized_marks.push(mark); - } - } if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, subscribers); } - for mark in sanitized_marks { - NemoRelayContextState::emit_event(&mark, subscribers); + for mark in marks { + if let Some(mark) = sanitize_event(mark).await { + NemoRelayContextState::emit_event(&mark, subscribers); + } } Ok(()) } diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 084a8900b..0f51a25c4 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -1465,6 +1465,22 @@ async fn invoke_native_async_callback( } }; unsafe { native_string_free(invocation as *mut NemoRelayNativeString) }; + let state = match NemoRelayNativeAsyncCallbackState::try_from(state) { + Ok(state) => state, + Err(()) => { + unsafe { + drop(Arc::from_raw( + completion_ref as *const NativeAsyncCompletion, + )); + if let Some(next_ref) = next_ref { + drop(Arc::from_raw(next_ref as *const NativeAsyncNext)); + } + } + return Err(FlowError::Internal( + "native async callback returned an invalid state".into(), + )); + } + }; if state == NemoRelayNativeAsyncCallbackState::Complete { unsafe { drop(Arc::from_raw( @@ -1935,7 +1951,7 @@ fn wrap_native_async_llm_stream_execution( unsafe extern "C" fn native_plugin_context_register_async_middleware( ctx: *mut NemoRelayNativePluginContext, - kind: NemoRelayNativeAsyncMiddlewareKind, + kind: u32, name: *const NemoRelayNativeString, priority: i32, break_chain: bool, @@ -1953,6 +1969,13 @@ unsafe extern "C" fn native_plugin_context_register_async_middleware( Ok(name) => name, Err(status) => return status, }; + let kind = match NemoRelayNativeAsyncMiddlewareKind::try_from(kind) { + Ok(kind) => kind, + Err(()) => { + set_native_last_error("invalid native async middleware kind"); + return NemoRelayStatus::InvalidArg; + } + }; let context = unsafe { &mut *host_ctx.ctx }; let registration = match kind { NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest => context diff --git a/crates/core/src/stream.rs b/crates/core/src/stream.rs index 2bae33344..a7cd81761 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -389,9 +389,9 @@ impl LlmStreamWrapper { Err(_) => None, } }; - if let Some(event) = event_snapshot - && let Some(sanitizers) = snapshot_event_sanitizers(&event, &self.scope_stack) - { + if let Some(event) = event_snapshot { + let sanitizers = + snapshot_event_sanitizers(&event, &self.scope_stack).unwrap_or_default(); let _ = subscriber_dispatcher::dispatch_sanitized_event( event, sanitizers, diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index d1928528e..00d417117 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -3,6 +3,7 @@ use std::ffi::c_void; use std::ptr; +use std::sync::atomic::{AtomicBool, Ordering}; use nemo_relay_plugin::{ CategoryProfile, ConfigDiagnostic, DiagnosticLevel, Event, EventCategory, EventSanitizeFields, @@ -17,6 +18,13 @@ use serde_json::{Map, json}; struct FixtureNativePlugin; +static ASYNC_PENDING_ENTERED: AtomicBool = AtomicBool::new(false); + +#[unsafe(no_mangle)] +pub extern "C" fn nemo_relay_fixture_async_pending_entered() -> bool { + ASYNC_PENDING_ENTERED.swap(false, Ordering::AcqRel) +} + impl NativePlugin for FixtureNativePlugin { fn plugin_kind(&self) -> &str { "fixture_native" @@ -634,7 +642,7 @@ unsafe extern "C" fn raw_register_async_tool_request( } let status = unsafe { (host.plugin_context_register_async_middleware)( - ctx, kind, name, 0, false, callback, user_data, None, + ctx, kind as u32, name, 0, false, callback, user_data, None, ) }; unsafe { (host.v1.string_free)(name) }; @@ -650,9 +658,9 @@ unsafe extern "C" fn raw_async_allow_callback( _invocation_json: *const NemoRelayNativeString, _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, -) -> NemoRelayNativeAsyncCallbackState { +) -> u32 { let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; }; let result = unsafe { raw_host_string(&host.v1, "null") }; if result.is_null() { @@ -663,7 +671,7 @@ unsafe extern "C" fn raw_async_allow_callback( (host.v1.string_free)(result); } } - NemoRelayNativeAsyncCallbackState::Complete + NemoRelayNativeAsyncCallbackState::Complete as u32 } unsafe extern "C" fn raw_async_passthrough_callback( @@ -671,9 +679,9 @@ unsafe extern "C" fn raw_async_passthrough_callback( invocation_json: *const NemoRelayNativeString, _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, -) -> NemoRelayNativeAsyncCallbackState { +) -> u32 { let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; }; let result = unsafe { raw_host_string_value(&host.v1, invocation_json) } .and_then(|value| serde_json::from_str::(&value).ok()) @@ -694,7 +702,7 @@ unsafe extern "C" fn raw_async_passthrough_callback( .and_then(|value| serde_json::to_string(&value).ok()); let Some(result) = result else { unsafe { reject_async_completion(host, completion, "invalid async passthrough invocation") }; - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; }; let result = unsafe { raw_host_string(&host.v1, &result) }; if result.is_null() { @@ -705,7 +713,7 @@ unsafe extern "C" fn raw_async_passthrough_callback( (host.v1.string_free)(result); } } - NemoRelayNativeAsyncCallbackState::Complete + NemoRelayNativeAsyncCallbackState::Complete as u32 } unsafe extern "C" fn raw_async_tool_request_callback( @@ -713,9 +721,9 @@ unsafe extern "C" fn raw_async_tool_request_callback( invocation_json: *const NemoRelayNativeString, _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, -) -> NemoRelayNativeAsyncCallbackState { +) -> u32 { let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; }; let invocation = unsafe { raw_host_string_value(&host.v1, invocation_json) } .and_then(|json| serde_json::from_str::(&json).ok()) @@ -737,9 +745,10 @@ unsafe extern "C" fn raw_async_tool_request_callback( }); let Some((result, pending, duplicate)) = invocation else { unsafe { reject_async_completion(host, completion, "invalid async tool request invocation") }; - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; }; if pending { + ASYNC_PENDING_ENTERED.store(true, Ordering::Release); let host = *host; let completion = completion as usize; std::thread::spawn(move || { @@ -769,7 +778,7 @@ unsafe extern "C" fn raw_async_tool_request_callback( } } }); - return NemoRelayNativeAsyncCallbackState::Pending; + return NemoRelayNativeAsyncCallbackState::Pending as u32; } let result = unsafe { raw_host_string(&host.v1, &result) }; if !result.is_null() { @@ -783,7 +792,7 @@ unsafe extern "C" fn raw_async_tool_request_callback( } else { unsafe { reject_async_completion(host, completion, "failed to allocate async tool request result") }; } - NemoRelayNativeAsyncCallbackState::Complete + NemoRelayNativeAsyncCallbackState::Complete as u32 } unsafe extern "C" fn raw_async_tool_execution_callback( @@ -791,16 +800,16 @@ unsafe extern "C" fn raw_async_tool_execution_callback( invocation_json: *const NemoRelayNativeString, next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, -) -> NemoRelayNativeAsyncCallbackState { +) -> u32 { let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; }; if next.is_null() || completion.is_null() { unsafe { reject_async_completion(host, completion, "async tool execution requires next and completion") }; if !next.is_null() { unsafe { (host.async_next_release)(next) }; } - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; } let value = unsafe { raw_host_string_value(&host.v1, invocation_json) } .and_then(|json| serde_json::from_str::(&json).ok()) @@ -816,13 +825,13 @@ unsafe extern "C" fn raw_async_tool_execution_callback( let Some(value) = value else { unsafe { reject_async_completion(host, completion, "invalid async tool execution invocation") }; unsafe { (host.async_next_release)(next) }; - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; }; let value = unsafe { raw_host_string(&host.v1, &value) }; if value.is_null() { unsafe { reject_async_completion(host, completion, "failed to allocate async tool execution invocation") }; unsafe { (host.async_next_release)(next) }; - return NemoRelayNativeAsyncCallbackState::Complete; + return NemoRelayNativeAsyncCallbackState::Complete as u32; } let status = unsafe { (host.async_next_invoke)(next, value, completion) }; unsafe { @@ -833,11 +842,11 @@ unsafe extern "C" fn raw_async_tool_execution_callback( (host.async_next_release)(next); (host.async_completion_release)(completion); } - NemoRelayNativeAsyncCallbackState::Pending + NemoRelayNativeAsyncCallbackState::Pending as u32 } else { unsafe { reject_async_completion(host, completion, "failed to invoke async tool execution next") }; unsafe { (host.async_next_release)(next) }; - NemoRelayNativeAsyncCallbackState::Complete + NemoRelayNativeAsyncCallbackState::Complete as u32 } } diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index 893c5ca15..91a07a4cf 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -673,6 +673,15 @@ async fn native_v3_async_registration_supports_all_middleware_kinds() { manifest_ref: manifest_ref.to_string_lossy().into_owned(), }]) .expect("v3 async native fixture should load"); + let fixture_library = unsafe { libloading::Library::new(&fixture.library_path) } + .expect("loaded v3 async native fixture should open for synchronization"); + let pending_entered = unsafe { + *fixture_library + .get:: bool>(b"nemo_relay_fixture_async_pending_entered\0") + .expect("v3 async native fixture should export its pending-entry signal") + }; + assert!(!unsafe { pending_entered() }); + drop(fixture_library); let mut cleanup = NativePluginTestCleanup::new(); let mut config = PluginConfig::default(); config.components.push(PluginComponentSpec { @@ -766,7 +775,13 @@ async fn native_v3_async_registration_supports_all_middleware_kinds() { let pending = tokio::spawn(async { tool_request_intercepts("async-pending", json!({"input": true})).await }); - tokio::task::yield_now().await; + tokio::time::timeout(std::time::Duration::from_secs(10), async { + while !unsafe { pending_entered() } { + tokio::task::yield_now().await; + } + }) + .await + .expect("native async callback should enter before plugin clear"); clear_plugin_configuration().expect("v3 async native fixture should clear while pending"); cleanup.plugin_configuration_active = false; let pending = pending diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index 5f5500072..1e08c05f7 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -800,6 +800,30 @@ pub enum NemoRelayNativeAsyncMiddlewareKind { ScopeSanitizeEnd = 13, } +impl TryFrom for NemoRelayNativeAsyncMiddlewareKind { + type Error = (); + + fn try_from(value: u32) -> std::result::Result { + match value { + 0 => Ok(Self::ToolSanitizeRequest), + 1 => Ok(Self::ToolSanitizeResponse), + 2 => Ok(Self::ToolConditionalExecution), + 3 => Ok(Self::ToolRequestIntercept), + 4 => Ok(Self::ToolExecutionIntercept), + 5 => Ok(Self::LlmSanitizeRequest), + 6 => Ok(Self::LlmSanitizeResponse), + 7 => Ok(Self::LlmConditionalExecution), + 8 => Ok(Self::LlmRequestIntercept), + 9 => Ok(Self::LlmExecutionIntercept), + 10 => Ok(Self::LlmStreamExecutionIntercept), + 11 => Ok(Self::MarkSanitize), + 12 => Ok(Self::ScopeSanitizeStart), + 13 => Ok(Self::ScopeSanitizeEnd), + _ => Err(()), + } + } +} + /// Indicates whether an asynchronous native callback settled before returning. #[repr(u32)] #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -810,6 +834,18 @@ pub enum NemoRelayNativeAsyncCallbackState { Pending = 1, } +impl TryFrom for NemoRelayNativeAsyncCallbackState { + type Error = (); + + fn try_from(value: u32) -> std::result::Result { + match value { + 0 => Ok(Self::Complete), + 1 => Ok(Self::Pending), + _ => Err(()), + } + } +} + /// Opaque one-shot completion retained by a pending native callback. #[repr(C)] pub struct NemoRelayNativeAsyncCompletion { @@ -827,18 +863,18 @@ pub struct NemoRelayNativeAsyncNext { /// Completion-based native middleware callback. /// /// `invocation_json` is borrowed for the call. A callback that returns -/// [`NemoRelayNativeAsyncCallbackState::Pending`] owns one completion -/// reference and must settle it then call the v3 `async_completion_release` -/// hook. When `next` is non-null, the callback owns that handle for the -/// invocation and must call `async_next_release` after its final use. `next` -/// is null for non-execution middleware. -pub type NemoRelayNativeAsyncMiddlewareCb = - unsafe extern "C" fn( - user_data: *mut c_void, - invocation_json: *const NemoRelayNativeString, - next: *const NemoRelayNativeAsyncNext, - completion: *const NemoRelayNativeAsyncCompletion, - ) -> NemoRelayNativeAsyncCallbackState; +/// [`NemoRelayNativeAsyncCallbackState::Pending`] as a `u32` owns one +/// completion reference and must settle it then call the v3 +/// `async_completion_release` hook. The host validates the returned +/// discriminant. When `next` is non-null, the callback owns that handle for +/// the invocation and must call `async_next_release` after its final use. +/// `next` is null for non-execution middleware. +pub type NemoRelayNativeAsyncMiddlewareCb = unsafe extern "C" fn( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, +) -> u32; /// ABI-v3 host extension appended to [`NemoRelayNativeHostApiV1`]. /// @@ -874,9 +910,12 @@ pub struct NemoRelayNativeHostApiV3 { /// Releases the callback-owned continuation reference for a pending callback. pub async_next_release: unsafe extern "C" fn(next: *const NemoRelayNativeAsyncNext), /// Registers any completion-based asynchronous middleware surface. + /// + /// `kind` must be a valid [`NemoRelayNativeAsyncMiddlewareKind`] + /// discriminant. The host rejects unknown `u32` values. pub plugin_context_register_async_middleware: unsafe extern "C" fn( ctx: *mut NemoRelayNativePluginContext, - kind: NemoRelayNativeAsyncMiddlewareKind, + kind: u32, name: *const NemoRelayNativeString, priority: i32, break_chain: bool, @@ -2396,7 +2435,7 @@ impl<'a> PluginContext<'a> { self.with_name(name, |_, name| unsafe { (host.plugin_context_register_async_middleware)( self.raw, - kind, + kind as u32, name, priority, break_chain, diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index 821b8015f..eccca24aa 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -16,7 +16,8 @@ use nemo_relay_plugin::{ AnnotatedLlmRequest, BuiltinLlmCodec, CategoryProfile, ConfigDiagnostic, DiagnosticLevel, Event, EventCategory, EventSanitizeFields, Json, LlmCodecIdentity, LlmJsonStream, LlmNext, LlmRequest, LlmRequestInterceptOutcome, LlmStream, LlmStreamNext, - NEMO_RELAY_NATIVE_ABI_VERSION, NativePlugin, NemoRelayNativeEventSanitizeCb, + NEMO_RELAY_NATIVE_ABI_VERSION, NativePlugin, NemoRelayNativeAsyncCallbackState, + NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmRequestCodec, @@ -32,6 +33,20 @@ use nemo_relay_plugin::{ }; use serde_json::{Map, json}; +#[test] +fn async_abi_discriminants_reject_unknown_values() { + assert_eq!( + NemoRelayNativeAsyncMiddlewareKind::try_from(13), + Ok(NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeEnd) + ); + assert!(NemoRelayNativeAsyncMiddlewareKind::try_from(14).is_err()); + assert_eq!( + NemoRelayNativeAsyncCallbackState::try_from(1), + Ok(NemoRelayNativeAsyncCallbackState::Pending) + ); + assert!(NemoRelayNativeAsyncCallbackState::try_from(2).is_err()); +} + struct TestString(Vec); struct RegisteredSubscriber { diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index bf6280c04..534f68c99 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1327,13 +1327,17 @@ fn tool_request_intercepts<'py>( .is_err() { let scope_stack = current_scope_stack_handle(); - let result = pyo3_async_runtimes::tokio::get_runtime() - .block_on(py_callable::PY_AWAITABLES_ALLOWED.scope( - false, - TASK_SCOPE_STACK.scope(scope_stack, async move { - core_tool_api::tool_request_intercepts(&name, args_json).await - }), - )) + let result = py + .detach(|| { + pyo3_async_runtimes::tokio::get_runtime().block_on( + py_callable::PY_AWAITABLES_ALLOWED.scope( + false, + TASK_SCOPE_STACK.scope(scope_stack, async move { + core_tool_api::tool_request_intercepts(&name, args_json).await + }), + ), + ) + }) .map_err(to_py_err)?; return json_to_py(py, &result).map(|value| value.into_bound(py)); } @@ -1370,14 +1374,17 @@ fn tool_conditional_execution<'py>( .is_err() { let scope_stack = current_scope_stack_handle(); - pyo3_async_runtimes::tokio::get_runtime() - .block_on(py_callable::PY_AWAITABLES_ALLOWED.scope( - false, - TASK_SCOPE_STACK.scope(scope_stack, async move { - core_tool_api::tool_conditional_execution(&name, &args_json).await - }), - )) - .map_err(to_py_err)?; + py.detach(|| { + pyo3_async_runtimes::tokio::get_runtime().block_on( + py_callable::PY_AWAITABLES_ALLOWED.scope( + false, + TASK_SCOPE_STACK.scope(scope_stack, async move { + core_tool_api::tool_conditional_execution(&name, &args_json).await + }), + ), + ) + }) + .map_err(to_py_err)?; return Ok(py.None().into_bound(py)); } let scope_stack = current_scope_stack_handle(); @@ -1413,13 +1420,17 @@ fn llm_request_intercepts<'py>( .is_err() { let scope_stack = current_scope_stack_handle(); - let result = pyo3_async_runtimes::tokio::get_runtime() - .block_on(py_callable::PY_AWAITABLES_ALLOWED.scope( - false, - TASK_SCOPE_STACK.scope(scope_stack, async move { - core_llm_api::llm_request_intercepts(&name, request.inner).await - }), - )) + let result = py + .detach(|| { + pyo3_async_runtimes::tokio::get_runtime().block_on( + py_callable::PY_AWAITABLES_ALLOWED.scope( + false, + TASK_SCOPE_STACK.scope(scope_stack, async move { + core_llm_api::llm_request_intercepts(&name, request.inner).await + }), + ), + ) + }) .map_err(to_py_err)?; return Py::new( py, @@ -1457,14 +1468,17 @@ fn llm_conditional_execution<'py>( .is_err() { let scope_stack = current_scope_stack_handle(); - pyo3_async_runtimes::tokio::get_runtime() - .block_on(py_callable::PY_AWAITABLES_ALLOWED.scope( - false, - TASK_SCOPE_STACK.scope(scope_stack, async move { - core_llm_api::llm_conditional_execution(&request.inner).await - }), - )) - .map_err(to_py_err)?; + py.detach(|| { + pyo3_async_runtimes::tokio::get_runtime().block_on( + py_callable::PY_AWAITABLES_ALLOWED.scope( + false, + TASK_SCOPE_STACK.scope(scope_stack, async move { + core_llm_api::llm_conditional_execution(&request.inner).await + }), + ), + ) + }) + .map_err(to_py_err)?; return Ok(py.None().into_bound(py)); } let scope_stack = current_scope_stack_handle(); diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index 3da6c3d68..b956d8f64 100644 --- a/python/tests/test_llm.py +++ b/python/tests/test_llm.py @@ -4,6 +4,7 @@ """Tests for NeMo Relay LLM lifecycle, guardrails, intercepts, and streaming.""" import asyncio +from collections.abc import AsyncIterator from typing import NoReturn, cast import pytest @@ -670,13 +671,13 @@ async def test_async_response_sanitizer_runs_during_stream_finalization(self): originating_loop = asyncio.get_running_loop() subscribers.register("py_llm_async_stream_sanitizer_sub", events.append) - async def sanitize_response(response, context): + async def sanitize_response(response, context) -> dict: del context await asyncio.sleep(0) assert asyncio.get_running_loop() is originating_loop return {"sanitized": response["raw"]} - async def stream_func(request): + async def stream_func(request) -> AsyncIterator[dict]: del request yield {"token": "hello"} @@ -686,7 +687,7 @@ async def stream_func(request): "stream_async_response_sanitizer", make_request(), stream_func, - lambda chunk: None, + lambda _chunk: None, lambda: {"raw": True}, ) assert [chunk async for chunk in stream] == [{"token": "hello"}] From 669a5296f516bd370796518370023c35c6a70678 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 19:39:41 -0400 Subject: [PATCH 24/52] fix: prevent async sanitizer flush deadlocks Signed-off-by: Will Killian --- crates/core/src/api/runtime.rs | 2 +- .../src/api/runtime/subscriber_dispatcher.rs | 39 +++++++++++++++++-- crates/core/src/api/subscriber.rs | 7 ++++ crates/ffi/nemo_relay.h | 6 +-- crates/ffi/src/api/llm_registry.rs | 8 ++-- crates/node/src/api/mod.rs | 8 ++-- crates/node/tests/event_sanitizers_tests.mjs | 19 +++++++++ crates/plugin/tests/typed_callbacks.rs | 26 +++++++++++-- crates/python/src/py_api/mod.rs | 8 ++-- 9 files changed, 100 insertions(+), 23 deletions(-) diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index 77670c804..b2878cead 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -27,4 +27,4 @@ pub use scope_stack::{ task_scope_top, with_active_event_uuid, with_scope_stack, }; pub use state::NemoRelayContextState; -pub use subscriber_dispatcher::flush_subscribers; +pub use subscriber_dispatcher::{flush_subscribers, flush_subscribers_from_binding}; diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 1d46f3fe6..8f9f1b8e3 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -52,11 +52,29 @@ mod native { OnceLock::new(); static DISPATCHER_FAILURE_LOGGED: AtomicBool = AtomicBool::new(false); static SANITIZER_RUNTIME_FAILURE_LOGGED: AtomicBool = AtomicBool::new(false); + static DISPATCH_IN_PROGRESS: AtomicBool = AtomicBool::new(false); thread_local! { static IN_DISPATCHER: Cell = const { Cell::new(false) }; } + struct DispatchGuard; + + impl DispatchGuard { + fn enter() -> Self { + debug_assert!(!DISPATCH_IN_PROGRESS.swap(true, Ordering::AcqRel)); + IN_DISPATCHER.with(|flag| flag.set(true)); + Self + } + } + + impl Drop for DispatchGuard { + fn drop(&mut self) { + IN_DISPATCHER.with(|flag| flag.set(false)); + DISPATCH_IN_PROGRESS.store(false, Ordering::Release); + } + } + fn sanitizer_runtime() -> std::result::Result<&'static tokio::runtime::Runtime, String> { SANITIZER_RUNTIME .get_or_init(|| { @@ -184,6 +202,13 @@ mod native { Ok(()) } + pub(super) fn flush_subscribers_from_binding() -> Result<()> { + if DISPATCH_IN_PROGRESS.load(Ordering::Acquire) { + return Ok(()); + } + flush_subscribers() + } + fn dispatcher_sender() -> std::result::Result, String> { DISPATCHER.get_or_init(start_dispatcher).clone() } @@ -288,9 +313,8 @@ mod native { ) { let previous_scope_stack = capture_thread_scope_stack(); set_thread_scope_stack(scope_stack); - IN_DISPATCHER.with(|flag| flag.set(true)); + let _dispatch_guard = DispatchGuard::enter(); let Some(event) = sanitize_event_snapshot(*event, transform, sanitizers) else { - IN_DISPATCHER.with(|flag| flag.set(false)); restore_thread_scope_stack(previous_scope_stack); return; }; @@ -303,7 +327,6 @@ mod native { ); } } - IN_DISPATCHER.with(|flag| flag.set(false)); restore_thread_scope_stack(previous_scope_stack); } @@ -417,3 +440,13 @@ pub(crate) fn register_async_publication() -> Option pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() } + +/// Wait for queued subscriber callbacks without creating a binding callback cycle. +/// +/// Language callbacks may resume on a different thread from the dispatcher. In that case the +/// thread-local reentrancy guard is insufficient, so bindings return early whenever an event is +/// actively being dispatched. +#[doc(hidden)] +pub fn flush_subscribers_from_binding() -> Result<()> { + native::flush_subscribers_from_binding() +} diff --git a/crates/core/src/api/subscriber.rs b/crates/core/src/api/subscriber.rs index 02ce9a322..a8c75c922 100644 --- a/crates/core/src/api/subscriber.rs +++ b/crates/core/src/api/subscriber.rs @@ -85,6 +85,13 @@ pub fn flush_subscribers() -> Result<()> { flush_runtime_subscribers() } +/// Binding-specific subscriber barrier that avoids cycles across asynchronous callback handoffs. +#[doc(hidden)] +pub fn flush_subscribers_from_binding() -> Result<()> { + ensure_runtime_owner()?; + crate::api::runtime::flush_subscribers_from_binding() +} + /// Register a scope-local lifecycle event subscriber. /// /// The subscriber remains active only while the target scope is still present diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 70c480b24..7e4ea851f 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -1258,9 +1258,9 @@ NemoRelayStatus nemo_relay_deregister_subscriber(const char *name); /** * Wait for subscriber callbacks queued before this call to finish. * - * Call this function outside native subscriber callbacks. A re-entrant call returns without - * waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can - * still run. + * If publication is currently executing, this function returns without waiting. This prevents + * subscriber and asynchronous event-sanitizer callbacks from creating a cycle with the serial + * dispatcher. */ NemoRelayStatus nemo_relay_flush_subscribers(void); diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index 38883ecee..486da1221 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -391,13 +391,13 @@ pub unsafe extern "C" fn nemo_relay_deregister_subscriber(name: *const c_char) - /// Wait for subscriber callbacks queued before this call to finish. /// -/// Call this function outside native subscriber callbacks. A re-entrant call returns without -/// waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can -/// still run. +/// If publication is currently executing, this function returns without waiting. This prevents +/// subscriber and asynchronous event-sanitizer callbacks from creating a cycle with the serial +/// dispatcher. #[unsafe(no_mangle)] pub extern "C" fn nemo_relay_flush_subscribers() -> NemoRelayStatus { clear_last_error(); - match core_subscriber_api::flush_subscribers() { + match core_subscriber_api::flush_subscribers_from_binding() { Ok(()) => NemoRelayStatus::Ok, Err(e) => status_from_error(&e), } diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index 2efa65892..daed96761 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -3172,9 +3172,9 @@ pub fn deregister_subscriber(name: String) -> Result { /// Return a Promise that resolves when native subscriber callbacks queued /// before this call finish. /// -/// Call this function outside native subscriber callbacks. A re-entrant call returns without -/// waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can -/// still run. +/// If publication is currently executing, this Promise resolves without waiting. This prevents +/// subscriber and asynchronous event-sanitizer callbacks from creating a cycle with the serial +/// dispatcher. /// /// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this /// Promise does not block the Node event loop while Promise-returning event sanitizers settle. @@ -3183,7 +3183,7 @@ pub fn deregister_subscriber(name: String) -> Result { /// Callers should handle errors when awaiting it. #[napi] pub async fn flush_subscribers() -> Result<()> { - tokio::task::spawn_blocking(core_subscriber_api::flush_subscribers) + tokio::task::spawn_blocking(core_subscriber_api::flush_subscribers_from_binding) .await .map_err(|error| to_napi_err(FlowError::Internal(error.to_string())))? .map_err(to_napi_err) diff --git a/crates/node/tests/event_sanitizers_tests.mjs b/crates/node/tests/event_sanitizers_tests.mjs index c5d6d75ea..e437d9574 100644 --- a/crates/node/tests/event_sanitizers_tests.mjs +++ b/crates/node/tests/event_sanitizers_tests.mjs @@ -129,6 +129,25 @@ describe('event sanitizer registries', () => { assert.deepEqual(events.at(-1).data, { sanitized: true }); }); + it('does not deadlock when an async sanitizer flushes subscribers', async () => { + const events = capture('node-event-sanitize-reentrant-flush-sub'); + let flushReturned = false; + lib.registerMarkSanitizeGuardrail('node-event-reentrant-flush', 0, async (_event, fields) => { + await lib.flushSubscribers(); + flushReturned = true; + return fields; + }); + try { + lib.event('reentrant-flush-checkpoint', null, { raw: true }); + await lib.flushSubscribers(); + await waitFor(events, 1); + } finally { + lib.deregisterMarkSanitizeGuardrail('node-event-reentrant-flush'); + lib.deregisterSubscriber('node-event-sanitize-reentrant-flush-sub'); + } + assert.equal(flushReturned, true); + }); + it('fails open and records invalid sanitizer results', async () => { const events = capture('node-event-sanitize-invalid-sub'); const invalidResults = { diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index eccca24aa..a56b2fdc6 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -35,10 +35,28 @@ use serde_json::{Map, json}; #[test] fn async_abi_discriminants_reject_unknown_values() { - assert_eq!( - NemoRelayNativeAsyncMiddlewareKind::try_from(13), - Ok(NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeEnd) - ); + use NemoRelayNativeAsyncMiddlewareKind as Kind; + + let middleware_kinds = [ + Kind::ToolSanitizeRequest, + Kind::ToolSanitizeResponse, + Kind::ToolConditionalExecution, + Kind::ToolRequestIntercept, + Kind::ToolExecutionIntercept, + Kind::LlmSanitizeRequest, + Kind::LlmSanitizeResponse, + Kind::LlmConditionalExecution, + Kind::LlmRequestIntercept, + Kind::LlmExecutionIntercept, + Kind::LlmStreamExecutionIntercept, + Kind::MarkSanitize, + Kind::ScopeSanitizeStart, + Kind::ScopeSanitizeEnd, + ]; + for (discriminant, kind) in middleware_kinds.into_iter().enumerate() { + assert_eq!(kind as u32, discriminant as u32); + assert_eq!(Kind::try_from(discriminant as u32), Ok(kind)); + } assert!(NemoRelayNativeAsyncMiddlewareKind::try_from(14).is_err()); assert_eq!( NemoRelayNativeAsyncCallbackState::try_from(1), diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index 534f68c99..678ed0e41 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1524,12 +1524,12 @@ fn deregister_subscriber(name: &str) -> PyResult { /// Wait for subscriber callbacks queued before this call to finish. /// -/// Call this function outside native subscriber callbacks. A re-entrant call returns without -/// waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can -/// still run. +/// If publication is currently executing, this function returns without waiting. This prevents +/// subscriber and asynchronous event-sanitizer callbacks from creating a cycle with the serial +/// dispatcher. #[pyfunction] fn flush_subscribers(py: Python<'_>) -> PyResult<()> { - py.detach(core_subscriber_api::flush_subscribers) + py.detach(core_subscriber_api::flush_subscribers_from_binding) .map_err(to_py_err) } From 6e194ff0576695daf6d5a488136a2542cde700c0 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 19:46:04 -0400 Subject: [PATCH 25/52] fix: scope reentrant flush guards to callbacks Signed-off-by: Will Killian --- crates/core/src/api/runtime.rs | 2 +- .../src/api/runtime/subscriber_dispatcher.rs | 21 ----------------- crates/core/src/api/subscriber.rs | 7 ------ crates/ffi/nemo_relay.h | 6 ++--- crates/ffi/src/api/llm_registry.rs | 8 +++---- crates/node/src/api/mod.rs | 10 ++++---- crates/node/src/callable.rs | 23 +++++++++++++++++++ crates/python/src/py_api/mod.rs | 10 ++++---- crates/python/src/py_callable.rs | 23 +++++++++++++++++++ 9 files changed, 66 insertions(+), 44 deletions(-) diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index b2878cead..77670c804 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -27,4 +27,4 @@ pub use scope_stack::{ task_scope_top, with_active_event_uuid, with_scope_stack, }; pub use state::NemoRelayContextState; -pub use subscriber_dispatcher::{flush_subscribers, flush_subscribers_from_binding}; +pub use subscriber_dispatcher::flush_subscribers; diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 8f9f1b8e3..47dace406 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -52,8 +52,6 @@ mod native { OnceLock::new(); static DISPATCHER_FAILURE_LOGGED: AtomicBool = AtomicBool::new(false); static SANITIZER_RUNTIME_FAILURE_LOGGED: AtomicBool = AtomicBool::new(false); - static DISPATCH_IN_PROGRESS: AtomicBool = AtomicBool::new(false); - thread_local! { static IN_DISPATCHER: Cell = const { Cell::new(false) }; } @@ -62,7 +60,6 @@ mod native { impl DispatchGuard { fn enter() -> Self { - debug_assert!(!DISPATCH_IN_PROGRESS.swap(true, Ordering::AcqRel)); IN_DISPATCHER.with(|flag| flag.set(true)); Self } @@ -71,7 +68,6 @@ mod native { impl Drop for DispatchGuard { fn drop(&mut self) { IN_DISPATCHER.with(|flag| flag.set(false)); - DISPATCH_IN_PROGRESS.store(false, Ordering::Release); } } @@ -202,13 +198,6 @@ mod native { Ok(()) } - pub(super) fn flush_subscribers_from_binding() -> Result<()> { - if DISPATCH_IN_PROGRESS.load(Ordering::Acquire) { - return Ok(()); - } - flush_subscribers() - } - fn dispatcher_sender() -> std::result::Result, String> { DISPATCHER.get_or_init(start_dispatcher).clone() } @@ -440,13 +429,3 @@ pub(crate) fn register_async_publication() -> Option pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() } - -/// Wait for queued subscriber callbacks without creating a binding callback cycle. -/// -/// Language callbacks may resume on a different thread from the dispatcher. In that case the -/// thread-local reentrancy guard is insufficient, so bindings return early whenever an event is -/// actively being dispatched. -#[doc(hidden)] -pub fn flush_subscribers_from_binding() -> Result<()> { - native::flush_subscribers_from_binding() -} diff --git a/crates/core/src/api/subscriber.rs b/crates/core/src/api/subscriber.rs index a8c75c922..02ce9a322 100644 --- a/crates/core/src/api/subscriber.rs +++ b/crates/core/src/api/subscriber.rs @@ -85,13 +85,6 @@ pub fn flush_subscribers() -> Result<()> { flush_runtime_subscribers() } -/// Binding-specific subscriber barrier that avoids cycles across asynchronous callback handoffs. -#[doc(hidden)] -pub fn flush_subscribers_from_binding() -> Result<()> { - ensure_runtime_owner()?; - crate::api::runtime::flush_subscribers_from_binding() -} - /// Register a scope-local lifecycle event subscriber. /// /// The subscriber remains active only while the target scope is still present diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 7e4ea851f..70c480b24 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -1258,9 +1258,9 @@ NemoRelayStatus nemo_relay_deregister_subscriber(const char *name); /** * Wait for subscriber callbacks queued before this call to finish. * - * If publication is currently executing, this function returns without waiting. This prevents - * subscriber and asynchronous event-sanitizer callbacks from creating a cycle with the serial - * dispatcher. + * Call this function outside native subscriber callbacks. A re-entrant call returns without + * waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can + * still run. */ NemoRelayStatus nemo_relay_flush_subscribers(void); diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index 486da1221..38883ecee 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -391,13 +391,13 @@ pub unsafe extern "C" fn nemo_relay_deregister_subscriber(name: *const c_char) - /// Wait for subscriber callbacks queued before this call to finish. /// -/// If publication is currently executing, this function returns without waiting. This prevents -/// subscriber and asynchronous event-sanitizer callbacks from creating a cycle with the serial -/// dispatcher. +/// Call this function outside native subscriber callbacks. A re-entrant call returns without +/// waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can +/// still run. #[unsafe(no_mangle)] pub extern "C" fn nemo_relay_flush_subscribers() -> NemoRelayStatus { clear_last_error(); - match core_subscriber_api::flush_subscribers_from_binding() { + match core_subscriber_api::flush_subscribers() { Ok(()) => NemoRelayStatus::Ok, Err(e) => status_from_error(&e), } diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index daed96761..6e22c3dc9 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -3172,9 +3172,8 @@ pub fn deregister_subscriber(name: String) -> Result { /// Return a Promise that resolves when native subscriber callbacks queued /// before this call finish. /// -/// If publication is currently executing, this Promise resolves without waiting. This prevents -/// subscriber and asynchronous event-sanitizer callbacks from creating a cycle with the serial -/// dispatcher. +/// When called from an event-sanitizer callback, this Promise resolves without waiting to prevent +/// a cycle with the serial dispatcher. /// /// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this /// Promise does not block the Node event loop while Promise-returning event sanitizers settle. @@ -3183,7 +3182,10 @@ pub fn deregister_subscriber(name: String) -> Result { /// Callers should handle errors when awaiting it. #[napi] pub async fn flush_subscribers() -> Result<()> { - tokio::task::spawn_blocking(core_subscriber_api::flush_subscribers_from_binding) + if crate::callable::event_sanitizer_callback_active() { + return Ok(()); + } + tokio::task::spawn_blocking(core_subscriber_api::flush_subscribers) .await .map_err(|error| to_napi_err(FlowError::Internal(error.to_string())))? .map_err(to_napi_err) diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index 4a7eb2a58..b541da082 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -12,6 +12,7 @@ use std::future::Future; use std::pin::Pin; use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; use napi::bindgen_prelude::ToNapiValue; use napi::threadsafe_function::{ErrorStrategy, ThreadsafeFunction, ThreadsafeFunctionCallMode}; @@ -44,6 +45,27 @@ use crate::convert::{callback_json, record_callback_error, to_napi_err}; use crate::promise_call::{JsonNextFn, JsonStreamNextFn, PromiseAwareFn}; use crate::types::{EventSanitizeFields, JsEvent, event_sanitize_fields_from_json}; +static ACTIVE_EVENT_SANITIZER_CALLBACKS: AtomicUsize = AtomicUsize::new(0); + +struct ActiveEventSanitizerCallback; + +impl ActiveEventSanitizerCallback { + fn enter() -> Self { + ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_add(1, Ordering::AcqRel); + Self + } +} + +impl Drop for ActiveEventSanitizerCallback { + fn drop(&mut self) { + ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_sub(1, Ordering::AcqRel); + } +} + +pub(crate) fn event_sanitizer_callback_active() -> bool { + ACTIVE_EVENT_SANITIZER_CALLBACKS.load(Ordering::Acquire) != 0 +} + /// Structured codec identity delivered to JavaScript LLM sanitizers. #[napi(object)] #[derive(Clone)] @@ -476,6 +498,7 @@ pub fn wrap_js_event_sanitize_promise_fn(func: Arc) -> EventSani Arc::new(move |event: Arc, fields: CoreEventSanitizeFields| { let func = func.clone(); Box::pin(async move { + let _active_callback = ActiveEventSanitizerCallback::enter(); let event_json = JsEvent::try_from_event(&event) .map(JsEvent::into_json) .map_err(|error| { diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index 678ed0e41..70d4c0742 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1524,12 +1524,14 @@ fn deregister_subscriber(name: &str) -> PyResult { /// Wait for subscriber callbacks queued before this call to finish. /// -/// If publication is currently executing, this function returns without waiting. This prevents -/// subscriber and asynchronous event-sanitizer callbacks from creating a cycle with the serial -/// dispatcher. +/// A call from an asynchronous event-sanitizer callback returns without waiting to prevent a +/// cycle with the serial dispatcher. #[pyfunction] fn flush_subscribers(py: Python<'_>) -> PyResult<()> { - py.detach(core_subscriber_api::flush_subscribers_from_binding) + if crate::py_callable::event_sanitizer_callback_active() { + return Ok(()); + } + py.detach(core_subscriber_api::flush_subscribers) .map_err(to_py_err) } diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index bea95c16d..9089da478 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -23,6 +23,7 @@ use std::future::Future; use std::pin::Pin; use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; use std::task::{Context, Poll}; use nemo_relay::api::runtime::{ @@ -43,6 +44,27 @@ use nemo_relay::api::event::{Event, EventSanitizeFields}; use nemo_relay::api::llm::LlmRequest; use nemo_relay::api::tool::ToolExecutionInterceptOutcome; use nemo_relay::codec::request::AnnotatedLlmRequest as AnnotatedLLMRequest; + +static ACTIVE_EVENT_SANITIZER_CALLBACKS: AtomicUsize = AtomicUsize::new(0); + +struct ActiveEventSanitizerCallback; + +impl ActiveEventSanitizerCallback { + fn enter() -> Self { + ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_add(1, Ordering::AcqRel); + Self + } +} + +impl Drop for ActiveEventSanitizerCallback { + fn drop(&mut self) { + ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_sub(1, Ordering::AcqRel); + } +} + +pub(crate) fn event_sanitizer_callback_active() -> bool { + ACTIVE_EVENT_SANITIZER_CALLBACKS.load(Ordering::Acquire) != 0 +} use nemo_relay::codec::response::AnnotatedLlmResponse as AnnotatedLLMResponse; use nemo_relay::codec::traits::{LlmCodec, LlmResponseCodec}; @@ -1142,6 +1164,7 @@ pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let py_fn = py_fn.clone(); let task_locals = task_locals.clone(); Box::pin(async move { + let _active_callback = ActiveEventSanitizerCallback::enter(); let result = Python::attach( |py| -> FlowResult, PyValueFuture>> { let py_event = match event.as_ref() { From 609afe81de8f7a762c4e21b99c1045fbee3192b0 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 20:24:34 -0400 Subject: [PATCH 26/52] fix(python): scope sanitizer flush reentrancy Signed-off-by: Will Killian --- crates/python/src/py_api/mod.rs | 7 +--- crates/python/src/py_callable.rs | 33 ++++------------ python/nemo_relay/_event_sanitizer_context.py | 38 +++++++++++++++++++ python/nemo_relay/subscribers.py | 9 +++-- python/tests/test_event_sanitizers.py | 30 +++++++++++++++ 5 files changed, 83 insertions(+), 34 deletions(-) create mode 100644 python/nemo_relay/_event_sanitizer_context.py diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index 70d4c0742..7bb90a042 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1524,13 +1524,10 @@ fn deregister_subscriber(name: &str) -> PyResult { /// Wait for subscriber callbacks queued before this call to finish. /// -/// A call from an asynchronous event-sanitizer callback returns without waiting to prevent a -/// cycle with the serial dispatcher. +/// Public Python wrappers prevent re-entrant event-sanitizer callbacks from waiting on the serial +/// dispatcher. #[pyfunction] fn flush_subscribers(py: Python<'_>) -> PyResult<()> { - if crate::py_callable::event_sanitizer_callback_active() { - return Ok(()); - } py.detach(core_subscriber_api::flush_subscribers) .map_err(to_py_err) } diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 9089da478..cc89236d4 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -23,7 +23,6 @@ use std::future::Future; use std::pin::Pin; use std::sync::Arc; -use std::sync::atomic::{AtomicUsize, Ordering}; use std::task::{Context, Poll}; use nemo_relay::api::runtime::{ @@ -44,27 +43,6 @@ use nemo_relay::api::event::{Event, EventSanitizeFields}; use nemo_relay::api::llm::LlmRequest; use nemo_relay::api::tool::ToolExecutionInterceptOutcome; use nemo_relay::codec::request::AnnotatedLlmRequest as AnnotatedLLMRequest; - -static ACTIVE_EVENT_SANITIZER_CALLBACKS: AtomicUsize = AtomicUsize::new(0); - -struct ActiveEventSanitizerCallback; - -impl ActiveEventSanitizerCallback { - fn enter() -> Self { - ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_add(1, Ordering::AcqRel); - Self - } -} - -impl Drop for ActiveEventSanitizerCallback { - fn drop(&mut self) { - ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_sub(1, Ordering::AcqRel); - } -} - -pub(crate) fn event_sanitizer_callback_active() -> bool { - ACTIVE_EVENT_SANITIZER_CALLBACKS.load(Ordering::Acquire) != 0 -} use nemo_relay::codec::response::AnnotatedLlmResponse as AnnotatedLLMResponse; use nemo_relay::codec::traits::{LlmCodec, LlmResponseCodec}; @@ -1164,7 +1142,6 @@ pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let py_fn = py_fn.clone(); let task_locals = task_locals.clone(); Box::pin(async move { - let _active_callback = ActiveEventSanitizerCallback::enter(); let result = Python::attach( |py| -> FlowResult, PyValueFuture>> { let py_event = match event.as_ref() { @@ -1201,10 +1178,14 @@ pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { return Err(FlowError::Internal(error.to_string())); } }; - let result = py_fn - .call1(py, (py_event, py_fields)) + let invoke = py + .import("nemo_relay._event_sanitizer_context") + .and_then(|module| module.getattr("invoke")) + .map_err(|error| FlowError::Internal(error.to_string()))?; + let result = invoke + .call1((py_fn.bind(py), py_event, py_fields)) .map_err(|error| FlowError::Internal(error.to_string()))?; - split_py_object_or_future_with_locals(py, result, task_locals.as_ref()) + split_py_object_or_future_with_locals(py, result.unbind(), task_locals.as_ref()) }, ); let result = resolve_py_object_or_future(result).await?; diff --git a/python/nemo_relay/_event_sanitizer_context.py b/python/nemo_relay/_event_sanitizer_context.py new file mode 100644 index 000000000..f2cc89a09 --- /dev/null +++ b/python/nemo_relay/_event_sanitizer_context.py @@ -0,0 +1,38 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Track re-entrant subscriber flushes from Python event sanitizers.""" + +from __future__ import annotations + +import inspect +from collections.abc import Awaitable, Callable +from contextvars import ContextVar +from typing import Any + +_ACTIVE: ContextVar[bool] = ContextVar("nemo_relay_event_sanitizer_active", default=False) + + +def callback_active() -> bool: + """Return whether the current Python context is running an event sanitizer.""" + return _ACTIVE.get() + + +async def _await_result(result: Awaitable[Any]) -> Any: + token = _ACTIVE.set(True) + try: + return await result + finally: + _ACTIVE.reset(token) + + +def invoke(callback: Callable[..., Any], *args: Any) -> Any: + """Invoke a sanitizer while marking its sync and async execution contexts.""" + token = _ACTIVE.set(True) + try: + result = callback(*args) + finally: + _ACTIVE.reset(token) + if inspect.isawaitable(result): + return _await_result(result) + return result diff --git a/python/nemo_relay/subscribers.py b/python/nemo_relay/subscribers.py index b54b42e2e..d5a60cc7c 100644 --- a/python/nemo_relay/subscribers.py +++ b/python/nemo_relay/subscribers.py @@ -25,6 +25,7 @@ def log_event(event): from collections.abc import Callable from typing import TYPE_CHECKING +from nemo_relay._event_sanitizer_context import callback_active as _event_sanitizer_callback_active from nemo_relay._native import ( deregister_subscriber as _native_deregister, ) @@ -94,10 +95,12 @@ def flush() -> None: waiting for observer work. Use this barrier in tests and shutdown paths when captured subscriber output must be complete before continuing. - Call this function outside subscriber callbacks. A re-entrant call returns - without waiting to avoid blocking the dispatcher, so callbacks later in the - same dispatch snapshot can still run. + Call this function outside subscriber and event-sanitizer callbacks. A + re-entrant call returns without waiting to avoid blocking the dispatcher, + so callbacks later in the same dispatch snapshot can still run. """ + if _event_sanitizer_callback_active(): + return None return _native_flush() diff --git a/python/tests/test_event_sanitizers.py b/python/tests/test_event_sanitizers.py index 4a91bf658..c0dc3c78e 100644 --- a/python/tests/test_event_sanitizers.py +++ b/python/tests/test_event_sanitizers.py @@ -97,6 +97,36 @@ async def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> Eve assert events[-1].data == {"async": True} +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_event_sanitizer_flush_is_reentrant(capture_events, asynchronous): + _capture_name, events = capture_events + flush_returned = False + + def sanitize_sync(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + nonlocal flush_returned + subscribers.flush() + flush_returned = True + return fields + + async def sanitize_async(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + await asyncio.sleep(0) + return sanitize_sync(_event, fields) + + guardrails.register_mark_sanitize( + "python-reentrant-mark", + 0, + sanitize_async if asynchronous else sanitize_sync, + ) + try: + scope.event("reentrant-checkpoint", data={"raw": True}) + await asyncio.to_thread(subscribers.flush) + finally: + guardrails.deregister_mark_sanitize("python-reentrant-mark") + + assert flush_returned is True + assert events[-1].data == {"raw": True} + + def test_scope_start_and_end_sanitizers_cover_category_profile(capture_events): _capture_name, events = capture_events From 82ee862f670b78f152f038149bc64e15f7df6390 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 21:07:00 -0400 Subject: [PATCH 27/52] fix: address async middleware review regressions Signed-off-by: Will Killian --- crates/core/src/api/llm.rs | 8 +++- crates/core/src/api/tool.rs | 8 +++- .../tests/integration/api_surface_tests.rs | 22 +++++----- .../tests/integration/native_plugin_tests.rs | 2 + crates/core/tests/unit/shared_tests.rs | 26 ++++++++--- crates/plugin/README.md | 5 ++- crates/python/src/py_callable.rs | 4 +- .../python/tests/coverage/coverage_tests.rs | 44 +++++++++++-------- .../coverage/py_callable_coverage_tests.rs | 31 ++++++++++++- 9 files changed, 106 insertions(+), 44 deletions(-) diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index 0301aeeda..b6a7ffa52 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -719,7 +719,9 @@ pub fn llm_call(params: LlmCallParams<'_>) -> Result { let handle = create_llm_handle(handle_params)?; let scope_stack = handle.captured_scope_stack().clone(); let (entries, subscribers, agent_is_fresh) = { - let mut scope_guard = scope_stack.write().expect("scope stack lock poisoned"); + let mut scope_guard = scope_stack + .write() + .map_err(|error| FlowError::Internal(error.to_string()))?; let scope_locals = scope_guard.collect_scope_local_registries(|registries| { ®istries.llm_sanitize_request_guardrails }); @@ -899,7 +901,9 @@ pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> { ensure_runtime_owner()?; let scope_stack = params.handle.captured_scope_stack().clone(); let (entries, subscribers) = { - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_guard = scope_stack + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; let scope_locals = scope_guard.collect_scope_local_registries(|registries| { ®istries.llm_sanitize_response_guardrails }); diff --git a/crates/core/src/api/tool.rs b/crates/core/src/api/tool.rs index 671471748..30a9249e1 100644 --- a/crates/core/src/api/tool.rs +++ b/crates/core/src/api/tool.rs @@ -231,7 +231,9 @@ pub fn tool_call(params: ToolCallParams<'_>) -> Result { ensure_runtime_owner()?; let scope_stack = current_scope_stack(); let (entries, subscribers) = { - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_guard = scope_stack + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; let scope_locals = scope_guard.collect_scope_local_registries(|registries| { ®istries.tool_sanitize_request_guardrails }); @@ -418,7 +420,9 @@ pub fn tool_call_end(params: ToolCallEndParams<'_>) -> Result<()> { ensure_runtime_owner()?; let scope_stack = current_scope_stack(); let (entries, subscribers) = { - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_guard = scope_stack + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; let scope_locals = scope_guard.collect_scope_local_registries(|registries| { ®istries.tool_sanitize_response_guardrails }); diff --git a/crates/core/tests/integration/api_surface_tests.rs b/crates/core/tests/integration/api_surface_tests.rs index 61f56251f..5b14fe82f 100644 --- a/crates/core/tests/integration/api_surface_tests.rs +++ b/crates/core/tests/integration/api_surface_tests.rs @@ -1921,21 +1921,19 @@ async fn test_llm_stream_api_covers_success_rejection_and_execution_error_paths( stream.close().await.unwrap(); let success_events = captured_events_snapshot(&events); - let success_start = success_events + let scope_events = success_events .iter() - .find(|event| { - event.kind() == "scope" && event.scope_category() == Some(ScopeCategory::Start) - }) - .expect("stream start event"); - let success_end = success_events - .iter() - .rev() - .find(|event| event.kind() == "scope" && event.scope_category() == Some(ScopeCategory::End)) - .expect("stream end event"); - assert_eq!(success_start.kind(), "scope"); + .filter(|event| event.kind() == "scope") + .collect::>(); + assert_eq!( + scope_events.len(), + 2, + "expected exactly one stream scope pair" + ); + let success_start = scope_events[0]; + let success_end = scope_events[1]; assert_eq!(success_start.scope_category(), Some(ScopeCategory::Start)); assert_eq!(success_start.category().unwrap().as_str(), "llm"); - assert_eq!(success_end.kind(), "scope"); assert_eq!(success_end.scope_category(), Some(ScopeCategory::End)); assert_eq!(success_end.category().unwrap().as_str(), "llm"); assert_eq!( diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index 91a07a4cf..8e9782def 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -680,6 +680,8 @@ async fn native_v3_async_registration_supports_all_middleware_kinds() { .get:: bool>(b"nemo_relay_fixture_async_pending_entered\0") .expect("v3 async native fixture should export its pending-entry signal") }; + // This pointer remains valid only while `activation` keeps the fixture + // library loaded; never call it after clearing the plugin configuration. assert!(!unsafe { pending_entered() }); drop(fixture_library); let mut cleanup = NativePluginTestCleanup::new(); diff --git a/crates/core/tests/unit/shared_tests.rs b/crates/core/tests/unit/shared_tests.rs index 1de4a9ed7..5375f4334 100644 --- a/crates/core/tests/unit/shared_tests.rs +++ b/crates/core/tests/unit/shared_tests.rs @@ -4,7 +4,7 @@ //! Unit tests for shared in the NeMo Relay core crate. use super::*; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; use serde_json::{Map, json}; @@ -174,13 +174,16 @@ async fn test_run_request_intercepts_with_codec_none_and_codec_paths() { let _guard = lock_runtime_owner(); reset_global(); + let observed_without_codec = Arc::new(Mutex::new(None)); + let callback_observed_without_codec = Arc::clone(&observed_without_codec); register_llm_request_intercept( "shared-none", 1, false, - Arc::new(|_name, mut request, annotated| { + Arc::new(move |_name, mut request, annotated| { + let callback_observed_without_codec = Arc::clone(&callback_observed_without_codec); Box::pin(async move { - assert!(annotated.is_none()); + *callback_observed_without_codec.lock().unwrap() = Some(annotated.is_none()); request.headers.insert("x-no-codec".into(), json!(true)); let mut annotated = SharedTestCodec.decode(&request)?; annotated.model = Some("interceptor-model".into()); @@ -201,6 +204,7 @@ async fn test_run_request_intercepts_with_codec_none_and_codec_paths() { ) .await .unwrap(); + assert_eq!(*observed_without_codec.lock().unwrap(), Some(true)); assert_eq!( request_without_codec.headers.get("x-no-codec"), Some(&json!(true)) @@ -214,13 +218,21 @@ async fn test_run_request_intercepts_with_codec_none_and_codec_paths() { assert!(pending_marks_without_codec.is_empty()); deregister_llm_request_intercept("shared-none").unwrap(); + let observed_with_codec = Arc::new(Mutex::new(None)); + let callback_observed_with_codec = Arc::clone(&observed_with_codec); register_llm_request_intercept( "shared-codec", 1, false, - Arc::new(|_name, mut request, annotated| { + Arc::new(move |_name, mut request, annotated| { + let callback_observed_with_codec = Arc::clone(&callback_observed_with_codec); Box::pin(async move { - let mut annotated = annotated.expect("codec should provide annotated request"); + *callback_observed_with_codec.lock().unwrap() = + Some(annotated.as_ref().and_then(|value| value.model.clone())); + let mut annotated = match annotated { + Some(value) => value, + None => SharedTestCodec.decode(&request)?, + }; annotated.model = Some("intercepted-model".into()); request.headers.insert("x-codec".into(), json!(true)); Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) @@ -242,6 +254,10 @@ async fn test_run_request_intercepts_with_codec_none_and_codec_paths() { .await .unwrap(); + assert_eq!( + *observed_with_codec.lock().unwrap(), + Some(Some("decoded-model".into())) + ); assert_eq!( request_with_codec.headers.get("x-codec"), Some(&json!(true)) diff --git a/crates/plugin/README.md b/crates/plugin/README.md index 0b68a888e..ecafb951f 100644 --- a/crates/plugin/README.md +++ b/crates/plugin/README.md @@ -43,8 +43,9 @@ the dynamic-library boundary on the stable C-compatible ABI. subscribers. - **`PluginRuntime`**: Typed helpers for Relay-owned scopes and marks. - **Stable native ABI v3**: C-compatible host and plugin tables behind the - safe Rust authoring interface, with a v2-compatible prefix for existing - plugins. + safe Rust authoring interface. The v3 tables preserve a v2-compatible field + prefix, but native plugins must still be rebuilt for v3 as described in the + [0.7 migration guide](../../docs/reference/migration-guides.mdx#upgrade-to-nemo-relay-07). - **Raw async middleware**: Completion-based raw registrations for plugins that need asynchronous guardrails, intercepts, or event sanitizers. Typed Rust callbacks remain synchronous convenience APIs. diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index cc89236d4..2c51cfef8 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -1072,7 +1072,7 @@ fn wrap_py_llm_sanitize_response_callback(py_fn: Py) -> LlmSanitizeRespon let task_locals = capture_python_task_locals(); Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { let py_fn = py_fn.clone(); - let task_locals = task_locals.clone(); + let task_locals = capture_python_task_locals().or_else(|| task_locals.clone()); Box::pin(async move { let result = resolve_py_object_or_future(Python::attach(|py| { let py_context = PyLlmSanitizeResponseContext { inner: context }; @@ -1140,7 +1140,7 @@ pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let task_locals = capture_python_task_locals(); Arc::new(move |event: Arc, fields: EventSanitizeFields| { let py_fn = py_fn.clone(); - let task_locals = task_locals.clone(); + let task_locals = capture_python_task_locals().or_else(|| task_locals.clone()); Box::pin(async move { let result = Python::attach( |py| -> FlowResult, PyValueFuture>> { diff --git a/crates/python/tests/coverage/coverage_tests.rs b/crates/python/tests/coverage/coverage_tests.rs index 2fa805abf..7eab65355 100644 --- a/crates/python/tests/coverage/coverage_tests.rs +++ b/crates/python/tests/coverage/coverage_tests.rs @@ -674,10 +674,12 @@ def event_fail(event): ); let tool_fail = wrap_py_tool_fn(module.getattr("tool_fail").unwrap().unbind()); + let error = runtime + .block_on(tool_fail("demo".to_string(), json!({"x": 1}))) + .unwrap_err(); assert!( - runtime - .block_on(tool_fail("demo".to_string(), json!({"x": 1}))) - .is_err() + error.to_string().contains("tool boom"), + "unexpected tool error: {error}" ); let tool_cond = @@ -694,13 +696,15 @@ def event_fail(event): let llm_sanitize = wrap_py_llm_sanitize_request_fn(module.getattr("llm_sanitize_bad").unwrap().unbind()) .unwrap(); + let error = runtime + .block_on(llm_sanitize( + request.clone(), + nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), + )) + .unwrap_err(); assert!( - runtime - .block_on(llm_sanitize( - request.clone(), - nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - )) - .is_err() + error.to_string().contains("unexpected type"), + "unexpected LLM request sanitizer error: {error}" ); let llm_cond = wrap_py_llm_conditional_fn(module.getattr("llm_cond_bad").unwrap().unbind()); @@ -730,22 +734,26 @@ def event_fail(event): let tool_req = wrap_py_tool_request_intercept_fn(module.getattr("tool_fail").unwrap().unbind()); + let error = runtime + .block_on(tool_req("demo".to_string(), json!({"x": 1}))) + .unwrap_err(); assert!( - runtime - .block_on(tool_req("demo".to_string(), json!({"x": 1}))) - .is_err() + error.to_string().contains("tool boom"), + "unexpected tool request intercept error: {error}" ); let llm_resp = wrap_py_llm_sanitize_response_fn(module.getattr("llm_resp_fail").unwrap().unbind()) .unwrap(); + let error = runtime + .block_on(llm_resp( + json!({"ok": true}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), + )) + .unwrap_err(); assert!( - runtime - .block_on(llm_resp( - json!({"ok": true}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - )) - .is_err() + error.to_string().contains("resp boom"), + "unexpected LLM response sanitizer error: {error}" ); let mut collector = diff --git a/crates/python/tests/coverage/py_callable_coverage_tests.rs b/crates/python/tests/coverage/py_callable_coverage_tests.rs index 53245133c..ccb6cef65 100644 --- a/crates/python/tests/coverage/py_callable_coverage_tests.rs +++ b/crates/python/tests/coverage/py_callable_coverage_tests.rs @@ -8,7 +8,7 @@ use super::*; use std::ffi::CString; use std::sync::Arc; -use pyo3::types::PyModule; +use pyo3::types::{PyDict, PyList, PyModule}; use serde_json::json; fn load_module<'py>(py: Python<'py>, code: &str) -> Bound<'py, PyModule> { @@ -18,6 +18,34 @@ fn load_module<'py>(py: Python<'py>, code: &str) -> Bound<'py, PyModule> { PyModule::from_code(py, &code, &file_name, &module_name).unwrap() } +fn install_event_sanitizer_context_module(py: Python<'_>) { + let code = CString::new(include_str!( + "../../../../python/nemo_relay/_event_sanitizer_context.py" + )) + .unwrap(); + let file_name = CString::new("_event_sanitizer_context.py").unwrap(); + let module_name = CString::new("nemo_relay._event_sanitizer_context").unwrap(); + let context = PyModule::from_code(py, &code, &file_name, &module_name).unwrap(); + let parent = PyModule::new(py, "nemo_relay").unwrap(); + parent + .setattr("__path__", PyList::empty(py)) + .expect("test package path should be writable"); + parent + .setattr("_event_sanitizer_context", &context) + .expect("test context module should be writable"); + let modules = py + .import("sys") + .unwrap() + .getattr("modules") + .unwrap() + .cast_into::() + .unwrap(); + modules.set_item("nemo_relay", parent).unwrap(); + modules + .set_item("nemo_relay._event_sanitizer_context", context) + .unwrap(); +} + fn make_request() -> LlmRequest { LlmRequest { headers: serde_json::Map::new(), @@ -673,6 +701,7 @@ fn event_sanitize_wrapper_covers_conversion_success_and_error_propagation() { let _python = crate::test_support::init_python_test(); Python::attach(|py| { + install_event_sanitizer_context_module(py); let module = load_module( py, r#" From e2ec4c72ec31b36fc61218f0e888bd02f3f09cbb Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 21:16:50 -0400 Subject: [PATCH 28/52] fix: address async binding review findings Signed-off-by: Will Killian --- .../adaptive/src/adaptive_hints_intercept.rs | 4 +- .../tests/fixtures/native_plugin/src/lib.rs | 232 +++++++++++++----- crates/ffi/src/callable.rs | 30 ++- crates/node/src/callable.rs | 1 + crates/python/src/py_callable.rs | 8 +- 5 files changed, 198 insertions(+), 77 deletions(-) diff --git a/crates/adaptive/src/adaptive_hints_intercept.rs b/crates/adaptive/src/adaptive_hints_intercept.rs index b4b649264..50437e4da 100644 --- a/crates/adaptive/src/adaptive_hints_intercept.rs +++ b/crates/adaptive/src/adaptive_hints_intercept.rs @@ -178,9 +178,9 @@ impl AdaptiveHintsIntercept { mut request: LlmRequest, mut annotated: Option| { let this = this.clone(); + let scope_path = extract_scope_path(); + let manual_ls = read_manual_latency_sensitivity(); Box::pin(async move { - let scope_path = extract_scope_path(); - let manual_ls = read_manual_latency_sensitivity(); let scope_depth = scope_path.len(); let call_index = this.call_counter.fetch_add(1, Ordering::Relaxed); diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index 00d417117..46c05c265 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -7,11 +7,11 @@ use std::sync::atomic::{AtomicBool, Ordering}; use nemo_relay_plugin::{ CategoryProfile, ConfigDiagnostic, DiagnosticLevel, Event, EventCategory, EventSanitizeFields, - Json, LlmJsonStream, LlmRequest, LlmRequestInterceptOutcome, NemoRelayNativeHostApiV1, - NemoRelayNativeHostApiV3, NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncMiddlewareCb, - NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativePluginContext, NemoRelayNativePluginV1, - NemoRelayNativeString, NemoRelayStatus, - NemoRelayNativeToolNextFn, NativePlugin, PendingMarkSpec, PluginContext, PluginRuntime, + Json, LlmJsonStream, LlmRequest, LlmRequestInterceptOutcome, NativePlugin, + NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncMiddlewareCb, + NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, + NemoRelayNativePluginContext, NemoRelayNativePluginV1, NemoRelayNativeString, + NemoRelayNativeToolNextFn, NemoRelayStatus, PendingMarkSpec, PluginContext, PluginRuntime, ScopeCategory, ScopeType, ToolExecutionInterceptOutcome, }; use serde_json::{Map, json}; @@ -65,11 +65,9 @@ impl NativePlugin for FixtureNativePlugin { 0, |_, fields| mark_event_fields(fields, "native_plugin_scope_start"), )?; - ctx.register_scope_sanitize_end_guardrail( - "fixture_scope_end_sanitize", - 0, - |_, fields| mark_event_fields(fields, "native_plugin_scope_end"), - )?; + ctx.register_scope_sanitize_end_guardrail("fixture_scope_end_sanitize", 0, |_, fields| { + mark_event_fields(fields, "native_plugin_scope_end") + })?; ctx.register_tool_sanitize_request_guardrail( "fixture_tool_sanitize_request", @@ -140,7 +138,12 @@ impl NativePlugin for FixtureNativePlugin { ctx.register_llm_sanitize_request_guardrail( "fixture_llm_sanitize_request", 0, - |request, _context| Some(mark_llm_request(request, "native_plugin_llm_sanitize_request")), + |request, _context| { + Some(mark_llm_request( + request, + "native_plugin_llm_sanitize_request", + )) + }, )?; ctx.register_llm_sanitize_response_guardrail( "fixture_llm_sanitize_response", @@ -175,13 +178,17 @@ impl NativePlugin for FixtureNativePlugin { )) }, )?; - ctx.register_llm_execution_intercept("fixture_llm_execution", 0, |_name, request, next| { - let response = next.call(mark_llm_request( - request, - "native_plugin_llm_execution_request", - ))?; - Ok(mark_json(response, "native_plugin_llm_execution")) - })?; + ctx.register_llm_execution_intercept( + "fixture_llm_execution", + 0, + |_name, request, next| { + let response = next.call(mark_llm_request( + request, + "native_plugin_llm_execution_request", + ))?; + Ok(mark_json(response, "native_plugin_llm_execution")) + }, + )?; ctx.register_llm_stream_execution_intercept( "fixture_llm_stream_execution", 0, @@ -616,24 +623,81 @@ unsafe extern "C" fn raw_register_async_tool_request( return NemoRelayStatus::NullPointer; } let host = unsafe { &*(user_data as *const NemoRelayNativeHostApiV3) }; - let registrations: [ - (NemoRelayNativeAsyncMiddlewareKind, &str, NemoRelayNativeAsyncMiddlewareCb); - 14 - ] = [ - (NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest, "fixture_async_tool_sanitize_request", raw_async_passthrough_callback), - (NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeResponse, "fixture_async_tool_sanitize_response", raw_async_passthrough_callback), - (NemoRelayNativeAsyncMiddlewareKind::ToolConditionalExecution, "fixture_async_tool_conditional", raw_async_allow_callback), - (NemoRelayNativeAsyncMiddlewareKind::ToolRequestIntercept, "fixture_async_request", raw_async_tool_request_callback), - (NemoRelayNativeAsyncMiddlewareKind::ToolExecutionIntercept, "fixture_async_execution", raw_async_tool_execution_callback), - (NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeRequest, "fixture_async_llm_sanitize_request", raw_async_passthrough_callback), - (NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeResponse, "fixture_async_llm_sanitize_response", raw_async_passthrough_callback), - (NemoRelayNativeAsyncMiddlewareKind::LlmConditionalExecution, "fixture_async_llm_conditional", raw_async_allow_callback), - (NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept, "fixture_async_llm_request", raw_async_passthrough_callback), - (NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept, "fixture_async_llm_execution", raw_async_tool_execution_callback), - (NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept, "fixture_async_llm_stream", raw_async_tool_execution_callback), - (NemoRelayNativeAsyncMiddlewareKind::MarkSanitize, "fixture_async_mark", raw_async_passthrough_callback), - (NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeStart, "fixture_async_scope_start", raw_async_passthrough_callback), - (NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeEnd, "fixture_async_scope_end", raw_async_passthrough_callback), + let registrations: [( + NemoRelayNativeAsyncMiddlewareKind, + &str, + NemoRelayNativeAsyncMiddlewareCb, + ); 14] = [ + ( + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest, + "fixture_async_tool_sanitize_request", + raw_async_passthrough_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeResponse, + "fixture_async_tool_sanitize_response", + raw_async_passthrough_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::ToolConditionalExecution, + "fixture_async_tool_conditional", + raw_async_allow_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::ToolRequestIntercept, + "fixture_async_request", + raw_async_tool_request_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::ToolExecutionIntercept, + "fixture_async_execution", + raw_async_tool_execution_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeRequest, + "fixture_async_llm_sanitize_request", + raw_async_passthrough_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeResponse, + "fixture_async_llm_sanitize_response", + raw_async_passthrough_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::LlmConditionalExecution, + "fixture_async_llm_conditional", + raw_async_allow_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept, + "fixture_async_llm_request", + raw_async_passthrough_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept, + "fixture_async_llm_execution", + raw_async_tool_execution_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept, + "fixture_async_llm_stream", + raw_async_tool_execution_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::MarkSanitize, + "fixture_async_mark", + raw_async_passthrough_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeStart, + "fixture_async_scope_start", + raw_async_passthrough_callback, + ), + ( + NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeEnd, + "fixture_async_scope_end", + raw_async_passthrough_callback, + ), ]; for (kind, registration_name, callback) in registrations { let name = unsafe { raw_host_string(&host.v1, registration_name) }; @@ -642,7 +706,14 @@ unsafe extern "C" fn raw_register_async_tool_request( } let status = unsafe { (host.plugin_context_register_async_middleware)( - ctx, kind as u32, name, 0, false, callback, user_data, None, + ctx, + kind as u32, + name, + 0, + false, + callback, + user_data, + None, ) }; unsafe { (host.v1.string_free)(name) }; @@ -664,7 +735,9 @@ unsafe extern "C" fn raw_async_allow_callback( }; let result = unsafe { raw_host_string(&host.v1, "null") }; if result.is_null() { - unsafe { reject_async_completion(host, completion, "failed to allocate async allow result") }; + unsafe { + reject_async_completion(host, completion, "failed to allocate async allow result") + }; } else { unsafe { (host.async_completion_resolve_json)(completion, result); @@ -686,27 +759,38 @@ unsafe extern "C" fn raw_async_passthrough_callback( let result = unsafe { raw_host_string_value(&host.v1, invocation_json) } .and_then(|value| serde_json::from_str::(&value).ok()) .and_then(|invocation| { - invocation.get("annotated").map(|annotated| { - json!({ - "request": invocation["request"], - "annotated_request": annotated, - "pending_marks": [], - "optimization_contributions": [], + invocation + .get("annotated") + .map(|annotated| { + json!({ + "request": invocation["request"], + "annotated_request": annotated, + "pending_marks": [], + "optimization_contributions": [], + }) + }) + .or_else(|| { + ["value", "request", "response", "fields"] + .into_iter() + .find_map(|key| invocation.get(key).cloned()) }) - }).or_else(|| { - ["value", "request", "response", "fields"] - .into_iter() - .find_map(|key| invocation.get(key).cloned()) - }) }) .and_then(|value| serde_json::to_string(&value).ok()); let Some(result) = result else { - unsafe { reject_async_completion(host, completion, "invalid async passthrough invocation") }; + unsafe { + reject_async_completion(host, completion, "invalid async passthrough invocation") + }; return NemoRelayNativeAsyncCallbackState::Complete as u32; }; let result = unsafe { raw_host_string(&host.v1, &result) }; if result.is_null() { - unsafe { reject_async_completion(host, completion, "failed to allocate async passthrough result") }; + unsafe { + reject_async_completion( + host, + completion, + "failed to allocate async passthrough result", + ) + }; } else { unsafe { (host.async_completion_resolve_json)(completion, result); @@ -744,7 +828,9 @@ unsafe extern "C" fn raw_async_tool_request_callback( .map(|value| (value, pending, duplicate)) }); let Some((result, pending, duplicate)) = invocation else { - unsafe { reject_async_completion(host, completion, "invalid async tool request invocation") }; + unsafe { + reject_async_completion(host, completion, "invalid async tool request invocation") + }; return NemoRelayNativeAsyncCallbackState::Complete as u32; }; if pending { @@ -790,7 +876,13 @@ unsafe extern "C" fn raw_async_tool_request_callback( (host.v1.string_free)(result); } } else { - unsafe { reject_async_completion(host, completion, "failed to allocate async tool request result") }; + unsafe { + reject_async_completion( + host, + completion, + "failed to allocate async tool request result", + ) + }; } NemoRelayNativeAsyncCallbackState::Complete as u32 } @@ -805,7 +897,13 @@ unsafe extern "C" fn raw_async_tool_execution_callback( return NemoRelayNativeAsyncCallbackState::Complete as u32; }; if next.is_null() || completion.is_null() { - unsafe { reject_async_completion(host, completion, "async tool execution requires next and completion") }; + unsafe { + reject_async_completion( + host, + completion, + "async tool execution requires next and completion", + ) + }; if !next.is_null() { unsafe { (host.async_next_release)(next) }; } @@ -823,13 +921,21 @@ unsafe extern "C" fn raw_async_tool_execution_callback( }) .and_then(|value| serde_json::to_string(&value).ok()); let Some(value) = value else { - unsafe { reject_async_completion(host, completion, "invalid async tool execution invocation") }; + unsafe { + reject_async_completion(host, completion, "invalid async tool execution invocation") + }; unsafe { (host.async_next_release)(next) }; return NemoRelayNativeAsyncCallbackState::Complete as u32; }; let value = unsafe { raw_host_string(&host.v1, &value) }; if value.is_null() { - unsafe { reject_async_completion(host, completion, "failed to allocate async tool execution invocation") }; + unsafe { + reject_async_completion( + host, + completion, + "failed to allocate async tool execution invocation", + ) + }; unsafe { (host.async_next_release)(next) }; return NemoRelayNativeAsyncCallbackState::Complete as u32; } @@ -844,7 +950,13 @@ unsafe extern "C" fn raw_async_tool_execution_callback( } NemoRelayNativeAsyncCallbackState::Pending as u32 } else { - unsafe { reject_async_completion(host, completion, "failed to invoke async tool execution next") }; + unsafe { + reject_async_completion( + host, + completion, + "failed to invoke async tool execution next", + ) + }; unsafe { (host.async_next_release)(next) }; NemoRelayNativeAsyncCallbackState::Complete as u32 } @@ -896,10 +1008,8 @@ unsafe extern "C" fn raw_tool_outcome_callback( } "fixture-status-error-outcome" => { unsafe { - *out_outcome_json = raw_host_string( - host, - r#"{"result":{"stale":true},"pending_marks":[]}"#, - ); + *out_outcome_json = + raw_host_string(host, r#"{"result":{"stale":true},"pending_marks":[]}"#); set_raw_last_error_from_user_data(user_data, "fixture tool execution failed"); } NemoRelayStatus::Internal diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 0fc04fa9f..58d4c3f60 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -324,13 +324,14 @@ pub fn wrap_tool_sanitize_fn( Arc::new(move |name: String, args: Json| { let ud = ud.clone(); Box::pin(async move { + clear_last_error(); let c_name = CString::new(name).unwrap_or_default(); let c_args = json_to_c_string(&args); let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; unsafe { nemo_relay_string_free_internal(c_args) }; - let result = ptr_to_json(result_ptr); + let result = json_result_from_ptr(result_ptr, "tool sanitize callback returned null"); unsafe { nemo_relay_string_free_internal(result_ptr) }; - Ok(result) + result }) }) } @@ -400,9 +401,9 @@ pub fn wrap_tool_exec_fn( let c_args = json_to_c_string(&args); let result_ptr = unsafe { cb(ud.ptr, c_args) }; unsafe { nemo_relay_string_free_internal(c_args) }; - let result = json_result_from_ptr(result_ptr, "tool execution callback failed")?; + let result = json_result_from_ptr(result_ptr, "tool execution callback failed"); unsafe { nemo_relay_string_free_internal(result_ptr) }; - Ok(result) + result }) }) } @@ -457,8 +458,9 @@ pub fn wrap_tool_exec_intercept_fn( unsafe { drop(Box::from_raw(next_ctx as *mut ToolExecutionNextFn)) }; unsafe { nemo_relay_string_free_internal(c_args) }; let outcome_json = - json_result_from_ptr(result_ptr, "tool execution intercept callback failed")?; + json_result_from_ptr(result_ptr, "tool execution intercept callback failed"); unsafe { nemo_relay_string_free_internal(result_ptr) }; + let outcome_json = outcome_json?; serde_json::from_value::(outcome_json).map_err(|error| { FlowError::Internal(format!( "invalid tool execution intercept outcome JSON: {error}" @@ -528,9 +530,9 @@ pub fn wrap_llm_exec_intercept_fn( unsafe { drop(Box::from_raw(next_ctx as *mut LlmExecutionNextFn)) }; unsafe { nemo_relay_string_free_internal(c_request) }; let result = - json_result_from_ptr(result_ptr, "LLM execution intercept callback failed")?; + json_result_from_ptr(result_ptr, "LLM execution intercept callback failed"); unsafe { nemo_relay_string_free_internal(result_ptr) }; - Ok(result) + result }) }, ) @@ -606,8 +608,9 @@ pub fn wrap_llm_stream_exec_intercept_fn( let result = json_result_from_ptr( result_ptr, "LLM stream execution intercept callback failed", - )?; + ); unsafe { nemo_relay_string_free_internal(result_ptr) }; + let result = result?; let stream = tokio_stream::once(Ok(result)); Ok(LlmJsonStream::new(stream)) }) @@ -843,9 +846,9 @@ pub fn wrap_llm_exec_fn( let c_request = json_to_c_string(&request_json); let result_ptr = unsafe { cb(ud.ptr, c_request) }; unsafe { nemo_relay_string_free_internal(c_request) }; - let result = json_result_from_ptr(result_ptr, "LLM execution callback failed")?; + let result = json_result_from_ptr(result_ptr, "LLM execution callback failed"); unsafe { nemo_relay_string_free_internal(result_ptr) }; - Ok(result) + result }) }) } @@ -868,8 +871,9 @@ pub fn wrap_llm_stream_exec_fn( let c_request = json_to_c_string(&request_json); let result_ptr = unsafe { cb(ud.ptr, c_request) }; unsafe { nemo_relay_string_free_internal(c_request) }; - let result = json_result_from_ptr(result_ptr, "LLM stream execution callback failed")?; + let result = json_result_from_ptr(result_ptr, "LLM stream execution callback failed"); unsafe { nemo_relay_string_free_internal(result_ptr) }; + let result = result?; // The C callback returns the full response as a single JSON value for stream // We emit it as a single-item stream let stream = tokio_stream::once(Ok(result)); @@ -1052,7 +1056,9 @@ fn json_result_from_ptr(ptr: *mut c_char, fallback: &str) -> Result { let message = last_error_message().unwrap_or_else(|| fallback.to_string()); return Err(FlowError::Internal(message)); } - Ok(ptr_to_json(ptr)) + let value = unsafe { CStr::from_ptr(ptr) }.to_string_lossy(); + serde_json::from_str(&value) + .map_err(|error| FlowError::Internal(format!("{fallback}: invalid JSON: {error}"))) } fn ptr_to_opt_string(ptr: *mut c_char) -> Option { diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index b541da082..b917f4042 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -1178,6 +1178,7 @@ pub fn wrap_js_event_sanitize_fn( Arc::new(move |event: Arc, fields: CoreEventSanitizeFields| { let func = func.clone(); Box::pin(async move { + let _active_callback = ActiveEventSanitizerCallback::enter(); let event_json = match JsEvent::try_from_event(&event) { Ok(event) => event.into_json(), Err(error) => { diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 2c51cfef8..b17f66f6b 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -854,9 +854,11 @@ pub fn wrap_py_llm_stream_exec_intercept_fn( /// Wrap a Python callable `(LlmRequest, LlmSanitizeRequestContext) -> Optional`. fn wrap_py_llm_sanitize_request_callback(py_fn: Py) -> LlmSanitizeRequestFn { let py_fn = Arc::new(py_fn); + let task_locals = capture_python_task_locals(); Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { let py_fn = py_fn.clone(); + let task_locals = capture_python_task_locals().or_else(|| task_locals.clone()); Box::pin(async move { let result = resolve_py_object_or_future(Python::attach(|py| { let result = py_fn @@ -868,7 +870,7 @@ fn wrap_py_llm_sanitize_request_callback(py_fn: Py) -> LlmSanitizeRequest ), ) .map_err(|e| FlowError::Internal(e.to_string()))?; - split_py_object_or_future(py, result) + split_py_object_or_future_with_locals(py, result, task_locals.as_ref()) })) .await?; Python::attach(|py| { @@ -893,14 +895,16 @@ fn wrap_py_llm_sanitize_request_callback(py_fn: Py) -> LlmSanitizeRequest /// Wrap a Python callable `(LlmRequest) -> Optional[str]` for LLM conditional guardrails. pub fn wrap_py_llm_conditional_fn(py_fn: Py) -> LlmConditionalFn { let py_fn = Arc::new(py_fn); + let task_locals = capture_python_task_locals(); Arc::new(move |request: LlmRequest| { let py_fn = py_fn.clone(); + let task_locals = capture_python_task_locals().or_else(|| task_locals.clone()); Box::pin(async move { let result = resolve_py_object_or_future(Python::attach(|py| { let result = py_fn .call1(py, (PyLLMRequest { inner: request },)) .map_err(|e| FlowError::Internal(e.to_string()))?; - split_py_object_or_future(py, result) + split_py_object_or_future_with_locals(py, result, task_locals.as_ref()) })) .await?; Python::attach(|py| { From dea33aa9e8071ba793b6b4f7994965a7a8f34930 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 21:54:53 -0400 Subject: [PATCH 29/52] fix: resolve async middleware review findings Signed-off-by: Will Killian --- .../adaptive/src/adaptive_hints_intercept.rs | 41 ++-- .../src/api/runtime/subscriber_dispatcher.rs | 14 ++ crates/ffi/nemo_relay.h | 13 +- crates/ffi/src/callable.rs | 58 ++++- crates/ffi/tests/unit/callable_tests.rs | 39 ++-- crates/node/src/api/mod.rs | 28 ++- crates/node/src/callable.rs | 199 ++++-------------- crates/node/src/callback_factory.rs | 83 ++++++-- crates/node/src/promise_call.rs | 60 +++++- crates/node/tests/event_sanitizers_tests.mjs | 85 ++++++++ crates/node/tests/llm_tests.mjs | 27 +++ crates/python/src/py_callable.rs | 69 ++++-- .../coverage/py_callable_coverage_tests.rs | 92 +++++++- python/nemo_relay/_event_sanitizer_context.py | 5 + python/tests/test_llm.py | 35 +++ 15 files changed, 588 insertions(+), 260 deletions(-) diff --git a/crates/adaptive/src/adaptive_hints_intercept.rs b/crates/adaptive/src/adaptive_hints_intercept.rs index 50437e4da..3a95f06ac 100644 --- a/crates/adaptive/src/adaptive_hints_intercept.rs +++ b/crates/adaptive/src/adaptive_hints_intercept.rs @@ -180,28 +180,25 @@ impl AdaptiveHintsIntercept { let this = this.clone(); let scope_path = extract_scope_path(); let manual_ls = read_manual_latency_sensitivity(); - Box::pin(async move { - let scope_depth = scope_path.len(); - let call_index = this.call_counter.fetch_add(1, Ordering::Relaxed); - - let effective_agent_id = this.effective_agent_id(); - let cached_hints = - this.load_hints(&scope_path, &effective_agent_id, call_index, scope_depth); - let final_hints = apply_manual_latency_override( - cached_hints, - manual_ls, - &effective_agent_id, - scope_depth, - ); - - if let Some(hints) = final_hints { - inject_agent_hints(&mut request, &mut annotated, &hints); - } - - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) - }) + let scope_depth = scope_path.len(); + let call_index = this.call_counter.fetch_add(1, Ordering::Relaxed); + let effective_agent_id = this.effective_agent_id(); + let cached_hints = + this.load_hints(&scope_path, &effective_agent_id, call_index, scope_depth); + let final_hints = apply_manual_latency_override( + cached_hints, + manual_ls, + &effective_agent_id, + scope_depth, + ); + + if let Some(hints) = final_hints { + inject_agent_hints(&mut request, &mut annotated, &hints); + } + + let outcome = + nemo_relay::api::llm::LlmRequestInterceptOutcome::new(request, annotated); + Box::pin(async move { Ok(outcome) }) }, ) } diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 47dace406..43412585b 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -198,6 +198,10 @@ mod native { Ok(()) } + pub(super) fn in_dispatcher_callback() -> bool { + IN_DISPATCHER.with(Cell::get) + } + fn dispatcher_sender() -> std::result::Result, String> { DISPATCHER.get_or_init(start_dispatcher).clone() } @@ -429,3 +433,13 @@ pub(crate) fn register_async_publication() -> Option pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() } + +/// Return whether the current callback was invoked by queued event publication. +/// +/// Bindings use this to make re-entrant flush operations non-blocking while +/// the serial dispatcher is awaiting middleware on another language runtime. +#[doc(hidden)] +#[must_use] +pub fn in_dispatcher_callback() -> bool { + native::in_dispatcher_callback() +} diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 70c480b24..accaa5e9e 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -252,6 +252,10 @@ typedef char *(*NemoRelayEventSanitizeCb)(void *user_data, /** * Optional destructor for user data passed to callbacks. * Called when the runtime no longer needs the associated callback. + * + * Middleware callbacks may run concurrently on Relay runtime or publication + * threads. Callers must keep `user_data` valid and thread-safe until this + * destructor runs. */ typedef void (*NemoRelayFreeFn)(void *user_data); @@ -361,6 +365,10 @@ typedef NemoRelayStatus (*NemoRelayLlmRequestInterceptCb)(void *user_data, /** * Runtime-provided "next" callback for LLM execution middleware chain. * Takes a native JSON C string, returns a response JSON C string. + * `next_ctx` is borrowed and valid only until the intercept callback returns; + * callers must not retain it or invoke `next_fn` asynchronously. The returned + * string belongs to the caller and must be released with + * `nemo_relay_string_free`. */ typedef char *(*NemoRelayLlmExecNextFn)(const char *native_json, void *next_ctx); @@ -405,7 +413,10 @@ typedef char *(*NemoRelayToolConditionalCb)(void *user_data, const char *name, c /** * Runtime-provided "next" callback for tool execution middleware chain. * Call this from an intercept to invoke the next layer (or original function). - * `next_ctx` is an opaque pointer managed by the runtime. + * `next_ctx` is borrowed and valid only until the intercept callback returns; + * callers must not retain it or invoke `next_fn` asynchronously. The returned + * string belongs to the caller and must be released with + * `nemo_relay_string_free`. */ typedef char *(*NemoRelayToolExecNextFn)(const char *args_json, void *next_ctx); diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 58d4c3f60..054c098bd 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -38,7 +38,7 @@ use nemo_relay::codec::request::AnnotatedLlmRequest as AnnotatedLLMRequest; use nemo_relay::codec::traits::LlmCodec; use nemo_relay::error::{FlowError, Result}; -use crate::convert::{c_str_to_json, json_to_c_string}; +use crate::convert::json_to_c_string; use crate::error::{NemoRelayStatus, clear_last_error, last_error_message, set_last_error}; use crate::types::{FfiEvent, FfiLLMRequest, FfiPluginContext}; @@ -48,6 +48,10 @@ use crate::types::{FfiEvent, FfiLLMRequest, FfiPluginContext}; /// Optional destructor for user data passed to callbacks. /// Called when the runtime no longer needs the associated callback. +/// +/// Middleware callbacks may run concurrently on Relay runtime or publication +/// threads. Callers must keep `user_data` valid and thread-safe until this +/// destructor runs. pub type NemoRelayFreeFn = Option; /// Callback for tool request/response sanitization guardrails and intercepts. @@ -76,7 +80,10 @@ pub type NemoRelayToolExecCb = /// Runtime-provided "next" callback for tool execution middleware chain. /// Call this from an intercept to invoke the next layer (or original function). -/// `next_ctx` is an opaque pointer managed by the runtime. +/// `next_ctx` is borrowed and valid only until the intercept callback returns; +/// callers must not retain it or invoke `next_fn` asynchronously. The returned +/// string belongs to the caller and must be released with +/// `nemo_relay_string_free`. pub type NemoRelayToolExecNextFn = unsafe extern "C" fn(args_json: *const c_char, next_ctx: *mut libc::c_void) -> *mut c_char; @@ -169,6 +176,10 @@ pub type NemoRelayLlmExecCb = /// Runtime-provided "next" callback for LLM execution middleware chain. /// Takes a native JSON C string, returns a response JSON C string. +/// `next_ctx` is borrowed and valid only until the intercept callback returns; +/// callers must not retain it or invoke `next_fn` asynchronously. The returned +/// string belongs to the caller and must be released with +/// `nemo_relay_string_free`. pub type NemoRelayLlmExecNextFn = unsafe extern "C" fn(native_json: *const c_char, next_ctx: *mut libc::c_void) -> *mut c_char; @@ -398,6 +409,7 @@ pub fn wrap_tool_exec_fn( Box::new(move |args: Json| { let ud = ud.clone(); Box::pin(async move { + clear_last_error(); let c_args = json_to_c_string(&args); let result_ptr = unsafe { cb(ud.ptr, c_args) }; unsafe { nemo_relay_string_free_internal(c_args) }; @@ -454,6 +466,7 @@ pub fn wrap_tool_exec_intercept_fn( } let c_args = json_to_c_string(&args); + clear_last_error(); let result_ptr = unsafe { cb(ud.ptr, c_args, tool_next_trampoline, next_ctx) }; unsafe { drop(Box::from_raw(next_ctx as *mut ToolExecutionNextFn)) }; unsafe { nemo_relay_string_free_internal(c_args) }; @@ -526,6 +539,7 @@ pub fn wrap_llm_exec_intercept_fn( let request_json = serde_json::to_value(&request).unwrap_or(Json::Null); let c_request = json_to_c_string(&request_json); + clear_last_error(); let result_ptr = unsafe { cb(ud.ptr, c_request, llm_next_trampoline, next_ctx) }; unsafe { drop(Box::from_raw(next_ctx as *mut LlmExecutionNextFn)) }; unsafe { nemo_relay_string_free_internal(c_request) }; @@ -601,6 +615,7 @@ pub fn wrap_llm_stream_exec_intercept_fn( let request_json = serde_json::to_value(&request).unwrap_or(Json::Null); let c_request = json_to_c_string(&request_json); + clear_last_error(); let result_ptr = unsafe { cb(ud.ptr, c_request, llm_stream_next_trampoline, next_ctx) }; unsafe { drop(Box::from_raw(next_ctx as *mut LlmStreamExecutionNextFn)) }; @@ -710,7 +725,7 @@ pub fn wrap_llm_sanitize_request_fn( Ok(identity) => identity, Err(error) => { set_last_error(&error.to_string()); - return Ok(None); + return Err(error); } }; let codec = context @@ -727,7 +742,10 @@ pub fn wrap_llm_sanitize_request_fn( let result_ptr = unsafe { cb(ud.ptr, ffi_req, ffi_context) }; if result_ptr.is_null() { unsafe { drop(Box::from_raw(ffi_req)) }; - return Ok(None); + return match last_error_message() { + Some(message) => Err(FlowError::Internal(message)), + None => Ok(None), + }; } if result_ptr == ffi_req { return Ok(Some(unsafe { Box::from_raw(ffi_req) }.0)); @@ -754,7 +772,7 @@ pub fn wrap_llm_sanitize_response_fn( Ok(identity) => identity, Err(error) => { set_last_error(&error.to_string()); - return Ok(None); + return Err(error); } }; let codec = context @@ -771,16 +789,32 @@ pub fn wrap_llm_sanitize_response_fn( let result_ptr = unsafe { cb(ud.ptr, response_json, ffi_context) }; if result_ptr.is_null() { unsafe { nemo_relay_string_free_internal(response_json) }; - return Ok(None); + return match last_error_message() { + Some(message) => Err(FlowError::Internal(message)), + None => Ok(None), + }; } - let result = c_str_to_json(result_ptr); + let result = unsafe { CStr::from_ptr(result_ptr) } + .to_str() + .map_err(|error| { + FlowError::Internal(format!( + "LLM response sanitizer returned invalid UTF-8: {error}" + )) + }) + .and_then(|value| { + serde_json::from_str(value).map_err(|error| { + FlowError::Internal(format!( + "LLM response sanitizer returned invalid JSON: {error}" + )) + }) + }); unsafe { nemo_relay_string_free_internal(response_json); if result_ptr != response_json { nemo_relay_string_free_internal(result_ptr); } } - Ok(result) + result.map(Some) }) }) } @@ -842,6 +876,7 @@ pub fn wrap_llm_exec_fn( Box::new(move |request: LlmRequest| { let ud = ud.clone(); Box::pin(async move { + clear_last_error(); let request_json = serde_json::to_value(&request).unwrap_or(Json::Null); let c_request = json_to_c_string(&request_json); let result_ptr = unsafe { cb(ud.ptr, c_request) }; @@ -867,6 +902,7 @@ pub fn wrap_llm_stream_exec_fn( Box::new(move |request: LlmRequest| { let ud = ud.clone(); Box::pin(async move { + clear_last_error(); let request_json = serde_json::to_value(&request).unwrap_or(Json::Null); let c_request = json_to_c_string(&request_json); let result_ptr = unsafe { cb(ud.ptr, c_request) }; @@ -1056,8 +1092,10 @@ fn json_result_from_ptr(ptr: *mut c_char, fallback: &str) -> Result { let message = last_error_message().unwrap_or_else(|| fallback.to_string()); return Err(FlowError::Internal(message)); } - let value = unsafe { CStr::from_ptr(ptr) }.to_string_lossy(); - serde_json::from_str(&value) + let value = unsafe { CStr::from_ptr(ptr) } + .to_str() + .map_err(|error| FlowError::Internal(format!("{fallback}: invalid UTF-8: {error}")))?; + serde_json::from_str(value) .map_err(|error| FlowError::Internal(format!("{fallback}: invalid JSON: {error}"))) } diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 70469b4f4..4542bb551 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -473,32 +473,37 @@ fn test_wrap_llm_request_response_and_conditional_callbacks() { for callback in [invalid_json_cb, invalid_utf8_cb] { let malformed_response = wrap_llm_sanitize_response_fn(callback, std::ptr::null_mut(), None); - assert_eq!( - resolve(malformed_response( - json!({"secret": "must be omitted"}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - )) - .unwrap(), - None + let error = resolve(malformed_response( + json!({"secret": "must be preserved"}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), + )) + .unwrap_err(); + assert!( + error.to_string().contains("invalid"), + "unexpected sanitizer error: {error}" ); } } #[test] -fn test_llm_sanitizers_fail_closed_for_runtime_codec_ids_with_embedded_nul() { +fn test_llm_sanitizers_report_runtime_codec_ids_with_embedded_nul() { let runtime_identity = nemo_relay::api::runtime::LlmCodecIdentity::Runtime("runtime\0codec".to_string()); let request_sanitizer = wrap_llm_sanitize_request_fn(llm_request_alias_cb, std::ptr::null_mut(), None); - let request_result = resolve(request_sanitizer( + let request_error = resolve(request_sanitizer( make_request(), nemo_relay::api::runtime::LlmSanitizeRequestContext::with_identity( runtime_identity.clone(), ), )) - .expect("legacy sanitizer wrappers report callback errors out of band"); - assert_eq!(request_result, None); + .unwrap_err(); + assert!( + request_error + .to_string() + .contains("runtime codec ID contains an embedded NUL") + ); assert!( last_error_message() .unwrap() @@ -507,12 +512,16 @@ fn test_llm_sanitizers_fail_closed_for_runtime_codec_ids_with_embedded_nul() { let response_sanitizer = wrap_llm_sanitize_response_fn(json_alias_cb, std::ptr::null_mut(), None); - let response_result = resolve(response_sanitizer( - json!({"secret": "must be omitted"}), + let response_error = resolve(response_sanitizer( + json!({"secret": "must be preserved"}), nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), )) - .expect("legacy sanitizer wrappers report callback errors out of band"); - assert_eq!(response_result, None); + .unwrap_err(); + assert!( + response_error + .to_string() + .contains("runtime codec ID contains an embedded NUL") + ); assert!( last_error_message() .unwrap() diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index 6e22c3dc9..9c00ee254 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -1413,7 +1413,9 @@ impl PersistentJsFunction { } fn node_event_sanitize_fn(env: &Env, func: &JsFunction) -> napi::Result { - let callback = Arc::new(crate::promise_call::PromiseAwareFn::new(env, func)?); + let callback = Arc::new(crate::promise_call::PromiseAwareFn::new_event_sanitizer( + env, func, + )?); Ok(callable::wrap_js_event_sanitize_promise_fn(callback)) } @@ -3180,15 +3182,21 @@ pub fn deregister_subscriber(name: String) -> Result { /// /// The Promise rejects if the blocking task fails or the core subscriber flush returns an error. /// Callers should handle errors when awaiting it. -#[napi] -pub async fn flush_subscribers() -> Result<()> { - if crate::callable::event_sanitizer_callback_active() { - return Ok(()); - } - tokio::task::spawn_blocking(core_subscriber_api::flush_subscribers) - .await - .map_err(|error| to_napi_err(FlowError::Internal(error.to_string())))? - .map_err(to_napi_err) +#[napi(ts_return_type = "Promise")] +pub fn flush_subscribers(env: Env) -> Result { + let reentrant = crate::callback_factory::event_sanitizer_callback_active(&env)?; + env.execute_tokio_future( + async move { + if reentrant { + return Ok(()); + } + tokio::task::spawn_blocking(core_subscriber_api::flush_subscribers) + .await + .map_err(|error| to_napi_err(FlowError::Internal(error.to_string())))? + .map_err(to_napi_err) + }, + |env, _| env.get_undefined(), + ) } // --------------------------------------------------------------------------- diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index b917f4042..905f77e49 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -12,7 +12,6 @@ use std::future::Future; use std::pin::Pin; use std::sync::Arc; -use std::sync::atomic::{AtomicUsize, Ordering}; use napi::bindgen_prelude::ToNapiValue; use napi::threadsafe_function::{ErrorStrategy, ThreadsafeFunction, ThreadsafeFunctionCallMode}; @@ -45,27 +44,6 @@ use crate::convert::{callback_json, record_callback_error, to_napi_err}; use crate::promise_call::{JsonNextFn, JsonStreamNextFn, PromiseAwareFn}; use crate::types::{EventSanitizeFields, JsEvent, event_sanitize_fields_from_json}; -static ACTIVE_EVENT_SANITIZER_CALLBACKS: AtomicUsize = AtomicUsize::new(0); - -struct ActiveEventSanitizerCallback; - -impl ActiveEventSanitizerCallback { - fn enter() -> Self { - ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_add(1, Ordering::AcqRel); - Self - } -} - -impl Drop for ActiveEventSanitizerCallback { - fn drop(&mut self) { - ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_sub(1, Ordering::AcqRel); - } -} - -pub(crate) fn event_sanitizer_callback_active() -> bool { - ACTIVE_EVENT_SANITIZER_CALLBACKS.load(Ordering::Acquire) != 0 -} - /// Structured codec identity delivered to JavaScript LLM sanitizers. #[napi(object)] #[derive(Clone)] @@ -321,6 +299,8 @@ pub fn wrap_js_llm_sanitize_request_promise_fn(func: Arc) -> Llm Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { let func = func.clone(); + let publication = + nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); Box::pin(async move { let request = serde_json::to_value(request).map_err(|error| { let error = FlowError::Internal(format!( @@ -330,26 +310,26 @@ pub fn wrap_js_llm_sanitize_request_promise_fn(func: Arc) -> Llm error })?; let context = js_llm_sanitize_request_context(&context); - let value = func - .call_spread_with_arg0(Box::new(move |env| { - let mut args = env.create_array_with_length(2)?; - let request = unsafe { - JsUnknown::from_raw_unchecked( - env.raw(), - Json::to_napi_value(env.raw(), request)?, - ) - }; - args.set_element(0, request)?; - args.set_element( - 1, - js_llm_sanitize_request_context_to_napi(env, context)?, - )?; - Ok(js_object_to_unknown(env, args)) - })) - .await - .inspect_err(|error| { - record_callback_error(error.to_string()); - })?; + let build_args: crate::promise_call::Arg0Builder = Box::new(move |env| { + let mut args = env.create_array_with_length(2)?; + let request = unsafe { + JsUnknown::from_raw_unchecked( + env.raw(), + Json::to_napi_value(env.raw(), request)?, + ) + }; + args.set_element(0, request)?; + args.set_element(1, js_llm_sanitize_request_context_to_napi(env, context)?)?; + Ok(js_object_to_unknown(env, args)) + }); + let value = if publication { + func.call_spread_with_arg0_for_publication(build_args).await + } else { + func.call_spread_with_arg0(build_args).await + } + .inspect_err(|error| { + record_callback_error(error.to_string()); + })?; if value.is_null() { Ok(None) } else { @@ -374,25 +354,29 @@ pub fn wrap_js_llm_sanitize_response_promise_fn( ) -> LlmSanitizeResponseFn { Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { let func = func.clone(); + let publication = nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); Box::pin(async move { let context = js_llm_sanitize_response_context(&context); - let value = func - .call_spread_with_arg0(Box::new(move |env| { - let mut args = env.create_array_with_length(2)?; - let response = unsafe { - JsUnknown::from_raw_unchecked( - env.raw(), - Json::to_napi_value(env.raw(), response)?, - ) - }; - args.set_element(0, response)?; - args.set_element(1, js_llm_sanitize_response_context_to_napi(env, context)?)?; - Ok(js_object_to_unknown(env, args)) - })) - .await - .inspect_err(|error| { - record_callback_error(error.to_string()); - })?; + let build_args: crate::promise_call::Arg0Builder = Box::new(move |env| { + let mut args = env.create_array_with_length(2)?; + let response = unsafe { + JsUnknown::from_raw_unchecked( + env.raw(), + Json::to_napi_value(env.raw(), response)?, + ) + }; + args.set_element(0, response)?; + args.set_element(1, js_llm_sanitize_response_context_to_napi(env, context)?)?; + Ok(js_object_to_unknown(env, args)) + }); + let value = if publication { + func.call_spread_with_arg0_for_publication(build_args).await + } else { + func.call_spread_with_arg0(build_args).await + } + .inspect_err(|error| { + record_callback_error(error.to_string()); + })?; Ok((!value.is_null()).then_some(value)) }) }) @@ -498,7 +482,6 @@ pub fn wrap_js_event_sanitize_promise_fn(func: Arc) -> EventSani Arc::new(move |event: Arc, fields: CoreEventSanitizeFields| { let func = func.clone(); Box::pin(async move { - let _active_callback = ActiveEventSanitizerCallback::enter(); let event_json = JsEvent::try_from_event(&event) .map(JsEvent::into_json) .map_err(|error| { @@ -1170,102 +1153,6 @@ pub fn wrap_js_event_subscriber( }) } -/// Wrap a JS event sanitizer: ``(event, fields) => fields``. -pub fn wrap_js_event_sanitize_fn( - func: ThreadsafeFunction<(Json, Json), ErrorStrategy::Fatal>, -) -> EventSanitizeFn { - let func = Arc::new(func); - Arc::new(move |event: Arc, fields: CoreEventSanitizeFields| { - let func = func.clone(); - Box::pin(async move { - let _active_callback = ActiveEventSanitizerCallback::enter(); - let event_json = match JsEvent::try_from_event(&event) { - Ok(event) => event.into_json(), - Err(error) => { - record_callback_error(format!( - "nemo_relay: failed to serialize JS event sanitizer context: {error}" - )); - return Err(FlowError::Internal(error.to_string())); - } - }; - let js_fields = EventSanitizeFields { - data: fields.data, - category_profile: fields - .category_profile - .as_ref() - .map(serde_json::to_value) - .transpose() - .map_err(|error| { - let error = FlowError::Internal(format!( - "failed to serialize JS event sanitizer category profile: {error}" - )); - record_callback_error(error.to_string()); - error - })?, - metadata: fields.metadata, - }; - let js_fields = serde_json::to_value(js_fields).map_err(|error| { - let error = FlowError::Internal(format!( - "failed to serialize JS event sanitizer fields: {error}" - )); - record_callback_error(error.to_string()); - error - })?; - let (tx, rx) = tokio::sync::oneshot::channel(); - let status = func.call_with_return_value( - (event_json, js_fields), - ThreadsafeFunctionCallMode::Blocking, - move |value: Option| { - let _ = tx.send(callback_json(value)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - record_callback_error(format!( - "nemo_relay: failed to queue JS event sanitizer callback: {status:?}" - )); - return Err(FlowError::Internal(format!( - "failed to queue JS event sanitizer callback: {status:?}" - ))); - } - let sanitized: Result = async { - let result = await_middleware_json_result( - rx, - "nemo_relay: JS event sanitizer callback failed", - ) - .await?; - let result = event_sanitize_fields_from_json(result).map_err(|error| { - FlowError::Internal(format!( - "nemo_relay: invalid JS event sanitizer result: {error}" - )) - })?; - let category_profile = result - .category_profile - .map(serde_json::from_value) - .transpose() - .map_err(|error| { - FlowError::Internal(format!( - "nemo_relay: invalid JS event sanitizer result: {error}" - )) - })?; - Ok(CoreEventSanitizeFields { - data: result.data, - category_profile, - metadata: result.metadata, - }) - } - .await; - match sanitized { - Ok(sanitized) => Ok(sanitized), - Err(error) => { - record_callback_error(error.to_string()); - Err(error) - } - } - }) - }) -} - // --------------------------------------------------------------------------- // Codec wrappers // --------------------------------------------------------------------------- diff --git a/crates/node/src/callback_factory.rs b/crates/node/src/callback_factory.rs index 5fcd45ed8..9314c544c 100644 --- a/crates/node/src/callback_factory.rs +++ b/crates/node/src/callback_factory.rs @@ -5,9 +5,12 @@ use napi::{Env, JsFunction, JsObject, JsUnknown, NapiRaw, NapiValue}; -const CALLBACK_FACTORIES_PROPERTY: &str = "__nemo_relay_callback_factories_v1"; +const CALLBACK_FACTORIES_PROPERTY: &str = "__nemo_relay_callback_factories_v2"; const CALLBACK_FACTORIES_SOURCE: &str = r#"(() => { + const { AsyncLocalStorage } = process.getBuiltinModule('node:async_hooks'); + const eventSanitizerContext = new AsyncLocalStorage(); + function jsonValue(value, seen = new Set()) { if (value === null || typeof value === 'string' || typeof value === 'boolean') { return value; @@ -49,6 +52,38 @@ const CALLBACK_FACTORIES_SOURCE: &str = r#"(() => { return result; } + function callPromise(fn, arg0, spread, next, resolve, reject, publication) { + const token = { active: publication }; + const invoke = () => { + Promise.resolve().then(() => ( + next === undefined + ? (spread ? fn(...arg0) : fn(arg0)) + : (spread ? fn(...arg0, next) : fn(arg0, next)) + )).then((value) => jsonValue(value === undefined ? null : value)).then((value) => { + token.active = false; + resolve(value); + }, (error) => { + token.active = false; + let message = 'unknown error'; + try { + if (typeof error === 'string') { + message = error; + } else if (error === null || (typeof error !== 'object' && typeof error !== 'function')) { + message = String(error); + } else if (error != null && typeof error.message === 'string') { + message = error.message; + } + } catch {} + reject(message); + }); + }; + if (publication) { + eventSanitizerContext.run(token, invoke); + } else { + invoke(); + } + } + return { execution(fn) { return function __nemo_relay_execution_wrapper(...args) { @@ -66,7 +101,7 @@ const CALLBACK_FACTORIES_SOURCE: &str = r#"(() => { }, promise(fn) { - return function __nemo_relay_promise_wrapper(error, arg0, spread, next, resolve, reject) { + return function __nemo_relay_promise_wrapper(error, arg0, spread, next, resolve, reject, publication) { if (error != null) { let message = 'unknown error'; try { @@ -75,25 +110,27 @@ const CALLBACK_FACTORIES_SOURCE: &str = r#"(() => { reject(message); return; } - Promise.resolve().then(() => ( - next === undefined - ? (spread ? fn(...arg0) : fn(arg0)) - : (spread ? fn(...arg0, next) : fn(arg0, next)) - )).then((value) => jsonValue(value === undefined ? null : value)).then(resolve, (error) => { + callPromise(fn, arg0, spread, next, resolve, reject, publication); + }; + }, + + eventSanitizerPromise(fn) { + return function __nemo_relay_event_sanitizer_promise_wrapper(error, arg0, spread, next, resolve, reject) { + if (error != null) { let message = 'unknown error'; try { - if (typeof error === 'string') { - message = error; - } else if (error === null || (typeof error !== 'object' && typeof error !== 'function')) { - message = String(error); - } else if (error != null && typeof error.message === 'string') { - message = error.message; - } + message = String(error?.message ?? error); } catch {} reject(message); - }); + return; + } + callPromise(fn, arg0, spread, next, resolve, reject, true); }; }, + + eventSanitizerCallbackActive() { + return eventSanitizerContext.getStore()?.active === true; + }, }; })()"#; @@ -140,3 +177,19 @@ pub(crate) fn wrap_execution_callback(env: &Env, func: &JsFunction) -> napi::Res pub(crate) fn wrap_promise_callback(env: &Env, func: &JsFunction) -> napi::Result { wrap_callback(env, func, "promise") } + +pub(crate) fn wrap_event_sanitizer_callback( + env: &Env, + func: &JsFunction, +) -> napi::Result { + wrap_callback(env, func, "eventSanitizerPromise") +} + +pub(crate) fn event_sanitizer_callback_active(env: &Env) -> napi::Result { + let factories = callback_factories(env)?; + let callback: JsFunction = factories.get_named_property("eventSanitizerCallbackActive")?; + callback + .call::(None, &[])? + .coerce_to_bool()? + .get_value() +} diff --git a/crates/node/src/promise_call.rs b/crates/node/src/promise_call.rs index cdc15bada..df9a7785a 100644 --- a/crates/node/src/promise_call.rs +++ b/crates/node/src/promise_call.rs @@ -55,6 +55,7 @@ struct CallArgs { arg0: PrimaryArg, spread: bool, next: Option, + publication: bool, completion: CallCompletion, } @@ -184,8 +185,19 @@ impl PromiseAwareFn { /// Must be called on the JS main thread (i.e., in a sync `#[napi]` function). pub fn new(env: &Env, func: &JsFunction) -> napi::Result { let wrapper = callback_factory::wrap_promise_callback(env, func)?; + Self::from_wrapper(env, &wrapper) + } + + /// Create a callback wrapper that marks only its JavaScript async context + /// as an active event sanitizer. + pub fn new_event_sanitizer(env: &Env, func: &JsFunction) -> napi::Result { + let wrapper = callback_factory::wrap_event_sanitizer_callback(env, func)?; + Self::from_wrapper(env, &wrapper) + } + + fn from_wrapper(env: &Env, wrapper: &JsFunction) -> napi::Result { let mut tsfn = - env.create_threadsafe_function(&wrapper, 0, |ctx: ThreadSafeCallContext| { + env.create_threadsafe_function(wrapper, 0, |ctx: ThreadSafeCallContext| { let next = match ctx.value.next { Some(next) => build_next_unknown(&ctx.env, next)?, None => undefined_to_unknown(&ctx.env)?, @@ -202,7 +214,13 @@ impl PromiseAwareFn { ctx.env.get_boolean(ctx.value.spread)?.raw(), ) }; - let args = vec![arg0, spread, next, resolve, reject]; + let publication = unsafe { + JsUnknown::from_raw_unchecked( + ctx.env.raw(), + ctx.env.get_boolean(ctx.value.publication)?.raw(), + ) + }; + let args = vec![arg0, spread, next, resolve, reject, publication]; Ok(args) })?; @@ -216,7 +234,8 @@ impl PromiseAwareFn { /// Call the JS function with the given args and await the result. pub async fn call(&self, args: Json) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), false, None).await + self.call_inner(PrimaryArg::Json(args), false, None, false) + .await } /// Call a JavaScript callback with several JSON arguments. @@ -225,7 +244,7 @@ impl PromiseAwareFn { /// guardrails, whose public contract is `(name, payload)` rather than a /// single envelope object. pub async fn call_spread(&self, args: Vec) -> FlowResult { - self.call_inner(PrimaryArg::Json(Json::Array(args)), true, None) + self.call_inner(PrimaryArg::Json(Json::Array(args)), true, None, false) .await } @@ -236,21 +255,35 @@ impl PromiseAwareFn { /// cannot cross the threadsafe-function boundary as plain JSON, such as a /// `#[napi]` class instance. pub async fn call_with_arg0(&self, build_arg0: Arg0Builder) -> FlowResult { - self.call_inner(PrimaryArg::Build(build_arg0), false, None) + self.call_inner(PrimaryArg::Build(build_arg0), false, None, false) .await } /// Call a JavaScript callback with builder-constructed spread arguments. pub async fn call_spread_with_arg0(&self, build_arg0: Arg0Builder) -> FlowResult { - self.call_inner(PrimaryArg::Build(build_arg0), true, None) + self.call_inner(PrimaryArg::Build(build_arg0), true, None, false) + .await + } + + /// Call a spread callback from queued event publication. + pub async fn call_spread_with_arg0_for_publication( + &self, + build_arg0: Arg0Builder, + ) -> FlowResult { + self.call_inner(PrimaryArg::Build(build_arg0), true, None, true) .await } /// Call the JS function with a middleware-style `next(arg)` callback that /// resolves to a JSON result. pub async fn call_with_json_next(&self, args: Json, next: JsonNextFn) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), false, Some(NextFn::Json(next))) - .await + self.call_inner( + PrimaryArg::Json(args), + false, + Some(NextFn::Json(next)), + false, + ) + .await } /// Call the JS function with a middleware-style `next(arg)` callback that @@ -260,8 +293,13 @@ impl PromiseAwareFn { args: Json, next: JsonStreamNextFn, ) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), false, Some(NextFn::Stream(next))) - .await + self.call_inner( + PrimaryArg::Json(args), + false, + Some(NextFn::Stream(next)), + false, + ) + .await } /// Release the underlying threadsafe function so it does not outlive its registration. @@ -276,6 +314,7 @@ impl PromiseAwareFn { arg0: PrimaryArg, spread: bool, next: Option, + publication: bool, ) -> FlowResult { let (sender, receiver) = tokio::sync::oneshot::channel(); let tsfn = self @@ -290,6 +329,7 @@ impl PromiseAwareFn { arg0, spread, next, + publication, completion: CallCompletion::new(sender), }), napi::threadsafe_function::ThreadsafeFunctionCallMode::NonBlocking, diff --git a/crates/node/tests/event_sanitizers_tests.mjs b/crates/node/tests/event_sanitizers_tests.mjs index e437d9574..2380f2cbc 100644 --- a/crates/node/tests/event_sanitizers_tests.mjs +++ b/crates/node/tests/event_sanitizers_tests.mjs @@ -148,6 +148,91 @@ describe('event sanitizer registries', () => { assert.equal(flushReturned, true); }); + it('does not treat an unrelated flush as sanitizer re-entrancy', async () => { + const events = capture('node-event-sanitize-independent-flush-sub'); + let releaseSanitizer; + let sanitizerEntered; + const entered = new Promise((resolve) => { + sanitizerEntered = resolve; + }); + const release = new Promise((resolve) => { + releaseSanitizer = resolve; + }); + lib.registerMarkSanitizeGuardrail('node-event-independent-flush', 0, async (_event, fields) => { + sanitizerEntered(); + await release; + return fields; + }); + try { + lib.event('independent-flush-checkpoint', null, { raw: true }); + await entered; + const flush = lib.flushSubscribers(); + const state = await Promise.race([ + flush.then(() => 'flushed'), + new Promise((resolve) => setImmediate(() => resolve('pending'))), + ]); + assert.equal(state, 'pending'); + releaseSanitizer(); + await flush; + await waitFor(events, 1); + } finally { + releaseSanitizer(); + lib.deregisterMarkSanitizeGuardrail('node-event-independent-flush'); + lib.deregisterSubscriber('node-event-sanitize-independent-flush-sub'); + } + }); + + it('clears sanitizer re-entrancy in async descendants after settlement', async () => { + const events = capture('node-event-sanitize-descendant-flush-sub'); + let secondSanitizerEntered; + const secondEntered = new Promise((resolve) => { + secondSanitizerEntered = resolve; + }); + let releaseSecondSanitizer; + const releaseSecond = new Promise((resolve) => { + releaseSecondSanitizer = resolve; + }); + let descendantFlushStarted; + const flushStarted = new Promise((resolve) => { + descendantFlushStarted = resolve; + }); + let descendantFlush; + const flushed = new Promise((resolve, reject) => { + descendantFlush = { resolve, reject }; + }); + lib.registerMarkSanitizeGuardrail('node-event-descendant-flush', 0, async (event, fields) => { + if (event.name === 'descendant-flush-origin') { + setTimeout(async () => { + await secondEntered; + descendantFlushStarted(); + lib.flushSubscribers().then(descendantFlush.resolve, descendantFlush.reject); + }, 0); + } else if (event.name === 'descendant-flush-blocked') { + secondSanitizerEntered(); + await releaseSecond; + } + return fields; + }); + try { + lib.event('descendant-flush-origin', null, { raw: true }); + lib.event('descendant-flush-blocked', null, { raw: true }); + await secondEntered; + await flushStarted; + const state = await Promise.race([ + flushed.then(() => 'flushed'), + new Promise((resolve) => setImmediate(() => resolve('pending'))), + ]); + assert.equal(state, 'pending'); + releaseSecondSanitizer(); + await flushed; + await waitFor(events, 2); + } finally { + releaseSecondSanitizer(); + lib.deregisterMarkSanitizeGuardrail('node-event-descendant-flush'); + lib.deregisterSubscriber('node-event-sanitize-descendant-flush-sub'); + } + }); + it('fails open and records invalid sanitizer results', async () => { const events = capture('node-event-sanitize-invalid-sub'); const invalidResults = { diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index 5408f6fa9..b54939a2d 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -645,6 +645,33 @@ describe('LLM guardrails', () => { deregisterLlmSanitizeRequestGuardrail('node_llm_san_req'); }); + it('manual async sanitizers can flush subscribers without deadlocking', async () => { + let requestFlushed = false; + let responseFlushed = false; + registerSubscriber('node_manual_flush_subscriber', () => {}); + registerLlmSanitizeRequestGuardrail('node_manual_flush_request', 10, async (request) => { + await flushSubscribers(); + requestFlushed = true; + return request; + }); + registerLlmSanitizeResponseGuardrail('node_manual_flush_response', 10, async (response) => { + await flushSubscribers(); + responseFlushed = true; + return response; + }); + try { + const handle = llmCall('node_manual_flush', makeNative()); + llmCallEnd(handle, { response: 'ok' }); + await flushSubscribers(); + } finally { + deregisterLlmSanitizeRequestGuardrail('node_manual_flush_request'); + deregisterLlmSanitizeResponseGuardrail('node_manual_flush_response'); + deregisterSubscriber('node_manual_flush_subscriber'); + } + assert.equal(requestFlushed, true); + assert.equal(responseFlushed, true); + }); + it('sanitize request guardrail rewrites start event payload', async () => { const events = []; registerSubscriber('node_llm_san_req_evt', (e) => events.push(e)); diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index b17f66f6b..ce23c5ec3 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -32,6 +32,7 @@ use nemo_relay::api::runtime::{ ToolConditionalFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, }; use nemo_relay::error::{FlowError, Result as FlowResult}; +use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; use pyo3::types::PyDict; use pyo3_async_runtimes::TaskLocals; @@ -128,7 +129,15 @@ fn split_py_object_or_future( py: Python<'_>, result: Py, ) -> FlowResult, PyValueFuture>> { - split_py_object_or_future_with_locals(py, result, None) + let bound = result.bind(py); + if bound.getattr("__await__").is_ok() { + reject_awaitable_from_sync_caller(bound)?; + let future = pyo3_async_runtimes::tokio::into_future(result.into_bound(py)) + .map_err(|error| FlowError::Internal(error.to_string()))?; + Ok(Err(Box::pin(future))) + } else { + Ok(Ok(result)) + } } fn split_py_object_or_future_with_locals( @@ -144,10 +153,21 @@ fn split_py_object_or_future_with_locals( pyo3_async_runtimes::into_future_with_locals(locals, result.into_bound(py)) .map_err(|e| FlowError::Internal(e.to_string()))?, ), - None => Box::pin( - pyo3_async_runtimes::tokio::into_future(result.into_bound(py)) - .map_err(|e| FlowError::Internal(e.to_string()))?, - ), + None => Box::pin(async move { + tokio::task::spawn_blocking(move || { + Python::attach(|py| { + let coroutine = py + .import("nemo_relay._event_sanitizer_context") + .and_then(|module| module.getattr("await_result")) + .and_then(|await_result| await_result.call1((result.bind(py),)))?; + py.import("asyncio") + .and_then(|asyncio| asyncio.call_method1("run", (coroutine,))) + .map(Bound::unbind) + }) + }) + .await + .map_err(|error| PyRuntimeError::new_err(error.to_string()))? + }), }; Ok(Err(future)) } else { @@ -859,18 +879,23 @@ fn wrap_py_llm_sanitize_request_callback(py_fn: Py) -> LlmSanitizeRequest move |request: LlmRequest, context: LlmSanitizeRequestContext| { let py_fn = py_fn.clone(); let task_locals = capture_python_task_locals().or_else(|| task_locals.clone()); + let publication = + nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); Box::pin(async move { let result = resolve_py_object_or_future(Python::attach(|py| { - let result = py_fn - .call1( - py, - ( - PyLLMRequest { inner: request }, - PyLlmSanitizeRequestContext { inner: context }, - ), - ) - .map_err(|e| FlowError::Internal(e.to_string()))?; - split_py_object_or_future_with_locals(py, result, task_locals.as_ref()) + let args = ( + PyLLMRequest { inner: request }, + PyLlmSanitizeRequestContext { inner: context }, + ); + let result = if publication { + py.import("nemo_relay._event_sanitizer_context") + .and_then(|module| module.getattr("invoke")) + .and_then(|invoke| invoke.call1((py_fn.bind(py), args.0, args.1))) + } else { + py_fn.bind(py).call1(args) + } + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_py_object_or_future_with_locals(py, result.unbind(), task_locals.as_ref()) })) .await?; Python::attach(|py| { @@ -1077,15 +1102,21 @@ fn wrap_py_llm_sanitize_response_callback(py_fn: Py) -> LlmSanitizeRespon Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { let py_fn = py_fn.clone(); let task_locals = capture_python_task_locals().or_else(|| task_locals.clone()); + let publication = nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); Box::pin(async move { let result = resolve_py_object_or_future(Python::attach(|py| { let py_context = PyLlmSanitizeResponseContext { inner: context }; let py_response = json_to_py(py, &response) .map_err(|error| FlowError::Internal(error.to_string()))?; - let result = py_fn - .call1(py, (py_response, py_context)) - .map_err(|error| FlowError::Internal(error.to_string()))?; - split_py_object_or_future_with_locals(py, result, task_locals.as_ref()) + let result = if publication { + py.import("nemo_relay._event_sanitizer_context") + .and_then(|module| module.getattr("invoke")) + .and_then(|invoke| invoke.call1((py_fn.bind(py), py_response, py_context))) + } else { + py_fn.bind(py).call1((py_response, py_context)) + } + .map_err(|error| FlowError::Internal(error.to_string()))?; + split_py_object_or_future_with_locals(py, result.unbind(), task_locals.as_ref()) })) .await?; Python::attach(|py| { diff --git a/crates/python/tests/coverage/py_callable_coverage_tests.rs b/crates/python/tests/coverage/py_callable_coverage_tests.rs index ccb6cef65..d8ece09e7 100644 --- a/crates/python/tests/coverage/py_callable_coverage_tests.rs +++ b/crates/python/tests/coverage/py_callable_coverage_tests.rs @@ -18,7 +18,41 @@ fn load_module<'py>(py: Python<'py>, code: &str) -> Bound<'py, PyModule> { PyModule::from_code(py, &code, &file_name, &module_name).unwrap() } -fn install_event_sanitizer_context_module(py: Python<'_>) { +struct InstalledContextModule { + previous_parent: Option>, + previous_context: Option>, +} + +impl Drop for InstalledContextModule { + fn drop(&mut self) { + Python::attach(|py| { + let Ok(modules) = py.import("sys").and_then(|sys| sys.getattr("modules")) else { + return; + }; + let Ok(modules) = modules.cast_into::() else { + return; + }; + for (name, previous) in [ + ("nemo_relay", self.previous_parent.take()), + ( + "nemo_relay._event_sanitizer_context", + self.previous_context.take(), + ), + ] { + match previous { + Some(module) => { + let _ = modules.set_item(name, module); + } + None => { + let _ = modules.del_item(name); + } + } + } + }); + } +} + +fn install_event_sanitizer_context_module(py: Python<'_>) -> InstalledContextModule { let code = CString::new(include_str!( "../../../../python/nemo_relay/_event_sanitizer_context.py" )) @@ -40,10 +74,19 @@ fn install_event_sanitizer_context_module(py: Python<'_>) { .unwrap() .cast_into::() .unwrap(); + let previous_parent = modules.get_item("nemo_relay").unwrap().map(Bound::unbind); + let previous_context = modules + .get_item("nemo_relay._event_sanitizer_context") + .unwrap() + .map(Bound::unbind); modules.set_item("nemo_relay", parent).unwrap(); modules .set_item("nemo_relay._event_sanitizer_context", context) .unwrap(); + InstalledContextModule { + previous_parent, + previous_context, + } } fn make_request() -> LlmRequest { @@ -701,16 +744,23 @@ fn event_sanitize_wrapper_covers_conversion_success_and_error_propagation() { let _python = crate::test_support::init_python_test(); Python::attach(|py| { - install_event_sanitizer_context_module(py); + let _context_module = install_event_sanitizer_context_module(py); let module = load_module( py, r#" +import asyncio + def sanitize(event, fields): assert event.kind == "mark" fields["data"] = {"safe": event.name} fields["metadata"] = None return fields +async def async_sanitize(event, fields): + await asyncio.sleep(0) + fields["data"] = {"async_safe": event.name} + return fields + def raises(event, fields): raise RuntimeError("sanitize boom") @@ -738,6 +788,16 @@ def invalid(event, fields): assert_eq!(sanitized.data, Some(json!({"safe": "checkpoint"}))); assert_eq!(sanitized.metadata, None); + let async_sanitized = runtime + .block_on(wrap_py_event_sanitize_fn( + module.getattr("async_sanitize").unwrap().unbind(), + )(Arc::new(event.clone()), fields.clone())) + .unwrap(); + assert_eq!( + async_sanitized.data, + Some(json!({"async_safe": "checkpoint"})) + ); + let raised = runtime .block_on(wrap_py_event_sanitize_fn( module.getattr("raises").unwrap().unbind(), @@ -810,3 +870,31 @@ async def llm_fail(request): }); }); } + +#[test] +fn background_middleware_accepts_custom_awaitables() { + let _python = crate::test_support::init_python_test(); + let (_context_module, llm_custom) = Python::attach(|py| { + let context_module = install_event_sanitizer_context_module(py); + let module = load_module( + py, + r#" +class CustomAwaitable: + def __await__(self): + async def resolve(): + return None + return resolve().__await__() + +def llm_custom_awaitable(request): + return CustomAwaitable() +"#, + ); + ( + context_module, + wrap_py_llm_conditional_fn(module.getattr("llm_custom_awaitable").unwrap().unbind()), + ) + }); + + let runtime = tokio::runtime::Runtime::new().unwrap(); + assert_eq!(runtime.block_on(llm_custom(make_request())).unwrap(), None); +} diff --git a/python/nemo_relay/_event_sanitizer_context.py b/python/nemo_relay/_event_sanitizer_context.py index f2cc89a09..7991ebf4a 100644 --- a/python/nemo_relay/_event_sanitizer_context.py +++ b/python/nemo_relay/_event_sanitizer_context.py @@ -26,6 +26,11 @@ async def _await_result(result: Awaitable[Any]) -> Any: _ACTIVE.reset(token) +async def await_result(result: Awaitable[Any]) -> Any: + """Await an arbitrary awaitable without changing sanitizer context.""" + return await result + + def invoke(callback: Callable[..., Any], *args: Any) -> Any: """Invoke a sanitizer while marking its sync and async execution contexts.""" token = _ACTIVE.set(True) diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index b956d8f64..22bc1aa18 100644 --- a/python/tests/test_llm.py +++ b/python/tests/test_llm.py @@ -179,6 +179,41 @@ def sanitize_response(response, context): assert context.codec.kind == "none" assert context.codec.id is None + async def test_manual_async_sanitizers_can_flush_subscribers(self): + request_flushed = False + response_flushed = False + + async def sanitize_request(request, context): + nonlocal request_flushed + del context + await asyncio.sleep(0) + subscribers.flush() + request_flushed = True + return request + + async def sanitize_response(response, context): + nonlocal response_flushed + del context + await asyncio.sleep(0) + subscribers.flush() + response_flushed = True + return response + + guardrails.register_llm_sanitize_request("py_manual_flush_request", 1, sanitize_request) + guardrails.register_llm_sanitize_response("py_manual_flush_response", 1, sanitize_response) + subscribers.register("py_manual_flush_subscriber", lambda _event: None) + try: + handle = llm.call("py_manual_flush", make_request()) + llm.call_end(handle, {"response": "ok"}) + await asyncio.wait_for(asyncio.to_thread(subscribers.flush), timeout=2) + finally: + guardrails.deregister_llm_sanitize_request("py_manual_flush_request") + guardrails.deregister_llm_sanitize_response("py_manual_flush_response") + subscribers.deregister("py_manual_flush_subscriber") + + assert request_flushed + assert response_flushed + async def test_sanitizers_resolve_active_builtin_codecs(self): request_codec_used = False response_codec_used = False From 4f47567e4b07f75405f31f97ca62f87438826919 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 22:00:25 -0400 Subject: [PATCH 30/52] test(go): encode streaming fixtures as valid JSON Signed-off-by: Will Killian --- go/nemo_relay/llm/llm_shorthand_test.go | 4 ++-- go/nemo_relay/llm_test.go | 9 ++++++--- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/go/nemo_relay/llm/llm_shorthand_test.go b/go/nemo_relay/llm/llm_shorthand_test.go index 40e3604bc..c8a5aa640 100644 --- a/go/nemo_relay/llm/llm_shorthand_test.go +++ b/go/nemo_relay/llm/llm_shorthand_test.go @@ -6,7 +6,6 @@ package llm_test import ( "encoding/json" "io" - "strings" "testing" "github.com/NVIDIA/NeMo-Relay/go/nemo_relay" @@ -121,7 +120,8 @@ func TestLlmShorthands(t *testing.T) { stream, err := llmpkg.StreamExecute("llm_stream", makeRequest(), func(nativeJSON json.RawMessage) (json.RawMessage, error) { - return json.RawMessage(`"` + strings.ReplaceAll("data: {\"chunk\": 1}\n\ndata: [DONE]\n\n", `"`, `\"`) + `"`), nil + encoded, err := json.Marshal("data: {\"chunk\": 1}\n\ndata: [DONE]\n\n") + return json.RawMessage(encoded), err }, nil, nil, ) diff --git a/go/nemo_relay/llm_test.go b/go/nemo_relay/llm_test.go index bf2a60fe4..5aa1eb837 100644 --- a/go/nemo_relay/llm_test.go +++ b/go/nemo_relay/llm_test.go @@ -1098,7 +1098,8 @@ func TestLlmStreamCallExecuteBasic(t *testing.T) { chunks := `data: {"chunk": 1}` + "\n\n" + `data: {"chunk": 2}` + "\n\n" + `data: [DONE]` + "\n\n" - return json.RawMessage(`"` + strings.ReplaceAll(chunks, `"`, `\"`) + `"`), nil + encoded, err := json.Marshal(chunks) + return json.RawMessage(encoded), err }, nil, nil, ) @@ -1146,7 +1147,8 @@ func TestLlmStreamCallExecuteWithCollectorFinalizer(t *testing.T) { func(nativeJSON json.RawMessage) (json.RawMessage, error) { chunks := `data: {"token": "hello"}` + "\n\n" + `data: [DONE]` + "\n\n" - return json.RawMessage(`"` + strings.ReplaceAll(chunks, `"`, `\"`) + `"`), nil + encoded, err := json.Marshal(chunks) + return json.RawMessage(encoded), err }, collector, finalizer, ) @@ -1412,7 +1414,8 @@ func TestLlmStreamCloseFinalizesPartialResponse(t *testing.T) { chunks := `data: {"chunk": 1}` + "\n\n" + `data: {"chunk": 2}` + "\n\n" + `data: [DONE]` + "\n\n" - return json.RawMessage(`"` + strings.ReplaceAll(chunks, `"`, `\"`) + `"`), nil + encoded, err := json.Marshal(chunks) + return json.RawMessage(encoded), err }, nil, func() string { finalizerCalls++ From 2f1aef2df7d3d906a67283b7f6565f7046b59889 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 22:41:03 -0400 Subject: [PATCH 31/52] fix: prevent queued sanitizer flush deadlocks Signed-off-by: Will Killian --- crates/core/src/api/llm.rs | 6 +- .../src/api/runtime/subscriber_dispatcher.rs | 73 +++++++++++++++++-- crates/core/src/logging/rotation.rs | 4 - crates/core/src/stream.rs | 1 + .../tests/coverage/logging_rotation_tests.rs | 34 --------- .../core/tests/coverage/logging_sink_tests.rs | 67 +---------------- .../core/tests/integration/pipeline_tests.rs | 42 +++++++++++ crates/node/src/api/mod.rs | 19 +++-- crates/node/src/callable.rs | 15 ++-- crates/node/src/promise_call.rs | 6 ++ crates/node/tests/llm_tests.mjs | 34 +++++++++ crates/node/tests/tools_tests.mjs | 34 +++++++++ crates/python/src/py_callable.rs | 24 ++++-- docs/about-nemo-relay/concepts/middleware.mdx | 20 +++-- .../advanced-guide.mdx | 20 +++-- docs/reference/event-sanitizers.mdx | 8 +- docs/reference/migration-guides.mdx | 19 +++-- python/nemo_relay/subscribers.py | 6 +- python/tests/test_tools.py | 43 +++++++++++ 19 files changed, 315 insertions(+), 160 deletions(-) delete mode 100644 crates/core/tests/coverage/logging_rotation_tests.rs diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index b6a7ffa52..1472cf66d 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -888,11 +888,13 @@ async fn build_llm_end_payload( /// the handle start time if the current time is not later. /// /// # Returns -/// A [`Result`] that is `Ok(())` when the end event has been emitted. +/// A [`Result`] that is `Ok(())` when the end event has been queued for +/// sanitization and publication. /// /// # Errors /// Returns an error when the runtime owner check fails, internal state cannot be -/// read safely, or response codec decoding fails. +/// read safely, or the event cannot be queued. Sanitizer and response-codec +/// errors discovered during queued publication are logged and fail open. /// /// # Notes /// Sanitize-response guardrails affect only the emitted end-event payload, not diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 43412585b..651d57848 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -18,6 +18,7 @@ pub(crate) type EventTransformFn = Box< mod native { use std::cell::Cell; + use std::collections::VecDeque; use std::panic::{AssertUnwindSafe, catch_unwind}; use std::sync::OnceLock; use std::sync::atomic::{AtomicBool, Ordering}; @@ -55,6 +56,9 @@ mod native { thread_local! { static IN_DISPATCHER: Cell = const { Cell::new(false) }; } + tokio::task_local! { + static IN_ASYNC_PUBLICATION: (); + } struct DispatchGuard; @@ -177,7 +181,7 @@ mod native { } pub(super) fn flush_subscribers() -> Result<()> { - if IN_DISPATCHER.with(Cell::get) { + if in_dispatcher_callback() { return Ok(()); } let Some(sender_result) = DISPATCHER.get() else { @@ -199,7 +203,15 @@ mod native { } pub(super) fn in_dispatcher_callback() -> bool { - IN_DISPATCHER.with(Cell::get) + IN_DISPATCHER.with(Cell::get) || IN_ASYNC_PUBLICATION.try_with(|_| ()).is_ok() + } + + pub(super) async fn with_async_publication_context(future: F) -> F::Output { + if IN_ASYNC_PUBLICATION.try_with(|_| ()).is_ok() { + future.await + } else { + IN_ASYNC_PUBLICATION.scope((), future).await + } } fn dispatcher_sender() -> std::result::Result, String> { @@ -248,10 +260,18 @@ mod native { } fn run_dispatcher(rx: Receiver) { - while let Ok(message) = rx.recv() { + let mut pending = VecDeque::new(); + loop { + let message = match pending.pop_front() { + Some(message) => message, + None => match rx.recv() { + Ok(message) => message, + Err(_) => break, + }, + }; match message { DispatcherMessage::Flush { done } => { - let pending_flushes = drain_pending_messages(&rx); + let pending_flushes = drain_pending_messages(&rx, &mut pending); let _ = done.send(()); for pending in pending_flushes { let _ = pending.send(()); @@ -265,13 +285,17 @@ mod native { } } - fn drain_pending_messages(rx: &Receiver) -> Vec> { + fn drain_pending_messages( + rx: &Receiver, + pending: &mut VecDeque, + ) -> Vec> { let mut pending_flushes = Vec::new(); while let Ok(message) = rx.try_recv() { match message { DispatcherMessage::Flush { done } => pending_flushes.push(done), - DispatcherMessage::Barrier { done } => { - let _ = done.recv(); + message @ DispatcherMessage::Barrier { .. } => { + pending.push_back(message); + break; } message => handle_message(message), } @@ -384,6 +408,33 @@ mod native { } } } + + #[cfg(test)] + mod tests { + use super::*; + use std::sync::Mutex; + use std::time::Duration; + + static TEST_MUTEX: Mutex<()> = Mutex::new(()); + + #[test] + fn flush_does_not_wait_for_a_later_publication_barrier() { + let _lock = TEST_MUTEX.lock().unwrap(); + let first = register_async_publication().expect("first publication barrier"); + let sender = dispatcher_sender().expect("dispatcher sender"); + let (flush_tx, flush_rx) = mpsc::channel(); + sender + .send(DispatcherMessage::Flush { done: flush_tx }) + .unwrap(); + let later = register_async_publication().expect("later publication barrier"); + + first.send(()).unwrap(); + flush_rx + .recv_timeout(Duration::from_secs(1)) + .expect("flush queued before the later barrier must complete"); + later.send(()).unwrap(); + } + } } #[cfg(test)] @@ -429,6 +480,14 @@ pub(crate) fn register_async_publication() -> Option native::register_async_publication() } +/// Run asynchronous middleware as part of an already-registered publication. +/// +/// Re-entrant subscriber flushes are no-ops in this context because the +/// publication's FIFO barrier cannot complete until the middleware returns. +pub(crate) async fn with_async_publication_context(future: F) -> F::Output { + native::with_async_publication_context(future).await +} + /// Wait for all queued subscriber callbacks submitted before this call. pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() diff --git a/crates/core/src/logging/rotation.rs b/crates/core/src/logging/rotation.rs index fa55780bb..3314377c1 100644 --- a/crates/core/src/logging/rotation.rs +++ b/crates/core/src/logging/rotation.rs @@ -146,7 +146,3 @@ pub(crate) fn rotated_log_path(base_path: &Path, index: usize) -> PathBuf { } base_path.with_file_name(file_name) } - -#[cfg(test)] -#[path = "../../tests/coverage/logging_rotation_tests.rs"] -mod tests; diff --git a/crates/core/src/stream.rs b/crates/core/src/stream.rs index a7cd81761..feaaace1e 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -337,6 +337,7 @@ impl LlmStreamWrapper { let _ = done.send(()); } }; + let finalize = subscriber_dispatcher::with_async_publication_context(finalize); if background_thread { // `Drop` can run while the current-thread Tokio executor is // synchronously flushing subscribers. Use a dedicated runtime so diff --git a/crates/core/tests/coverage/logging_rotation_tests.rs b/crates/core/tests/coverage/logging_rotation_tests.rs deleted file mode 100644 index 842fa7927..000000000 --- a/crates/core/tests/coverage/logging_rotation_tests.rs +++ /dev/null @@ -1,34 +0,0 @@ -// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -// SPDX-License-Identifier: Apache-2.0 - -use super::*; - -#[test] -fn rotating_writer_rotates_retains_and_reports_closed_file_errors() { - let directory = tempfile::tempdir().unwrap(); - let path = directory.path().join("nested").join("relay.log"); - let mut writer = SizeRotatingFileWriter::new(path.clone(), 4, 2).unwrap(); - assert_eq!(writer.write(b"abcd").unwrap(), 4); - writer.flush().unwrap(); - assert_eq!(writer.write(b"e").unwrap(), 1); - writer.flush().unwrap(); - - assert_eq!(std::fs::read(rotated_log_path(&path, 1)).unwrap(), b"abcd"); - assert_eq!(std::fs::read(&path).unwrap(), b"e"); - - writer.file = None; - assert!(writer.write(b"x").is_err()); - assert!(writer.flush().is_err()); -} - -#[test] -fn rotation_helpers_handle_empty_and_missing_files() { - let directory = tempfile::tempdir().unwrap(); - let path = directory.path().join("missing.log"); - rotate_files(&path, 2).unwrap(); - assert_eq!( - rotated_log_path(&path, 2), - directory.path().join("missing.2.log") - ); - create_parent_directory(std::path::Path::new("plain.log")).unwrap(); -} diff --git a/crates/core/tests/coverage/logging_sink_tests.rs b/crates/core/tests/coverage/logging_sink_tests.rs index 88f142655..044ce2b47 100644 --- a/crates/core/tests/coverage/logging_sink_tests.rs +++ b/crates/core/tests/coverage/logging_sink_tests.rs @@ -2,16 +2,10 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - DROP_REPORT_INTERVAL_MILLIS, DropNoticeRateLimiter, build_logger, dropped_record_error_handler, - log_level_filter, now_millis, reserved_sink_paths, resolve_log_path, rotated_log_path, - spdlog_level, stderr_error_handler, + DROP_REPORT_INTERVAL_MILLIS, DropNoticeRateLimiter, dropped_record_error_handler, + log_level_filter, now_millis, spdlog_level, stderr_error_handler, }; use crate::logging::LogLevel; -use crate::logging::{ - FileLogRotationConfig, FileLogSinkConfig, LogFormat, LogSinkConfig, LoggingConfig, - MAX_FILE_SINK_QUEUE_ENTRIES, -}; -use std::path::PathBuf; #[test] fn drop_notice_rate_limiter_reports_immediately_then_once_per_interval() { @@ -39,60 +33,3 @@ fn sink_helpers_cover_boundary_levels_time_and_emergency_handlers() { "expected test error", ))); } - -#[test] -fn logger_builder_rejects_duplicate_conflicting_and_invalid_file_sinks() { - let directory = tempfile::tempdir().unwrap(); - let log_path = directory.path().join("relay.log"); - let file_sink = |path: PathBuf, rotation| { - LogSinkConfig::File(FileLogSinkConfig { - path, - level: LogLevel::Info, - format: LogFormat::Jsonl, - queue_capacity: 8, - rotation, - }) - }; - - assert!(resolve_log_path(std::path::Path::new("")).is_err()); - let rotation = FileLogRotationConfig::new(32, 1).unwrap(); - assert_eq!(reserved_sink_paths(&log_path, Some(rotation)).len(), 2); - - let duplicate = LoggingConfig { - sinks: vec![ - file_sink(log_path.clone(), None), - file_sink(log_path.clone(), None), - ], - ..LoggingConfig::default() - }; - let error = match build_logger(&duplicate, "root".into()) { - Ok(_) => panic!("duplicate file sinks must be rejected"), - Err(error) => error, - }; - assert!(error.to_string().contains("duplicate logging sink path")); - - let conflict = LoggingConfig { - sinks: vec![ - file_sink(log_path.clone(), Some(rotation)), - file_sink(rotated_log_path(&log_path, 1), None), - ], - ..LoggingConfig::default() - }; - let error = match build_logger(&conflict, "root".into()) { - Ok(_) => panic!("active and rotated file paths must not overlap"), - Err(error) => error, - }; - assert!(error.to_string().contains("conflicts")); - - let mut invalid_capacity = LoggingConfig { - sinks: vec![file_sink(log_path, None)], - ..LoggingConfig::default() - }; - let LogSinkConfig::File(file_sink) = &mut invalid_capacity.sinks[0]; - file_sink.queue_capacity = MAX_FILE_SINK_QUEUE_ENTRIES + 1; - let error = match build_logger(&invalid_capacity, "root".into()) { - Ok(_) => panic!("oversized async queues must be rejected"), - Err(error) => error, - }; - assert!(error.to_string().contains("queue_capacity")); -} diff --git a/crates/core/tests/integration/pipeline_tests.rs b/crates/core/tests/integration/pipeline_tests.rs index 9711974de..ca8303bef 100644 --- a/crates/core/tests/integration/pipeline_tests.rs +++ b/crates/core/tests/integration/pipeline_tests.rs @@ -1972,3 +1972,45 @@ async fn test_stream_response_codec_annotation_uses_sanitized_aggregated_respons deregister_subscriber("stream_sanitized_resp_codec_sub").unwrap(); deregister_llm_sanitize_response_guardrail("stream_sanitize_resp_codec_annotation").unwrap(); } + +#[tokio::test] +async fn test_stream_response_sanitizer_can_flush_subscribers() { + let _lock = TEST_MUTEX.lock().unwrap(); + reset_global(); + setup_isolated_thread(); + + register_subscriber("stream_reentrant_flush_subscriber", Arc::new(|_| {})).unwrap(); + register_llm_sanitize_response_guardrail( + "stream_reentrant_flush_sanitizer", + 1, + Arc::new(|response, _context| { + Box::pin(async move { + flush_subscribers()?; + Ok(Some(response)) + }) + }), + ) + .unwrap(); + + let mut stream = llm_stream_call_execute( + LlmStreamCallExecuteParams::builder() + .name("stream_reentrant_flush") + .request(make_openai_chat_request("stream me")) + .func(noop_stream_exec_fn()) + .collector(Box::new(|_chunk| Ok(()))) + .finalizer(Box::new(|| make_openai_chat_response("done"))) + .build(), + ) + .await + .unwrap(); + + while stream.next().await.is_some() {} + tokio::time::timeout(std::time::Duration::from_secs(2), stream.close()) + .await + .expect("stream close deadlocked in response sanitizer") + .unwrap(); + flush_subscribers().unwrap(); + + deregister_llm_sanitize_response_guardrail("stream_reentrant_flush_sanitizer").unwrap(); + deregister_subscriber("stream_reentrant_flush_subscriber").unwrap(); +} diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index 9c00ee254..40e9a9075 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -2918,8 +2918,8 @@ pub fn deregister_tool_execution_intercept(name: String) -> Result { /// /// The `guardrail` callback receives `(request, context)` and must return the sanitized request, /// or `null` to omit the observability payload. Lower `priority` values run first. Throws if a -/// guardrail with the same `name` already exists. If the callback throws, Relay omits the payload -/// and records the error for `getLastCallbackError()`. +/// guardrail with the same `name` already exists. If the callback throws, Relay preserves the last +/// valid payload, continues publication, and records the error for `getLastCallbackError()`. #[napi] pub fn register_llm_sanitize_request_guardrail( env: Env, @@ -2952,8 +2952,8 @@ pub fn deregister_llm_sanitize_request_guardrail(name: String) -> Result { /// /// The `guardrail` callback receives `(response, context)` and must return the sanitized response, /// or `null` to omit the observability payload. Lower `priority` values run first. Throws if a -/// guardrail with the same `name` already exists. If the callback throws, Relay omits the payload -/// and records the error for `getLastCallbackError()`. +/// guardrail with the same `name` already exists. If the callback throws, Relay preserves the last +/// valid payload, continues publication, and records the error for `getLastCallbackError()`. #[napi] pub fn register_llm_sanitize_response_guardrail( env: Env, @@ -3174,8 +3174,9 @@ pub fn deregister_subscriber(name: String) -> Result { /// Return a Promise that resolves when native subscriber callbacks queued /// before this call finish. /// -/// When called from an event-sanitizer callback, this Promise resolves without waiting to prevent -/// a cycle with the serial dispatcher. +/// When called from a queued publication sanitizer callback (including event and manual tool/LLM +/// sanitizers), this Promise resolves without waiting to prevent a cycle with the serial +/// dispatcher. /// /// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this /// Promise does not block the Node event loop while Promise-returning event sanitizers settle. @@ -3502,7 +3503,8 @@ pub fn scope_deregister_tool_execution_intercept(scope_uuid: String, name: Strin /// The `guardrail` callback receives `(request, context)` and must return the sanitized request, /// or `null` to omit the observability payload. Lower `priority` values run first. Throws if a /// guardrail with the same `name` already exists on the specified scope. If the callback throws, -/// Relay omits the payload and records the error for `getLastCallbackError()`. +/// Relay preserves the last valid payload, continues publication, and records the error for +/// `getLastCallbackError()`. #[napi] pub fn scope_register_llm_sanitize_request_guardrail( env: Env, @@ -3546,7 +3548,8 @@ pub fn scope_deregister_llm_sanitize_request_guardrail( /// The `guardrail` callback receives `(response, context)` and must return the sanitized response, /// or `null` to omit the observability payload. Lower `priority` values run first. Throws if a /// guardrail with the same `name` already exists on the specified scope. If the callback throws, -/// Relay omits the payload and records the error for `getLastCallbackError()`. +/// Relay preserves the last valid payload, continues publication, and records the error for +/// `getLastCallbackError()`. #[napi] pub fn scope_register_llm_sanitize_response_guardrail( env: Env, diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index 905f77e49..9d08f1ba0 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -284,12 +284,17 @@ pub fn wrap_js_tool_request_intercept_promise_fn(func: Arc) -> T pub fn wrap_js_tool_sanitize_promise_fn(func: Arc) -> ToolSanitizeFn { Arc::new(move |name: String, value: Json| { let func = func.clone(); + let publication = nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); Box::pin(async move { - func.call_spread(vec![Json::String(name), value]) - .await - .inspect_err(|error| { - record_callback_error(error.to_string()); - }) + let args = vec![Json::String(name), value]; + let result = if publication { + func.call_spread_for_publication(args).await + } else { + func.call_spread(args).await + }; + result.inspect_err(|error| { + record_callback_error(error.to_string()); + }) }) }) } diff --git a/crates/node/src/promise_call.rs b/crates/node/src/promise_call.rs index df9a7785a..40f780bc0 100644 --- a/crates/node/src/promise_call.rs +++ b/crates/node/src/promise_call.rs @@ -248,6 +248,12 @@ impl PromiseAwareFn { .await } + /// Call a spread callback from queued event publication. + pub async fn call_spread_for_publication(&self, args: Vec) -> FlowResult { + self.call_inner(PrimaryArg::Json(Json::Array(args)), true, None, true) + .await + } + /// Call the JS function with a builder-constructed first argument and await /// the result. /// diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index b54939a2d..c2bad301d 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -568,6 +568,40 @@ describe('LLM guardrails', () => { } }); + it('stream response sanitizers can flush subscribers without deadlocking', async () => { + let responseFlushed = false; + registerSubscriber('node_stream_flush_subscriber', () => {}); + registerLlmSanitizeResponseGuardrail('node_stream_flush_response', 10, async (response) => { + await flushSubscribers(); + responseFlushed = true; + return response; + }); + try { + const stream = await llmStreamCallExecute( + 'node_stream_flush', + makeNative(), + (wrapper) => { + lib.pushStreamChunk(wrapper.__nemo_relay_stream_id, { delta: 'ok' }); + lib.endStream(wrapper.__nemo_relay_stream_id); + }, + null, + () => ({ response: 'ok' }), + null, + null, + null, + null, + null, + ); + assert.deepEqual(await stream.next(), { delta: 'ok' }); + assert.equal(await stream.next(), null); + await flushSubscribers(); + } finally { + deregisterLlmSanitizeResponseGuardrail('node_stream_flush_response'); + deregisterSubscriber('node_stream_flush_subscriber'); + } + assert.equal(responseFlushed, true); + }); + it('releases custom stream codec references safely after early garbage collection', () => { const modulePath = path.join(nodeDir, 'index.js'); const script = ` diff --git a/crates/node/tests/tools_tests.mjs b/crates/node/tests/tools_tests.mjs index c9884bc05..4576d8253 100644 --- a/crates/node/tests/tools_tests.mjs +++ b/crates/node/tests/tools_tests.mjs @@ -709,6 +709,40 @@ describe('Tool guardrails', () => { } }); + it('manual async sanitizers can flush subscribers without deadlocking', async () => { + const events = []; + let requestFlushed = false; + let responseFlushed = false; + registerSubscriber('node_manual_tool_flush_subscriber', (event) => events.push(event)); + registerToolSanitizeRequestGuardrail('node_manual_tool_flush_request', 10, async (_name, args) => { + await flushSubscribers(); + requestFlushed = true; + return { ...args, requestSanitized: true }; + }); + registerToolSanitizeResponseGuardrail('node_manual_tool_flush_response', 10, async (_name, response) => { + await flushSubscribers(); + responseFlushed = true; + return { ...response, responseSanitized: true }; + }); + try { + const handle = toolCall('node_manual_tool_flush', { original: true }); + toolCallEnd(handle, { ok: true }); + await flushSubscribers(); + } finally { + deregisterToolSanitizeRequestGuardrail('node_manual_tool_flush_request'); + deregisterToolSanitizeResponseGuardrail('node_manual_tool_flush_response'); + deregisterSubscriber('node_manual_tool_flush_subscriber'); + } + assert.equal(requestFlushed, true); + assert.equal(responseFlushed, true); + const start = events.find( + (event) => event.name === 'node_manual_tool_flush' && event.scope_category === 'start', + ); + const end = events.find((event) => event.name === 'node_manual_tool_flush' && event.scope_category === 'end'); + assert.deepEqual(start.data, { original: true, requestSanitized: true }); + assert.deepEqual(end.data, { ok: true, responseSanitized: true }); + }); + it('conditional guardrail (block)', () => { registerToolConditionalExecutionGuardrail('node_tool_block', 10, (name, args) => 'blocked'); deregisterToolConditionalExecutionGuardrail('node_tool_block'); diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index ce23c5ec3..279040388 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -469,18 +469,30 @@ fn stream_from_async_iter(async_iter: Py) -> FlowResult { /// Wrap a Python callable `(str, Json) -> Json` for tool sanitize/intercept fns. pub fn wrap_py_tool_fn(py_fn: Py) -> ToolSanitizeFn { let py_fn = Arc::new(py_fn); + let task_locals = capture_python_task_locals(); Arc::new(move |name: String, args: Json| { let py_fn = py_fn.clone(); + let task_locals = capture_python_task_locals().or_else(|| task_locals.clone()); + let publication = nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); Box::pin(async move { - resolve_json_or_future(Python::attach(|py| { + let result = resolve_py_object_or_future(Python::attach(|py| { let py_args = json_to_py(py, &args) .map_err(|e| FlowError::Internal(format!("tool json_to_py failed: {e}")))?; - let result = py_fn.call1(py, (name, py_args)).map_err(|e| { - FlowError::Internal(format!("Python tool callback failed: {e}")) - })?; - split_json_or_future(py, result) + let result = if publication { + py.import("nemo_relay._event_sanitizer_context") + .and_then(|module| module.getattr("invoke")) + .and_then(|invoke| invoke.call1((py_fn.bind(py), name, py_args))) + } else { + py_fn.bind(py).call1((name, py_args)) + } + .map_err(|e| FlowError::Internal(format!("Python tool callback failed: {e}")))?; + split_py_object_or_future_with_locals(py, result.unbind(), task_locals.as_ref()) })) - .await + .await?; + Python::attach(|py| { + py_to_json(result.bind(py)) + .map_err(|e| FlowError::Internal(format!("tool py_to_json failed: {e}"))) + }) }) }) } diff --git a/docs/about-nemo-relay/concepts/middleware.mdx b/docs/about-nemo-relay/concepts/middleware.mdx index 1af705901..f74a3bb89 100644 --- a/docs/about-nemo-relay/concepts/middleware.mdx +++ b/docs/about-nemo-relay/concepts/middleware.mdx @@ -328,14 +328,18 @@ register_llm_sanitize_request_guardrail( "redact-openai-chat", 10, Arc::new(|mut request, context| { - if context.codec() == &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) - && let Some(codec) = context.resolve_codec() - && let Ok(mut annotated) = codec.decode(&request) - { - annotated.messages.clear(); - request = codec.encode(&annotated, &request).ok()?; - } - Some(request) + Box::pin(async move { + if context.codec() == &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + && let Some(codec) = context.resolve_codec() + && let Ok(mut annotated) = codec.decode(&request) + { + annotated.messages.clear(); + if let Ok(encoded) = codec.encode(&annotated, &request) { + request = encoded; + } + } + Ok(Some(request)) + }) }), )?; ``` diff --git a/docs/instrument-applications/advanced-guide.mdx b/docs/instrument-applications/advanced-guide.mdx index 87c41c380..38e05d264 100644 --- a/docs/instrument-applications/advanced-guide.mdx +++ b/docs/instrument-applications/advanced-guide.mdx @@ -124,12 +124,14 @@ register_tool_sanitize_request_guardrail( "search.redact_api_key", 10, Arc::new(|_tool_name, mut args| { - if let Some(object) = args.as_object_mut() { - if object.contains_key("api_key") { - object.insert("api_key".into(), json!("")); + Box::pin(async move { + if let Some(object) = args.as_object_mut() { + if object.contains_key("api_key") { + object.insert("api_key".into(), json!("")); + } } - } - args + Ok(args) + }) }), )?; @@ -137,9 +139,11 @@ register_tool_conditional_execution_guardrail( "search.require_query", 20, Arc::new(|_tool_name, args| { - Ok(match args.get("query").and_then(|value| value.as_str()) { - Some(query) if !query.is_empty() => None, - _ => Some("search.query is required".into()), + Box::pin(async move { + Ok(match args.get("query").and_then(|value| value.as_str()) { + Some(query) if !query.is_empty() => None, + _ => Some("search.query is required".into()), + }) }) }), )?; diff --git a/docs/reference/event-sanitizers.mdx b/docs/reference/event-sanitizers.mdx index 6122394b6..9ebeb5eae 100644 --- a/docs/reference/event-sanitizers.mdx +++ b/docs/reference/event-sanitizers.mdx @@ -135,9 +135,11 @@ register_mark_sanitize_guardrail( "safe-marks", 100, Arc::new(|event, mut fields| { - fields.data = Some(json!({"checkpoint": event.name()})); - fields.metadata = None; - fields + Box::pin(async move { + fields.data = Some(json!({"checkpoint": event.name()})); + fields.metadata = None; + Ok(fields) + }) }), )?; diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index e344cf738..4fc5529e0 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -72,11 +72,15 @@ Python and Node.js registration names are unchanged. Mark a Python callback `async def`, or return a Promise from Node.js, only when it needs asynchronous work; existing direct-value callbacks remain supported. -Scope lifecycle and mark emission remain synchronous. `push_scope`, -`pop_scope`, and mark APIs snapshot the event and visible sanitizer/subscriber -chain, then enqueue sanitization and publication on a serial dispatcher. Event -subscribers and exporters therefore receive sanitized events later, in emission -order. Do not add `await` to scope or mark emission calls. +Scope, mark, and manual tool/LLM lifecycle APIs remain synchronous. This +includes `push_scope`, `pop_scope`, mark APIs, `tool_call`, `tool_call_end`, +`llm_call`, and `llm_call_end`. These APIs snapshot the event and visible +sanitizer/subscriber chain, then enqueue sanitization and publication on a +serial dispatcher. Event subscribers and exporters therefore receive sanitized +events later, in emission order. Do not add `await` to these lifecycle calls. +Only enqueue-time validation and runtime-state errors are returned directly; +middleware and codec errors discovered during queued publication are logged and +handled according to their fail-open contracts. ### Update LLM Sanitizer Callbacks @@ -171,8 +175,9 @@ Check every callback for an implicit empty return. In particular: - A Python function that reaches the end without `return` omits the payload. - A JavaScript callback that returns `null` or `undefined` omits the payload. - A Rust callback must return `Some(payload)` to retain the payload. -- A sanitizer error reported through a plugin or binding boundary also omits - the payload and annotation. +- A sanitizer error or panic fails open: Relay preserves the last valid payload + and annotation snapshot, logs or records the callback error, and continues + publication. Use omission only when recording the payload would be unsafe. diff --git a/python/nemo_relay/subscribers.py b/python/nemo_relay/subscribers.py index d5a60cc7c..ebd6b2975 100644 --- a/python/nemo_relay/subscribers.py +++ b/python/nemo_relay/subscribers.py @@ -95,9 +95,9 @@ def flush() -> None: waiting for observer work. Use this barrier in tests and shutdown paths when captured subscriber output must be complete before continuing. - Call this function outside subscriber and event-sanitizer callbacks. A - re-entrant call returns without waiting to avoid blocking the dispatcher, - so callbacks later in the same dispatch snapshot can still run. + Call this function outside subscriber and queued publication sanitizer + callbacks. A re-entrant call returns without waiting to avoid blocking the + dispatcher, so callbacks later in the same dispatch snapshot can still run. """ if _event_sanitizer_callback_active(): return None diff --git a/python/tests/test_tools.py b/python/tests/test_tools.py index a401eb73b..f6d74c084 100644 --- a/python/tests/test_tools.py +++ b/python/tests/test_tools.py @@ -3,6 +3,7 @@ """Tests for NeMo Relay tool lifecycle, guardrails, and intercepts.""" +import asyncio from collections import UserDict, UserList from dataclasses import dataclass from typing import cast @@ -315,6 +316,48 @@ def test_deregister_nonexistent(self): class TestToolGuardrailsAsync: + async def test_manual_async_sanitizers_publish_transformed_payloads_and_can_flush(self): + events = [] + request_flushed = False + response_flushed = False + + async def sanitize_request(name, args): + nonlocal request_flushed + await asyncio.sleep(0) + subscribers.flush() + request_flushed = True + return {**args, "request_sanitized": True} + + async def sanitize_response(name, response): + nonlocal response_flushed + await asyncio.sleep(0) + subscribers.flush() + response_flushed = True + return {**response, "response_sanitized": True} + + subscribers.register("py_manual_tool_flush_subscriber", events.append) + guardrails.register_tool_sanitize_request("py_manual_tool_flush_request", 1, sanitize_request) + guardrails.register_tool_sanitize_response("py_manual_tool_flush_response", 1, sanitize_response) + try: + handle = tools.call("py_manual_tool_flush", {"original": True}) + tools.call_end(handle, {"ok": True}) + await asyncio.wait_for(asyncio.to_thread(subscribers.flush), timeout=2) + finally: + guardrails.deregister_tool_sanitize_request("py_manual_tool_flush_request") + guardrails.deregister_tool_sanitize_response("py_manual_tool_flush_response") + subscribers.deregister("py_manual_tool_flush_subscriber") + + assert request_flushed + assert response_flushed + assert _tool_event(events, "py_manual_tool_flush", "start").data == { + "original": True, + "request_sanitized": True, + } + assert _tool_event(events, "py_manual_tool_flush", "end").data == { + "ok": True, + "response_sanitized": True, + } + async def test_conditional_blocks_execution(self): guardrails.register_tool_conditional_execution("py_async_blocker", 1, lambda name, args: "blocked by policy") From b7c5e1c9cfaeefd1b18242224aeb61499e285c55 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 23:11:17 -0400 Subject: [PATCH 32/52] fix: make publication barriers flush-safe Signed-off-by: Will Killian --- crates/core/src/api/llm.rs | 13 +++-- .../src/api/runtime/subscriber_dispatcher.rs | 49 +++++++++++++++---- crates/core/src/api/subscriber.rs | 7 +-- crates/core/src/api/tool.rs | 20 +++++--- .../core/tests/integration/pipeline_tests.rs | 16 +++--- crates/node/src/api/mod.rs | 7 +-- crates/python/src/py_api/mod.rs | 5 +- python/nemo_relay/subscribers.py | 6 +-- 8 files changed, 83 insertions(+), 40 deletions(-) diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index 1472cf66d..dc2e68846 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -691,11 +691,13 @@ fn emit_optimization_marks_with( /// the emitted start event. When `None`, the current UTC time is used. /// /// # Returns -/// A [`Result`] containing the created [`LlmHandle`]. +/// A [`Result`] containing the created [`LlmHandle`] after its start-event +/// snapshot has been submitted for queued publication. /// /// # Errors /// Returns an error when the runtime owner check fails or when internal state -/// cannot be read safely. +/// cannot be read safely. Dispatcher submission failures are logged because +/// observability publication is best effort. /// /// # Notes /// The runtime removes standard credential headers (`authorization`, @@ -892,9 +894,10 @@ async fn build_llm_end_payload( /// sanitization and publication. /// /// # Errors -/// Returns an error when the runtime owner check fails, internal state cannot be -/// read safely, or the event cannot be queued. Sanitizer and response-codec -/// errors discovered during queued publication are logged and fail open. +/// Returns an error when the runtime owner check fails or internal state cannot +/// be read safely. Dispatcher submission failures are logged because +/// observability publication is best effort. Sanitizer and response-codec errors +/// discovered during queued publication are also logged and fail open. /// /// # Notes /// Sanitize-response guardrails affect only the emitted end-event payload, not diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 651d57848..8e16e927a 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -23,6 +23,7 @@ mod native { use std::sync::OnceLock; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::mpsc::{self, Receiver, Sender}; + use std::time::Duration; use super::*; use crate::api::runtime::scope_stack::{ @@ -278,13 +279,37 @@ mod native { } } DispatcherMessage::Barrier { done } => { - let _ = done.recv(); + wait_for_barrier(done, &rx, &mut pending); } message => handle_message(message), } } } + /// Preserve FIFO delivery behind an asynchronous publication boundary while + /// allowing flush requests to return. A flush cannot wait for the current + /// publication without risking a cycle when middleware spawned the caller. + fn wait_for_barrier( + done: Receiver<()>, + rx: &Receiver, + pending: &mut VecDeque, + ) { + loop { + match done.recv_timeout(Duration::from_millis(10)) { + Ok(()) | Err(mpsc::RecvTimeoutError::Disconnected) => return, + Err(mpsc::RecvTimeoutError::Timeout) => {} + } + while let Ok(message) = rx.try_recv() { + match message { + DispatcherMessage::Flush { done } => { + let _ = done.send(()); + } + message => pending.push_back(message), + } + } + } + } + fn drain_pending_messages( rx: &Receiver, pending: &mut VecDeque, @@ -412,14 +437,13 @@ mod native { #[cfg(test)] mod tests { use super::*; - use std::sync::Mutex; - use std::time::Duration; - - static TEST_MUTEX: Mutex<()> = Mutex::new(()); #[test] - fn flush_does_not_wait_for_a_later_publication_barrier() { - let _lock = TEST_MUTEX.lock().unwrap(); + fn flush_does_not_wait_for_active_or_later_publication_barriers() { + let _lock = crate::shared_runtime::runtime_owner_test_mutex() + .lock() + .unwrap_or_else(|error| error.into_inner()); + flush_subscribers().unwrap(); let first = register_async_publication().expect("first publication barrier"); let sender = dispatcher_sender().expect("dispatcher sender"); let (flush_tx, flush_rx) = mpsc::channel(); @@ -428,11 +452,12 @@ mod native { .unwrap(); let later = register_async_publication().expect("later publication barrier"); - first.send(()).unwrap(); flush_rx .recv_timeout(Duration::from_secs(1)) - .expect("flush queued before the later barrier must complete"); + .expect("flush must not wait for an active publication barrier"); + first.send(()).unwrap(); later.send(()).unwrap(); + flush_subscribers().unwrap(); } } } @@ -488,7 +513,11 @@ pub(crate) async fn with_async_publication_context(future: F) -> F::O native::with_async_publication_context(future).await } -/// Wait for all queued subscriber callbacks submitted before this call. +/// Wait for queued subscriber callbacks submitted before this call. +/// +/// If an asynchronous publication boundary is still active, this returns +/// without waiting for that publication or work queued behind it. This avoids +/// a cycle when publication middleware spawns or offloads the caller. pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() } diff --git a/crates/core/src/api/subscriber.rs b/crates/core/src/api/subscriber.rs index 02ce9a322..1c8cf2b95 100644 --- a/crates/core/src/api/subscriber.rs +++ b/crates/core/src/api/subscriber.rs @@ -72,9 +72,10 @@ pub fn deregister_subscriber(name: &str) -> Result { /// Wait for all subscriber callbacks queued before this call to finish. /// -/// Call this helper outside native subscriber callbacks. A re-entrant call returns without -/// waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can -/// still run. +/// A re-entrant call returns without waiting. The same applies while an +/// asynchronous publication boundary is active, including calls spawned or +/// offloaded by publication middleware. Call again after that middleware +/// settles to wait for its event and work queued behind it. /// /// Native targets deliver subscriber callbacks on a background dispatcher so /// event-producing APIs do not wait for observer work. Call this helper from diff --git a/crates/core/src/api/tool.rs b/crates/core/src/api/tool.rs index 30a9249e1..14b974651 100644 --- a/crates/core/src/api/tool.rs +++ b/crates/core/src/api/tool.rs @@ -202,8 +202,8 @@ pub struct ToolCallEndParams<'a> { /// Start a manual tool lifecycle span. /// -/// This emits a tool-start event after applying sanitize-request guardrails to -/// the payload recorded for observability. +/// This submits a tool-start event for queued sanitize-request guardrails and +/// publication without waiting for that work. /// /// # Parameters /// - `name`: Tool name recorded on the emitted lifecycle event. @@ -218,11 +218,13 @@ pub struct ToolCallEndParams<'a> { /// the emitted start event. When `None`, the current UTC time is used. /// /// # Returns -/// A [`Result`] containing the created [`ToolHandle`]. +/// A [`Result`] containing the created [`ToolHandle`] after its start-event +/// snapshot has been submitted for queued publication. /// /// # Errors /// Returns an error when the runtime owner check fails or when internal state -/// cannot be read safely. +/// cannot be read safely. Dispatcher submission failures are logged because +/// observability publication is best effort. /// /// # Notes /// Sanitize-request guardrails affect only the emitted start-event payload, not @@ -392,8 +394,8 @@ async fn tool_call_with_subscriber_snapshot( /// Finish a manual tool lifecycle span. /// -/// This emits a tool-end event for a handle previously returned by -/// [`tool_call`]. +/// This submits a tool-end event for queued sanitization and publication for a +/// handle previously returned by [`tool_call`]. /// /// # Parameters /// - `handle`: Tool handle to close. @@ -407,11 +409,13 @@ async fn tool_call_with_subscriber_snapshot( /// the handle start time if the current time is not later. /// /// # Returns -/// A [`Result`] that is `Ok(())` when the end event has been emitted. +/// A [`Result`] that is `Ok(())` when the end-event snapshot has been submitted +/// for queued publication. /// /// # Errors /// Returns an error when the runtime owner check fails or when internal state -/// cannot be read safely. +/// cannot be read safely. Dispatcher submission failures are logged because +/// observability publication is best effort. /// /// # Notes /// Sanitize-response guardrails affect only the emitted end-event payload, not diff --git a/crates/core/tests/integration/pipeline_tests.rs b/crates/core/tests/integration/pipeline_tests.rs index ca8303bef..bff17d6a7 100644 --- a/crates/core/tests/integration/pipeline_tests.rs +++ b/crates/core/tests/integration/pipeline_tests.rs @@ -1985,7 +1985,9 @@ async fn test_stream_response_sanitizer_can_flush_subscribers() { 1, Arc::new(|response, _context| { Box::pin(async move { - flush_subscribers()?; + tokio::task::spawn_blocking(flush_subscribers) + .await + .map_err(|error| FlowError::Internal(error.to_string()))??; Ok(Some(response)) }) }), @@ -2004,11 +2006,13 @@ async fn test_stream_response_sanitizer_can_flush_subscribers() { .await .unwrap(); - while stream.next().await.is_some() {} - tokio::time::timeout(std::time::Duration::from_secs(2), stream.close()) - .await - .expect("stream close deadlocked in response sanitizer") - .unwrap(); + tokio::time::timeout(std::time::Duration::from_secs(2), async { + while stream.next().await.is_some() {} + stream.close().await + }) + .await + .expect("stream finalization deadlocked in response sanitizer") + .unwrap(); flush_subscribers().unwrap(); deregister_llm_sanitize_response_guardrail("stream_reentrant_flush_sanitizer").unwrap(); diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index 40e9a9075..c1e6ecdd3 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -3174,9 +3174,10 @@ pub fn deregister_subscriber(name: String) -> Result { /// Return a Promise that resolves when native subscriber callbacks queued /// before this call finish. /// -/// When called from a queued publication sanitizer callback (including event and manual tool/LLM -/// sanitizers), this Promise resolves without waiting to prevent a cycle with the serial -/// dispatcher. +/// When called from queued publication middleware, or while an asynchronous +/// publication boundary is active, this Promise resolves without waiting. +/// Call it again after the middleware settles to wait for its event and later +/// work. /// /// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this /// Promise does not block the Node event loop while Promise-returning event sanitizers settle. diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index 7bb90a042..0e00bab51 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1524,8 +1524,9 @@ fn deregister_subscriber(name: &str) -> PyResult { /// Wait for subscriber callbacks queued before this call to finish. /// -/// Public Python wrappers prevent re-entrant event-sanitizer callbacks from waiting on the serial -/// dispatcher. +/// Re-entrant calls and calls observed while an asynchronous publication +/// boundary is active return without waiting. Call again after middleware +/// settles to wait for its event and later work. #[pyfunction] fn flush_subscribers(py: Python<'_>) -> PyResult<()> { py.detach(core_subscriber_api::flush_subscribers) diff --git a/python/nemo_relay/subscribers.py b/python/nemo_relay/subscribers.py index ebd6b2975..1137fc829 100644 --- a/python/nemo_relay/subscribers.py +++ b/python/nemo_relay/subscribers.py @@ -95,9 +95,9 @@ def flush() -> None: waiting for observer work. Use this barrier in tests and shutdown paths when captured subscriber output must be complete before continuing. - Call this function outside subscriber and queued publication sanitizer - callbacks. A re-entrant call returns without waiting to avoid blocking the - dispatcher, so callbacks later in the same dispatch snapshot can still run. + A re-entrant call, or a call observed while an asynchronous publication + boundary is active, returns without waiting. Call ``flush()`` again after + that middleware settles to wait for its event and later work. """ if _event_sanitizer_callback_active(): return None From c21131ee9d01ac70e735f3a5bf07986854d7e5b7 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 23:29:52 -0400 Subject: [PATCH 33/52] fix: preserve async publication flush ordering Signed-off-by: Will Killian --- crates/core/src/api/llm.rs | 22 +- .../src/api/runtime/subscriber_dispatcher.rs | 251 ++++++++++++------ crates/core/src/api/subscriber.rs | 8 +- crates/core/src/stream.rs | 12 +- .../core/tests/integration/pipeline_tests.rs | 123 ++++++++- crates/node/src/api/mod.rs | 8 +- crates/python/src/py_api/mod.rs | 6 +- python/nemo_relay/subscribers.py | 7 +- 8 files changed, 327 insertions(+), 110 deletions(-) diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index dc2e68846..cfd107173 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -21,7 +21,7 @@ use crate::api::runtime::LlmCodecIdentity; use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::global_context; use crate::api::runtime::subscriber_dispatcher::{ - dispatch_sanitized_event, dispatch_transformed_event, + dispatch_reserved_sanitized_event, dispatch_sanitized_event, dispatch_transformed_event, }; use crate::api::runtime::{ EventSubscriberFn, LlmCollectorFn, LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, @@ -547,6 +547,26 @@ pub(crate) async fn emit_optimization_marks(handle: &LlmHandle, subscribers: &[E .await; } +pub(crate) async fn emit_reserved_optimization_marks( + handle: &LlmHandle, + subscribers: &[EventSubscriberFn], +) { + emit_optimization_marks_with_async( + handle, + subscribers, + |event| sanitize_event_with_scope_stack(event, handle.captured_scope_stack()), + |event, subscribers| { + dispatch_reserved_sanitized_event( + event.clone(), + Vec::new(), + subscribers, + handle.captured_scope_stack().clone(), + ) + }, + ) + .await; +} + /// Queue optimization marks from a synchronous lifecycle API. /// /// The public manual lifecycle APIs must not await middleware. Capture each diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index 8e16e927a..fc77a9bcc 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -17,13 +17,12 @@ pub(crate) type EventTransformFn = Box< >; mod native { - use std::cell::Cell; + use std::cell::{Cell, RefCell}; use std::collections::VecDeque; use std::panic::{AssertUnwindSafe, catch_unwind}; use std::sync::OnceLock; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::mpsc::{self, Receiver, Sender}; - use std::time::Duration; use super::*; use crate::api::runtime::scope_stack::{ @@ -44,7 +43,7 @@ mod native { done: Sender<()>, }, Barrier { - done: Receiver<()>, + publications: Receiver>, }, } @@ -58,11 +57,15 @@ mod native { static IN_DISPATCHER: Cell = const { Cell::new(false) }; } tokio::task_local! { - static IN_ASYNC_PUBLICATION: (); + static ASYNC_PUBLICATION_MESSAGES: RefCell>>; } struct DispatchGuard; + pub(crate) struct AsyncPublication { + sender: Sender>, + } + impl DispatchGuard { fn enter() -> Self { IN_DISPATCHER.with(|flag| flag.set(true)); @@ -106,31 +109,7 @@ mod native { subscribers: subscribers.to_vec(), scope_stack: current_scope_stack(), }; - match dispatcher_sender() { - Ok(sender) => { - if sender.send(message).is_err() { - log::warn!( - target: "nemo_relay.runtime", - event = "subscriber_event_dropped", - reason = "dispatcher_disconnected"; - "Subscriber event was dropped because the dispatcher stopped" - ); - false - } else { - true - } - } - Err(_error) if !DISPATCHER_FAILURE_LOGGED.swap(true, Ordering::AcqRel) => { - log::error!( - target: "nemo_relay.runtime", - event = "subscriber_dispatcher_failed", - error_kind = "initialization"; - "Subscriber dispatcher failed to start" - ); - false - } - Err(_) => false, - } + send_dispatch_message(message) } pub(super) fn dispatch_sanitized_event( @@ -152,6 +131,39 @@ mod native { send_dispatch_message(message) } + pub(super) fn dispatch_reserved_sanitized_event( + event: Event, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, + ) -> bool { + if subscribers.is_empty() { + return true; + } + let message = DispatcherMessage::Deliver { + event: Box::new(event), + transform: None, + sanitizers, + subscribers: subscribers.to_vec(), + scope_stack, + }; + let buffer_active = ASYNC_PUBLICATION_MESSAGES + .try_with(|messages| messages.borrow().is_some()) + .unwrap_or(false); + if buffer_active { + ASYNC_PUBLICATION_MESSAGES.with(|messages| { + messages + .borrow_mut() + .as_mut() + .expect("publication buffer checked above") + .push(message); + }); + true + } else { + send_dispatch_message(message) + } + } + pub(super) fn dispatch_transformed_event( event: Event, transform: EventTransformFn, @@ -169,16 +181,20 @@ mod native { send_dispatch_message(message) } - /// Insert a FIFO barrier for work that will enqueue a publication from an - /// async task. A later flush waits for the task to signal completion, then - /// drains the event it queued before acknowledging the flush. - pub(super) fn register_async_publication() -> Option> { + /// Reserve a FIFO position for publications produced by an async task. + /// A later flush waits for the task and drains its buffered publications + /// at the reserved position before acknowledging the flush. + pub(super) fn register_async_publication() -> Option { let sender = dispatcher_sender().ok()?; - let (done_tx, done_rx) = mpsc::channel(); + let (publication_tx, publication_rx) = mpsc::channel(); sender - .send(DispatcherMessage::Barrier { done: done_rx }) + .send(DispatcherMessage::Barrier { + publications: publication_rx, + }) .ok() - .map(|_| done_tx) + .map(|_| AsyncPublication { + sender: publication_tx, + }) } pub(super) fn flush_subscribers() -> Result<()> { @@ -204,14 +220,31 @@ mod native { } pub(super) fn in_dispatcher_callback() -> bool { - IN_DISPATCHER.with(Cell::get) || IN_ASYNC_PUBLICATION.try_with(|_| ()).is_ok() + IN_DISPATCHER.with(Cell::get) || ASYNC_PUBLICATION_MESSAGES.try_with(|_| ()).is_ok() } - pub(super) async fn with_async_publication_context(future: F) -> F::Output { - if IN_ASYNC_PUBLICATION.try_with(|_| ()).is_ok() { + pub(super) async fn with_async_publication_context( + publication: Option, + future: F, + ) -> F::Output { + if ASYNC_PUBLICATION_MESSAGES.try_with(|_| ()).is_ok() { future.await } else { - IN_ASYNC_PUBLICATION.scope((), future).await + let (output, publications) = ASYNC_PUBLICATION_MESSAGES + .scope( + RefCell::new(publication.as_ref().map(|_| Vec::new())), + async { + let output = future.await; + let publications = ASYNC_PUBLICATION_MESSAGES + .with(|messages| messages.borrow_mut().take()); + (output, publications) + }, + ) + .await; + if let (Some(publication), Some(publications)) = (publication, publications) { + let _ = publication.sender.send(publications); + } + output } } @@ -278,34 +311,14 @@ mod native { let _ = pending.send(()); } } - DispatcherMessage::Barrier { done } => { - wait_for_barrier(done, &rx, &mut pending); - } - message => handle_message(message), - } - } - } - - /// Preserve FIFO delivery behind an asynchronous publication boundary while - /// allowing flush requests to return. A flush cannot wait for the current - /// publication without risking a cycle when middleware spawned the caller. - fn wait_for_barrier( - done: Receiver<()>, - rx: &Receiver, - pending: &mut VecDeque, - ) { - loop { - match done.recv_timeout(Duration::from_millis(10)) { - Ok(()) | Err(mpsc::RecvTimeoutError::Disconnected) => return, - Err(mpsc::RecvTimeoutError::Timeout) => {} - } - while let Ok(message) = rx.try_recv() { - match message { - DispatcherMessage::Flush { done } => { - let _ = done.send(()); + DispatcherMessage::Barrier { publications } => { + if let Ok(publications) = publications.recv() { + for publication in publications { + handle_message(publication); + } } - message => pending.push_back(message), } + message => handle_message(message), } } } @@ -340,8 +353,12 @@ mod native { DispatcherMessage::Flush { done } => { let _ = done.send(()); } - DispatcherMessage::Barrier { done } => { - let _ = done.recv(); + DispatcherMessage::Barrier { publications } => { + if let Ok(publications) = publications.recv() { + for publication in publications { + handle_message(publication); + } + } } } } @@ -439,24 +456,79 @@ mod native { use super::*; #[test] - fn flush_does_not_wait_for_active_or_later_publication_barriers() { + fn flush_waits_for_active_but_not_later_publication_barriers() { let _lock = crate::shared_runtime::runtime_owner_test_mutex() .lock() .unwrap_or_else(|error| error.into_inner()); flush_subscribers().unwrap(); let first = register_async_publication().expect("first publication barrier"); let sender = dispatcher_sender().expect("dispatcher sender"); + let delivered = std::sync::Arc::new(std::sync::Mutex::new(Vec::new())); + let subscriber: EventSubscriberFn = { + let delivered = delivered.clone(); + std::sync::Arc::new(move |event| { + delivered + .lock() + .unwrap_or_else(|error| error.into_inner()) + .push(event.name().to_string()); + }) + }; + let queued_event = serde_json::from_value(serde_json::json!({ + "kind": "mark", + "atof_version": "0.1", + "uuid": "019c1df6-4a57-7000-8000-000000000001", + "timestamp": "2026-07-28T00:00:00Z", + "name": "queued-before-flush" + })) + .expect("valid event"); + sender + .send(DispatcherMessage::Deliver { + event: Box::new(queued_event), + transform: None, + sanitizers: Vec::new(), + subscribers: vec![subscriber.clone()], + scope_stack: current_scope_stack(), + }) + .unwrap(); let (flush_tx, flush_rx) = mpsc::channel(); sender .send(DispatcherMessage::Flush { done: flush_tx }) .unwrap(); let later = register_async_publication().expect("later publication barrier"); + assert!( + flush_rx + .recv_timeout(std::time::Duration::from_millis(50)) + .is_err(), + "flush must wait for an active publication barrier" + ); + let deferred_event = serde_json::from_value(serde_json::json!({ + "kind": "mark", + "atof_version": "0.1", + "uuid": "019c1df6-4a57-7000-8000-000000000002", + "timestamp": "2026-07-28T00:00:00Z", + "name": "deferred-at-barrier" + })) + .expect("valid event"); + first + .sender + .send(vec![DispatcherMessage::Deliver { + event: Box::new(deferred_event), + transform: None, + sanitizers: Vec::new(), + subscribers: vec![subscriber], + scope_stack: current_scope_stack(), + }]) + .unwrap(); flush_rx - .recv_timeout(Duration::from_secs(1)) - .expect("flush must not wait for an active publication barrier"); - first.send(()).unwrap(); - later.send(()).unwrap(); + .recv_timeout(std::time::Duration::from_secs(1)) + .expect("flush queued before the later barrier must complete"); + assert_eq!( + *delivered.lock().unwrap_or_else(|error| error.into_inner()), + ["deferred-at-barrier", "queued-before-flush"], + "the barrier must publish deferred work at its reserved FIFO position" + ); + later.sender.send(Vec::new()).unwrap(); flush_subscribers().unwrap(); } } @@ -485,6 +557,16 @@ pub(crate) fn dispatch_sanitized_event( native::dispatch_sanitized_event(event, sanitizers, subscribers, scope_stack) } +/// Publish a stream-finalization event at its reserved FIFO position. +pub(crate) fn dispatch_reserved_sanitized_event( + event: Event, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, +) -> bool { + native::dispatch_reserved_sanitized_event(event, sanitizers, subscribers, scope_stack) +} + /// Queue a snapshot for a middleware-specific asynchronous transformation, /// followed by event sanitization and subscriber delivery. pub(crate) fn dispatch_transformed_event( @@ -499,25 +581,26 @@ pub(crate) fn dispatch_transformed_event( /// Register a FIFO barrier for async work that will queue a subscriber event. /// -/// Dropping the returned sender releases the barrier, so error paths cannot -/// leave the dispatcher blocked. -pub(crate) fn register_async_publication() -> Option> { +/// Dropping the returned publication handle releases the barrier, so error +/// paths cannot leave the dispatcher blocked. +pub(crate) fn register_async_publication() -> Option { native::register_async_publication() } -/// Run asynchronous middleware as part of an already-registered publication. +/// Run asynchronous middleware as part of an already-registered publication, +/// buffering the finalization publications explicitly assigned to its reserved +/// FIFO position. /// /// Re-entrant subscriber flushes are no-ops in this context because the /// publication's FIFO barrier cannot complete until the middleware returns. -pub(crate) async fn with_async_publication_context(future: F) -> F::Output { - native::with_async_publication_context(future).await +pub(crate) async fn with_async_publication_context( + publication: Option, + future: F, +) -> F::Output { + native::with_async_publication_context(publication, future).await } -/// Wait for queued subscriber callbacks submitted before this call. -/// -/// If an asynchronous publication boundary is still active, this returns -/// without waiting for that publication or work queued behind it. This avoids -/// a cycle when publication middleware spawns or offloads the caller. +/// Wait for all queued subscriber callbacks submitted before this call. pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() } diff --git a/crates/core/src/api/subscriber.rs b/crates/core/src/api/subscriber.rs index 1c8cf2b95..337544167 100644 --- a/crates/core/src/api/subscriber.rs +++ b/crates/core/src/api/subscriber.rs @@ -72,10 +72,10 @@ pub fn deregister_subscriber(name: &str) -> Result { /// Wait for all subscriber callbacks queued before this call to finish. /// -/// A re-entrant call returns without waiting. The same applies while an -/// asynchronous publication boundary is active, including calls spawned or -/// offloaded by publication middleware. Call again after that middleware -/// settles to wait for its event and work queued behind it. +/// A direct re-entrant call from queued publication middleware returns without +/// waiting. Publication middleware must not move such a flush into +/// `tokio::spawn`, `tokio::task::spawn_blocking`, or another unmarked task or +/// thread because the publication cannot complete while awaiting that flush. /// /// Native targets deliver subscriber callbacks on a background dispatcher so /// event-producing APIs do not wait for observer work. Call this helper from diff --git a/crates/core/src/stream.rs b/crates/core/src/stream.rs index feaaace1e..a1c02179b 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -34,7 +34,7 @@ use tokio_stream::Stream; use crate::api::event::{BaseEvent, MarkEvent}; use crate::api::llm::LlmHandle; -use crate::api::llm::emit_optimization_marks; +use crate::api::llm::emit_reserved_optimization_marks; use crate::api::optimization::finalize_optimization_summary; use crate::api::runtime::LlmSanitizeResponseContext; use crate::api::runtime::NemoRelayContextState; @@ -295,7 +295,7 @@ impl LlmStreamWrapper { handle .optimization_recorder .close_for_finalization(interruption); - emit_optimization_marks(&handle, &subscribers).await; + emit_reserved_optimization_marks(&handle, &subscribers).await; let pricing = crate::codec::response::active_pricing_resolver(); let summary = finalize_optimization_summary( &handle.optimization_recorder, @@ -326,18 +326,16 @@ impl LlmStreamWrapper { if let Some(event) = event_snapshot && let Some(event) = sanitize_event_with_scope_stack(event, &scope_stack).await { - let _ = subscriber_dispatcher::dispatch_sanitized_event( + let _ = subscriber_dispatcher::dispatch_reserved_sanitized_event( event, Vec::new(), &subscribers, scope_stack.clone(), ); } - if let Some(done) = publication_barrier { - let _ = done.send(()); - } }; - let finalize = subscriber_dispatcher::with_async_publication_context(finalize); + let finalize = + subscriber_dispatcher::with_async_publication_context(publication_barrier, finalize); if background_thread { // `Drop` can run while the current-thread Tokio executor is // synchronously flushing subscribers. Use a dedicated runtime so diff --git a/crates/core/tests/integration/pipeline_tests.rs b/crates/core/tests/integration/pipeline_tests.rs index bff17d6a7..999369412 100644 --- a/crates/core/tests/integration/pipeline_tests.rs +++ b/crates/core/tests/integration/pipeline_tests.rs @@ -30,7 +30,7 @@ use nemo_relay::api::runtime::NemoRelayContextState; use nemo_relay::api::runtime::global_context; use nemo_relay::api::runtime::{LlmExecutionNextFn, LlmJsonStream, LlmStreamExecutionNextFn}; use nemo_relay::api::runtime::{create_scope_stack, set_thread_scope_stack}; -use nemo_relay::api::scope::ScopeType; +use nemo_relay::api::scope::{EmitMarkEventParams, ScopeType, event}; use nemo_relay::api::subscriber::{deregister_subscriber, flush_subscribers, register_subscriber}; use nemo_relay::codec::anthropic::AnthropicMessagesCodec; use nemo_relay::codec::openai_chat::OpenAIChatCodec; @@ -1985,9 +1985,7 @@ async fn test_stream_response_sanitizer_can_flush_subscribers() { 1, Arc::new(|response, _context| { Box::pin(async move { - tokio::task::spawn_blocking(flush_subscribers) - .await - .map_err(|error| FlowError::Internal(error.to_string()))??; + flush_subscribers()?; Ok(Some(response)) }) }), @@ -2018,3 +2016,120 @@ async fn test_stream_response_sanitizer_can_flush_subscribers() { deregister_llm_sanitize_response_guardrail("stream_reentrant_flush_sanitizer").unwrap(); deregister_subscriber("stream_reentrant_flush_subscriber").unwrap(); } + +#[tokio::test] +async fn test_dropped_stream_end_keeps_fifo_position_before_later_mark() { + let _lock = TEST_MUTEX.lock().unwrap(); + reset_global(); + setup_isolated_thread(); + + let sanitizer_started = Arc::new(tokio::sync::Notify::new()); + let sanitizer_release = Arc::new(tokio::sync::Notify::new()); + let events = Arc::new(Mutex::new(Vec::new())); + let captured_events = events.clone(); + register_subscriber( + "stream_fifo_subscriber", + Arc::new(move |event| { + captured_events.lock().unwrap().push(event.clone()); + }), + ) + .unwrap(); + register_llm_sanitize_response_guardrail( + "stream_fifo_sanitizer", + 1, + Arc::new({ + let sanitizer_started = sanitizer_started.clone(); + let sanitizer_release = sanitizer_release.clone(); + move |response, _context| { + let sanitizer_started = sanitizer_started.clone(); + let sanitizer_release = sanitizer_release.clone(); + Box::pin(async move { + sanitizer_started.notify_one(); + sanitizer_release.notified().await; + Ok(Some(response)) + }) + } + }), + ) + .unwrap(); + + let stream = llm_stream_call_execute( + LlmStreamCallExecuteParams::builder() + .name("stream_fifo") + .request(make_openai_chat_request("stream me")) + .func(Arc::new(|_| { + Box::pin(async { + assert!(record_llm_optimization_contribution( + routed_model_contribution() + )); + Ok(LlmJsonStream::new(tokio_stream::empty())) + }) + })) + .collector(Box::new(|_chunk| Ok(()))) + .finalizer(Box::new(|| make_openai_chat_response("done"))) + .build(), + ) + .await + .unwrap(); + drop(stream); + + tokio::time::timeout( + std::time::Duration::from_secs(2), + sanitizer_started.notified(), + ) + .await + .expect("stream response sanitizer did not start"); + event( + EmitMarkEventParams::builder() + .name("mark-after-stream-drop") + .build(), + ) + .unwrap(); + + let (flush_done_tx, flush_done_rx) = std::sync::mpsc::channel(); + std::thread::spawn(move || { + let result = flush_subscribers(); + let _ = flush_done_tx.send(result); + }); + let flush_waited_for_end = flush_done_rx + .recv_timeout(std::time::Duration::from_millis(50)) + .is_err(); + sanitizer_release.notify_one(); + assert!( + flush_waited_for_end, + "flush must wait for the pending stream END" + ); + flush_done_rx + .recv_timeout(std::time::Duration::from_secs(2)) + .expect("flush did not finish after sanitizer release") + .unwrap(); + + let events = events.lock().unwrap(); + let end_index = events + .iter() + .position(|event| { + event.name() == "stream_fifo" + && is_scope_event(event, ScopeType::Llm, ScopeCategory::End) + }) + .expect("stream END event"); + let optimization_index = events + .iter() + .position(|event| event.name() == "nemo_relay.llm.optimization") + .expect("stream optimization mark"); + let mark_index = events + .iter() + .position(|event| event.name() == "mark-after-stream-drop") + .expect("later mark event"); + assert!( + optimization_index < end_index, + "optimization marks must retain their position before stream END" + ); + assert!( + end_index < mark_index, + "stream END must retain its FIFO position before the later mark" + ); + + drop(events); + deregister_llm_sanitize_response_guardrail("stream_fifo_sanitizer").unwrap(); + deregister_subscriber("stream_fifo_subscriber").unwrap(); +} diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index c1e6ecdd3..6f10cb2f6 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -3174,10 +3174,10 @@ pub fn deregister_subscriber(name: String) -> Result { /// Return a Promise that resolves when native subscriber callbacks queued /// before this call finish. /// -/// When called from queued publication middleware, or while an asynchronous -/// publication boundary is active, this Promise resolves without waiting. -/// Call it again after the middleware settles to wait for its event and later -/// work. +/// When called from a queued publication sanitizer callback (including event and manual tool/LLM +/// sanitizers), this Promise resolves without waiting to prevent a cycle with the serial +/// dispatcher. Publication middleware must not move such a re-entrant flush to +/// an unmarked worker thread. /// /// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this /// Promise does not block the Node event loop while Promise-returning event sanitizers settle. diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index 0e00bab51..53f48f33e 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1524,9 +1524,9 @@ fn deregister_subscriber(name: &str) -> PyResult { /// Wait for subscriber callbacks queued before this call to finish. /// -/// Re-entrant calls and calls observed while an asynchronous publication -/// boundary is active return without waiting. Call again after middleware -/// settles to wait for its event and later work. +/// Public Python wrappers prevent re-entrant event-sanitizer callbacks from +/// waiting on the serial dispatcher. Publication middleware must not move such +/// a re-entrant flush to an unmarked worker thread. #[pyfunction] fn flush_subscribers(py: Python<'_>) -> PyResult<()> { py.detach(core_subscriber_api::flush_subscribers) diff --git a/python/nemo_relay/subscribers.py b/python/nemo_relay/subscribers.py index 1137fc829..aeabaf557 100644 --- a/python/nemo_relay/subscribers.py +++ b/python/nemo_relay/subscribers.py @@ -95,9 +95,10 @@ def flush() -> None: waiting for observer work. Use this barrier in tests and shutdown paths when captured subscriber output must be complete before continuing. - A re-entrant call, or a call observed while an asynchronous publication - boundary is active, returns without waiting. Call ``flush()`` again after - that middleware settles to wait for its event and later work. + Call this function outside subscriber and queued publication sanitizer + callbacks. A re-entrant call returns without waiting to avoid blocking the + dispatcher. Publication middleware must not move such a call to an unmarked + worker thread. """ if _event_sanitizer_callback_active(): return None From 7af6674678e2993a28dc50892ad2965a94420930 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 14:05:07 -0400 Subject: [PATCH 34/52] feat: add async middleware C and Go APIs Signed-off-by: Will Killian --- crates/ffi/build.rs | 150 +++- crates/ffi/nemo_relay.h | 128 ++++ crates/ffi/src/api/event_registry.rs | 133 +++- crates/ffi/src/api/llm_registry.rs | 105 ++- crates/ffi/src/api/mod.rs | 20 +- crates/ffi/src/api/scope_registry.rs | 148 +++- crates/ffi/src/api/tool_registry.rs | 83 ++- crates/ffi/src/callable.rs | 649 +++++++++++++++++- crates/ffi/tests/integration/api_tests.rs | 162 +++++ .../tests/unit/api/coverage_sweeps_tests.rs | 273 ++++++++ crates/ffi/tests/unit/api/registry_tests.rs | 4 +- .../ffi/tests/unit/callable_private_tests.rs | 247 +++++++ crates/ffi/tests/unit/callable_tests.rs | 244 ++++++- go/nemo_relay/adaptive_runtime_test.go | 55 +- go/nemo_relay/async_middleware_test.go | 224 ++++++ go/nemo_relay/callbacks.go | 180 +++++ go/nemo_relay/nemo_relay.go | 300 ++++++++ go/nemo_relay/optimization_test.go | 20 + 18 files changed, 3088 insertions(+), 37 deletions(-) create mode 100644 go/nemo_relay/async_middleware_test.go diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index 013b53194..d7f0d3d32 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -5,6 +5,7 @@ fn main() { let crate_dir = std::env::var("CARGO_MANIFEST_DIR").unwrap(); + validate_async_registration_parity(&crate_dir); let config = cbindgen::Config::from_file(format!("{crate_dir}/cbindgen.toml")) .expect("Unable to read cbindgen.toml"); @@ -13,6 +14,153 @@ fn main() { .with_config(config) .generate() { - bindings.write_to_file(format!("{crate_dir}/nemo_relay.h")); + let header_path = format!("{crate_dir}/nemo_relay.h"); + bindings.write_to_file(&header_path); + // cbindgen intentionally does not expand declarative macros. Keep the + // macro-generated async registration functions in the generated C ABI. + let header = std::fs::read_to_string(&header_path).expect("read generated FFI header"); + let marker = "\n#endif /* NEMO_RELAY_H */\n"; + assert!( + header.contains(marker), + "generated FFI header is missing its NEMO_RELAY_H closing guard" + ); + let header = header.replacen( + marker, + &format!("\n{}\n#endif /* NEMO_RELAY_H */\n", ASYNC_REGISTRATIONS), + 1, + ); + std::fs::write(header_path, header).expect("write generated FFI header"); } } + +#[derive(Debug, PartialEq, Eq)] +struct AsyncPrototype<'a> { + name: &'a str, + parameters: Vec<&'a str>, +} + +/// cbindgen does not expand the declarative registration macros. Keep the +/// handwritten C declarations checked against macro-generated exports, their +/// complete parameter lists, and ordering. +fn validate_async_registration_parity(crate_dir: &str) { + const REGISTRATION_SOURCES: &[&str] = &[ + "src/api/event_registry.rs", + "src/api/llm_registry.rs", + "src/api/scope_registry.rs", + "src/api/tool_registry.rs", + ]; + + let mut exported = Vec::new(); + for source in REGISTRATION_SOURCES { + println!("cargo:rerun-if-changed={source}"); + let source_path = format!("{crate_dir}/{source}"); + let contents = std::fs::read_to_string(&source_path) + .unwrap_or_else(|error| panic!("read {source_path}: {error}")); + exported.extend( + contents + .split(|character: char| !character.is_ascii_alphanumeric() && character != '_') + .filter(|token| token.starts_with("nemo_relay_") && token.ends_with("_async")) + .map(str::to_owned), + ); + } + exported.sort(); + exported.dedup(); + + let mut declared = ASYNC_REGISTRATIONS + .lines() + .filter_map(parse_async_prototype) + .collect::>(); + declared.sort_by(|left, right| left.name.cmp(right.name)); + for duplicates in declared.windows(2) { + assert_ne!( + duplicates[0].name, duplicates[1].name, + "ASYNC_REGISTRATIONS contains duplicate declaration for {}", + duplicates[0].name + ); + } + let declared_names = declared + .iter() + .map(|prototype| prototype.name.to_owned()) + .collect::>(); + assert_eq!( + declared_names, exported, + "ASYNC_REGISTRATIONS must declare exactly the async Rust FFI exports" + ); + + for prototype in declared { + assert!( + exported + .binary_search_by(|name| name.as_str().cmp(prototype.name)) + .is_ok(), + "async declaration for {} is not a Rust FFI export", + prototype.name + ); + assert_eq!( + prototype, + expected_async_prototype(prototype.name), + "async declaration for {} has a mismatched C prototype", + prototype.name + ); + } +} + +fn parse_async_prototype(line: &str) -> Option> { + let line = line.strip_prefix("NemoRelayStatus ")?; + let (name, parameters) = line.split_once('(')?; + let parameters = parameters.strip_suffix(");")?; + Some(AsyncPrototype { + name, + parameters: parameters.split(", ").collect(), + }) +} + +fn expected_async_prototype(name: &str) -> AsyncPrototype<'_> { + let mut parameters = Vec::new(); + if name.starts_with("nemo_relay_scope_") { + parameters.push("const char *scope_uuid"); + } + parameters.extend(["const char *name", "int32_t priority"]); + if name.contains("request_intercept_async") { + parameters.push("bool break_chain"); + } + parameters.push(if name.contains("execution_intercept_async") { + "NemoRelayAsyncInterceptCb cb" + } else { + "NemoRelayAsyncJsonCb cb" + }); + parameters.extend(["void *user_data", "NemoRelayFreeFn free_fn"]); + AsyncPrototype { name, parameters } +} + +const ASYNC_REGISTRATIONS: &str = r#" +/* Completion-based async middleware registrations generated from Rust macros. */ +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion); +NemoRelayStatus nemo_relay_register_mark_sanitize_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_scope_sanitize_start_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_scope_sanitize_end_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_mark_sanitize_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_scope_sanitize_start_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_scope_sanitize_end_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_sanitize_request_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_sanitize_response_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_sanitize_request_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_sanitize_response_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_stream_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_stream_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +"#; diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index accaa5e9e..83af886eb 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -132,6 +132,21 @@ enum NemoRelayScopeType { }; typedef int32_t NemoRelayScopeType; +/** + * Indicates whether an async callback settled its completion before returning. + */ +enum NemoRelayAsyncCallbackState { + /** + * The callback called a resolve/reject function before returning. + */ + NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0, + /** + * The callback retained the completion and will settle it later. + */ + NEMO_RELAY_ASYNC_CALLBACK_STATE_PENDING = 1, +}; +typedef uint32_t NemoRelayAsyncCallbackState; + /** * Opaque owned adaptive runtime handle. */ @@ -235,6 +250,16 @@ typedef struct FfiThreadScopeStackBinding FfiThreadScopeStackBinding; */ typedef struct FfiToolHandle FfiToolHandle; +/** + * One-shot completion passed to asynchronous C callbacks. + */ +typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; + +/** + * Runtime-owned asynchronous `next` continuation for execution intercepts. + */ +typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; + typedef struct Option_NemoRelayCollectorCb Option_NemoRelayCollectorCb; typedef struct Option_NemoRelayFinalizerCb Option_NemoRelayFinalizerCb; @@ -445,6 +470,26 @@ typedef char *(*NemoRelayToolExecInterceptCb)(void *user_data, */ typedef char *(*NemoRelayToolExecCb)(void *user_data, const char *args_json); +/** + * Completion-based execution-intercept callback. + */ +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, + const char *invocation_json, + const struct NemoRelayAsyncNext *next, + const struct NemoRelayAsyncCompletion *completion); + +/** + * Result callback used by channel/future-style async `next` wrappers. + * + * Invoked on a Tokio runtime worker thread, not necessarily the thread that + * called `nemo_relay_async_next_invoke_callback`; `user_data` must therefore + * be safe for cross-thread use. `value_json` and `error_message` are borrowed + * for the duration of the callback only. + */ +typedef void (*NemoRelayAsyncNextResultCb)(void *user_data, + const char *value_json, + const char *error_message); + /** * Run the registered tool request intercept chain on the given arguments. * @@ -2484,6 +2529,15 @@ NemoRelayStatus nemo_relay_tool_call_execute(const char *name, const char *metadata_json, char **out); +/** + * Register a completion-based asynchronous tool execution intercept. + */ +NemoRelayStatus nemo_relay_register_tool_execution_intercept_async(const char *name, + int32_t priority, + NemoRelayAsyncInterceptCb cb, + void *user_data, + NemoRelayFreeFn free_fn); + /** * Register a tool conditional execution guardrail. The callback decides whether * a tool call should proceed. Returns an error message to reject, or null to allow. @@ -2546,6 +2600,48 @@ NemoRelayStatus nemo_relay_register_tool_execution_intercept(const char *name, */ NemoRelayStatus nemo_relay_deregister_tool_execution_intercept(const char *name); +/** + * Release the callback-owned async `next` reference after a pending intercept. + */ +void nemo_relay_async_next_release(const struct NemoRelayAsyncNext *next); + +/** + * Invoke the next execution layer and settle `completion` with its result. + */ +NemoRelayStatus nemo_relay_async_next_invoke(const struct NemoRelayAsyncNext *next, + const char *invocation_json, + const struct NemoRelayAsyncCompletion *completion); + +/** + * Invoke the next execution layer and report its result through a callback. + */ +NemoRelayStatus nemo_relay_async_next_invoke_callback(const struct NemoRelayAsyncNext *next, + const char *invocation_json, + NemoRelayAsyncNextResultCb callback, + void *user_data); + +/** + * Resolve an async C callback with owned JSON. + */ +NemoRelayStatus nemo_relay_async_completion_resolve_json(const struct NemoRelayAsyncCompletion *completion, + const char *value_json); + +/** + * Reject an async C callback with an error message. + */ +NemoRelayStatus nemo_relay_async_completion_reject(const struct NemoRelayAsyncCompletion *completion, + const char *message); + +/** + * Returns whether an async completion's invocation has been cancelled. + */ +bool nemo_relay_async_completion_is_cancelled(const struct NemoRelayAsyncCompletion *completion); + +/** + * Release the callback-owned completion reference after a pending invocation. + */ +void nemo_relay_async_completion_release(const struct NemoRelayAsyncCompletion *completion); + /** * Free a C string previously returned by any `nemo_relay_*` accessor function. * Passing null is a safe no-op. @@ -3039,4 +3135,36 @@ char *nemo_relay_event_annotated_request(const struct FfiEvent *ptr); */ char *nemo_relay_event_annotated_response(const struct FfiEvent *ptr); + +/* Completion-based async middleware registrations generated from Rust macros. */ +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion); +NemoRelayStatus nemo_relay_register_mark_sanitize_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_scope_sanitize_start_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_scope_sanitize_end_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_mark_sanitize_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_scope_sanitize_start_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_scope_sanitize_end_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_sanitize_request_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_sanitize_response_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_sanitize_request_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_sanitize_response_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_stream_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_stream_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); + #endif /* NEMO_RELAY_H */ diff --git a/crates/ffi/src/api/event_registry.rs b/crates/ffi/src/api/event_registry.rs index 13af2b8d8..4badaa381 100644 --- a/crates/ffi/src/api/event_registry.rs +++ b/crates/ffi/src/api/event_registry.rs @@ -2,8 +2,9 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayEventSanitizeCb, NemoRelayFreeFn, NemoRelayStatus, c_char, c_str_to_string, - clear_last_error, core_registry_api, set_last_error, status_from_error, wrap_event_sanitize_fn, + NemoRelayAsyncJsonCb, NemoRelayEventSanitizeCb, NemoRelayFreeFn, NemoRelayStatus, c_char, + c_str_to_string, clear_last_error, core_registry_api, set_last_error, status_from_error, + wrap_async_event_sanitize_fn, wrap_event_sanitize_fn, }; #[derive(Clone, Copy)] @@ -43,6 +44,66 @@ unsafe fn register_global( .unwrap_or_else(|error| status_from_error(&error)) } +unsafe fn register_global_async( + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + surface: Surface, +) -> NemoRelayStatus { + clear_last_error(); + let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + let result = match surface { + Surface::Mark => { + core_registry_api::register_mark_sanitize_guardrail(&name, priority, callback) + } + Surface::Start => { + core_registry_api::register_scope_sanitize_start_guardrail(&name, priority, callback) + } + Surface::End => { + core_registry_api::register_scope_sanitize_end_guardrail(&name, priority, callback) + } + }; + result + .map(|()| NemoRelayStatus::Ok) + .unwrap_or_else(|error| status_from_error(&error)) +} + +macro_rules! async_event_registration { + ($name:ident, $surface:expr) => { + /// Register a completion-based asynchronous event sanitizer. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name( + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + unsafe { register_global_async(name, priority, cb, user_data, free_fn, $surface) } + } + }; +} + +async_event_registration!( + nemo_relay_register_mark_sanitize_guardrail_async, + Surface::Mark +); +async_event_registration!( + nemo_relay_register_scope_sanitize_start_guardrail_async, + Surface::Start +); +async_event_registration!( + nemo_relay_register_scope_sanitize_end_guardrail_async, + Surface::End +); + unsafe fn deregister_global(name: *const c_char, surface: Surface) -> NemoRelayStatus { clear_last_error(); let name = match c_str_to_string(name) { @@ -102,6 +163,74 @@ unsafe fn register_scope( .unwrap_or_else(|error| status_from_error(&error)) } +unsafe fn register_scope_async( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + surface: Surface, +) -> NemoRelayStatus { + clear_last_error(); + let uuid = match parse_scope_uuid(scope_uuid) { + Ok(uuid) => uuid, + Err(status) => return status, + }; + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); + let result = match surface { + Surface::Mark => core_registry_api::scope_register_mark_sanitize_guardrail( + &uuid, &name, priority, callback, + ), + Surface::Start => core_registry_api::scope_register_scope_sanitize_start_guardrail( + &uuid, &name, priority, callback, + ), + Surface::End => core_registry_api::scope_register_scope_sanitize_end_guardrail( + &uuid, &name, priority, callback, + ), + }; + result + .map(|()| NemoRelayStatus::Ok) + .unwrap_or_else(|error| status_from_error(&error)) +} + +macro_rules! scope_async_event_registration { + ($name:ident, $surface:expr) => { + /// Register a scope-local completion-based asynchronous event sanitizer. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + unsafe { + register_scope_async(scope_uuid, name, priority, cb, user_data, free_fn, $surface) + } + } + }; +} + +scope_async_event_registration!( + nemo_relay_scope_register_mark_sanitize_guardrail_async, + Surface::Mark +); +scope_async_event_registration!( + nemo_relay_scope_register_scope_sanitize_start_guardrail_async, + Surface::Start +); +scope_async_event_registration!( + nemo_relay_scope_register_scope_sanitize_end_guardrail_async, + Surface::End +); + unsafe fn deregister_scope( scope_uuid: *const c_char, name: *const c_char, diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index 38883ecee..60d1d0053 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -2,14 +2,107 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayEventSubscriberCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, - NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, - NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, c_char, c_str_to_string, clear_last_error, - core_registry_api, core_subscriber_api, status_from_error, wrap_event_subscriber, - wrap_llm_conditional_fn, wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, - wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_intercept_fn, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayEventSubscriberCb, NemoRelayFreeFn, + NemoRelayLlmConditionalCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, + NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, c_char, + c_str_to_string, clear_last_error, core_registry_api, core_subscriber_api, status_from_error, + wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, + wrap_async_llm_request_intercept_fn, wrap_async_llm_sanitize_request_fn, + wrap_async_llm_sanitize_response_fn, wrap_async_llm_stream_execution_intercept_fn, + wrap_event_subscriber, wrap_llm_conditional_fn, wrap_llm_exec_intercept_fn, + wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, + wrap_llm_stream_exec_intercept_fn, }; +macro_rules! async_llm_registration { + ($fn_name:ident, $register:path, $wrapper:path $(, $break_chain:ident)?) => { + /// Register a completion-based asynchronous LLM middleware callback. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $fn_name( + name: *const c_char, + priority: i32, + $( $break_chain: bool, )? + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register( + &name, + priority, + $( $break_chain, )? + $wrapper(cb, user_data, free_fn), + ) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +async_llm_registration!( + nemo_relay_register_llm_sanitize_request_guardrail_async, + core_registry_api::register_llm_sanitize_request_guardrail, + wrap_async_llm_sanitize_request_fn +); + +macro_rules! async_llm_execution_registration { + ($name:ident, $register:path, $wrapper:path) => { + /// Register a completion-based asynchronous LLM execution intercept. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name( + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register(&name, priority, $wrapper(cb, user_data, free_fn)) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +async_llm_execution_registration!( + nemo_relay_register_llm_execution_intercept_async, + core_registry_api::register_llm_execution_intercept, + wrap_async_llm_execution_intercept_fn +); +async_llm_execution_registration!( + nemo_relay_register_llm_stream_execution_intercept_async, + core_registry_api::register_llm_stream_execution_intercept, + wrap_async_llm_stream_execution_intercept_fn +); +async_llm_registration!( + nemo_relay_register_llm_sanitize_response_guardrail_async, + core_registry_api::register_llm_sanitize_response_guardrail, + wrap_async_llm_sanitize_response_fn +); +async_llm_registration!( + nemo_relay_register_llm_conditional_execution_guardrail_async, + core_registry_api::register_llm_conditional_execution_guardrail, + wrap_async_llm_conditional_fn +); +async_llm_registration!( + nemo_relay_register_llm_request_intercept_async, + core_registry_api::register_llm_request_intercept, + wrap_async_llm_request_intercept_fn, + break_chain +); + // --------------------------------------------------------------------------- // LLM guardrail registrations // --------------------------------------------------------------------------- diff --git a/crates/ffi/src/api/mod.rs b/crates/ffi/src/api/mod.rs index 1fc8977fd..fa5ff2bc3 100644 --- a/crates/ffi/src/api/mod.rs +++ b/crates/ffi/src/api/mod.rs @@ -14,17 +14,23 @@ use std::sync::{Arc, OnceLock}; use std::time::Duration; use crate::callable::{ - NemoRelayCodecDecodeFn, NemoRelayCodecEncodeFn, NemoRelayCollectorCb, NemoRelayEventSanitizeCb, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayCodecDecodeFn, + NemoRelayCodecEncodeFn, NemoRelayCollectorCb, NemoRelayEventSanitizeCb, NemoRelayEventSubscriberCb, NemoRelayFinalizerCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, NemoRelayLlmExecCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayPluginRegisterCb, NemoRelayPluginValidateCb, NemoRelayToolConditionalCb, NemoRelayToolExecCb, - NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, wrap_codec_fn, wrap_collector_fn, - wrap_event_sanitize_fn, wrap_event_subscriber, wrap_finalizer_fn, wrap_llm_conditional_fn, - wrap_llm_exec_fn, wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, - wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_fn, - wrap_llm_stream_exec_intercept_fn, wrap_tool_conditional_fn, wrap_tool_exec_fn, - wrap_tool_exec_intercept_fn, wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, + NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, wrap_async_event_sanitize_fn, + wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, + wrap_async_llm_request_intercept_fn, wrap_async_llm_sanitize_request_fn, + wrap_async_llm_sanitize_response_fn, wrap_async_llm_stream_execution_intercept_fn, + wrap_async_tool_conditional_fn, wrap_async_tool_execution_intercept_fn, + wrap_async_tool_json_fn, wrap_codec_fn, wrap_collector_fn, wrap_event_sanitize_fn, + wrap_event_subscriber, wrap_finalizer_fn, wrap_llm_conditional_fn, wrap_llm_exec_fn, + wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, + wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_fn, wrap_llm_stream_exec_intercept_fn, + wrap_tool_conditional_fn, wrap_tool_exec_fn, wrap_tool_exec_intercept_fn, + wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, }; use crate::convert::{ c_str_to_json, c_str_to_opt_json, c_str_to_string, json_to_c_string, nemo_relay_string_free, diff --git a/crates/ffi/src/api/scope_registry.rs b/crates/ffi/src/api/scope_registry.rs index 50efd644b..aebf09cd1 100644 --- a/crates/ffi/src/api/scope_registry.rs +++ b/crates/ffi/src/api/scope_registry.rs @@ -2,15 +2,19 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayEventSubscriberCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, - NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, - NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, NemoRelayToolConditionalCb, - NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, c_char, c_str_to_string, - clear_last_error, core_registry_api, core_subscriber_api, set_last_error, status_from_error, - wrap_event_subscriber, wrap_llm_conditional_fn, wrap_llm_exec_intercept_fn, - wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, - wrap_llm_stream_exec_intercept_fn, wrap_tool_conditional_fn, wrap_tool_exec_intercept_fn, - wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayEventSubscriberCb, NemoRelayFreeFn, + NemoRelayLlmConditionalCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, + NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, + NemoRelayToolConditionalCb, NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, c_char, + c_str_to_string, clear_last_error, core_registry_api, core_subscriber_api, set_last_error, + status_from_error, wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, + wrap_async_llm_request_intercept_fn, wrap_async_llm_sanitize_request_fn, + wrap_async_llm_sanitize_response_fn, wrap_async_llm_stream_execution_intercept_fn, + wrap_async_tool_conditional_fn, wrap_async_tool_execution_intercept_fn, + wrap_async_tool_json_fn, wrap_event_subscriber, wrap_llm_conditional_fn, + wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, + wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_intercept_fn, wrap_tool_conditional_fn, + wrap_tool_exec_intercept_fn, wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, }; // --------------------------------------------------------------------------- @@ -26,6 +30,132 @@ fn parse_scope_uuid(scope_uuid: *const c_char) -> Result { + /// Register a scope-local completion-based asynchronous middleware callback. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $fn_name( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + $( $break_chain: bool, )? + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let uuid = match parse_scope_uuid(scope_uuid) { + Ok(uuid) => uuid, + Err(status) => return status, + }; + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register( + &uuid, + &name, + priority, + $( $break_chain, )? + $wrapper(cb, user_data, free_fn), + ) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +scope_async_registration!( + nemo_relay_scope_register_tool_sanitize_request_guardrail_async, + core_registry_api::scope_register_tool_sanitize_request_guardrail, + wrap_async_tool_json_fn +); +scope_async_registration!( + nemo_relay_scope_register_tool_sanitize_response_guardrail_async, + core_registry_api::scope_register_tool_sanitize_response_guardrail, + wrap_async_tool_json_fn +); +scope_async_registration!( + nemo_relay_scope_register_tool_conditional_execution_guardrail_async, + core_registry_api::scope_register_tool_conditional_execution_guardrail, + wrap_async_tool_conditional_fn +); +scope_async_registration!( + nemo_relay_scope_register_tool_request_intercept_async, + core_registry_api::scope_register_tool_request_intercept, + wrap_async_tool_json_fn, + break_chain +); +scope_async_registration!( + nemo_relay_scope_register_llm_sanitize_request_guardrail_async, + core_registry_api::scope_register_llm_sanitize_request_guardrail, + wrap_async_llm_sanitize_request_fn +); +scope_async_registration!( + nemo_relay_scope_register_llm_sanitize_response_guardrail_async, + core_registry_api::scope_register_llm_sanitize_response_guardrail, + wrap_async_llm_sanitize_response_fn +); +scope_async_registration!( + nemo_relay_scope_register_llm_conditional_execution_guardrail_async, + core_registry_api::scope_register_llm_conditional_execution_guardrail, + wrap_async_llm_conditional_fn +); +scope_async_registration!( + nemo_relay_scope_register_llm_request_intercept_async, + core_registry_api::scope_register_llm_request_intercept, + wrap_async_llm_request_intercept_fn, + break_chain +); + +macro_rules! scope_async_execution_registration { + ($fn_name:ident, $register:path, $wrapper:path) => { + /// Register a scope-local completion-based asynchronous execution intercept. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $fn_name( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let uuid = match parse_scope_uuid(scope_uuid) { + Ok(uuid) => uuid, + Err(status) => return status, + }; + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register(&uuid, &name, priority, $wrapper(cb, user_data, free_fn)) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +scope_async_execution_registration!( + nemo_relay_scope_register_tool_execution_intercept_async, + core_registry_api::scope_register_tool_execution_intercept, + wrap_async_tool_execution_intercept_fn +); +scope_async_execution_registration!( + nemo_relay_scope_register_llm_execution_intercept_async, + core_registry_api::scope_register_llm_execution_intercept, + wrap_async_llm_execution_intercept_fn +); +scope_async_execution_registration!( + nemo_relay_scope_register_llm_stream_execution_intercept_async, + core_registry_api::scope_register_llm_stream_execution_intercept, + wrap_async_llm_stream_execution_intercept_fn +); + macro_rules! ffi_scope_guardrail_tool_api { ($(#[$reg_doc:meta])* $register_name:ident, $(#[$dereg_doc:meta])* $deregister_name:ident, diff --git a/crates/ffi/src/api/tool_registry.rs b/crates/ffi/src/api/tool_registry.rs index 5d5cefb40..fd1f2e580 100644 --- a/crates/ffi/src/api/tool_registry.rs +++ b/crates/ffi/src/api/tool_registry.rs @@ -2,12 +2,89 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayFreeFn, NemoRelayStatus, NemoRelayToolConditionalCb, NemoRelayToolExecInterceptCb, - NemoRelayToolSanitizeCb, c_char, c_str_to_string, clear_last_error, core_registry_api, - status_from_error, wrap_tool_conditional_fn, wrap_tool_exec_intercept_fn, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayFreeFn, NemoRelayStatus, + NemoRelayToolConditionalCb, NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, c_char, + c_str_to_string, clear_last_error, core_registry_api, status_from_error, + wrap_async_tool_conditional_fn, wrap_async_tool_execution_intercept_fn, + wrap_async_tool_json_fn, wrap_tool_conditional_fn, wrap_tool_exec_intercept_fn, wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, }; +macro_rules! async_tool_json_registration { + ($name:ident, $register:path, $wrapper:path $(, $break_chain:ident)?) => { + /// Register a completion-based asynchronous tool middleware callback. + #[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name( + name: *const c_char, + priority: i32, + $( $break_chain: bool, )? + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + let callback = $wrapper(cb, user_data, free_fn); + match $register(&name, priority, $( $break_chain, )? callback) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +async_tool_json_registration!( + nemo_relay_register_tool_sanitize_request_guardrail_async, + core_registry_api::register_tool_sanitize_request_guardrail, + wrap_async_tool_json_fn +); +async_tool_json_registration!( + nemo_relay_register_tool_sanitize_response_guardrail_async, + core_registry_api::register_tool_sanitize_response_guardrail, + wrap_async_tool_json_fn +); +async_tool_json_registration!( + nemo_relay_register_tool_conditional_execution_guardrail_async, + core_registry_api::register_tool_conditional_execution_guardrail, + wrap_async_tool_conditional_fn +); + +async_tool_json_registration!( + nemo_relay_register_tool_request_intercept_async, + core_registry_api::register_tool_request_intercept, + wrap_async_tool_json_fn, + break_chain +); + +/// Register a completion-based asynchronous tool execution intercept. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_register_tool_execution_intercept_async( + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> NemoRelayStatus { + clear_last_error(); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match core_registry_api::register_tool_execution_intercept( + &name, + priority, + wrap_async_tool_execution_intercept_fn(cb, user_data, free_fn), + ) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } +} + // --------------------------------------------------------------------------- // Tool guardrail registrations // --------------------------------------------------------------------------- diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 054c098bd..9414d31ac 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -19,14 +19,17 @@ use std::ffi::{CStr, CString}; use std::future::Future; use std::pin::Pin; +use std::ptr; use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; use libc::c_char; use nemo_relay::api::runtime::{ - EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmConditionalFn, LlmExecutionNextFn, - LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, - LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionNextFn, ToolConditionalFn, - ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmConditionalFn, LlmExecutionFn, + LlmExecutionNextFn, LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, + LlmSanitizeRequestFn, LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, + LlmStreamExecutionNextFn, ToolConditionalFn, ToolExecutionFn, ToolExecutionNextFn, + ToolInterceptFn, ToolSanitizeFn, }; use serde_json::Value as Json; use tokio_stream::StreamExt; @@ -54,6 +57,401 @@ use crate::types::{FfiEvent, FfiLLMRequest, FfiPluginContext}; /// destructor runs. pub type NemoRelayFreeFn = Option; +/// Indicates whether an async callback settled its completion before returning. +#[repr(u32)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum NemoRelayAsyncCallbackState { + /// The callback called a resolve/reject function before returning. + Complete = 0, + /// The callback retained the completion and will settle it later. + Pending = 1, +} + +/// One-shot completion passed to asynchronous C callbacks. +pub struct NemoRelayAsyncCompletion { + sender: std::sync::Mutex>>>, + cancelled: AtomicBool, +} + +/// Generic completion-based middleware callback. +/// +/// `invocation_json` is borrowed for the duration of the call. The completion +/// has one callback-owned reference. A callback returning `Complete` need not +/// release it; a callback returning `Pending` must eventually settle and call +/// `nemo_relay_async_completion_release`. +pub type NemoRelayAsyncJsonCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState; + +/// Runtime-owned asynchronous `next` continuation for execution intercepts. +pub struct NemoRelayAsyncNext { + inner: AsyncNextInner, + runtime: tokio::runtime::Handle, +} + +enum AsyncNextInner { + Tool(ToolExecutionNextFn), + Llm(LlmExecutionNextFn), + LlmStream(LlmStreamExecutionNextFn), +} + +/// Completion-based execution-intercept callback. +pub type NemoRelayAsyncInterceptCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState; + +/// Result callback used by channel/future-style async `next` wrappers. +/// +/// Invoked on a Tokio runtime worker thread, not necessarily the thread that +/// called `nemo_relay_async_next_invoke_callback`; `user_data` must therefore +/// be safe for cross-thread use. `value_json` and `error_message` are borrowed +/// for the duration of the callback only. +pub type NemoRelayAsyncNextResultCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + value_json: *const c_char, + error_message: *const c_char, +); + +struct CompletionWait { + completion: Arc, + receiver: tokio::sync::oneshot::Receiver>, +} + +impl Drop for CompletionWait { + fn drop(&mut self) { + self.completion.cancelled.store(true, Ordering::Release); + } +} + +async fn invoke_async_json( + cb: NemoRelayAsyncJsonCb, + user_data: Arc, + invocation: Json, +) -> Result { + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let callback_ref = Arc::into_raw(completion.clone()); + let invocation = json_to_c_string(&invocation); + let state = unsafe { cb(user_data.ptr, invocation, callback_ref) }; + unsafe { nemo_relay_string_free_internal(invocation) }; + if state == NemoRelayAsyncCallbackState::Complete { + unsafe { drop(Arc::from_raw(callback_ref)) }; + if completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .is_some() + { + return Err(FlowError::Internal( + "async C callback returned Complete without settling".into(), + )); + } + } + let mut wait = CompletionWait { + completion, + receiver, + }; + (&mut wait.receiver) + .await + .map_err(|_| FlowError::Internal("async C callback dropped without settling".into()))? +} + +async fn invoke_async_intercept( + cb: NemoRelayAsyncInterceptCb, + user_data: Arc, + invocation: Json, + next: AsyncNextInner, +) -> Result { + let runtime = tokio::runtime::Handle::try_current().map_err(|error| { + FlowError::Internal(format!( + "async C intercept requires a Tokio runtime: {error}" + )) + })?; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let callback_ref = Arc::into_raw(completion.clone()); + let next = Arc::new(NemoRelayAsyncNext { + inner: next, + runtime, + }); + let next_ref = Arc::into_raw(next); + let invocation = json_to_c_string(&invocation); + let state = unsafe { cb(user_data.ptr, invocation, next_ref, callback_ref) }; + unsafe { nemo_relay_string_free_internal(invocation) }; + if state == NemoRelayAsyncCallbackState::Complete { + unsafe { drop(Arc::from_raw(callback_ref)) }; + unsafe { drop(Arc::from_raw(next_ref)) }; + if completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .is_some() + { + return Err(FlowError::Internal( + "async C intercept returned Complete without settling".into(), + )); + } + } + let mut wait = CompletionWait { + completion, + receiver, + }; + (&mut wait.receiver) + .await + .map_err(|_| FlowError::Internal("async C intercept dropped without settling".into()))? +} + +/// Release the callback-owned async `next` reference after a pending intercept. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_next_release(next: *const NemoRelayAsyncNext) { + if !next.is_null() { + unsafe { drop(Arc::from_raw(next)) }; + } +} + +/// Invoke the next execution layer and settle `completion` with its result. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_next_invoke( + next: *const NemoRelayAsyncNext, + invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayStatus { + let Some(next) = (unsafe { next.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if completion.is_null() { + return NemoRelayStatus::NullPointer; + } + let Some(invocation) = c_str_to_json(invocation_json) else { + return NemoRelayStatus::InvalidJson; + }; + unsafe { Arc::increment_strong_count(completion) }; + let completion = unsafe { Arc::from_raw(completion) }; + let future: Pin> + Send>> = match &next.inner { + AsyncNextInner::Tool(next) => { + let next = next.clone(); + Box::pin(async move { + let outcome = next(invocation).await?; + serde_json::to_value(outcome) + .map_err(|error| FlowError::Internal(error.to_string())) + }) + } + AsyncNextInner::Llm(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(error) => { + return { + let _ = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); + NemoRelayStatus::InvalidJson + }; + } + }; + let next = next.clone(); + Box::pin(async move { next(request).await }) + } + AsyncNextInner::LlmStream(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(error) => { + return { + let _ = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); + NemoRelayStatus::InvalidJson + }; + } + }; + let next = next.clone(); + Box::pin(async move { + let mut stream = next(request).await?; + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk?); + } + Ok(Json::Array(chunks)) + }) + } + }; + next.runtime.spawn(async move { + let result = future.await; + if let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + { + let _ = sender.send(result); + } + }); + NemoRelayStatus::Ok +} + +/// Invoke the next execution layer and report its result through a callback. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_next_invoke_callback( + next: *const NemoRelayAsyncNext, + invocation_json: *const c_char, + callback: NemoRelayAsyncNextResultCb, + user_data: *mut libc::c_void, +) -> NemoRelayStatus { + let Some(next) = (unsafe { next.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + let Some(invocation) = c_str_to_json(invocation_json) else { + return NemoRelayStatus::InvalidJson; + }; + let future: Pin> + Send>> = match &next.inner { + AsyncNextInner::Tool(next) => { + let next = next.clone(); + Box::pin(async move { + serde_json::to_value(next(invocation).await?) + .map_err(|error| FlowError::Internal(error.to_string())) + }) + } + AsyncNextInner::Llm(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(_) => return NemoRelayStatus::InvalidJson, + }; + let next = next.clone(); + Box::pin(async move { next(request).await }) + } + AsyncNextInner::LlmStream(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(_) => return NemoRelayStatus::InvalidJson, + }; + let next = next.clone(); + Box::pin(async move { + let mut stream = next(request).await?; + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk?); + } + Ok(Json::Array(chunks)) + }) + } + }; + let user_data = user_data as usize; + next.runtime.spawn(async move { + match future.await { + Ok(value) => { + let value = json_to_c_string(&value); + unsafe { callback(user_data as *mut libc::c_void, value, ptr::null()) }; + unsafe { nemo_relay_string_free_internal(value) }; + } + Err(error) => { + let error = CString::new(error.to_string()).unwrap_or_default(); + unsafe { callback(user_data as *mut libc::c_void, ptr::null(), error.as_ptr()) }; + } + } + }); + NemoRelayStatus::Ok +} + +/// Resolve an async C callback with owned JSON. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_completion_resolve_json( + completion: *const NemoRelayAsyncCompletion, + value_json: *const c_char, +) -> NemoRelayStatus { + let Some(completion) = (unsafe { completion.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if completion.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let Some(value) = c_str_to_json(value_json) else { + return NemoRelayStatus::InvalidJson; + }; + let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + else { + return NemoRelayStatus::InvalidArg; + }; + let _ = sender.send(Ok(value)); + NemoRelayStatus::Ok +} + +/// Reject an async C callback with an error message. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_completion_reject( + completion: *const NemoRelayAsyncCompletion, + message: *const c_char, +) -> NemoRelayStatus { + let Some(completion) = (unsafe { completion.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if completion.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let message = if message.is_null() { + "async C callback rejected".to_string() + } else { + unsafe { CStr::from_ptr(message) } + .to_string_lossy() + .into_owned() + }; + let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + else { + return NemoRelayStatus::InvalidArg; + }; + let _ = sender.send(Err(FlowError::Internal(message))); + NemoRelayStatus::Ok +} + +/// Returns whether an async completion's invocation has been cancelled. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_completion_is_cancelled( + completion: *const NemoRelayAsyncCompletion, +) -> bool { + unsafe { completion.as_ref() } + .is_none_or(|completion| completion.cancelled.load(Ordering::Acquire)) +} + +/// Release the callback-owned completion reference after a pending invocation. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_completion_release( + completion: *const NemoRelayAsyncCompletion, +) { + if !completion.is_null() { + unsafe { drop(Arc::from_raw(completion)) }; + } +} + /// Callback for tool request/response sanitization guardrails and intercepts. /// Receives tool name and arguments as JSON, returns sanitized arguments as JSON. /// The returned string must be allocated with `malloc` or equivalent. @@ -321,6 +719,249 @@ fn make_user_data( }) } +/// Wrap a completion-based C tool sanitizer or request intercept. +pub fn wrap_async_tool_json_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> ToolSanitizeFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |name: String, value: Json| { + let user_data = user_data.clone(); + Box::pin(invoke_async_json( + cb, + user_data, + serde_json::json!({"name": name, "value": value}), + )) + }) +} + +/// Wrap a completion-based C tool conditional guardrail. +pub fn wrap_async_tool_conditional_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> ToolConditionalFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |name: String, value: Json| { + let user_data = user_data.clone(); + Box::pin(async move { + match invoke_async_json( + cb, + user_data, + serde_json::json!({"name": name, "value": value}), + ) + .await? + { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => Err(FlowError::Internal(format!( + "async conditional callback returned {other}; expected string or null" + ))), + } + }) + }) +} + +/// Wrap a completion-based C event sanitizer. +pub fn wrap_async_event_sanitize_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> EventSanitizeFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |event: Event, fields: EventSanitizeFields| { + let user_data = user_data.clone(); + Box::pin(async move { + let value = invoke_async_json( + cb, + user_data, + serde_json::json!({"event": event, "fields": fields}), + ) + .await?; + serde_json::from_value(value) + .map_err(|error| FlowError::Internal(format!("invalid event fields: {error}"))) + }) + }) +} + +/// Wrap a completion-based C LLM conditional guardrail. +pub fn wrap_async_llm_conditional_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmConditionalFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |request: LlmRequest| { + let user_data = user_data.clone(); + Box::pin(async move { + match invoke_async_json(cb, user_data, serde_json::json!({"request": request})).await? { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => Err(FlowError::Internal(format!( + "async conditional callback returned {other}; expected string or null" + ))), + } + }) + }) +} + +/// Wrap a completion-based C LLM request sanitizer. +pub fn wrap_async_llm_sanitize_request_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmSanitizeRequestFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new( + move |request: LlmRequest, context: LlmSanitizeRequestContext| { + let user_data = user_data.clone(); + let codec = format!("{:?}", context.codec()); + Box::pin(async move { + let value = invoke_async_json( + cb, + user_data, + serde_json::json!({"request": request, "context": {"codec": codec}}), + ) + .await?; + if value.is_null() { + Ok(None) + } else { + serde_json::from_value(value) + .map(Some) + .map_err(|error| FlowError::Internal(error.to_string())) + } + }) + }, + ) +} + +/// Wrap a completion-based C LLM response sanitizer. +pub fn wrap_async_llm_sanitize_response_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmSanitizeResponseFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { + let user_data = user_data.clone(); + let codec = format!("{:?}", context.codec()); + Box::pin(async move { + let value = invoke_async_json( + cb, + user_data, + serde_json::json!({"response": response, "context": {"codec": codec}}), + ) + .await?; + Ok((!value.is_null()).then_some(value)) + }) + }) +} + +/// Wrap a completion-based C LLM request intercept. +pub fn wrap_async_llm_request_intercept_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmRequestInterceptFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new( + move |name: String, request: LlmRequest, annotated: Option| { + let user_data = user_data.clone(); + Box::pin(async move { + let value = invoke_async_json( + cb, + user_data, + serde_json::json!({ + "name": name, + "request": request, + "annotated": annotated, + }), + ) + .await?; + serde_json::from_value(value).map_err(|error| { + FlowError::Internal(format!("invalid LLM request intercept outcome: {error}")) + }) + }) + }, + ) +} + +/// Wrap a completion-based C tool execution intercept. +pub fn wrap_async_tool_execution_intercept_fn( + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> ToolExecutionFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |name: &str, args: Json, next: ToolExecutionNextFn| { + let user_data = user_data.clone(); + let invocation = serde_json::json!({"name": name, "value": args}); + Box::pin(async move { + let value = + invoke_async_intercept(cb, user_data, invocation, AsyncNextInner::Tool(next)) + .await?; + serde_json::from_value(value) + .map_err(|error| FlowError::Internal(format!("invalid tool outcome: {error}"))) + }) + }) +} + +/// Wrap a completion-based C LLM execution intercept. +pub fn wrap_async_llm_execution_intercept_fn( + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmExecutionFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new( + move |name: &str, request: LlmRequest, next: LlmExecutionNextFn| { + let user_data = user_data.clone(); + let invocation = serde_json::json!({"name": name, "request": request}); + Box::pin(invoke_async_intercept( + cb, + user_data, + invocation, + AsyncNextInner::Llm(next), + )) + }, + ) +} + +/// Wrap a completion-based C LLM stream execution intercept. +/// +/// The completion ABI resolves one JSON value, so a stream intercept must +/// resolve to an array of chunks. Relay replays that array as a stream after +/// completion; incremental chunk delivery is not available through this ABI. +pub fn wrap_async_llm_stream_execution_intercept_fn( + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmStreamExecutionFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new( + move |name: &str, request: LlmRequest, next: LlmStreamExecutionNextFn| { + let user_data = user_data.clone(); + let invocation = serde_json::json!({"name": name, "request": request}); + Box::pin(async move { + let value = invoke_async_intercept( + cb, + user_data, + invocation, + AsyncNextInner::LlmStream(next), + ) + .await?; + let chunks = value.as_array().cloned().ok_or_else(|| { + FlowError::Internal("async stream intercept must resolve to an array".into()) + })?; + Ok(LlmJsonStream::new(tokio_stream::iter( + chunks.into_iter().map(Ok), + ))) + }) + }, + ) +} + // --------------------------------------------------------------------------- // Wrapper functions: C callback -> core trait objects // --------------------------------------------------------------------------- diff --git a/crates/ffi/tests/integration/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index 140f089ce..4f6e90e10 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -661,6 +661,168 @@ fn scope_stack_api_round_trip() { unsafe { nemo_relay_scope_stack_free(stack) }; } +#[test] +fn scope_stack_propagation_and_thread_binding_validate_all_ffi_inputs() { + let _guard = TEST_MUTEX.lock().unwrap(); + + assert_eq!( + unsafe { nemo_relay_capture_propagation_context_json(ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { + nemo_relay_capture_propagation_context_with_root_json(ptr::null(), ptr::null_mut()) + }, + NemoRelayStatus::NullPointer + ); + + let mut inherited_context = ptr::null_mut(); + assert_eq!( + unsafe { nemo_relay_capture_propagation_context_json(&mut inherited_context) }, + NemoRelayStatus::Ok + ); + let inherited_context = unsafe { take_string(inherited_context) }.unwrap(); + assert_eq!( + serde_json::from_str::(&inherited_context).unwrap()["version"], + json!(1) + ); + + let root_uuid = cstring("018f13f0-7c1a-7a80-8000-000000000001"); + let mut rooted_context = ptr::null_mut(); + assert_eq!( + unsafe { + nemo_relay_capture_propagation_context_with_root_json( + root_uuid.as_ptr(), + &mut rooted_context, + ) + }, + NemoRelayStatus::Ok + ); + let rooted_context = unsafe { take_string(rooted_context) }.unwrap(); + assert_eq!( + serde_json::from_str::(&rooted_context).unwrap()["root_uuid"], + json!("018f13f0-7c1a-7a80-8000-000000000001") + ); + + let invalid_root = cstring("not-a-uuid"); + let mut context_json = ptr::null_mut(); + assert_eq!( + unsafe { + nemo_relay_capture_propagation_context_with_root_json( + invalid_root.as_ptr(), + &mut context_json, + ) + }, + NemoRelayStatus::InvalidArg + ); + assert!(context_json.is_null()); + + assert_eq!( + unsafe { + nemo_relay_scope_stack_create_from_propagation_json(ptr::null(), ptr::null_mut()) + }, + NemoRelayStatus::NullPointer + ); + let invalid_context = cstring("not-json"); + let mut stack = ptr::null_mut(); + assert_eq!( + unsafe { + nemo_relay_scope_stack_create_from_propagation_json( + invalid_context.as_ptr(), + &mut stack, + ) + }, + NemoRelayStatus::InvalidJson + ); + assert!(stack.is_null()); + + let rooted_context = cstring(&rooted_context); + assert_eq!( + unsafe { + nemo_relay_scope_stack_create_from_propagation_json(rooted_context.as_ptr(), &mut stack) + }, + NemoRelayStatus::Ok + ); + assert!(!stack.is_null()); + unsafe { nemo_relay_scope_stack_free(stack) }; + + assert_eq!( + unsafe { nemo_relay_scope_stack_set_thread(ptr::null()) }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { nemo_relay_scope_stack_capture_thread(ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { nemo_relay_scope_stack_restore_thread(ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + + let stack = unsafe { fresh_scope_stack() }; + let mut binding = ptr::null_mut(); + assert_eq!( + unsafe { nemo_relay_scope_stack_capture_thread(&mut binding) }, + NemoRelayStatus::Ok + ); + assert!(!binding.is_null()); + assert_eq!( + unsafe { nemo_relay_scope_stack_restore_thread(binding) }, + NemoRelayStatus::Ok + ); + unsafe { nemo_relay_scope_stack_free(stack) }; +} + +#[test] +fn observability_component_helpers_serialize_defaults_and_validate_inputs() { + let _guard = TEST_MUTEX.lock().unwrap(); + + let kind = unsafe { take_string(api::nemo_relay_observability_plugin_kind()) }.unwrap(); + assert_eq!(kind, "observability"); + + assert_eq!( + unsafe { api::nemo_relay_observability_default_config_json(ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + let mut default_config = ptr::null_mut(); + assert_eq!( + unsafe { api::nemo_relay_observability_default_config_json(&mut default_config) }, + NemoRelayStatus::Ok + ); + assert!(unsafe { returned_json(default_config) }.is_object()); + + assert_eq!( + unsafe { + api::nemo_relay_observability_component_spec_json(ptr::null(), true, ptr::null_mut()) + }, + NemoRelayStatus::NullPointer + ); + let mut component = ptr::null_mut(); + assert_eq!( + unsafe { + api::nemo_relay_observability_component_spec_json(ptr::null(), true, &mut component) + }, + NemoRelayStatus::Ok + ); + let component = unsafe { returned_json(component) }; + assert_eq!(component["kind"], json!("observability")); + assert_eq!(component["enabled"], json!(true)); + + let invalid_config = cstring("not-json"); + let mut rejected = ptr::null_mut(); + assert_eq!( + unsafe { + api::nemo_relay_observability_component_spec_json( + invalid_config.as_ptr(), + false, + &mut rejected, + ) + }, + NemoRelayStatus::InvalidJson + ); + assert!(rejected.is_null()); +} + #[test] fn llm_request_accessors_round_trip() { let headers = cstring(r#"{"x-trace":"1"}"#); diff --git a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs index b778b9a2c..7c7e012d0 100644 --- a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs +++ b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs @@ -13,6 +13,279 @@ struct EnvGuard { original: Option, } +unsafe extern "C" fn async_json_registration_callback( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _completion: *const callable::NemoRelayAsyncCompletion, +) -> callable::NemoRelayAsyncCallbackState { + callable::NemoRelayAsyncCallbackState::Pending +} + +unsafe extern "C" fn async_intercept_registration_callback( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _next: *const callable::NemoRelayAsyncNext, + _completion: *const callable::NemoRelayAsyncCompletion, +) -> callable::NemoRelayAsyncCallbackState { + callable::NemoRelayAsyncCallbackState::Pending +} + +#[test] +fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { + let _lock = TEST_MUTEX.lock().unwrap_or_else(|e| e.into_inner()); + reset_globals(); + + macro_rules! global_json { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + name.as_ptr(), + 0, + async_json_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(unsafe { $deregister(name.as_ptr()) }, NemoRelayStatus::Ok); + }}; + } + macro_rules! global_json_with_break_chain { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + name.as_ptr(), + 0, + false, + async_json_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(unsafe { $deregister(name.as_ptr()) }, NemoRelayStatus::Ok); + }}; + } + macro_rules! global_intercept { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + name.as_ptr(), + 0, + async_intercept_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(unsafe { $deregister(name.as_ptr()) }, NemoRelayStatus::Ok); + }}; + } + + global_json!( + nemo_relay_register_mark_sanitize_guardrail_async, + nemo_relay_deregister_mark_sanitize_guardrail + ); + global_json!( + nemo_relay_register_scope_sanitize_start_guardrail_async, + nemo_relay_deregister_scope_sanitize_start_guardrail + ); + global_json!( + nemo_relay_register_scope_sanitize_end_guardrail_async, + nemo_relay_deregister_scope_sanitize_end_guardrail + ); + global_json!( + nemo_relay_register_tool_sanitize_request_guardrail_async, + nemo_relay_deregister_tool_sanitize_request_guardrail + ); + global_json!( + nemo_relay_register_tool_sanitize_response_guardrail_async, + nemo_relay_deregister_tool_sanitize_response_guardrail + ); + global_json!( + nemo_relay_register_tool_conditional_execution_guardrail_async, + nemo_relay_deregister_tool_conditional_execution_guardrail + ); + global_json_with_break_chain!( + nemo_relay_register_tool_request_intercept_async, + nemo_relay_deregister_tool_request_intercept + ); + global_intercept!( + nemo_relay_register_tool_execution_intercept_async, + nemo_relay_deregister_tool_execution_intercept + ); + global_json!( + nemo_relay_register_llm_sanitize_request_guardrail_async, + nemo_relay_deregister_llm_sanitize_request_guardrail + ); + global_json!( + nemo_relay_register_llm_sanitize_response_guardrail_async, + nemo_relay_deregister_llm_sanitize_response_guardrail + ); + global_json!( + nemo_relay_register_llm_conditional_execution_guardrail_async, + nemo_relay_deregister_llm_conditional_execution_guardrail + ); + global_json_with_break_chain!( + nemo_relay_register_llm_request_intercept_async, + nemo_relay_deregister_llm_request_intercept + ); + global_intercept!( + nemo_relay_register_llm_execution_intercept_async, + nemo_relay_deregister_llm_execution_intercept + ); + global_intercept!( + nemo_relay_register_llm_stream_execution_intercept_async, + nemo_relay_deregister_llm_stream_execution_intercept + ); + + let _stack = unsafe { fresh_scope_stack() }; + let mut scope = ptr::null_mut(); + assert_eq!( + unsafe { nemo_relay_get_handle(&mut scope) }, + NemoRelayStatus::Ok + ); + let scope_uuid = cstring( + &unsafe { take_string(nemo_relay_scope_handle_uuid(scope)) } + .expect("root scope UUID should exist"), + ); + + macro_rules! scope_json { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + scope_uuid.as_ptr(), + name.as_ptr(), + 0, + async_json_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { $deregister(scope_uuid.as_ptr(), name.as_ptr()) }, + NemoRelayStatus::Ok + ); + }}; + } + macro_rules! scope_json_with_break_chain { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + scope_uuid.as_ptr(), + name.as_ptr(), + 0, + false, + async_json_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { $deregister(scope_uuid.as_ptr(), name.as_ptr()) }, + NemoRelayStatus::Ok + ); + }}; + } + macro_rules! scope_intercept { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + scope_uuid.as_ptr(), + name.as_ptr(), + 0, + async_intercept_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { $deregister(scope_uuid.as_ptr(), name.as_ptr()) }, + NemoRelayStatus::Ok + ); + }}; + } + + scope_json!( + nemo_relay_scope_register_mark_sanitize_guardrail_async, + nemo_relay_scope_deregister_mark_sanitize_guardrail + ); + scope_json!( + nemo_relay_scope_register_scope_sanitize_start_guardrail_async, + nemo_relay_scope_deregister_scope_sanitize_start_guardrail + ); + scope_json!( + nemo_relay_scope_register_scope_sanitize_end_guardrail_async, + nemo_relay_scope_deregister_scope_sanitize_end_guardrail + ); + scope_json!( + nemo_relay_scope_register_tool_sanitize_request_guardrail_async, + nemo_relay_scope_deregister_tool_sanitize_request_guardrail + ); + scope_json!( + nemo_relay_scope_register_tool_sanitize_response_guardrail_async, + nemo_relay_scope_deregister_tool_sanitize_response_guardrail + ); + scope_json!( + nemo_relay_scope_register_tool_conditional_execution_guardrail_async, + nemo_relay_scope_deregister_tool_conditional_execution_guardrail + ); + scope_json_with_break_chain!( + nemo_relay_scope_register_tool_request_intercept_async, + nemo_relay_scope_deregister_tool_request_intercept + ); + scope_intercept!( + nemo_relay_scope_register_tool_execution_intercept_async, + nemo_relay_scope_deregister_tool_execution_intercept + ); + scope_json!( + nemo_relay_scope_register_llm_sanitize_request_guardrail_async, + nemo_relay_scope_deregister_llm_sanitize_request_guardrail + ); + scope_json!( + nemo_relay_scope_register_llm_sanitize_response_guardrail_async, + nemo_relay_scope_deregister_llm_sanitize_response_guardrail + ); + scope_json!( + nemo_relay_scope_register_llm_conditional_execution_guardrail_async, + nemo_relay_scope_deregister_llm_conditional_execution_guardrail + ); + scope_json_with_break_chain!( + nemo_relay_scope_register_llm_request_intercept_async, + nemo_relay_scope_deregister_llm_request_intercept + ); + scope_intercept!( + nemo_relay_scope_register_llm_execution_intercept_async, + nemo_relay_scope_deregister_llm_execution_intercept + ); + scope_intercept!( + nemo_relay_scope_register_llm_stream_execution_intercept_async, + nemo_relay_scope_deregister_llm_stream_execution_intercept + ); + unsafe { nemo_relay_scope_handle_free(scope) }; +} + impl EnvGuard { fn set(key: &'static str, value: &str) -> Self { let original = std::env::var_os(key); diff --git a/crates/ffi/tests/unit/api/registry_tests.rs b/crates/ffi/tests/unit/api/registry_tests.rs index 4ae23ce05..79cc9c296 100644 --- a/crates/ffi/tests/unit/api/registry_tests.rs +++ b/crates/ffi/tests/unit/api/registry_tests.rs @@ -243,7 +243,8 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { nemo_relay_deregister_mark_sanitize_guardrail(invalid_guard.as_ptr()), NemoRelayStatus::Ok ); - // The queued event retains its sanitizer snapshot after deregistration. + // A queued event owns its sanitizer snapshot until publication. Flush + // before observing the callback-data destructor. assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); assert_eq!(*lock_unpoisoned(plugin_frees()), 4); @@ -408,6 +409,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ); nemo_relay_scope_handle_free(owner); assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_eq!(*lock_unpoisoned(plugin_frees()), 7); let events = lock_unpoisoned(event_log()); let invalid_callback_event = events .iter() diff --git a/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index a2684dd2f..f2aced5ee 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -5,6 +5,35 @@ use super::*; +unsafe extern "C" fn complete_without_settling( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + NemoRelayAsyncCallbackState::Complete +} + +unsafe extern "C" fn send_next_result( + user_data: *mut libc::c_void, + value_json: *const c_char, + error_message: *const c_char, +) { + let sender = unsafe { + Box::from_raw( + user_data.cast::>>(), + ) + }; + let result = if error_message.is_null() { + serde_json::from_str(unsafe { CStr::from_ptr(value_json) }.to_str().unwrap()) + .map_err(|error| error.to_string()) + } else { + Err(unsafe { CStr::from_ptr(error_message) } + .to_string_lossy() + .into_owned()) + }; + let _ = sender.send(result); +} + #[test] fn test_callable_private_helper_paths() { clear_last_error(); @@ -17,3 +46,221 @@ fn test_callable_private_helper_paths() { assert_eq!(ptr_to_opt_string(raw), Some("ffi-string".into())); unsafe { nemo_relay_string_free_internal(raw) }; } + +#[test] +fn async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlements() { + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + let invalid_json = CString::new("not-json").unwrap(); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion_ref, invalid_json.as_ptr()) }, + NemoRelayStatus::InvalidJson + ); + let value = CString::new(r#"{"ok":true}"#).unwrap(); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion_ref, value.as_ptr()) }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion_ref, value.as_ptr()) }, + NemoRelayStatus::InvalidArg + ); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + assert_eq!( + runtime.block_on(receiver).unwrap().unwrap(), + serde_json::json!({"ok": true}) + ); + unsafe { nemo_relay_async_completion_release(completion_ref) }; + + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + assert_eq!( + unsafe { nemo_relay_async_completion_reject(completion_ref, std::ptr::null()) }, + NemoRelayStatus::Ok + ); + assert!( + runtime + .block_on(receiver) + .unwrap() + .unwrap_err() + .to_string() + .contains("async C callback rejected") + ); + unsafe { nemo_relay_async_completion_release(completion_ref) }; + + let (sender, _receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(true), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + assert!(unsafe { nemo_relay_async_completion_is_cancelled(completion_ref) }); + assert!(unsafe { nemo_relay_async_completion_is_cancelled(std::ptr::null()) }); + assert_eq!( + unsafe { nemo_relay_async_completion_reject(completion_ref, std::ptr::null()) }, + NemoRelayStatus::InvalidArg + ); + unsafe { nemo_relay_async_completion_release(completion_ref) }; +} + +#[test] +fn async_callback_wrappers_reject_complete_callbacks_without_settlement() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + + let error = runtime + .block_on(invoke_async_json( + complete_without_settling, + Arc::new(UserData { + ptr: std::ptr::null_mut(), + free_fn: None, + }), + serde_json::json!({}), + )) + .unwrap_err(); + assert!( + error + .to_string() + .contains("returned Complete without settling") + ); +} + +#[test] +fn async_next_invocation_supports_tool_llm_and_stream_continuations() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + + let cases: Vec<(AsyncNextInner, CString, serde_json::Value)> = vec![ + ( + AsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + CString::new(r#"{"tool":true}"#).unwrap(), + serde_json::json!({"tool": true}), + ), + ( + AsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + CString::new( + serde_json::to_string(&LlmRequest { + headers: serde_json::Map::new(), + content: serde_json::json!({"llm": true}), + }) + .unwrap(), + ) + .unwrap(), + serde_json::json!({"llm": true}), + ), + ( + AsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![ + Ok(serde_json::json!({"chunk": 1})), + Ok(serde_json::json!({"chunk": 2})), + ]))) + }) + })), + CString::new( + serde_json::to_string(&LlmRequest { + headers: serde_json::Map::new(), + content: serde_json::json!({"stream": true}), + }) + .unwrap(), + ) + .unwrap(), + serde_json::json!([{"chunk": 1}, {"chunk": 2}]), + ), + ]; + + for (inner, invocation, expected) in cases { + let next = Arc::new(NemoRelayAsyncNext { + inner, + runtime: runtime.handle().clone(), + }); + let next_ref = Arc::into_raw(next); + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + assert_eq!( + unsafe { nemo_relay_async_next_invoke(next_ref, invocation.as_ptr(), completion_ref) }, + NemoRelayStatus::Ok + ); + assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); + unsafe { + nemo_relay_async_next_release(next_ref); + nemo_relay_async_completion_release(completion_ref); + } + } +} + +#[test] +fn async_next_callback_reports_tool_llm_and_stream_results() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let cases: Vec<(AsyncNextInner, CString, Json)> = vec![ + ( + AsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + CString::new(r#"{"tool":true}"#).unwrap(), + serde_json::json!({"tool": true}), + ), + ( + AsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + CString::new(r#"{"headers":{},"content":{"llm":true}}"#).unwrap(), + serde_json::json!({"llm": true}), + ), + ( + AsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( + serde_json::json!({"stream": true}), + )]))) + }) + })), + CString::new(r#"{"headers":{},"content":{}}"#).unwrap(), + serde_json::json!([{ "stream": true }]), + ), + ]; + for (inner, invocation, expected) in cases { + let next = Arc::new(NemoRelayAsyncNext { + inner, + runtime: runtime.handle().clone(), + }); + let next_ref = Arc::into_raw(next); + let (sender, receiver) = + tokio::sync::oneshot::channel::>(); + assert_eq!( + unsafe { + nemo_relay_async_next_invoke_callback( + next_ref, + invocation.as_ptr(), + send_next_result, + Box::into_raw(Box::new(sender)).cast(), + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); + unsafe { nemo_relay_async_next_release(next_ref) }; + } +} diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 4542bb551..65a326f4c 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -18,6 +18,246 @@ extern "C" fn free_arc_counter(user_data: *mut libc::c_void) { counter.fetch_add(1, Ordering::SeqCst); } +unsafe extern "C" fn async_json_passthrough_callback( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + let kind = unsafe { *(user_data.cast::()) }; + let invocation: Json = serde_json::from_str( + unsafe { CStr::from_ptr(invocation_json) } + .to_str() + .expect("invocation must be UTF-8"), + ) + .expect("invocation must be JSON"); + let value = match kind { + 0 => invocation["value"].clone(), + 1 | 3 => Json::Null, + 2 | 4 => invocation["request"].clone(), + 5 => invocation["response"].clone(), + 6 => json!({ + "request": invocation["request"], + "annotated_request": invocation["annotated"], + "pending_marks": [], + "optimization_contributions": [], + }), + 7 => invocation["fields"].clone(), + _ => unreachable!("test callback kind must be known"), + }; + let value = CString::new(value.to_string()).expect("JSON has no NUL"); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, + NemoRelayStatus::Ok + ); + NemoRelayAsyncCallbackState::Complete +} + +fn async_callback_user_data(kind: usize) -> *mut libc::c_void { + Box::into_raw(Box::new(kind)).cast() +} + +unsafe extern "C" fn free_async_callback_user_data(user_data: *mut libc::c_void) { + unsafe { drop(Box::from_raw(user_data.cast::())) }; +} + +unsafe extern "C" fn async_next_callback( + _user_data: *mut libc::c_void, + invocation_json: *const c_char, + next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + let invocation: Json = serde_json::from_str( + unsafe { CStr::from_ptr(invocation_json) } + .to_str() + .expect("invocation must be UTF-8"), + ) + .expect("invocation must be JSON"); + let value = invocation + .get("value") + .or_else(|| invocation.get("request")) + .expect("intercept invocation must carry a value") + .to_string(); + let value = CString::new(value).expect("JSON has no NUL"); + assert_eq!( + unsafe { nemo_relay_async_next_invoke(next, value.as_ptr(), completion) }, + NemoRelayStatus::Ok + ); + unsafe { + nemo_relay_async_next_release(next); + nemo_relay_async_completion_release(completion); + } + NemoRelayAsyncCallbackState::Pending +} + +unsafe extern "C" fn async_tool_outcome_callback( + _user_data: *mut libc::c_void, + invocation_json: *const c_char, + _next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + let invocation: Json = serde_json::from_str( + unsafe { CStr::from_ptr(invocation_json) } + .to_str() + .expect("invocation must be UTF-8"), + ) + .expect("invocation must be JSON"); + let value = + CString::new(json!({"result": invocation["value"], "pending_marks": []}).to_string()) + .expect("JSON has no NUL"); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, + NemoRelayStatus::Ok + ); + NemoRelayAsyncCallbackState::Complete +} + +#[test] +fn async_callback_wrappers_cover_all_middleware_shapes() { + let tool_json = wrap_async_tool_json_fn( + async_json_passthrough_callback, + async_callback_user_data(0), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(tool_json("tool".into(), json!({"value": true}))).unwrap(), + json!({"value": true}) + ); + + let tool_conditional = wrap_async_tool_conditional_fn( + async_json_passthrough_callback, + async_callback_user_data(1), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(tool_conditional("tool".into(), json!({}))).unwrap(), + None + ); + + let llm_conditional = wrap_async_llm_conditional_fn( + async_json_passthrough_callback, + async_callback_user_data(3), + Some(free_async_callback_user_data), + ); + assert_eq!(resolve(llm_conditional(make_request())).unwrap(), None); + + let request_sanitizer = wrap_async_llm_sanitize_request_fn( + async_json_passthrough_callback, + async_callback_user_data(4), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(request_sanitizer( + make_request(), + nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), + )) + .unwrap(), + Some(make_request()) + ); + + let response_sanitizer = wrap_async_llm_sanitize_response_fn( + async_json_passthrough_callback, + async_callback_user_data(5), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(response_sanitizer( + json!({"response": true}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), + )) + .unwrap(), + Some(json!({"response": true})) + ); + + let request_intercept = wrap_async_llm_request_intercept_fn( + async_json_passthrough_callback, + async_callback_user_data(6), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(request_intercept("llm".into(), make_request(), None)) + .unwrap() + .request, + make_request() + ); + + let event = Event::Scope(nemo_relay::api::event::ScopeEvent::new( + nemo_relay::api::event::BaseEvent::builder() + .name("async-event") + .build(), + nemo_relay::api::event::ScopeCategory::Start, + Vec::new(), + nemo_relay::api::event::EventCategory::llm(), + None, + )); + let fields = EventSanitizeFields::builder() + .data(json!({"safe": true})) + .build(); + let event_sanitizer = wrap_async_event_sanitize_fn( + async_json_passthrough_callback, + async_callback_user_data(7), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(event_sanitizer(event, fields.clone())).unwrap(), + fields + ); +} + +#[test] +fn async_execution_wrappers_continue_tool_and_llm_calls() { + let tool_intercept = wrap_async_tool_execution_intercept_fn( + async_tool_outcome_callback, + std::ptr::null_mut(), + None, + ); + let tool_next: ToolExecutionNextFn = Arc::new(|args| Box::pin(async move { Ok(args) })); + assert_eq!( + resolve(tool_intercept("tool", json!({"ok": true}), tool_next)) + .unwrap() + .result, + json!({"ok": true}) + ); + + let llm_intercept = + wrap_async_llm_execution_intercept_fn(async_next_callback, std::ptr::null_mut(), None); + let llm_next: LlmExecutionNextFn = + Arc::new(|request| Box::pin(async move { Ok(json!({"model": request.content["model"]})) })); + assert_eq!( + resolve(llm_intercept("llm", make_request(), llm_next)).unwrap(), + json!({"model": "test-model"}) + ); +} + +#[test] +fn async_stream_execution_wrapper_collects_the_continued_stream() { + let intercept = wrap_async_llm_stream_execution_intercept_fn( + async_next_callback, + std::ptr::null_mut(), + None, + ); + let next: LlmStreamExecutionNextFn = Arc::new(|request| { + Box::pin(async move { + Ok(nemo_relay::api::runtime::LlmJsonStream::new( + tokio_stream::iter(vec![ + Ok(json!({"model": request.content["model"], "chunk": 1})), + Ok(json!({"chunk": 2})), + ]), + )) + }) + }); + + let mut stream = resolve(intercept("llm", make_request(), next)).unwrap(); + assert_eq!( + resolve(async { stream.next().await.unwrap().unwrap() }), + json!({"model": "test-model", "chunk": 1}) + ); + assert_eq!( + resolve(async { stream.next().await.unwrap().unwrap() }), + json!({"chunk": 2}) + ); + assert!(resolve(async { stream.next().await }).is_none()); +} + fn user_data_counter() -> (*mut libc::c_void, Arc) { let counter = Arc::new(AtomicUsize::new(0)); let ptr = Box::into_raw(Box::new(counter.clone())) as *mut libc::c_void; @@ -498,7 +738,7 @@ fn test_llm_sanitizers_report_runtime_codec_ids_with_embedded_nul() { runtime_identity.clone(), ), )) - .unwrap_err(); + .expect_err("an embedded runtime codec ID must fail the async callback wrapper"); assert!( request_error .to_string() @@ -516,7 +756,7 @@ fn test_llm_sanitizers_report_runtime_codec_ids_with_embedded_nul() { json!({"secret": "must be preserved"}), nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), )) - .unwrap_err(); + .expect_err("an embedded runtime codec ID must fail the async callback wrapper"); assert!( response_error .to_string() diff --git a/go/nemo_relay/adaptive_runtime_test.go b/go/nemo_relay/adaptive_runtime_test.go index d63dbcda5..018ea3d11 100644 --- a/go/nemo_relay/adaptive_runtime_test.go +++ b/go/nemo_relay/adaptive_runtime_test.go @@ -172,12 +172,36 @@ func TestSetLatencySensitivityRejectsInvalidValue(t *testing.T) { } } +func TestAdaptiveRuntimeRejectsNilHandles(t *testing.T) { + var nilRuntime *AdaptiveRuntime + if err := nilRuntime.Register(); err == nil { + t.Fatal("expected nil Register to fail") + } + if err := nilRuntime.Deregister(); err == nil { + t.Fatal("expected nil Deregister to fail") + } + if err := nilRuntime.Shutdown(); err == nil { + t.Fatal("expected nil Shutdown to fail") + } + if err := nilRuntime.WaitForIdle(); err == nil { + t.Fatal("expected nil WaitForIdle to fail") + } + if _, err := nilRuntime.Report(); err == nil { + t.Fatal("expected nil Report to fail") + } + if err := nilRuntime.BindScope(nil); err == nil { + t.Fatal("expected nil BindScope to fail") + } + if _, err := nilRuntime.BuildCacheRequestFacts(CacheRequestFactsInput{}); err == nil { + t.Fatal("expected nil BuildCacheRequestFacts to fail") + } +} + func TestAdaptiveRuntimeLifecycleRejectsUseAfterShutdown(t *testing.T) { - runtime, err := NewAdaptiveRuntime(testAdaptiveRuntimeConfig("openai")) + runtime, err := NewAdaptiveRuntime(NewAdaptiveConfig()) if err != nil { t.Fatalf(newAdaptiveRuntimeFailedMsg, err) } - if err := runtime.Shutdown(); err != nil { t.Fatalf("Shutdown failed: %v", err) } @@ -213,6 +237,33 @@ func assertAdaptiveRuntimeClosed(t *testing.T, runtime *AdaptiveRuntime) { } } +func TestAdaptiveRuntimeHelpersRejectInvalidInputs(t *testing.T) { + if _, err := BuildCacheTelemetryEvent(CacheTelemetryEventInput{ + Provider: "unsupported", + RequestID: "not-a-uuid", + }); err == nil { + t.Fatal("expected invalid telemetry input to fail") + } + + runtime, err := NewAdaptiveRuntime(testAdaptiveRuntimeConfig("openai")) + if err != nil { + t.Fatalf(newAdaptiveRuntimeFailedMsg, err) + } + defer runtime.Shutdown() + if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{ + Provider: "unsupported", + RequestID: "not-a-uuid", + AnnotatedRequest: json.RawMessage(`{}`), + }); err == nil { + t.Fatal("expected invalid cache request facts input to fail") + } + if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{ + AnnotatedRequest: json.RawMessage(`not-json`), + }); err == nil { + t.Fatal("expected malformed annotated request JSON to fail before the FFI call") + } +} + func TestAdaptiveRuntimePublicHelpersPropagateJSONMarshalFailures(t *testing.T) { oldMarshal := jsonMarshal t.Cleanup(func() { jsonMarshal = oldMarshal }) diff --git a/go/nemo_relay/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go new file mode 100644 index 000000000..ee501908b --- /dev/null +++ b/go/nemo_relay/async_middleware_test.go @@ -0,0 +1,224 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +package nemo_relay + +import ( + "context" + "encoding/json" + "errors" + "strings" + "testing" +) + +func asyncMiddlewareNoop(context.Context, json.RawMessage) (any, error) { + return nil, nil +} + +func asyncExecutionNoop(context.Context, json.RawMessage, AsyncNext) (any, error) { + return nil, nil +} + +func TestAsyncMiddlewareGlobalRegistrationParity(t *testing.T) { + registrations := []struct { + name string + register func(string) error + deregister func(string) error + }{ + {"mark", func(name string) error { return RegisterMarkSanitizeGuardrailAsync(name, 0, asyncMiddlewareNoop) }, DeregisterMarkSanitizeGuardrail}, + {"scope-start", func(name string) error { return RegisterScopeSanitizeStartGuardrailAsync(name, 0, asyncMiddlewareNoop) }, DeregisterScopeSanitizeStartGuardrail}, + {"scope-end", func(name string) error { return RegisterScopeSanitizeEndGuardrailAsync(name, 0, asyncMiddlewareNoop) }, DeregisterScopeSanitizeEndGuardrail}, + {"tool-sanitize-request", func(name string) error { + return RegisterToolSanitizeRequestGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterToolSanitizeRequestGuardrail}, + {"tool-sanitize-response", func(name string) error { + return RegisterToolSanitizeResponseGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterToolSanitizeResponseGuardrail}, + {"tool-conditional", func(name string) error { + return RegisterToolConditionalExecutionGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterToolConditionalExecutionGuardrail}, + {"tool-request", func(name string) error { return RegisterToolRequestInterceptAsync(name, 0, false, asyncMiddlewareNoop) }, DeregisterToolRequestIntercept}, + {"tool-execution", func(name string) error { return RegisterToolExecutionInterceptAsync(name, 0, asyncExecutionNoop) }, DeregisterToolExecutionIntercept}, + {"llm-sanitize-request", func(name string) error { return RegisterLlmSanitizeRequestGuardrailAsync(name, 0, asyncMiddlewareNoop) }, DeregisterLlmSanitizeRequestGuardrail}, + {"llm-sanitize-response", func(name string) error { + return RegisterLlmSanitizeResponseGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterLlmSanitizeResponseGuardrail}, + {"llm-conditional", func(name string) error { + return RegisterLlmConditionalExecutionGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterLlmConditionalExecutionGuardrail}, + {"llm-request", func(name string) error { return RegisterLlmRequestInterceptAsync(name, 0, false, asyncMiddlewareNoop) }, DeregisterLlmRequestIntercept}, + {"llm-execution", func(name string) error { return RegisterLlmExecutionInterceptAsync(name, 0, asyncExecutionNoop) }, DeregisterLlmExecutionIntercept}, + {"llm-stream-execution", func(name string) error { return RegisterLlmStreamExecutionInterceptAsync(name, 0, asyncExecutionNoop) }, DeregisterLlmStreamExecutionIntercept}, + } + + for _, registration := range registrations { + t.Run(registration.name, func(t *testing.T) { + name := "go-async-global-" + registration.name + if err := registration.register(name); err != nil { + t.Fatalf("register: %v", err) + } + if err := registration.deregister(name); err != nil { + t.Fatalf("deregister: %v", err) + } + }) + } +} + +func TestAsyncMiddlewareScopeLocalRegistrationParity(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + handle, err := PushScope("async-registration-owner", ScopeTypeAgent) + if err != nil { + t.Fatalf("push scope: %v", err) + } + defer func() { + if err := PopScope(handle); err != nil { + t.Fatalf("pop scope: %v", err) + } + }() + + scopeUUID := handle.UUID() + registrations := []struct { + name string + register func(string) error + deregister func(string) error + }{ + {"mark", func(name string) error { + return ScopeRegisterMarkSanitizeGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterMarkSanitizeGuardrail(scopeUUID, name) }}, + {"scope-start", func(name string) error { + return ScopeRegisterScopeSanitizeStartGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterScopeSanitizeStartGuardrail(scopeUUID, name) }}, + {"scope-end", func(name string) error { + return ScopeRegisterScopeSanitizeEndGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterScopeSanitizeEndGuardrail(scopeUUID, name) }}, + {"tool-sanitize-request", func(name string) error { + return ScopeRegisterToolSanitizeRequestGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterToolSanitizeRequestGuardrail(scopeUUID, name) }}, + {"tool-sanitize-response", func(name string) error { + return ScopeRegisterToolSanitizeResponseGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterToolSanitizeResponseGuardrail(scopeUUID, name) }}, + {"tool-conditional", func(name string) error { + return ScopeRegisterToolConditionalExecutionGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterToolConditionalExecutionGuardrail(scopeUUID, name) }}, + {"tool-request", func(name string) error { + return ScopeRegisterToolRequestInterceptAsync(scopeUUID, name, 0, false, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterToolRequestIntercept(scopeUUID, name) }}, + {"tool-execution", func(name string) error { + return ScopeRegisterToolExecutionInterceptAsync(scopeUUID, name, 0, asyncExecutionNoop) + }, func(name string) error { return ScopeDeregisterToolExecutionIntercept(scopeUUID, name) }}, + {"llm-sanitize-request", func(name string) error { + return ScopeRegisterLlmSanitizeRequestGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterLlmSanitizeRequestGuardrail(scopeUUID, name) }}, + {"llm-sanitize-response", func(name string) error { + return ScopeRegisterLlmSanitizeResponseGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterLlmSanitizeResponseGuardrail(scopeUUID, name) }}, + {"llm-conditional", func(name string) error { + return ScopeRegisterLlmConditionalExecutionGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterLlmConditionalExecutionGuardrail(scopeUUID, name) }}, + {"llm-request", func(name string) error { + return ScopeRegisterLlmRequestInterceptAsync(scopeUUID, name, 0, false, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterLlmRequestIntercept(scopeUUID, name) }}, + {"llm-execution", func(name string) error { + return ScopeRegisterLlmExecutionInterceptAsync(scopeUUID, name, 0, asyncExecutionNoop) + }, func(name string) error { return ScopeDeregisterLlmExecutionIntercept(scopeUUID, name) }}, + {"llm-stream-execution", func(name string) error { + return ScopeRegisterLlmStreamExecutionInterceptAsync(scopeUUID, name, 0, asyncExecutionNoop) + }, func(name string) error { return ScopeDeregisterLlmStreamExecutionIntercept(scopeUUID, name) }}, + } + + for _, registration := range registrations { + name := "go-async-local-" + registration.name + if err := registration.register(name); err != nil { + t.Fatalf("register %s: %v", registration.name, err) + } + if err := registration.deregister(name); err != nil { + t.Fatalf("deregister %s: %v", registration.name, err) + } + } + }) +} + +func TestAsyncToolMiddlewareCompletionAndNext(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + if err := RegisterToolConditionalExecutionGuardrailAsync("go-async-tool-conditional", 0, asyncMiddlewareNoop); err != nil { + t.Fatalf("register conditional: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolConditionalExecutionGuardrail("go-async-tool-conditional") }) + + if err := RegisterToolExecutionInterceptAsync("go-async-tool-execution", 0, + func(ctx context.Context, invocation json.RawMessage, next AsyncNext) (any, error) { + var payload struct { + Value json.RawMessage `json:"value"` + } + if err := json.Unmarshal(invocation, &payload); err != nil { + return nil, err + } + result, err := next(ctx, payload.Value) + if err != nil { + return nil, err + } + return map[string]json.RawMessage{"result": result}, nil + }, + ); err != nil { + t.Fatalf("register execution intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolExecutionIntercept("go-async-tool-execution") }) + + result, err := ToolCallExecute("go-async-tool", json.RawMessage(`{"value":1}`), func(args json.RawMessage) (json.RawMessage, error) { + return args, nil + }) + if err != nil { + t.Fatalf("tool call execute: %v", err) + } + if string(result) != `{"value":1}` { + t.Fatalf("tool result = %s, want original result", result) + } + }) +} + +func TestAsyncToolMiddlewarePropagatesCallbackAndNextErrors(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + const conditionalName = "go-async-tool-conditional-error" + if err := RegisterToolConditionalExecutionGuardrailAsync(conditionalName, 0, + func(context.Context, json.RawMessage) (any, error) { + return nil, errors.New("conditional callback failed") + }, + ); err != nil { + t.Fatalf("register conditional: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolConditionalExecutionGuardrail(conditionalName) }) + + _, err := ToolCallExecute("go-async-tool-conditional-error", json.RawMessage(`{}`), func(json.RawMessage) (json.RawMessage, error) { + return json.RawMessage(`{}`), nil + }) + if err == nil || !strings.Contains(err.Error(), "conditional callback failed") { + t.Fatalf("conditional error = %v, want callback failure", err) + } + + if err := DeregisterToolConditionalExecutionGuardrail(conditionalName); err != nil { + t.Fatalf("deregister conditional: %v", err) + } + const executionName = "go-async-tool-next-error" + if err := RegisterToolExecutionInterceptAsync(executionName, 0, + func(ctx context.Context, invocation json.RawMessage, next AsyncNext) (any, error) { + var payload struct { + Value json.RawMessage `json:"value"` + } + if err := json.Unmarshal(invocation, &payload); err != nil { + return nil, err + } + return next(ctx, payload.Value) + }, + ); err != nil { + t.Fatalf("register execution intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolExecutionIntercept(executionName) }) + + _, err = ToolCallExecute("go-async-tool-next-error", json.RawMessage(`{}`), func(json.RawMessage) (json.RawMessage, error) { + return nil, errors.New("tool implementation failed") + }) + if err == nil || !strings.Contains(err.Error(), "tool implementation failed") { + t.Fatalf("next error = %v, want implementation failure", err) + } + }) +} diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index 8096c7447..f7837ab2b 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -50,6 +50,16 @@ typedef char* (*NemoRelayLlmSanitizeResponseCb)(void* user_data, const char* res typedef void (*NemoRelayEventSubscriberFn)(void* user_data, const FfiEvent* event); typedef char* (*NemoRelayEventSanitizeFn)(void* user_data, const FfiEvent* event, const char* fields_json); typedef struct FfiPluginContext FfiPluginContext; +typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; +typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; +typedef void (*NemoRelayAsyncNextResultCb)(void*, const char*, const char*); +extern int32_t nemo_relay_async_completion_resolve_json(const NemoRelayAsyncCompletion*, const char*); +extern int32_t nemo_relay_async_completion_reject(const NemoRelayAsyncCompletion*, const char*); +extern bool nemo_relay_async_completion_is_cancelled(const NemoRelayAsyncCompletion*); +extern void nemo_relay_async_completion_release(const NemoRelayAsyncCompletion*); +extern int32_t nemo_relay_async_next_invoke_callback(const NemoRelayAsyncNext*, const char*, NemoRelayAsyncNextResultCb, void*); +extern void nemo_relay_async_next_release(const NemoRelayAsyncNext*); +extern void goAsyncNextResultTrampoline(void*, char*, char*); // Middleware chain next function types typedef char* (*NemoRelayToolExecNextFn)(const char* args_json, void* next_ctx); @@ -87,10 +97,12 @@ typedef NemoRelayCodecEncodeCb NemoRelayCodecEncodeFn; import "C" import ( + "context" "encoding/json" "errors" "sync" "sync/atomic" + "time" "unsafe" ) @@ -167,6 +179,43 @@ type ToolSanitizeFunc func(name string, args json.RawMessage) json.RawMessage // message string to reject the call. type ToolConditionalFunc func(name string, args json.RawMessage) *string +// AsyncMiddlewareFunc is the common completion-based middleware callback. +// The JSON envelope identifies the middleware family and invocation fields. +type AsyncMiddlewareFunc func(ctx context.Context, invocation json.RawMessage) (any, error) + +// AsyncNext invokes the remaining execution chain and returns its eventual result. +type AsyncNext func(ctx context.Context, invocation json.RawMessage) (json.RawMessage, error) + +// AsyncExecutionInterceptFunc is an asynchronous execution intercept with an awaitable next helper. +type AsyncExecutionInterceptFunc func(ctx context.Context, invocation json.RawMessage, next AsyncNext) (any, error) + +const asyncCallbackPending = C.uint32_t(1) + +func contextForCompletion(completion *C.NemoRelayAsyncCompletion) (context.Context, func()) { + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + var doneOnce sync.Once + go func() { + ticker := time.NewTicker(10 * time.Millisecond) + defer ticker.Stop() + for { + select { + case <-done: + return + case <-ticker.C: + if bool(C.nemo_relay_async_completion_is_cancelled(completion)) { + cancel() + return + } + } + } + }() + return ctx, func() { + doneOnce.Do(func() { close(done) }) + cancel() + } +} + // ToolExecutionFunc is a callback that executes a tool call, receiving the // arguments as JSON and returning the result JSON or an error. type ToolExecutionFunc func(args json.RawMessage) (json.RawMessage, error) @@ -590,6 +639,137 @@ func goToolSanitizeTrampoline(userData unsafe.Pointer, name *C.char, argsJSON *C return C.CString(string(result)) } +//export goAsyncMiddlewareTrampoline +func goAsyncMiddlewareTrampoline(userData unsafe.Pointer, invocationJSON *C.char, completion *C.NemoRelayAsyncCompletion) C.uint32_t { + fn, ok := lookupClosure(userData).(AsyncMiddlewareFunc) + if !ok { + message := C.CString("nemo_relay: async middleware callback is not registered") + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) + C.nemo_relay_async_completion_release(completion) + return asyncCallbackPending + } + invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) + go func() { + defer C.nemo_relay_async_completion_release(completion) + ctx, cancel := contextForCompletion(completion) + defer cancel() + value, err := fn(ctx, invocation) + if err != nil { + message := C.CString(err.Error()) + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) + return + } + encoded, err := json.Marshal(value) + if err != nil { + message := C.CString(err.Error()) + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) + return + } + result := C.CString(string(encoded)) + defer C.free(unsafe.Pointer(result)) + C.nemo_relay_async_completion_resolve_json(completion, result) + }() + return asyncCallbackPending +} + +type asyncNextResult struct { + value json.RawMessage + err error +} + +//export goAsyncNextResultTrampoline +func goAsyncNextResultTrampoline(userData unsafe.Pointer, valueJSON *C.char, errorMessage *C.char) { + ch, ok := lookupClosure(userData).(chan asyncNextResult) + if !ok { + return + } + defer unregisterClosure(userData) + if errorMessage != nil { + select { + case ch <- asyncNextResult{err: errors.New(C.GoString(errorMessage))}: + default: + } + return + } + select { + case ch <- asyncNextResult{value: append(json.RawMessage(nil), []byte(C.GoString(valueJSON))...)}: + default: + } +} + +//export goAsyncExecutionInterceptTrampoline +func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON *C.char, next *C.NemoRelayAsyncNext, completion *C.NemoRelayAsyncCompletion) C.uint32_t { + fn, ok := lookupClosure(userData).(AsyncExecutionInterceptFunc) + if !ok { + message := C.CString("nemo_relay: async execution intercept callback is not registered") + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) + C.nemo_relay_async_completion_release(completion) + C.nemo_relay_async_next_release(next) + return asyncCallbackPending + } + invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) + go func() { + defer C.nemo_relay_async_completion_release(completion) + ctx, cancel := contextForCompletion(completion) + defer cancel() + var nextMu sync.RWMutex + nextOpen := true + defer func() { + nextMu.Lock() + nextOpen = false + nextMu.Unlock() + C.nemo_relay_async_next_release(next) + }() + nextFn := func(ctx context.Context, payload json.RawMessage) (json.RawMessage, error) { + nextMu.RLock() + defer nextMu.RUnlock() + if !nextOpen { + return nil, context.Canceled + } + ch := make(chan asyncNextResult, 1) + token := registerClosure(ch) + cPayload := C.CString(string(payload)) + status := C.nemo_relay_async_next_invoke_callback( + next, cPayload, + (C.NemoRelayAsyncNextResultCb)(C.goAsyncNextResultTrampoline), token, + ) + C.free(unsafe.Pointer(cPayload)) + if err := checkStatus(status); err != nil { + unregisterClosure(token) + return nil, err + } + select { + case result := <-ch: + return result.value, result.err + case <-ctx.Done(): + return nil, ctx.Err() + } + } + value, err := fn(ctx, invocation, nextFn) + if err != nil { + message := C.CString(err.Error()) + C.nemo_relay_async_completion_reject(completion, message) + C.free(unsafe.Pointer(message)) + return + } + encoded, err := json.Marshal(value) + if err != nil { + message := C.CString(err.Error()) + C.nemo_relay_async_completion_reject(completion, message) + C.free(unsafe.Pointer(message)) + return + } + result := C.CString(string(encoded)) + C.nemo_relay_async_completion_resolve_json(completion, result) + C.free(unsafe.Pointer(result)) + }() + return asyncCallbackPending +} + //export goToolConditionalTrampoline func goToolConditionalTrampoline(userData unsafe.Pointer, name *C.char, argsJSON *C.char) *C.char { fn := lookupClosure(userData).(ToolConditionalFunc) diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index 978e40228..52ab649f0 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -44,6 +44,10 @@ typedef struct NemoRelayLlmSanitizeRequestContext { uint32_t codec_kind; const c typedef struct NemoRelayLlmSanitizeResponseContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeResponseCodec* codec; } NemoRelayLlmSanitizeResponseContext; typedef void (*NemoRelayFreeFn)(void* user_data); +typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; +typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; +typedef uint32_t (*NemoRelayAsyncJsonCb)(void*, const char*, const NemoRelayAsyncCompletion*); +typedef uint32_t (*NemoRelayAsyncInterceptCb)(void*, const char*, const NemoRelayAsyncNext*, const NemoRelayAsyncCompletion*); // Core API extern int32_t nemo_relay_get_handle(FfiScopeHandle** out); @@ -121,46 +125,57 @@ extern void nemo_relay_set_last_error_message(const char* msg); // Tool guardrails typedef char* (*NemoRelayToolSanitizeFn)(void* user_data, const char* name, const char* args_json); extern int32_t nemo_relay_register_tool_sanitize_request_guardrail(const char* name, int32_t priority, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_tool_sanitize_request_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_sanitize_request_guardrail(const char* name); extern int32_t nemo_relay_register_tool_sanitize_response_guardrail(const char* name, int32_t priority, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_tool_sanitize_response_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_sanitize_response_guardrail(const char* name); typedef char* (*NemoRelayToolConditionalFn)(void* user_data, const char* name, const char* args_json); extern int32_t nemo_relay_register_tool_conditional_execution_guardrail(const char* name, int32_t priority, NemoRelayToolConditionalFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_tool_conditional_execution_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_conditional_execution_guardrail(const char* name); // Tool intercepts extern int32_t nemo_relay_register_tool_request_intercept(const char* name, int32_t priority, _Bool break_chain, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_tool_request_intercept_async(const char*, int32_t, _Bool, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_request_intercept(const char* name); // Middleware chain intercept callback types (must be declared before use in externs) typedef char* (*NemoRelayToolExecNextFn)(const char* args_json, void* next_ctx); typedef char* (*NemoRelayToolExecInterceptCb)(void* user_data, const char* args_json, NemoRelayToolExecNextFn next_fn, void* next_ctx); extern int32_t nemo_relay_register_tool_execution_intercept(const char* name, int32_t priority, NemoRelayToolExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_register_tool_execution_intercept_async(const char*, int32_t, NemoRelayAsyncInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_execution_intercept(const char* name); // LLM guardrails typedef FfiLLMRequest* (*NemoRelayLlmSanitizeRequestCb)(void* user_data, const FfiLLMRequest* request, NemoRelayLlmSanitizeRequestContext context); extern int32_t nemo_relay_register_llm_sanitize_request_guardrail(const char* name, int32_t priority, NemoRelayLlmSanitizeRequestCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_llm_sanitize_request_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_sanitize_request_guardrail(const char* name); typedef char* (*NemoRelayLlmSanitizeResponseCb)(void* user_data, const char* response_json, NemoRelayLlmSanitizeResponseContext context); extern int32_t nemo_relay_register_llm_sanitize_response_guardrail(const char* name, int32_t priority, NemoRelayLlmSanitizeResponseCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_llm_sanitize_response_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_sanitize_response_guardrail(const char* name); typedef char* (*NemoRelayLlmConditionalCb)(void* user_data, const FfiLLMRequest* request); extern int32_t nemo_relay_register_llm_conditional_execution_guardrail(const char* name, int32_t priority, NemoRelayLlmConditionalCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_llm_conditional_execution_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_conditional_execution_guardrail(const char* name); // LLM intercepts typedef int32_t (*NemoRelayLlmRequestInterceptCb)(void* user_data, const char* name, const FfiLLMRequest* request, const char* annotated_json, char** out_outcome_json); extern int32_t nemo_relay_register_llm_request_intercept(const char* name, int32_t priority, _Bool break_chain, NemoRelayLlmRequestInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_llm_request_intercept_async(const char*, int32_t, _Bool, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_request_intercept(const char* name); typedef char* (*NemoRelayLlmExecNextFn)(const char* native_json, void* next_ctx); typedef char* (*NemoRelayLlmExecInterceptCb)(void* user_data, const char* native_json, NemoRelayLlmExecNextFn next_fn, void* next_ctx); extern int32_t nemo_relay_register_llm_execution_intercept(const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_register_llm_execution_intercept_async(const char*, int32_t, NemoRelayAsyncInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_execution_intercept(const char* name); extern int32_t nemo_relay_register_llm_stream_execution_intercept(const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_register_llm_stream_execution_intercept_async(const char*, int32_t, NemoRelayAsyncInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_stream_execution_intercept(const char* name); // Subscribers @@ -170,46 +185,63 @@ extern int32_t nemo_relay_register_subscriber(const char* name, NemoRelayEventSu extern int32_t nemo_relay_deregister_subscriber(const char* name); extern int32_t nemo_relay_flush_subscribers(void); extern int32_t nemo_relay_register_mark_sanitize_guardrail(const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_mark_sanitize_guardrail_async(const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_deregister_mark_sanitize_guardrail(const char* name); extern int32_t nemo_relay_register_scope_sanitize_start_guardrail(const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_scope_sanitize_start_guardrail_async(const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_deregister_scope_sanitize_start_guardrail(const char* name); extern int32_t nemo_relay_register_scope_sanitize_end_guardrail(const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_scope_sanitize_end_guardrail_async(const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_deregister_scope_sanitize_end_guardrail(const char* name); // Scope-local tool guardrails extern int32_t nemo_relay_scope_register_mark_sanitize_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_mark_sanitize_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_mark_sanitize_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_scope_sanitize_start_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_scope_sanitize_start_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_scope_sanitize_start_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_scope_sanitize_end_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_scope_sanitize_end_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_scope_sanitize_end_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_tool_sanitize_request_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_tool_sanitize_request_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_sanitize_request_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_tool_sanitize_response_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_tool_sanitize_response_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_sanitize_response_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_tool_conditional_execution_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayToolConditionalFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_tool_conditional_execution_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_conditional_execution_guardrail(const char* scope_uuid, const char* name); // Scope-local tool intercepts extern int32_t nemo_relay_scope_register_tool_request_intercept(const char* scope_uuid, const char* name, int32_t priority, _Bool break_chain, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_tool_request_intercept_async(const char* scope_uuid, const char* name, int32_t priority, _Bool break_chain, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_request_intercept(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_tool_execution_intercept(const char* scope_uuid, const char* name, int32_t priority, NemoRelayToolExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_scope_register_tool_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_execution_intercept(const char* scope_uuid, const char* name); // Scope-local LLM guardrails extern int32_t nemo_relay_scope_register_llm_sanitize_request_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmSanitizeRequestCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_sanitize_request_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_sanitize_request_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_sanitize_response_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmSanitizeResponseCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_sanitize_response_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_sanitize_response_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_conditional_execution_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmConditionalCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_conditional_execution_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_conditional_execution_guardrail(const char* scope_uuid, const char* name); // Scope-local LLM intercepts extern int32_t nemo_relay_scope_register_llm_request_intercept(const char* scope_uuid, const char* name, int32_t priority, _Bool break_chain, NemoRelayLlmRequestInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_request_intercept_async(const char* scope_uuid, const char* name, int32_t priority, _Bool break_chain, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_request_intercept(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_execution_intercept(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_scope_register_llm_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_execution_intercept(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_stream_execution_intercept(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_scope_register_llm_stream_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_stream_execution_intercept(const char* scope_uuid, const char* name); // Scope-local subscribers @@ -267,6 +299,8 @@ extern void nemo_relay_otel_subscriber_free(void*); // Go trampoline forward declarations (defined via //export in callbacks.go) extern char* goToolSanitizeTrampoline(void*, const char*, const char*); +extern uint32_t goAsyncMiddlewareTrampoline(void*, const char*, const NemoRelayAsyncCompletion*); +extern uint32_t goAsyncExecutionInterceptTrampoline(void*, const char*, const NemoRelayAsyncNext*, const NemoRelayAsyncCompletion*); extern char* goEventSanitizeTrampoline(void*, const FfiEvent*, const char*); extern char* goToolConditionalTrampoline(void*, const char*, const char*); extern char* goToolExecTrampoline(void*, const char*); @@ -1198,11 +1232,34 @@ func registerEventSanitizer(name string, priority int32, fn EventSanitizeFunc, k return checkStatus(status) } +func registerAsyncEventSanitizer(name string, priority int32, fn AsyncMiddlewareFunc, kind int) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + callback := C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline) + free := C.NemoRelayFreeFn(C.goFreeTrampoline) + var status C.int32_t + switch kind { + case 0: + status = C.nemo_relay_register_mark_sanitize_guardrail_async(cName, C.int32_t(priority), callback, id, free) + case 1: + status = C.nemo_relay_register_scope_sanitize_start_guardrail_async(cName, C.int32_t(priority), callback, id, free) + default: + status = C.nemo_relay_register_scope_sanitize_end_guardrail_async(cName, C.int32_t(priority), callback, id, free) + } + return checkStatus(status) +} + // RegisterMarkSanitizeGuardrail registers a global mark event sanitizer. func RegisterMarkSanitizeGuardrail(name string, priority int32, fn EventSanitizeFunc) error { return registerEventSanitizer(name, priority, fn, 0) } +// RegisterMarkSanitizeGuardrailAsync registers an asynchronous global mark sanitizer. +func RegisterMarkSanitizeGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + return registerAsyncEventSanitizer(name, priority, fn, 0) +} + // DeregisterMarkSanitizeGuardrail removes a global mark event sanitizer. func DeregisterMarkSanitizeGuardrail(name string) error { cName := C.CString(name) @@ -1215,6 +1272,11 @@ func RegisterScopeSanitizeStartGuardrail(name string, priority int32, fn EventSa return registerEventSanitizer(name, priority, fn, 1) } +// RegisterScopeSanitizeStartGuardrailAsync registers an asynchronous global scope-start sanitizer. +func RegisterScopeSanitizeStartGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + return registerAsyncEventSanitizer(name, priority, fn, 1) +} + // DeregisterScopeSanitizeStartGuardrail removes a global scope-start event sanitizer. func DeregisterScopeSanitizeStartGuardrail(name string) error { cName := C.CString(name) @@ -1227,6 +1289,11 @@ func RegisterScopeSanitizeEndGuardrail(name string, priority int32, fn EventSani return registerEventSanitizer(name, priority, fn, 2) } +// RegisterScopeSanitizeEndGuardrailAsync registers an asynchronous global scope-end sanitizer. +func RegisterScopeSanitizeEndGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + return registerAsyncEventSanitizer(name, priority, fn, 2) +} + // DeregisterScopeSanitizeEndGuardrail removes a global scope-end event sanitizer. func DeregisterScopeSanitizeEndGuardrail(name string) error { cName := C.CString(name) @@ -1252,6 +1319,17 @@ func RegisterToolSanitizeRequestGuardrail(name string, priority int32, fn ToolSa )) } +// RegisterToolSanitizeRequestGuardrailAsync registers an asynchronous tool request sanitizer. +func RegisterToolSanitizeRequestGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_sanitize_request_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolSanitizeRequestGuardrail removes a previously registered tool // sanitize-request guardrail by name. Returns a NotFound error if no guardrail // with the given name is registered. @@ -1277,6 +1355,17 @@ func RegisterToolSanitizeResponseGuardrail(name string, priority int32, fn ToolS )) } +// RegisterToolSanitizeResponseGuardrailAsync registers an asynchronous tool response sanitizer. +func RegisterToolSanitizeResponseGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_sanitize_response_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolSanitizeResponseGuardrail removes a previously registered tool // sanitize-response guardrail by name. Returns a NotFound error if no guardrail // with the given name is registered. @@ -1304,6 +1393,17 @@ func RegisterToolConditionalExecutionGuardrail(name string, priority int32, fn T )) } +// RegisterToolConditionalExecutionGuardrailAsync registers an asynchronous tool guardrail. +func RegisterToolConditionalExecutionGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_conditional_execution_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolConditionalExecutionGuardrail removes a previously registered // tool conditional-execution guardrail by name. Returns a NotFound error if no // guardrail with the given name is registered. @@ -1330,6 +1430,18 @@ func RegisterToolRequestIntercept(name string, priority int32, breakChain bool, )) } +// RegisterToolRequestInterceptAsync registers an asynchronous tool request intercept. +func RegisterToolRequestInterceptAsync(name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_request_intercept_async( + cName, C.int32_t(priority), C._Bool(breakChain), + C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolRequestIntercept removes a previously registered tool request // intercept by name. func DeregisterToolRequestIntercept(name string) error { @@ -1354,6 +1466,18 @@ func RegisterToolExecutionIntercept(name string, priority int32, execFn ToolExec )) } +// RegisterToolExecutionInterceptAsync registers an asynchronous tool execution intercept. +func RegisterToolExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_execution_intercept_async( + cName, C.int32_t(priority), + C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolExecutionIntercept removes a previously registered tool // execution intercept by name. func DeregisterToolExecutionIntercept(name string) error { @@ -1380,6 +1504,17 @@ func RegisterLlmSanitizeRequestGuardrail(name string, priority int32, fn LLMRequ )) } +// RegisterLlmSanitizeRequestGuardrailAsync registers an asynchronous LLM request sanitizer. +func RegisterLlmSanitizeRequestGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_sanitize_request_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmSanitizeRequestGuardrail removes a previously registered LLM // sanitize-request guardrail by name. func DeregisterLlmSanitizeRequestGuardrail(name string) error { @@ -1402,6 +1537,17 @@ func RegisterLlmSanitizeResponseGuardrail(name string, priority int32, fn LLMRes )) } +// RegisterLlmSanitizeResponseGuardrailAsync registers an asynchronous LLM response sanitizer. +func RegisterLlmSanitizeResponseGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_sanitize_response_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmSanitizeResponseGuardrail removes a previously registered LLM // sanitize-response guardrail by name. func DeregisterLlmSanitizeResponseGuardrail(name string) error { @@ -1428,6 +1574,17 @@ func RegisterLlmConditionalExecutionGuardrail(name string, priority int32, fn LL )) } +// RegisterLlmConditionalExecutionGuardrailAsync registers an asynchronous LLM guardrail. +func RegisterLlmConditionalExecutionGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_conditional_execution_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmConditionalExecutionGuardrail removes a previously registered // LLM conditional-execution guardrail by name. func DeregisterLlmConditionalExecutionGuardrail(name string) error { @@ -1454,6 +1611,18 @@ func RegisterLlmRequestIntercept(name string, priority int32, breakChain bool, f )) } +// RegisterLlmRequestInterceptAsync registers an asynchronous LLM request intercept. +func RegisterLlmRequestInterceptAsync(name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_request_intercept_async( + cName, C.int32_t(priority), C._Bool(breakChain), + C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmRequestIntercept removes a previously registered LLM request // intercept by name. func DeregisterLlmRequestIntercept(name string) error { @@ -1478,6 +1647,18 @@ func RegisterLlmExecutionIntercept(name string, priority int32, execFn LLMExecut )) } +// RegisterLlmExecutionInterceptAsync registers an asynchronous LLM execution intercept. +func RegisterLlmExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_execution_intercept_async( + cName, C.int32_t(priority), + C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmExecutionIntercept removes a previously registered LLM // execution intercept by name. func DeregisterLlmExecutionIntercept(name string) error { @@ -1503,6 +1684,18 @@ func RegisterLlmStreamExecutionIntercept(name string, priority int32, execFn LLM )) } +// RegisterLlmStreamExecutionInterceptAsync registers an asynchronous streaming LLM intercept. +func RegisterLlmStreamExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_stream_execution_intercept_async( + cName, C.int32_t(priority), + C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmStreamExecutionIntercept removes a previously registered LLM // stream execution intercept by name. func DeregisterLlmStreamExecutionIntercept(name string) error { @@ -2166,6 +2359,71 @@ func (s *OpenTelemetrySubscriber) Close() { // Scope-local guardrail/intercept registration (Tool) // --------------------------------------------------------------------------- +func withScopeAsyncMiddleware(scopeUUID, name string, priority int32, fn any, call func(*C.char, *C.char, C.int32_t, unsafe.Pointer) C.int32_t) error { + id := registerClosure(fn) + cScopeUUID := C.CString(scopeUUID) + defer C.free(unsafe.Pointer(cScopeUUID)) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(call(cScopeUUID, cName, C.int32_t(priority), id)) +} + +// ScopeRegisterMarkSanitizeGuardrailAsync registers an asynchronous scope-local mark sanitizer. +func ScopeRegisterMarkSanitizeGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_mark_sanitize_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterScopeSanitizeStartGuardrailAsync registers an asynchronous scope-local start sanitizer. +func ScopeRegisterScopeSanitizeStartGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_scope_sanitize_start_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterScopeSanitizeEndGuardrailAsync registers an asynchronous scope-local end sanitizer. +func ScopeRegisterScopeSanitizeEndGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_scope_sanitize_end_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolSanitizeRequestGuardrailAsync registers an asynchronous scope-local tool request sanitizer. +func ScopeRegisterToolSanitizeRequestGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_sanitize_request_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolSanitizeResponseGuardrailAsync registers an asynchronous scope-local tool response sanitizer. +func ScopeRegisterToolSanitizeResponseGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_sanitize_response_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolConditionalExecutionGuardrailAsync registers an asynchronous scope-local tool guardrail. +func ScopeRegisterToolConditionalExecutionGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_conditional_execution_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolRequestInterceptAsync registers an asynchronous scope-local tool request intercept. +func ScopeRegisterToolRequestInterceptAsync(scopeUUID, name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_request_intercept_async(scope, name, priority, C._Bool(breakChain), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolExecutionInterceptAsync registers an asynchronous scope-local tool execution intercept. +func ScopeRegisterToolExecutionInterceptAsync(scopeUUID, name string, priority int32, fn AsyncExecutionInterceptFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_execution_intercept_async(scope, name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + func registerScopeEventSanitizer(scopeUUID, name string, priority int32, fn EventSanitizeFunc, kind int) error { id := registerClosure(fn) cScopeUUID := C.CString(scopeUUID) @@ -2364,6 +2622,48 @@ func ScopeDeregisterToolExecutionIntercept(scopeUUID, name string) error { // Scope-local guardrail/intercept registration (LLM) // --------------------------------------------------------------------------- +// ScopeRegisterLlmSanitizeRequestGuardrailAsync registers an asynchronous scope-local LLM request sanitizer. +func ScopeRegisterLlmSanitizeRequestGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_sanitize_request_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmSanitizeResponseGuardrailAsync registers an asynchronous scope-local LLM response sanitizer. +func ScopeRegisterLlmSanitizeResponseGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_sanitize_response_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmConditionalExecutionGuardrailAsync registers an asynchronous scope-local LLM guardrail. +func ScopeRegisterLlmConditionalExecutionGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_conditional_execution_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmRequestInterceptAsync registers an asynchronous scope-local LLM request intercept. +func ScopeRegisterLlmRequestInterceptAsync(scopeUUID, name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_request_intercept_async(scope, name, priority, C._Bool(breakChain), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmExecutionInterceptAsync registers an asynchronous scope-local LLM execution intercept. +func ScopeRegisterLlmExecutionInterceptAsync(scopeUUID, name string, priority int32, fn AsyncExecutionInterceptFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_execution_intercept_async(scope, name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmStreamExecutionInterceptAsync registers an asynchronous scope-local streaming LLM intercept. +func ScopeRegisterLlmStreamExecutionInterceptAsync(scopeUUID, name string, priority int32, fn AsyncExecutionInterceptFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_stream_execution_intercept_async(scope, name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + // ScopeRegisterLlmSanitizeRequestGuardrail registers a scope-local guardrail // that sanitizes LLM request data. func ScopeRegisterLlmSanitizeRequestGuardrail(scopeUUID, name string, priority int32, fn LLMRequestFunc) error { diff --git a/go/nemo_relay/optimization_test.go b/go/nemo_relay/optimization_test.go index 60712a8da..491899ce5 100644 --- a/go/nemo_relay/optimization_test.go +++ b/go/nemo_relay/optimization_test.go @@ -98,6 +98,26 @@ func TestLLMOptimizationContributionOmittedAppliedIsNonApplied(t *testing.T) { } } +func TestLLMOptimizationContributionRejectsMalformedAndUnknownWireShapes(t *testing.T) { + var contribution LLMOptimizationContribution + if err := json.Unmarshal([]byte(`not-json`), &contribution); err == nil { + t.Fatal("expected malformed optimization contribution JSON to fail") + } + if err := json.Unmarshal([]byte(`[]`), &contribution); err == nil { + t.Fatal("expected non-object optimization contribution JSON to fail") + } + + contribution = LLMOptimizationContribution{ + Producer: "test", + Kind: "custom", + PayloadSchema: &LLMOptimizationDataSchema{Name: "test", Version: "v1"}, + Payload: json.RawMessage(`not-json`), + } + if _, err := json.Marshal(contribution); err == nil { + t.Fatal("expected malformed payload JSON to fail") + } +} + func TestLLMRequestInterceptOptimizationContributionsRoundTrip(t *testing.T) { fixture, contribution := optimizationContributionFixture(t) const interceptName = "go_optimization_fixture" From 3bbbe9e89061cfa3b5d719d662d5ff0c09cc19c3 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 16:04:42 -0400 Subject: [PATCH 35/52] fix: address async FFI and Go review feedback Signed-off-by: Will Killian --- crates/ffi/build.rs | 6 + crates/ffi/nemo_relay.h | 41 +---- crates/ffi/src/api/event_registry.rs | 2 +- crates/ffi/src/api/llm_registry.rs | 74 ++------ crates/ffi/src/api/mod.rs | 63 +++++++ crates/ffi/src/api/scope_registry.rs | 84 ++------- crates/ffi/src/api/tool_registry.rs | 69 ++------ crates/ffi/src/callable.rs | 104 +++++++----- .../tests/unit/api/coverage_sweeps_tests.rs | 3 +- crates/ffi/tests/unit/callable_tests.rs | 2 +- go/nemo_relay/async_middleware_test.go | 47 ++++++ go/nemo_relay/callbacks.go | 103 ++++++++---- go/nemo_relay/nemo_relay.go | 159 +++++++----------- 13 files changed, 363 insertions(+), 394 deletions(-) diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index d7f0d3d32..b6d6d335b 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -134,7 +134,13 @@ fn expected_async_prototype(name: &str) -> AsyncPrototype<'_> { const ASYNC_REGISTRATIONS: &str = r#" /* Completion-based async middleware registrations generated from Rust macros. */ +typedef uint32_t NemoRelayAsyncCallbackState; +enum { + NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0, + NEMO_RELAY_ASYNC_CALLBACK_STATE_PENDING = 1, +}; typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion); +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncCompletion *completion); NemoRelayStatus nemo_relay_register_mark_sanitize_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_scope_sanitize_start_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_scope_sanitize_end_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 83af886eb..fa52c428d 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -132,21 +132,6 @@ enum NemoRelayScopeType { }; typedef int32_t NemoRelayScopeType; -/** - * Indicates whether an async callback settled its completion before returning. - */ -enum NemoRelayAsyncCallbackState { - /** - * The callback called a resolve/reject function before returning. - */ - NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0, - /** - * The callback retained the completion and will settle it later. - */ - NEMO_RELAY_ASYNC_CALLBACK_STATE_PENDING = 1, -}; -typedef uint32_t NemoRelayAsyncCallbackState; - /** * Opaque owned adaptive runtime handle. */ @@ -470,14 +455,6 @@ typedef char *(*NemoRelayToolExecInterceptCb)(void *user_data, */ typedef char *(*NemoRelayToolExecCb)(void *user_data, const char *args_json); -/** - * Completion-based execution-intercept callback. - */ -typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, - const char *invocation_json, - const struct NemoRelayAsyncNext *next, - const struct NemoRelayAsyncCompletion *completion); - /** * Result callback used by channel/future-style async `next` wrappers. * @@ -2529,15 +2506,6 @@ NemoRelayStatus nemo_relay_tool_call_execute(const char *name, const char *metadata_json, char **out); -/** - * Register a completion-based asynchronous tool execution intercept. - */ -NemoRelayStatus nemo_relay_register_tool_execution_intercept_async(const char *name, - int32_t priority, - NemoRelayAsyncInterceptCb cb, - void *user_data, - NemoRelayFreeFn free_fn); - /** * Register a tool conditional execution guardrail. The callback decides whether * a tool call should proceed. Returns an error message to reject, or null to allow. @@ -2607,6 +2575,9 @@ void nemo_relay_async_next_release(const struct NemoRelayAsyncNext *next); /** * Invoke the next execution layer and settle `completion` with its result. + * + * A non-`Ok` return means invocation was not scheduled and never settles + * `completion`; the caller remains responsible for rejecting or releasing it. */ NemoRelayStatus nemo_relay_async_next_invoke(const struct NemoRelayAsyncNext *next, const char *invocation_json, @@ -3137,7 +3108,13 @@ char *nemo_relay_event_annotated_response(const struct FfiEvent *ptr); /* Completion-based async middleware registrations generated from Rust macros. */ +typedef uint32_t NemoRelayAsyncCallbackState; +enum { + NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0, + NEMO_RELAY_ASYNC_CALLBACK_STATE_PENDING = 1, +}; typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion); +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncCompletion *completion); NemoRelayStatus nemo_relay_register_mark_sanitize_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_scope_sanitize_start_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_scope_sanitize_end_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); diff --git a/crates/ffi/src/api/event_registry.rs b/crates/ffi/src/api/event_registry.rs index 4badaa381..273b98df3 100644 --- a/crates/ffi/src/api/event_registry.rs +++ b/crates/ffi/src/api/event_registry.rs @@ -173,6 +173,7 @@ unsafe fn register_scope_async( surface: Surface, ) -> NemoRelayStatus { clear_last_error(); + let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); let uuid = match parse_scope_uuid(scope_uuid) { Ok(uuid) => uuid, Err(status) => return status, @@ -181,7 +182,6 @@ unsafe fn register_scope_async( Ok(name) => name, Err(status) => return status, }; - let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); let result = match surface { Surface::Mark => core_registry_api::scope_register_mark_sanitize_guardrail( &uuid, &name, priority, callback, diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index 60d1d0053..218ca1138 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -14,90 +14,40 @@ use super::{ wrap_llm_stream_exec_intercept_fn, }; -macro_rules! async_llm_registration { - ($fn_name:ident, $register:path, $wrapper:path $(, $break_chain:ident)?) => { - /// Register a completion-based asynchronous LLM middleware callback. - #[allow(clippy::missing_safety_doc)] - #[unsafe(no_mangle)] - pub unsafe extern "C" fn $fn_name( - name: *const c_char, - priority: i32, - $( $break_chain: bool, )? - cb: NemoRelayAsyncJsonCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - ) -> NemoRelayStatus { - clear_last_error(); - let name = match c_str_to_string(name) { - Ok(name) => name, - Err(status) => return status, - }; - match $register( - &name, - priority, - $( $break_chain, )? - $wrapper(cb, user_data, free_fn), - ) { - Ok(()) => NemoRelayStatus::Ok, - Err(error) => status_from_error(&error), - } - } - }; -} - -async_llm_registration!( +global_async_registration!( nemo_relay_register_llm_sanitize_request_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::register_llm_sanitize_request_guardrail, wrap_async_llm_sanitize_request_fn ); -macro_rules! async_llm_execution_registration { - ($name:ident, $register:path, $wrapper:path) => { - /// Register a completion-based asynchronous LLM execution intercept. - #[allow(clippy::missing_safety_doc)] - #[unsafe(no_mangle)] - pub unsafe extern "C" fn $name( - name: *const c_char, - priority: i32, - cb: NemoRelayAsyncInterceptCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - ) -> NemoRelayStatus { - clear_last_error(); - let name = match c_str_to_string(name) { - Ok(name) => name, - Err(status) => return status, - }; - match $register(&name, priority, $wrapper(cb, user_data, free_fn)) { - Ok(()) => NemoRelayStatus::Ok, - Err(error) => status_from_error(&error), - } - } - }; -} - -async_llm_execution_registration!( +global_async_registration!( nemo_relay_register_llm_execution_intercept_async, + NemoRelayAsyncInterceptCb, core_registry_api::register_llm_execution_intercept, wrap_async_llm_execution_intercept_fn ); -async_llm_execution_registration!( +global_async_registration!( nemo_relay_register_llm_stream_execution_intercept_async, + NemoRelayAsyncInterceptCb, core_registry_api::register_llm_stream_execution_intercept, wrap_async_llm_stream_execution_intercept_fn ); -async_llm_registration!( +global_async_registration!( nemo_relay_register_llm_sanitize_response_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::register_llm_sanitize_response_guardrail, wrap_async_llm_sanitize_response_fn ); -async_llm_registration!( +global_async_registration!( nemo_relay_register_llm_conditional_execution_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::register_llm_conditional_execution_guardrail, wrap_async_llm_conditional_fn ); -async_llm_registration!( +global_async_registration!( nemo_relay_register_llm_request_intercept_async, + NemoRelayAsyncJsonCb, core_registry_api::register_llm_request_intercept, wrap_async_llm_request_intercept_fn, break_chain diff --git a/crates/ffi/src/api/mod.rs b/crates/ffi/src/api/mod.rs index fa5ff2bc3..0b7005975 100644 --- a/crates/ffi/src/api/mod.rs +++ b/crates/ffi/src/api/mod.rs @@ -71,6 +71,69 @@ use nemo_relay::plugin::{ use nemo_relay_adaptive::plugin_component::register_adaptive_component; use tokio::runtime::Runtime; +macro_rules! global_async_registration { + ($fn_name:ident, $callback_ty:ty, $register:path, $wrapper:path $(, $break_chain:ident)?) => { + /// Register a completion-based asynchronous middleware callback. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $fn_name( + name: *const c_char, + priority: i32, + $( $break_chain: bool, )? + cb: $callback_ty, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + // The wrapper assumes ownership before validation so every return + // path invokes free_fn exactly once. + let callback = $wrapper(cb, user_data, free_fn); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register(&name, priority, $( $break_chain, )? callback) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +macro_rules! scope_async_registration { + ($fn_name:ident, $callback_ty:ty, $register:path, $wrapper:path $(, $break_chain:ident)?) => { + /// Register a scope-local completion-based asynchronous middleware callback. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $fn_name( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + $( $break_chain: bool, )? + cb: $callback_ty, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + // The wrapper assumes ownership before validation so every return + // path invokes free_fn exactly once. + let callback = $wrapper(cb, user_data, free_fn); + let uuid = match parse_scope_uuid(scope_uuid) { + Ok(uuid) => uuid, + Err(status) => return status, + }; + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register(&uuid, &name, priority, $( $break_chain, )? callback) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + mod adaptive; mod event_registry; mod llm; diff --git a/crates/ffi/src/api/scope_registry.rs b/crates/ffi/src/api/scope_registry.rs index aebf09cd1..21d9767e9 100644 --- a/crates/ffi/src/api/scope_registry.rs +++ b/crates/ffi/src/api/scope_registry.rs @@ -30,128 +30,72 @@ fn parse_scope_uuid(scope_uuid: *const c_char) -> Result { - /// Register a scope-local completion-based asynchronous middleware callback. - #[allow(clippy::missing_safety_doc)] - #[unsafe(no_mangle)] - pub unsafe extern "C" fn $fn_name( - scope_uuid: *const c_char, - name: *const c_char, - priority: i32, - $( $break_chain: bool, )? - cb: NemoRelayAsyncJsonCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - ) -> NemoRelayStatus { - clear_last_error(); - let uuid = match parse_scope_uuid(scope_uuid) { - Ok(uuid) => uuid, - Err(status) => return status, - }; - let name = match c_str_to_string(name) { - Ok(name) => name, - Err(status) => return status, - }; - match $register( - &uuid, - &name, - priority, - $( $break_chain, )? - $wrapper(cb, user_data, free_fn), - ) { - Ok(()) => NemoRelayStatus::Ok, - Err(error) => status_from_error(&error), - } - } - }; -} - scope_async_registration!( nemo_relay_scope_register_tool_sanitize_request_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::scope_register_tool_sanitize_request_guardrail, wrap_async_tool_json_fn ); scope_async_registration!( nemo_relay_scope_register_tool_sanitize_response_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::scope_register_tool_sanitize_response_guardrail, wrap_async_tool_json_fn ); scope_async_registration!( nemo_relay_scope_register_tool_conditional_execution_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::scope_register_tool_conditional_execution_guardrail, wrap_async_tool_conditional_fn ); scope_async_registration!( nemo_relay_scope_register_tool_request_intercept_async, + NemoRelayAsyncJsonCb, core_registry_api::scope_register_tool_request_intercept, wrap_async_tool_json_fn, break_chain ); scope_async_registration!( nemo_relay_scope_register_llm_sanitize_request_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::scope_register_llm_sanitize_request_guardrail, wrap_async_llm_sanitize_request_fn ); scope_async_registration!( nemo_relay_scope_register_llm_sanitize_response_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::scope_register_llm_sanitize_response_guardrail, wrap_async_llm_sanitize_response_fn ); scope_async_registration!( nemo_relay_scope_register_llm_conditional_execution_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::scope_register_llm_conditional_execution_guardrail, wrap_async_llm_conditional_fn ); scope_async_registration!( nemo_relay_scope_register_llm_request_intercept_async, + NemoRelayAsyncJsonCb, core_registry_api::scope_register_llm_request_intercept, wrap_async_llm_request_intercept_fn, break_chain ); -macro_rules! scope_async_execution_registration { - ($fn_name:ident, $register:path, $wrapper:path) => { - /// Register a scope-local completion-based asynchronous execution intercept. - #[allow(clippy::missing_safety_doc)] - #[unsafe(no_mangle)] - pub unsafe extern "C" fn $fn_name( - scope_uuid: *const c_char, - name: *const c_char, - priority: i32, - cb: NemoRelayAsyncInterceptCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - ) -> NemoRelayStatus { - clear_last_error(); - let uuid = match parse_scope_uuid(scope_uuid) { - Ok(uuid) => uuid, - Err(status) => return status, - }; - let name = match c_str_to_string(name) { - Ok(name) => name, - Err(status) => return status, - }; - match $register(&uuid, &name, priority, $wrapper(cb, user_data, free_fn)) { - Ok(()) => NemoRelayStatus::Ok, - Err(error) => status_from_error(&error), - } - } - }; -} - -scope_async_execution_registration!( +scope_async_registration!( nemo_relay_scope_register_tool_execution_intercept_async, + NemoRelayAsyncInterceptCb, core_registry_api::scope_register_tool_execution_intercept, wrap_async_tool_execution_intercept_fn ); -scope_async_execution_registration!( +scope_async_registration!( nemo_relay_scope_register_llm_execution_intercept_async, + NemoRelayAsyncInterceptCb, core_registry_api::scope_register_llm_execution_intercept, wrap_async_llm_execution_intercept_fn ); -scope_async_execution_registration!( +scope_async_registration!( nemo_relay_scope_register_llm_stream_execution_intercept_async, + NemoRelayAsyncInterceptCb, core_registry_api::scope_register_llm_stream_execution_intercept, wrap_async_llm_stream_execution_intercept_fn ); diff --git a/crates/ffi/src/api/tool_registry.rs b/crates/ffi/src/api/tool_registry.rs index fd1f2e580..b25621e1f 100644 --- a/crates/ffi/src/api/tool_registry.rs +++ b/crates/ffi/src/api/tool_registry.rs @@ -10,80 +10,39 @@ use super::{ wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, }; -macro_rules! async_tool_json_registration { - ($name:ident, $register:path, $wrapper:path $(, $break_chain:ident)?) => { - /// Register a completion-based asynchronous tool middleware callback. - #[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. - #[unsafe(no_mangle)] - pub unsafe extern "C" fn $name( - name: *const c_char, - priority: i32, - $( $break_chain: bool, )? - cb: NemoRelayAsyncJsonCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - ) -> NemoRelayStatus { - clear_last_error(); - let name = match c_str_to_string(name) { - Ok(name) => name, - Err(status) => return status, - }; - let callback = $wrapper(cb, user_data, free_fn); - match $register(&name, priority, $( $break_chain, )? callback) { - Ok(()) => NemoRelayStatus::Ok, - Err(error) => status_from_error(&error), - } - } - }; -} - -async_tool_json_registration!( +global_async_registration!( nemo_relay_register_tool_sanitize_request_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::register_tool_sanitize_request_guardrail, wrap_async_tool_json_fn ); -async_tool_json_registration!( +global_async_registration!( nemo_relay_register_tool_sanitize_response_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::register_tool_sanitize_response_guardrail, wrap_async_tool_json_fn ); -async_tool_json_registration!( +global_async_registration!( nemo_relay_register_tool_conditional_execution_guardrail_async, + NemoRelayAsyncJsonCb, core_registry_api::register_tool_conditional_execution_guardrail, wrap_async_tool_conditional_fn ); -async_tool_json_registration!( +global_async_registration!( nemo_relay_register_tool_request_intercept_async, + NemoRelayAsyncJsonCb, core_registry_api::register_tool_request_intercept, wrap_async_tool_json_fn, break_chain ); -/// Register a completion-based asynchronous tool execution intercept. -#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. -#[unsafe(no_mangle)] -pub unsafe extern "C" fn nemo_relay_register_tool_execution_intercept_async( - name: *const c_char, - priority: i32, - cb: NemoRelayAsyncInterceptCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, -) -> NemoRelayStatus { - clear_last_error(); - let name = match c_str_to_string(name) { - Ok(name) => name, - Err(status) => return status, - }; - match core_registry_api::register_tool_execution_intercept( - &name, - priority, - wrap_async_tool_execution_intercept_fn(cb, user_data, free_fn), - ) { - Ok(()) => NemoRelayStatus::Ok, - Err(error) => status_from_error(&error), - } -} +global_async_registration!( + nemo_relay_register_tool_execution_intercept_async, + NemoRelayAsyncInterceptCb, + core_registry_api::register_tool_execution_intercept, + wrap_async_tool_execution_intercept_fn +); // --------------------------------------------------------------------------- // Tool guardrail registrations diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 9414d31ac..9fd1e136a 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -97,6 +97,36 @@ enum AsyncNextInner { LlmStream(LlmStreamExecutionNextFn), } +const ASYNC_STREAM_MAX_CHUNKS: usize = 4096; +const ASYNC_STREAM_MAX_SERIALIZED_BYTES: usize = 16 * 1024 * 1024; + +async fn collect_async_stream_for_completion(mut stream: LlmJsonStream) -> Result { + let mut chunks = Vec::new(); + let mut serialized_bytes = 0usize; + while let Some(chunk) = stream.next().await { + let chunk = chunk?; + if chunks.len() >= ASYNC_STREAM_MAX_CHUNKS { + return Err(FlowError::Internal(format!( + "async stream continuation exceeded the {ASYNC_STREAM_MAX_CHUNKS}-chunk completion limit" + ))); + } + serialized_bytes = serialized_bytes.saturating_add( + serde_json::to_vec(&chunk) + .map_err(|error| { + FlowError::Internal(format!("failed to measure async stream chunk: {error}")) + })? + .len(), + ); + if serialized_bytes > ASYNC_STREAM_MAX_SERIALIZED_BYTES { + return Err(FlowError::Internal(format!( + "async stream continuation exceeded the {ASYNC_STREAM_MAX_SERIALIZED_BYTES}-byte completion limit" + ))); + } + chunks.push(chunk); + } + Ok(Json::Array(chunks)) +} + /// Completion-based execution-intercept callback. pub type NemoRelayAsyncInterceptCb = unsafe extern "C" fn( user_data: *mut libc::c_void, @@ -222,6 +252,9 @@ pub unsafe extern "C" fn nemo_relay_async_next_release(next: *const NemoRelayAsy } /// Invoke the next execution layer and settle `completion` with its result. +/// +/// A non-`Ok` return means invocation was not scheduled and never settles +/// `completion`; the caller remains responsible for rejecting or releasing it. #[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. #[unsafe(no_mangle)] pub unsafe extern "C" fn nemo_relay_async_next_invoke( @@ -252,17 +285,7 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke( AsyncNextInner::Llm(next) => { let request = match serde_json::from_value(invocation) { Ok(request) => request, - Err(error) => { - return { - let _ = completion - .sender - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); - NemoRelayStatus::InvalidJson - }; - } + Err(_) => return NemoRelayStatus::InvalidJson, }; let next = next.clone(); Box::pin(async move { next(request).await }) @@ -270,27 +293,10 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke( AsyncNextInner::LlmStream(next) => { let request = match serde_json::from_value(invocation) { Ok(request) => request, - Err(error) => { - return { - let _ = completion - .sender - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); - NemoRelayStatus::InvalidJson - }; - } + Err(_) => return NemoRelayStatus::InvalidJson, }; let next = next.clone(); - Box::pin(async move { - let mut stream = next(request).await?; - let mut chunks = Vec::new(); - while let Some(chunk) = stream.next().await { - chunks.push(chunk?); - } - Ok(Json::Array(chunks)) - }) + Box::pin(async move { collect_async_stream_for_completion(next(request).await?).await }) } }; next.runtime.spawn(async move { @@ -344,14 +350,7 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke_callback( Err(_) => return NemoRelayStatus::InvalidJson, }; let next = next.clone(); - Box::pin(async move { - let mut stream = next(request).await?; - let mut chunks = Vec::new(); - while let Some(chunk) = stream.next().await { - chunks.push(chunk?); - } - Ok(Json::Array(chunks)) - }) + Box::pin(async move { collect_async_stream_for_completion(next(request).await?).await }) } }; let user_data = user_data as usize; @@ -770,7 +769,7 @@ pub fn wrap_async_event_sanitize_fn( free_fn: NemoRelayFreeFn, ) -> EventSanitizeFn { let user_data = make_user_data(user_data, free_fn); - Arc::new(move |event: Event, fields: EventSanitizeFields| { + Arc::new(move |event: Arc, fields: EventSanitizeFields| { let user_data = user_data.clone(); Box::pin(async move { let value = invoke_async_json( @@ -816,12 +815,13 @@ pub fn wrap_async_llm_sanitize_request_fn( Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { let user_data = user_data.clone(); - let codec = format!("{:?}", context.codec()); + let codec = ffi_codec_identity_json(context.codec()); Box::pin(async move { + let codec = codec?; let value = invoke_async_json( cb, user_data, - serde_json::json!({"request": request, "context": {"codec": codec}}), + serde_json::json!({"request": request, "context": codec}), ) .await?; if value.is_null() { @@ -845,12 +845,13 @@ pub fn wrap_async_llm_sanitize_response_fn( let user_data = make_user_data(user_data, free_fn); Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { let user_data = user_data.clone(); - let codec = format!("{:?}", context.codec()); + let codec = ffi_codec_identity_json(context.codec()); Box::pin(async move { + let codec = codec?; let value = invoke_async_json( cb, user_data, - serde_json::json!({"response": response, "context": {"codec": codec}}), + serde_json::json!({"response": response, "context": codec}), ) .await?; Ok((!value.is_null()).then_some(value)) @@ -933,6 +934,7 @@ pub fn wrap_async_llm_execution_intercept_fn( /// The completion ABI resolves one JSON value, so a stream intercept must /// resolve to an array of chunks. Relay replays that array as a stream after /// completion; incremental chunk delivery is not available through this ABI. +/// Relay rejects more than 4096 chunks or 16 MiB of serialized chunk data. pub fn wrap_async_llm_stream_execution_intercept_fn( cb: NemoRelayAsyncInterceptCb, user_data: *mut libc::c_void, @@ -1479,6 +1481,22 @@ fn ffi_codec_identity( }) } +fn ffi_codec_identity_json(identity: &LlmCodecIdentity) -> Result { + let (kind, id) = ffi_codec_identity(identity)?; + let id = id + .as_ref() + .map(|id| { + id.to_str() + .map(str::to_owned) + .map_err(|error| FlowError::Internal(error.to_string())) + }) + .transpose()?; + Ok(serde_json::json!({ + "codec_kind": kind as u32, + "codec_id": id, + })) +} + /// Wrap a C LLM conditional callback into a Rust closure. pub fn wrap_llm_conditional_fn( cb: NemoRelayLlmConditionalCb, diff --git a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs index 7c7e012d0..73343b59c 100644 --- a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs +++ b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs @@ -148,7 +148,7 @@ fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { nemo_relay_deregister_llm_stream_execution_intercept ); - let _stack = unsafe { fresh_scope_stack() }; + let stack = unsafe { fresh_scope_stack() }; let mut scope = ptr::null_mut(); assert_eq!( unsafe { nemo_relay_get_handle(&mut scope) }, @@ -284,6 +284,7 @@ fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { nemo_relay_scope_deregister_llm_stream_execution_intercept ); unsafe { nemo_relay_scope_handle_free(scope) }; + unsafe { nemo_relay_scope_stack_free(stack) }; } impl EnvGuard { diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 65a326f4c..5fb38c3d1 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -198,7 +198,7 @@ fn async_callback_wrappers_cover_all_middleware_shapes() { Some(free_async_callback_user_data), ); assert_eq!( - resolve(event_sanitizer(event, fields.clone())).unwrap(), + resolve(event_sanitizer(Arc::new(event), fields.clone())).unwrap(), fields ); } diff --git a/go/nemo_relay/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go index ee501908b..d34bfdcd9 100644 --- a/go/nemo_relay/async_middleware_test.go +++ b/go/nemo_relay/async_middleware_test.go @@ -57,9 +57,15 @@ func TestAsyncMiddlewareGlobalRegistrationParity(t *testing.T) { if err := registration.register(name); err != nil { t.Fatalf("register: %v", err) } + if err := registration.register(name); err == nil { + t.Fatal("duplicate registration unexpectedly succeeded") + } if err := registration.deregister(name); err != nil { t.Fatalf("deregister: %v", err) } + if err := registration.deregister(name); err != nil { + t.Fatalf("idempotent deregister: %v", err) + } }) } } @@ -131,9 +137,50 @@ func TestAsyncMiddlewareScopeLocalRegistrationParity(t *testing.T) { if err := registration.register(name); err != nil { t.Fatalf("register %s: %v", registration.name, err) } + if err := registration.register(name); err == nil { + t.Fatalf("duplicate registration %s unexpectedly succeeded", registration.name) + } if err := registration.deregister(name); err != nil { t.Fatalf("deregister %s: %v", registration.name, err) } + if err := registration.deregister(name); err != nil { + t.Fatalf("idempotent deregistration %s: %v", registration.name, err) + } + } + }) +} + +func TestAsyncToolRequestInterceptPriorityOrdering(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + register := func(name string, priority int32, marker string) { + t.Helper() + err := RegisterToolRequestInterceptAsync(name, priority, false, + func(_ context.Context, invocation json.RawMessage) (any, error) { + var envelope struct { + Value map[string]any `json:"value"` + } + if err := json.Unmarshal(invocation, &envelope); err != nil { + return nil, err + } + order, _ := envelope.Value["order"].(string) + envelope.Value["order"] = order + marker + return envelope.Value, nil + }, + ) + if err != nil { + t.Fatalf("register %s: %v", name, err) + } + t.Cleanup(func() { _ = DeregisterToolRequestIntercept(name) }) + } + register("go-async-priority-late", 10, "B") + register("go-async-priority-early", 0, "A") + + result, err := ToolRequestIntercepts("priority", json.RawMessage(`{"order":""}`)) + if err != nil { + t.Fatalf("tool request intercepts: %v", err) + } + if string(result) != `{"order":"AB"}` { + t.Fatalf("result = %s, want priority order AB", result) } }) } diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index f7837ab2b..3b5615f4f 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -114,6 +114,7 @@ import ( var ( closureRegistryMu sync.Mutex closureRegistry = make(map[uintptr]interface{}) + closureTokens = make(map[uintptr]uintptr) closureNextID atomic.Uint64 closureTokenAlloc = func() unsafe.Pointer { return C.malloc(C.size_t(unsafe.Sizeof(uintptr(0)))) @@ -131,10 +132,6 @@ func setLastErrorMessage(msg string) { // suitable for passing as void* user_data to C callbacks. func registerClosure(fn interface{}) unsafe.Pointer { id := uintptr(closureNextID.Add(1)) - closureRegistryMu.Lock() - closureRegistry[id] = fn - closureRegistryMu.Unlock() - // Allocate the callback token in C-owned memory so we don't pass a Go // pointer through C and can release it explicitly on deregistration. p := (*uintptr)(closureTokenAlloc()) @@ -142,27 +139,39 @@ func registerClosure(fn interface{}) unsafe.Pointer { panic("nemo_relay: failed to allocate callback token") } *p = id - return unsafe.Pointer(p) -} - -func closureID(userData unsafe.Pointer) uintptr { - return *(*uintptr)(userData) + token := unsafe.Pointer(p) + closureRegistryMu.Lock() + closureRegistry[id] = fn + closureTokens[uintptr(token)] = id + closureRegistryMu.Unlock() + return token } func lookupClosure(userData unsafe.Pointer) interface{} { - id := closureID(userData) closureRegistryMu.Lock() + id := closureTokens[uintptr(userData)] fn := closureRegistry[id] closureRegistryMu.Unlock() return fn } +func closureID(userData unsafe.Pointer) uintptr { + closureRegistryMu.Lock() + defer closureRegistryMu.Unlock() + return closureTokens[uintptr(userData)] +} + func unregisterClosure(userData unsafe.Pointer) { - id := closureID(userData) closureRegistryMu.Lock() + id, registered := closureTokens[uintptr(userData)] + if registered { + delete(closureTokens, uintptr(userData)) + } delete(closureRegistry, id) closureRegistryMu.Unlock() - C.free(userData) + if registered { + C.free(userData) + } } // --------------------------------------------------------------------------- @@ -191,28 +200,57 @@ type AsyncExecutionInterceptFunc func(ctx context.Context, invocation json.RawMe const asyncCallbackPending = C.uint32_t(1) +const asyncCancellationPollInterval = 10 * time.Millisecond + +type completionCancellationWatch struct { + completion *C.NemoRelayAsyncCompletion + cancel context.CancelFunc +} + +var ( + completionCancellationMu sync.Mutex + completionCancellationWatches = make(map[uint64]completionCancellationWatch) + completionCancellationNextID atomic.Uint64 + completionCancellationOnce sync.Once +) + +func startCompletionCancellationMonitor() { + completionCancellationOnce.Do(func() { + go func() { + ticker := time.NewTicker(asyncCancellationPollInterval) + defer ticker.Stop() + for range ticker.C { + completionCancellationMu.Lock() + for id, watch := range completionCancellationWatches { + if bool(C.nemo_relay_async_completion_is_cancelled(watch.completion)) { + delete(completionCancellationWatches, id) + watch.cancel() + } + } + completionCancellationMu.Unlock() + } + }() + }) +} + func contextForCompletion(completion *C.NemoRelayAsyncCompletion) (context.Context, func()) { ctx, cancel := context.WithCancel(context.Background()) - done := make(chan struct{}) + startCompletionCancellationMonitor() + id := completionCancellationNextID.Add(1) + completionCancellationMu.Lock() + completionCancellationWatches[id] = completionCancellationWatch{ + completion: completion, + cancel: cancel, + } + completionCancellationMu.Unlock() var doneOnce sync.Once - go func() { - ticker := time.NewTicker(10 * time.Millisecond) - defer ticker.Stop() - for { - select { - case <-done: - return - case <-ticker.C: - if bool(C.nemo_relay_async_completion_is_cancelled(completion)) { - cancel() - return - } - } - } - }() return ctx, func() { - doneOnce.Do(func() { close(done) }) - cancel() + doneOnce.Do(func() { + completionCancellationMu.Lock() + delete(completionCancellationWatches, id) + completionCancellationMu.Unlock() + cancel() + }) } } @@ -684,6 +722,7 @@ type asyncNextResult struct { func goAsyncNextResultTrampoline(userData unsafe.Pointer, valueJSON *C.char, errorMessage *C.char) { ch, ok := lookupClosure(userData).(chan asyncNextResult) if !ok { + unregisterClosure(userData) return } defer unregisterClosure(userData) @@ -719,6 +758,8 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON var nextMu sync.RWMutex nextOpen := true defer func() { + // Unblock in-flight next calls before waiting for their read locks. + cancel() nextMu.Lock() nextOpen = false nextMu.Unlock() @@ -732,6 +773,7 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON } ch := make(chan asyncNextResult, 1) token := registerClosure(ch) + defer unregisterClosure(token) cPayload := C.CString(string(payload)) status := C.nemo_relay_async_next_invoke_callback( next, cPayload, @@ -739,7 +781,6 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON ) C.free(unsafe.Pointer(cPayload)) if err := checkStatus(status); err != nil { - unregisterClosure(token) return nil, err } select { diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index 52ab649f0..bb931c9a5 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -1233,21 +1233,18 @@ func registerEventSanitizer(name string, priority int32, fn EventSanitizeFunc, k } func registerAsyncEventSanitizer(name string, priority int32, fn AsyncMiddlewareFunc, kind int) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - callback := C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline) - free := C.NemoRelayFreeFn(C.goFreeTrampoline) - var status C.int32_t - switch kind { - case 0: - status = C.nemo_relay_register_mark_sanitize_guardrail_async(cName, C.int32_t(priority), callback, id, free) - case 1: - status = C.nemo_relay_register_scope_sanitize_start_guardrail_async(cName, C.int32_t(priority), callback, id, free) - default: - status = C.nemo_relay_register_scope_sanitize_end_guardrail_async(cName, C.int32_t(priority), callback, id, free) - } - return checkStatus(status) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + callback := C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline) + free := C.NemoRelayFreeFn(C.goFreeTrampoline) + switch kind { + case 0: + return C.nemo_relay_register_mark_sanitize_guardrail_async(name, priority, callback, id, free) + case 1: + return C.nemo_relay_register_scope_sanitize_start_guardrail_async(name, priority, callback, id, free) + default: + return C.nemo_relay_register_scope_sanitize_end_guardrail_async(name, priority, callback, id, free) + } + }) } // RegisterMarkSanitizeGuardrail registers a global mark event sanitizer. @@ -1321,13 +1318,9 @@ func RegisterToolSanitizeRequestGuardrail(name string, priority int32, fn ToolSa // RegisterToolSanitizeRequestGuardrailAsync registers an asynchronous tool request sanitizer. func RegisterToolSanitizeRequestGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_tool_sanitize_request_guardrail_async( - cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_tool_sanitize_request_guardrail_async(name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterToolSanitizeRequestGuardrail removes a previously registered tool @@ -1357,13 +1350,9 @@ func RegisterToolSanitizeResponseGuardrail(name string, priority int32, fn ToolS // RegisterToolSanitizeResponseGuardrailAsync registers an asynchronous tool response sanitizer. func RegisterToolSanitizeResponseGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_tool_sanitize_response_guardrail_async( - cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_tool_sanitize_response_guardrail_async(name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterToolSanitizeResponseGuardrail removes a previously registered tool @@ -1395,13 +1384,9 @@ func RegisterToolConditionalExecutionGuardrail(name string, priority int32, fn T // RegisterToolConditionalExecutionGuardrailAsync registers an asynchronous tool guardrail. func RegisterToolConditionalExecutionGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_tool_conditional_execution_guardrail_async( - cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_tool_conditional_execution_guardrail_async(name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterToolConditionalExecutionGuardrail removes a previously registered @@ -1432,14 +1417,9 @@ func RegisterToolRequestIntercept(name string, priority int32, breakChain bool, // RegisterToolRequestInterceptAsync registers an asynchronous tool request intercept. func RegisterToolRequestInterceptAsync(name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_tool_request_intercept_async( - cName, C.int32_t(priority), C._Bool(breakChain), - C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_tool_request_intercept_async(name, priority, C._Bool(breakChain), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterToolRequestIntercept removes a previously registered tool request @@ -1468,14 +1448,9 @@ func RegisterToolExecutionIntercept(name string, priority int32, execFn ToolExec // RegisterToolExecutionInterceptAsync registers an asynchronous tool execution intercept. func RegisterToolExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_tool_execution_intercept_async( - cName, C.int32_t(priority), - C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_tool_execution_intercept_async(name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterToolExecutionIntercept removes a previously registered tool @@ -1506,13 +1481,9 @@ func RegisterLlmSanitizeRequestGuardrail(name string, priority int32, fn LLMRequ // RegisterLlmSanitizeRequestGuardrailAsync registers an asynchronous LLM request sanitizer. func RegisterLlmSanitizeRequestGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_llm_sanitize_request_guardrail_async( - cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_llm_sanitize_request_guardrail_async(name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterLlmSanitizeRequestGuardrail removes a previously registered LLM @@ -1539,13 +1510,9 @@ func RegisterLlmSanitizeResponseGuardrail(name string, priority int32, fn LLMRes // RegisterLlmSanitizeResponseGuardrailAsync registers an asynchronous LLM response sanitizer. func RegisterLlmSanitizeResponseGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_llm_sanitize_response_guardrail_async( - cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_llm_sanitize_response_guardrail_async(name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterLlmSanitizeResponseGuardrail removes a previously registered LLM @@ -1576,13 +1543,9 @@ func RegisterLlmConditionalExecutionGuardrail(name string, priority int32, fn LL // RegisterLlmConditionalExecutionGuardrailAsync registers an asynchronous LLM guardrail. func RegisterLlmConditionalExecutionGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_llm_conditional_execution_guardrail_async( - cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_llm_conditional_execution_guardrail_async(name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterLlmConditionalExecutionGuardrail removes a previously registered @@ -1613,14 +1576,9 @@ func RegisterLlmRequestIntercept(name string, priority int32, breakChain bool, f // RegisterLlmRequestInterceptAsync registers an asynchronous LLM request intercept. func RegisterLlmRequestInterceptAsync(name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_llm_request_intercept_async( - cName, C.int32_t(priority), C._Bool(breakChain), - C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_llm_request_intercept_async(name, priority, C._Bool(breakChain), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterLlmRequestIntercept removes a previously registered LLM request @@ -1649,14 +1607,9 @@ func RegisterLlmExecutionIntercept(name string, priority int32, execFn LLMExecut // RegisterLlmExecutionInterceptAsync registers an asynchronous LLM execution intercept. func RegisterLlmExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_llm_execution_intercept_async( - cName, C.int32_t(priority), - C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_llm_execution_intercept_async(name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterLlmExecutionIntercept removes a previously registered LLM @@ -1686,14 +1639,9 @@ func RegisterLlmStreamExecutionIntercept(name string, priority int32, execFn LLM // RegisterLlmStreamExecutionInterceptAsync registers an asynchronous streaming LLM intercept. func RegisterLlmStreamExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { - id := registerClosure(fn) - cName := C.CString(name) - defer C.free(unsafe.Pointer(cName)) - return checkStatus(C.nemo_relay_register_llm_stream_execution_intercept_async( - cName, C.int32_t(priority), - C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), - id, C.NemoRelayFreeFn(C.goFreeTrampoline), - )) + return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_register_llm_stream_execution_intercept_async(name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) } // DeregisterLlmStreamExecutionIntercept removes a previously registered LLM @@ -2359,12 +2307,27 @@ func (s *OpenTelemetrySubscriber) Close() { // Scope-local guardrail/intercept registration (Tool) // --------------------------------------------------------------------------- -func withScopeAsyncMiddleware(scopeUUID, name string, priority int32, fn any, call func(*C.char, *C.char, C.int32_t, unsafe.Pointer) C.int32_t) error { +type asyncMiddlewareCallback interface { + AsyncMiddlewareFunc | AsyncExecutionInterceptFunc +} + +func withGlobalAsyncMiddleware[T asyncMiddlewareCallback](name string, priority int32, fn T, call func(*C.char, C.int32_t, unsafe.Pointer) C.int32_t) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + // The C registration entry point owns id on every return path and invokes + // goFreeTrampoline exactly once if registration fails. + return checkStatus(call(cName, C.int32_t(priority), id)) +} + +func withScopeAsyncMiddleware[T asyncMiddlewareCallback](scopeUUID, name string, priority int32, fn T, call func(*C.char, *C.char, C.int32_t, unsafe.Pointer) C.int32_t) error { id := registerClosure(fn) cScopeUUID := C.CString(scopeUUID) defer C.free(unsafe.Pointer(cScopeUUID)) cName := C.CString(name) defer C.free(unsafe.Pointer(cName)) + // The C registration entry point owns id on every return path and invokes + // goFreeTrampoline exactly once if registration fails. return checkStatus(call(cScopeUUID, cName, C.int32_t(priority), id)) } From e1c29135fbe775914a084f09d8d83c4b40d4307d Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 16:32:18 -0400 Subject: [PATCH 36/52] fix: address follow-up FFI and Go review feedback Signed-off-by: Will Killian --- crates/ffi/build.rs | 115 +++++++++++++------------ crates/ffi/src/api/event_registry.rs | 22 +++-- crates/ffi/src/callable.rs | 18 +++- go/nemo_relay/adaptive_runtime_test.go | 16 ++-- go/nemo_relay/callbacks.go | 55 +++++++++--- 5 files changed, 145 insertions(+), 81 deletions(-) diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index b6d6d335b..1e2cead74 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -34,9 +34,9 @@ fn main() { } #[derive(Debug, PartialEq, Eq)] -struct AsyncPrototype<'a> { - name: &'a str, - parameters: Vec<&'a str>, +struct AsyncPrototype { + name: String, + parameters: Vec, } /// cbindgen does not expand the declarative registration macros. Keep the @@ -50,27 +50,21 @@ fn validate_async_registration_parity(crate_dir: &str) { "src/api/tool_registry.rs", ]; - let mut exported = Vec::new(); + let mut expected = Vec::new(); for source in REGISTRATION_SOURCES { println!("cargo:rerun-if-changed={source}"); let source_path = format!("{crate_dir}/{source}"); let contents = std::fs::read_to_string(&source_path) .unwrap_or_else(|error| panic!("read {source_path}: {error}")); - exported.extend( - contents - .split(|character: char| !character.is_ascii_alphanumeric() && character != '_') - .filter(|token| token.starts_with("nemo_relay_") && token.ends_with("_async")) - .map(str::to_owned), - ); + expected.extend(parse_async_macro_invocations(&contents)); } - exported.sort(); - exported.dedup(); + expected.sort_by(|left, right| left.name.cmp(&right.name)); let mut declared = ASYNC_REGISTRATIONS .lines() .filter_map(parse_async_prototype) .collect::>(); - declared.sort_by(|left, right| left.name.cmp(right.name)); + declared.sort_by(|left, right| left.name.cmp(&right.name)); for duplicates in declared.windows(2) { assert_ne!( duplicates[0].name, duplicates[1].name, @@ -78,58 +72,73 @@ fn validate_async_registration_parity(crate_dir: &str) { duplicates[0].name ); } - let declared_names = declared - .iter() - .map(|prototype| prototype.name.to_owned()) - .collect::>(); assert_eq!( - declared_names, exported, - "ASYNC_REGISTRATIONS must declare exactly the async Rust FFI exports" + declared, expected, + "ASYNC_REGISTRATIONS must exactly match the macro-generated Rust FFI exports" ); - - for prototype in declared { - assert!( - exported - .binary_search_by(|name| name.as_str().cmp(prototype.name)) - .is_ok(), - "async declaration for {} is not a Rust FFI export", - prototype.name - ); - assert_eq!( - prototype, - expected_async_prototype(prototype.name), - "async declaration for {} has a mismatched C prototype", - prototype.name - ); - } } -fn parse_async_prototype(line: &str) -> Option> { +fn parse_async_prototype(line: &str) -> Option { let line = line.strip_prefix("NemoRelayStatus ")?; let (name, parameters) = line.split_once('(')?; let parameters = parameters.strip_suffix(");")?; Some(AsyncPrototype { - name, - parameters: parameters.split(", ").collect(), + name: name.to_owned(), + parameters: parameters.split(", ").map(str::to_owned).collect(), }) } -fn expected_async_prototype(name: &str) -> AsyncPrototype<'_> { - let mut parameters = Vec::new(); - if name.starts_with("nemo_relay_scope_") { - parameters.push("const char *scope_uuid"); - } - parameters.extend(["const char *name", "int32_t priority"]); - if name.contains("request_intercept_async") { - parameters.push("bool break_chain"); +fn parse_async_macro_invocations(source: &str) -> Vec { + const MACROS: &[(&str, bool)] = &[ + ("global_async_registration!(", false), + ("scope_async_registration!(", true), + ("global_async_event_registration!(", false), + ("scope_async_event_registration!(", true), + ]; + + let mut prototypes = Vec::new(); + for (prefix, scope_local) in MACROS { + let mut remaining = source; + while let Some(start) = remaining.find(prefix) { + let invocation = &remaining[start + prefix.len()..]; + let Some(end) = invocation.find(");") else { + break; + }; + let arguments = invocation[..end] + .split(',') + .map(str::trim) + .collect::>(); + remaining = &invocation[end + 2..]; + + let Some(name) = arguments.first().copied() else { + continue; + }; + if !name.starts_with("nemo_relay_") || !name.ends_with("_async") { + continue; + } + let callback_type = arguments + .get(1) + .unwrap_or_else(|| panic!("{name} macro invocation is missing its callback type")); + let mut parameters = Vec::new(); + if *scope_local { + parameters.push("const char *scope_uuid".to_owned()); + } + parameters.extend(["const char *name".to_owned(), "int32_t priority".to_owned()]); + if arguments.contains(&"break_chain") { + parameters.push("bool break_chain".to_owned()); + } + parameters.extend([ + format!("{callback_type} cb"), + "void *user_data".to_owned(), + "NemoRelayFreeFn free_fn".to_owned(), + ]); + prototypes.push(AsyncPrototype { + name: name.to_owned(), + parameters, + }); + } } - parameters.push(if name.contains("execution_intercept_async") { - "NemoRelayAsyncInterceptCb cb" - } else { - "NemoRelayAsyncJsonCb cb" - }); - parameters.extend(["void *user_data", "NemoRelayFreeFn free_fn"]); - AsyncPrototype { name, parameters } + prototypes } const ASYNC_REGISTRATIONS: &str = r#" diff --git a/crates/ffi/src/api/event_registry.rs b/crates/ffi/src/api/event_registry.rs index 273b98df3..a1e471865 100644 --- a/crates/ffi/src/api/event_registry.rs +++ b/crates/ffi/src/api/event_registry.rs @@ -74,15 +74,15 @@ unsafe fn register_global_async( .unwrap_or_else(|error| status_from_error(&error)) } -macro_rules! async_event_registration { - ($name:ident, $surface:expr) => { +macro_rules! global_async_event_registration { + ($name:ident, $callback_ty:ty, $surface:expr) => { /// Register a completion-based asynchronous event sanitizer. #[allow(clippy::missing_safety_doc)] #[unsafe(no_mangle)] pub unsafe extern "C" fn $name( name: *const c_char, priority: i32, - cb: NemoRelayAsyncJsonCb, + cb: $callback_ty, user_data: *mut libc::c_void, free_fn: NemoRelayFreeFn, ) -> NemoRelayStatus { @@ -91,16 +91,19 @@ macro_rules! async_event_registration { }; } -async_event_registration!( +global_async_event_registration!( nemo_relay_register_mark_sanitize_guardrail_async, + NemoRelayAsyncJsonCb, Surface::Mark ); -async_event_registration!( +global_async_event_registration!( nemo_relay_register_scope_sanitize_start_guardrail_async, + NemoRelayAsyncJsonCb, Surface::Start ); -async_event_registration!( +global_async_event_registration!( nemo_relay_register_scope_sanitize_end_guardrail_async, + NemoRelayAsyncJsonCb, Surface::End ); @@ -199,7 +202,7 @@ unsafe fn register_scope_async( } macro_rules! scope_async_event_registration { - ($name:ident, $surface:expr) => { + ($name:ident, $callback_ty:ty, $surface:expr) => { /// Register a scope-local completion-based asynchronous event sanitizer. #[allow(clippy::missing_safety_doc)] #[unsafe(no_mangle)] @@ -207,7 +210,7 @@ macro_rules! scope_async_event_registration { scope_uuid: *const c_char, name: *const c_char, priority: i32, - cb: NemoRelayAsyncJsonCb, + cb: $callback_ty, user_data: *mut libc::c_void, free_fn: NemoRelayFreeFn, ) -> NemoRelayStatus { @@ -220,14 +223,17 @@ macro_rules! scope_async_event_registration { scope_async_event_registration!( nemo_relay_scope_register_mark_sanitize_guardrail_async, + NemoRelayAsyncJsonCb, Surface::Mark ); scope_async_event_registration!( nemo_relay_scope_register_scope_sanitize_start_guardrail_async, + NemoRelayAsyncJsonCb, Surface::Start ); scope_async_event_registration!( nemo_relay_scope_register_scope_sanitize_end_guardrail_async, + NemoRelayAsyncJsonCb, Surface::End ); diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 9fd1e136a..0e4da5eda 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -147,6 +147,18 @@ pub type NemoRelayAsyncNextResultCb = unsafe extern "C" fn( error_message: *const c_char, ); +struct SendUserData(*mut libc::c_void); + +// SAFETY: NemoRelayAsyncNextResultCb requires callers to keep user_data valid +// and safe to access until the asynchronously invoked callback runs. +unsafe impl Send for SendUserData {} + +impl SendUserData { + fn as_ptr(&self) -> *mut libc::c_void { + self.0 + } +} + struct CompletionWait { completion: Arc, receiver: tokio::sync::oneshot::Receiver>, @@ -353,17 +365,17 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke_callback( Box::pin(async move { collect_async_stream_for_completion(next(request).await?).await }) } }; - let user_data = user_data as usize; + let user_data = SendUserData(user_data); next.runtime.spawn(async move { match future.await { Ok(value) => { let value = json_to_c_string(&value); - unsafe { callback(user_data as *mut libc::c_void, value, ptr::null()) }; + unsafe { callback(user_data.as_ptr(), value, ptr::null()) }; unsafe { nemo_relay_string_free_internal(value) }; } Err(error) => { let error = CString::new(error.to_string()).unwrap_or_default(); - unsafe { callback(user_data as *mut libc::c_void, ptr::null(), error.as_ptr()) }; + unsafe { callback(user_data.as_ptr(), ptr::null(), error.as_ptr()) }; } } }); diff --git a/go/nemo_relay/adaptive_runtime_test.go b/go/nemo_relay/adaptive_runtime_test.go index 018ea3d11..7e96328ab 100644 --- a/go/nemo_relay/adaptive_runtime_test.go +++ b/go/nemo_relay/adaptive_runtime_test.go @@ -240,9 +240,15 @@ func assertAdaptiveRuntimeClosed(t *testing.T, runtime *AdaptiveRuntime) { func TestAdaptiveRuntimeHelpersRejectInvalidInputs(t *testing.T) { if _, err := BuildCacheTelemetryEvent(CacheTelemetryEventInput{ Provider: "unsupported", + RequestID: "018f13f0-7c1a-7a80-8000-000000000001", + }); err == nil || !strings.Contains(err.Error(), "provider") { + t.Fatalf("expected unsupported provider rejection, got %v", err) + } + if _, err := BuildCacheTelemetryEvent(CacheTelemetryEventInput{ + Provider: "openai", RequestID: "not-a-uuid", - }); err == nil { - t.Fatal("expected invalid telemetry input to fail") + }); err == nil || !strings.Contains(err.Error(), "request_id") { + t.Fatalf("expected invalid request ID rejection, got %v", err) } runtime, err := NewAdaptiveRuntime(testAdaptiveRuntimeConfig("openai")) @@ -251,11 +257,11 @@ func TestAdaptiveRuntimeHelpersRejectInvalidInputs(t *testing.T) { } defer runtime.Shutdown() if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{ - Provider: "unsupported", + Provider: "openai", RequestID: "not-a-uuid", AnnotatedRequest: json.RawMessage(`{}`), - }); err == nil { - t.Fatal("expected invalid cache request facts input to fail") + }); err == nil || !strings.Contains(err.Error(), "request_id") { + t.Fatalf("expected invalid request ID rejection, got %v", err) } if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{ AnnotatedRequest: json.RawMessage(`not-json`), diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index 3b5615f4f..c91e20e69 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -205,50 +205,81 @@ const asyncCancellationPollInterval = 10 * time.Millisecond type completionCancellationWatch struct { completion *C.NemoRelayAsyncCompletion cancel context.CancelFunc + probeMu *sync.Mutex } var ( completionCancellationMu sync.Mutex completionCancellationWatches = make(map[uint64]completionCancellationWatch) completionCancellationNextID atomic.Uint64 - completionCancellationOnce sync.Once + completionCancellationRunning bool ) func startCompletionCancellationMonitor() { - completionCancellationOnce.Do(func() { - go func() { - ticker := time.NewTicker(asyncCancellationPollInterval) - defer ticker.Stop() - for range ticker.C { + completionCancellationMu.Lock() + if completionCancellationRunning { + completionCancellationMu.Unlock() + return + } + completionCancellationRunning = true + completionCancellationMu.Unlock() + go func() { + ticker := time.NewTicker(asyncCancellationPollInterval) + defer ticker.Stop() + for range ticker.C { + completionCancellationMu.Lock() + if len(completionCancellationWatches) == 0 { + completionCancellationRunning = false + completionCancellationMu.Unlock() + return + } + snapshot := make(map[uint64]completionCancellationWatch, len(completionCancellationWatches)) + for id, watch := range completionCancellationWatches { + snapshot[id] = watch + } + completionCancellationMu.Unlock() + + for id, watch := range snapshot { + watch.probeMu.Lock() completionCancellationMu.Lock() - for id, watch := range completionCancellationWatches { - if bool(C.nemo_relay_async_completion_is_cancelled(watch.completion)) { + _, live := completionCancellationWatches[id] + completionCancellationMu.Unlock() + if live && bool(C.nemo_relay_async_completion_is_cancelled(watch.completion)) { + completionCancellationMu.Lock() + if _, live = completionCancellationWatches[id]; live { delete(completionCancellationWatches, id) + } + completionCancellationMu.Unlock() + if live { watch.cancel() } } - completionCancellationMu.Unlock() + watch.probeMu.Unlock() } - }() - }) + } + }() } func contextForCompletion(completion *C.NemoRelayAsyncCompletion) (context.Context, func()) { ctx, cancel := context.WithCancel(context.Background()) - startCompletionCancellationMonitor() id := completionCancellationNextID.Add(1) + probeMu := &sync.Mutex{} completionCancellationMu.Lock() completionCancellationWatches[id] = completionCancellationWatch{ completion: completion, cancel: cancel, + probeMu: probeMu, } completionCancellationMu.Unlock() + startCompletionCancellationMonitor() var doneOnce sync.Once return ctx, func() { doneOnce.Do(func() { completionCancellationMu.Lock() delete(completionCancellationWatches, id) completionCancellationMu.Unlock() + probeMu.Lock() + probeMu.Unlock() cancel() }) } From 7f672f78bf8f5a551b6c1f4074f85fb0a3c47afc Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 16:37:15 -0400 Subject: [PATCH 37/52] fix: retain async next callback ownership Signed-off-by: Will Killian --- go/nemo_relay/async_middleware_test.go | 42 ++++++++++++++++++++++++++ go/nemo_relay/callbacks.go | 14 ++++++--- 2 files changed, 52 insertions(+), 4 deletions(-) diff --git a/go/nemo_relay/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go index d34bfdcd9..ea8bfd2d0 100644 --- a/go/nemo_relay/async_middleware_test.go +++ b/go/nemo_relay/async_middleware_test.go @@ -9,6 +9,7 @@ import ( "errors" "strings" "testing" + "time" ) func asyncMiddlewareNoop(context.Context, json.RawMessage) (any, error) { @@ -269,3 +270,44 @@ func TestAsyncToolMiddlewarePropagatesCallbackAndNextErrors(t *testing.T) { } }) } + +func TestAsyncNextObservesOuterCancellationWithDetachedContext(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + const name = "go-async-detached-next" + nextStarted := make(chan struct{}) + releaseNext := make(chan struct{}) + if err := RegisterToolExecutionInterceptAsync(name, 0, + func(_ context.Context, invocation json.RawMessage, next AsyncNext) (any, error) { + go func() { + _, _ = next(context.Background(), invocation) + }() + <-nextStarted + return nil, errors.New("intercept returned early") + }, + ); err != nil { + t.Fatalf("register execution intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolExecutionIntercept(name) }) + + result := make(chan error, 1) + go func() { + _, err := ToolCallExecute(name, json.RawMessage(`{}`), func(args json.RawMessage) (json.RawMessage, error) { + close(nextStarted) + <-releaseNext + return args, nil + }) + result <- err + }() + + select { + case err := <-result: + if err == nil || !strings.Contains(err.Error(), "intercept returned early") { + t.Fatalf("execution error = %v, want intercept failure", err) + } + case <-time.After(time.Second): + close(releaseNext) + t.Fatal("detached next context prevented intercept cleanup") + } + close(releaseNext) + }) +} diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index c91e20e69..a424279a5 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -796,7 +796,8 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON nextMu.Unlock() C.nemo_relay_async_next_release(next) }() - nextFn := func(ctx context.Context, payload json.RawMessage) (json.RawMessage, error) { + outerCtx := ctx + nextFn := func(nextCtx context.Context, payload json.RawMessage) (json.RawMessage, error) { nextMu.RLock() defer nextMu.RUnlock() if !nextOpen { @@ -804,7 +805,6 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON } ch := make(chan asyncNextResult, 1) token := registerClosure(ch) - defer unregisterClosure(token) cPayload := C.CString(string(payload)) status := C.nemo_relay_async_next_invoke_callback( next, cPayload, @@ -812,13 +812,19 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON ) C.free(unsafe.Pointer(cPayload)) if err := checkStatus(status); err != nil { + // Rust did not retain user_data when invocation was rejected. + unregisterClosure(token) return nil, err } + // Successful invocation transfers token ownership to the one-shot + // result trampoline, even if this waiter is cancelled first. select { case result := <-ch: return result.value, result.err - case <-ctx.Done(): - return nil, ctx.Err() + case <-nextCtx.Done(): + return nil, nextCtx.Err() + case <-outerCtx.Done(): + return nil, outerCtx.Err() } } value, err := fn(ctx, invocation, nextFn) From 31cd595b204f13f2d20cc18ebff665e47664b24a Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 17:32:09 -0400 Subject: [PATCH 38/52] fix: contain Go async callback panics Signed-off-by: Will Killian --- crates/ffi/build.rs | 4 ++- crates/ffi/src/api/event_registry.rs | 4 +++ crates/ffi/src/callable.rs | 12 ++++++-- go/nemo_relay/adaptive_runtime_test.go | 2 +- go/nemo_relay/async_middleware_test.go | 41 ++++++++++++++++++++++++++ go/nemo_relay/callbacks.go | 13 ++++++++ 6 files changed, 71 insertions(+), 5 deletions(-) diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index 1e2cead74..26400e815 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -50,9 +50,11 @@ fn validate_async_registration_parity(crate_dir: &str) { "src/api/tool_registry.rs", ]; + println!("cargo:rerun-if-changed=cbindgen.toml"); + println!("cargo:rerun-if-changed=src"); + let mut expected = Vec::new(); for source in REGISTRATION_SOURCES { - println!("cargo:rerun-if-changed={source}"); let source_path = format!("{crate_dir}/{source}"); let contents = std::fs::read_to_string(&source_path) .unwrap_or_else(|error| panic!("read {source_path}: {error}")); diff --git a/crates/ffi/src/api/event_registry.rs b/crates/ffi/src/api/event_registry.rs index a1e471865..414d537c5 100644 --- a/crates/ffi/src/api/event_registry.rs +++ b/crates/ffi/src/api/event_registry.rs @@ -53,6 +53,8 @@ unsafe fn register_global_async( surface: Surface, ) -> NemoRelayStatus { clear_last_error(); + // The wrapper assumes ownership before validation so every return path + // invokes free_fn exactly once. let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); let name = match c_str_to_string(name) { Ok(name) => name, @@ -176,6 +178,8 @@ unsafe fn register_scope_async( surface: Surface, ) -> NemoRelayStatus { clear_last_error(); + // The wrapper assumes ownership before validation so every return path + // invokes free_fn exactly once. let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); let uuid = match parse_scope_uuid(scope_uuid) { Ok(uuid) => uuid, diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 0e4da5eda..e82443100 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -76,9 +76,10 @@ pub struct NemoRelayAsyncCompletion { /// Generic completion-based middleware callback. /// /// `invocation_json` is borrowed for the duration of the call. The completion -/// has one callback-owned reference. A callback returning `Complete` need not -/// release it; a callback returning `Pending` must eventually settle and call -/// `nemo_relay_async_completion_release`. +/// has one callback-owned reference. A callback returning `Complete` must not +/// release it because the runtime does so; a callback returning `Pending` must +/// eventually settle and call `nemo_relay_async_completion_release` exactly +/// once. pub type NemoRelayAsyncJsonCb = unsafe extern "C" fn( user_data: *mut libc::c_void, invocation_json: *const c_char, @@ -128,6 +129,11 @@ async fn collect_async_stream_for_completion(mut stream: LlmJsonStream) -> Resul } /// Completion-based execution-intercept callback. +/// +/// A callback returning `Complete` must not release either `completion` or +/// `next` because the runtime does so. A callback returning `Pending` must +/// eventually settle and release its callback-owned `completion` and `next` +/// references exactly once. pub type NemoRelayAsyncInterceptCb = unsafe extern "C" fn( user_data: *mut libc::c_void, invocation_json: *const c_char, diff --git a/go/nemo_relay/adaptive_runtime_test.go b/go/nemo_relay/adaptive_runtime_test.go index 7e96328ab..99f2ba220 100644 --- a/go/nemo_relay/adaptive_runtime_test.go +++ b/go/nemo_relay/adaptive_runtime_test.go @@ -44,7 +44,7 @@ func TestValidateAdaptiveConfigAndOwnedRuntime(t *testing.T) { if err != nil { t.Fatalf(newAdaptiveRuntimeFailedMsg, err) } - defer runtime.Shutdown() + defer func() { _ = runtime.Shutdown() }() if err := runtime.Register(); err != nil { t.Fatalf("Register failed: %v", err) } diff --git a/go/nemo_relay/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go index ea8bfd2d0..4261f0d08 100644 --- a/go/nemo_relay/async_middleware_test.go +++ b/go/nemo_relay/async_middleware_test.go @@ -271,6 +271,47 @@ func TestAsyncToolMiddlewarePropagatesCallbackAndNextErrors(t *testing.T) { }) } +func TestAsyncMiddlewarePanicsBecomeInvocationErrors(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + const conditionalName = "go-async-tool-conditional-panic" + if err := RegisterToolConditionalExecutionGuardrailAsync(conditionalName, 0, + func(context.Context, json.RawMessage) (any, error) { + panic("conditional callback panicked") + }, + ); err != nil { + t.Fatalf("register conditional: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolConditionalExecutionGuardrail(conditionalName) }) + + _, err := ToolCallExecute(conditionalName, json.RawMessage(`{}`), func(json.RawMessage) (json.RawMessage, error) { + return json.RawMessage(`{}`), nil + }) + if err == nil || !strings.Contains(err.Error(), "conditional callback panicked") { + t.Fatalf("conditional error = %v, want recovered panic", err) + } + if err := DeregisterToolConditionalExecutionGuardrail(conditionalName); err != nil { + t.Fatalf("deregister conditional: %v", err) + } + + const executionName = "go-async-tool-execution-panic" + if err := RegisterToolExecutionInterceptAsync(executionName, 0, + func(context.Context, json.RawMessage, AsyncNext) (any, error) { + panic("execution intercept panicked") + }, + ); err != nil { + t.Fatalf("register execution intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolExecutionIntercept(executionName) }) + + _, err = ToolCallExecute(executionName, json.RawMessage(`{}`), func(json.RawMessage) (json.RawMessage, error) { + return json.RawMessage(`{}`), nil + }) + if err == nil || !strings.Contains(err.Error(), "execution intercept panicked") { + t.Fatalf("execution error = %v, want recovered panic", err) + } + }) +} + func TestAsyncNextObservesOuterCancellationWithDetachedContext(t *testing.T) { runTestWithScopeStack(t, func(t *testing.T) { const name = "go-async-detached-next" diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index a424279a5..0cd23bbb3 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -100,6 +100,7 @@ import ( "context" "encoding/json" "errors" + "fmt" "sync" "sync/atomic" "time" @@ -721,6 +722,7 @@ func goAsyncMiddlewareTrampoline(userData unsafe.Pointer, invocationJSON *C.char invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) go func() { defer C.nemo_relay_async_completion_release(completion) + defer rejectAsyncCallbackPanic(completion) ctx, cancel := contextForCompletion(completion) defer cancel() value, err := fn(ctx, invocation) @@ -744,6 +746,16 @@ func goAsyncMiddlewareTrampoline(userData unsafe.Pointer, invocationJSON *C.char return asyncCallbackPending } +func rejectAsyncCallbackPanic(completion *C.NemoRelayAsyncCompletion) { + recovered := recover() + if recovered == nil { + return + } + message := C.CString(fmt.Sprintf("panic in async middleware callback: %v", recovered)) + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) +} + type asyncNextResult struct { value json.RawMessage err error @@ -784,6 +796,7 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) go func() { defer C.nemo_relay_async_completion_release(completion) + defer rejectAsyncCallbackPanic(completion) ctx, cancel := contextForCompletion(completion) defer cancel() var nextMu sync.RWMutex From 1a4d041d6fac72c5f6db7e99b2a835d5c0cfb84e Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 17:52:07 -0400 Subject: [PATCH 39/52] test: isolate null propagation context coverage Signed-off-by: Will Killian --- crates/ffi/tests/integration/api_tests.rs | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/crates/ffi/tests/integration/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index 4f6e90e10..a457c4238 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -723,6 +723,18 @@ fn scope_stack_propagation_and_thread_binding_validate_all_ffi_inputs() { }, NemoRelayStatus::NullPointer ); + let mut null_context_stack = ptr::null_mut(); + assert_eq!( + unsafe { + nemo_relay_scope_stack_create_from_propagation_json( + ptr::null(), + &mut null_context_stack, + ) + }, + NemoRelayStatus::NullPointer + ); + assert!(null_context_stack.is_null()); + let invalid_context = cstring("not-json"); let mut stack = ptr::null_mut(); assert_eq!( From c0a6907a60ef1e98baff8b7080012e3b37b8b54d Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:05:40 -0400 Subject: [PATCH 40/52] test: cover completion cancellation lifecycle Signed-off-by: Will Killian --- crates/ffi/src/callable.rs | 3 +- .../ffi/tests/unit/callable_private_tests.rs | 45 ++++++++++++++++--- crates/ffi/tests/unit/callable_tests.rs | 4 +- 3 files changed, 44 insertions(+), 8 deletions(-) diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index e82443100..b7378936d 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -952,7 +952,8 @@ pub fn wrap_async_llm_execution_intercept_fn( /// The completion ABI resolves one JSON value, so a stream intercept must /// resolve to an array of chunks. Relay replays that array as a stream after /// completion; incremental chunk delivery is not available through this ABI. -/// Relay rejects more than 4096 chunks or 16 MiB of serialized chunk data. +/// When the callback invokes `next`, Relay rejects more than 4096 chunks or +/// 16 MiB of serialized chunk data while collecting that continuation. pub fn wrap_async_llm_stream_execution_intercept_fn( cb: NemoRelayAsyncInterceptCb, user_data: *mut libc::c_void, diff --git a/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index f2aced5ee..13d0cc2d0 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -13,6 +13,16 @@ unsafe extern "C" fn complete_without_settling( NemoRelayAsyncCallbackState::Complete } +unsafe extern "C" fn retain_pending_completion( + user_data: *mut libc::c_void, + _invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + let slot = unsafe { &*user_data.cast::() }; + slot.store(completion as usize, Ordering::Release); + NemoRelayAsyncCallbackState::Pending +} + unsafe extern "C" fn send_next_result( user_data: *mut libc::c_void, value_json: *const c_char, @@ -99,14 +109,39 @@ fn async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlements() { ); unsafe { nemo_relay_async_completion_release(completion_ref) }; - let (sender, _receiver) = tokio::sync::oneshot::channel(); - let completion = Arc::new(NemoRelayAsyncCompletion { - sender: std::sync::Mutex::new(Some(sender)), - cancelled: AtomicBool::new(true), + let retained_completion = std::sync::atomic::AtomicUsize::new(0); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + runtime.block_on(async { + let mut invocation = Box::pin(invoke_async_json( + retain_pending_completion, + Arc::new(UserData { + ptr: (&retained_completion as *const std::sync::atomic::AtomicUsize) + .cast_mut() + .cast(), + free_fn: None, + }), + serde_json::json!({}), + )); + tokio::select! { + biased; + result = &mut invocation => panic!("pending callback unexpectedly settled: {result:?}"), + _ = tokio::task::yield_now() => {} + } + drop(invocation); }); - let completion_ref = Arc::into_raw(Arc::clone(&completion)); + let completion_ref = + retained_completion.load(Ordering::Acquire) as *const NemoRelayAsyncCompletion; + assert!(!completion_ref.is_null()); assert!(unsafe { nemo_relay_async_completion_is_cancelled(completion_ref) }); assert!(unsafe { nemo_relay_async_completion_is_cancelled(std::ptr::null()) }); + let value = CString::new(r#"{"late":true}"#).unwrap(); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion_ref, value.as_ptr()) }, + NemoRelayStatus::InvalidArg + ); assert_eq!( unsafe { nemo_relay_async_completion_reject(completion_ref, std::ptr::null()) }, NemoRelayStatus::InvalidArg diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 5fb38c3d1..75323d30d 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -738,7 +738,7 @@ fn test_llm_sanitizers_report_runtime_codec_ids_with_embedded_nul() { runtime_identity.clone(), ), )) - .expect_err("an embedded runtime codec ID must fail the async callback wrapper"); + .expect_err("an embedded runtime codec ID must fail the request sanitizer wrapper"); assert!( request_error .to_string() @@ -756,7 +756,7 @@ fn test_llm_sanitizers_report_runtime_codec_ids_with_embedded_nul() { json!({"secret": "must be preserved"}), nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), )) - .expect_err("an embedded runtime codec ID must fail the async callback wrapper"); + .expect_err("an embedded runtime codec ID must fail the response sanitizer wrapper"); assert!( response_error .to_string() From d5b0cbd74ca424cd82c70fa61cf90f1a85834e7b Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:12:43 -0400 Subject: [PATCH 41/52] fix: fail fast on FFI header generation Signed-off-by: Will Killian --- crates/ffi/build.rs | 38 +++++++++++++++++++------------------- 1 file changed, 19 insertions(+), 19 deletions(-) diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index 26400e815..ae871818b 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -8,29 +8,29 @@ fn main() { validate_async_registration_parity(&crate_dir); let config = cbindgen::Config::from_file(format!("{crate_dir}/cbindgen.toml")) .expect("Unable to read cbindgen.toml"); + let include_guard = config + .include_guard + .clone() + .expect("cbindgen.toml must configure an include guard"); - if let Ok(bindings) = cbindgen::Builder::new() + let bindings = cbindgen::Builder::new() .with_crate(&crate_dir) .with_config(config) .generate() - { - let header_path = format!("{crate_dir}/nemo_relay.h"); - bindings.write_to_file(&header_path); - // cbindgen intentionally does not expand declarative macros. Keep the - // macro-generated async registration functions in the generated C ABI. - let header = std::fs::read_to_string(&header_path).expect("read generated FFI header"); - let marker = "\n#endif /* NEMO_RELAY_H */\n"; - assert!( - header.contains(marker), - "generated FFI header is missing its NEMO_RELAY_H closing guard" - ); - let header = header.replacen( - marker, - &format!("\n{}\n#endif /* NEMO_RELAY_H */\n", ASYNC_REGISTRATIONS), - 1, - ); - std::fs::write(header_path, header).expect("write generated FFI header"); - } + .expect("Unable to generate FFI header"); + let header_path = format!("{crate_dir}/nemo_relay.h"); + bindings.write_to_file(&header_path); + // cbindgen intentionally does not expand declarative macros. Keep the + // macro-generated async registration functions in the generated C ABI. + let header = std::fs::read_to_string(&header_path).expect("read generated FFI header"); + let marker = format!("\n#endif /* {include_guard} */\n"); + assert!( + header.contains(&marker), + "generated FFI header is missing its configured closing guard" + ); + let replacement = format!("\n{}\n#endif /* {include_guard} */\n", ASYNC_REGISTRATIONS); + let header = header.replacen(&marker, &replacement, 1); + std::fs::write(header_path, header).expect("write generated FFI header"); } #[derive(Debug, PartialEq, Eq)] From fda6bf9cdb8f60a5e9ee810ecdd375e6cbb23ef6 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:15:43 -0400 Subject: [PATCH 42/52] refactor: consolidate async event registrations Signed-off-by: Will Killian --- crates/ffi/src/api/event_registry.rs | 136 ++++--------------------- go/nemo_relay/adaptive_runtime_test.go | 2 +- 2 files changed, 19 insertions(+), 119 deletions(-) diff --git a/crates/ffi/src/api/event_registry.rs b/crates/ffi/src/api/event_registry.rs index 414d537c5..45524edf7 100644 --- a/crates/ffi/src/api/event_registry.rs +++ b/crates/ffi/src/api/event_registry.rs @@ -44,69 +44,23 @@ unsafe fn register_global( .unwrap_or_else(|error| status_from_error(&error)) } -unsafe fn register_global_async( - name: *const c_char, - priority: i32, - cb: NemoRelayAsyncJsonCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - surface: Surface, -) -> NemoRelayStatus { - clear_last_error(); - // The wrapper assumes ownership before validation so every return path - // invokes free_fn exactly once. - let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); - let name = match c_str_to_string(name) { - Ok(name) => name, - Err(status) => return status, - }; - let result = match surface { - Surface::Mark => { - core_registry_api::register_mark_sanitize_guardrail(&name, priority, callback) - } - Surface::Start => { - core_registry_api::register_scope_sanitize_start_guardrail(&name, priority, callback) - } - Surface::End => { - core_registry_api::register_scope_sanitize_end_guardrail(&name, priority, callback) - } - }; - result - .map(|()| NemoRelayStatus::Ok) - .unwrap_or_else(|error| status_from_error(&error)) -} - -macro_rules! global_async_event_registration { - ($name:ident, $callback_ty:ty, $surface:expr) => { - /// Register a completion-based asynchronous event sanitizer. - #[allow(clippy::missing_safety_doc)] - #[unsafe(no_mangle)] - pub unsafe extern "C" fn $name( - name: *const c_char, - priority: i32, - cb: $callback_ty, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - ) -> NemoRelayStatus { - unsafe { register_global_async(name, priority, cb, user_data, free_fn, $surface) } - } - }; -} - -global_async_event_registration!( +global_async_registration!( nemo_relay_register_mark_sanitize_guardrail_async, NemoRelayAsyncJsonCb, - Surface::Mark + core_registry_api::register_mark_sanitize_guardrail, + wrap_async_event_sanitize_fn ); -global_async_event_registration!( +global_async_registration!( nemo_relay_register_scope_sanitize_start_guardrail_async, NemoRelayAsyncJsonCb, - Surface::Start + core_registry_api::register_scope_sanitize_start_guardrail, + wrap_async_event_sanitize_fn ); -global_async_event_registration!( +global_async_registration!( nemo_relay_register_scope_sanitize_end_guardrail_async, NemoRelayAsyncJsonCb, - Surface::End + core_registry_api::register_scope_sanitize_end_guardrail, + wrap_async_event_sanitize_fn ); unsafe fn deregister_global(name: *const c_char, surface: Surface) -> NemoRelayStatus { @@ -168,77 +122,23 @@ unsafe fn register_scope( .unwrap_or_else(|error| status_from_error(&error)) } -unsafe fn register_scope_async( - scope_uuid: *const c_char, - name: *const c_char, - priority: i32, - cb: NemoRelayAsyncJsonCb, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - surface: Surface, -) -> NemoRelayStatus { - clear_last_error(); - // The wrapper assumes ownership before validation so every return path - // invokes free_fn exactly once. - let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); - let uuid = match parse_scope_uuid(scope_uuid) { - Ok(uuid) => uuid, - Err(status) => return status, - }; - let name = match c_str_to_string(name) { - Ok(name) => name, - Err(status) => return status, - }; - let result = match surface { - Surface::Mark => core_registry_api::scope_register_mark_sanitize_guardrail( - &uuid, &name, priority, callback, - ), - Surface::Start => core_registry_api::scope_register_scope_sanitize_start_guardrail( - &uuid, &name, priority, callback, - ), - Surface::End => core_registry_api::scope_register_scope_sanitize_end_guardrail( - &uuid, &name, priority, callback, - ), - }; - result - .map(|()| NemoRelayStatus::Ok) - .unwrap_or_else(|error| status_from_error(&error)) -} - -macro_rules! scope_async_event_registration { - ($name:ident, $callback_ty:ty, $surface:expr) => { - /// Register a scope-local completion-based asynchronous event sanitizer. - #[allow(clippy::missing_safety_doc)] - #[unsafe(no_mangle)] - pub unsafe extern "C" fn $name( - scope_uuid: *const c_char, - name: *const c_char, - priority: i32, - cb: $callback_ty, - user_data: *mut libc::c_void, - free_fn: NemoRelayFreeFn, - ) -> NemoRelayStatus { - unsafe { - register_scope_async(scope_uuid, name, priority, cb, user_data, free_fn, $surface) - } - } - }; -} - -scope_async_event_registration!( +scope_async_registration!( nemo_relay_scope_register_mark_sanitize_guardrail_async, NemoRelayAsyncJsonCb, - Surface::Mark + core_registry_api::scope_register_mark_sanitize_guardrail, + wrap_async_event_sanitize_fn ); -scope_async_event_registration!( +scope_async_registration!( nemo_relay_scope_register_scope_sanitize_start_guardrail_async, NemoRelayAsyncJsonCb, - Surface::Start + core_registry_api::scope_register_scope_sanitize_start_guardrail, + wrap_async_event_sanitize_fn ); -scope_async_event_registration!( +scope_async_registration!( nemo_relay_scope_register_scope_sanitize_end_guardrail_async, NemoRelayAsyncJsonCb, - Surface::End + core_registry_api::scope_register_scope_sanitize_end_guardrail, + wrap_async_event_sanitize_fn ); unsafe fn deregister_scope( diff --git a/go/nemo_relay/adaptive_runtime_test.go b/go/nemo_relay/adaptive_runtime_test.go index 99f2ba220..1051e053a 100644 --- a/go/nemo_relay/adaptive_runtime_test.go +++ b/go/nemo_relay/adaptive_runtime_test.go @@ -159,7 +159,7 @@ func TestAdaptiveRuntimeBindScopeRejectsNilScope(t *testing.T) { if err != nil { t.Fatalf(newAdaptiveRuntimeFailedMsg, err) } - defer runtime.Shutdown() + defer func() { _ = runtime.Shutdown() }() if runtime.BindScope(nil) == nil { t.Fatal("expected BindScope to reject nil scope") From 32ac7f1a0cc5421470d46af91373ae7c4378261e Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:25:28 -0400 Subject: [PATCH 43/52] test: cover async conditional rejection paths Signed-off-by: Will Killian --- crates/ffi/tests/unit/callable_tests.rs | 81 +++++++++++++++++++++++++ 1 file changed, 81 insertions(+) diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 75323d30d..eb26c2333 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -42,6 +42,8 @@ unsafe extern "C" fn async_json_passthrough_callback( "optimization_contributions": [], }), 7 => invocation["fields"].clone(), + 8 => Json::String("blocked by async guardrail".into()), + 9 => json!({"invalid": true}), _ => unreachable!("test callback kind must be known"), }; let value = CString::new(value.to_string()).expect("JSON has no NUL"); @@ -52,6 +54,20 @@ unsafe extern "C" fn async_json_passthrough_callback( NemoRelayAsyncCallbackState::Complete } +unsafe extern "C" fn async_invalid_stream_callback( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + let value = CString::new(json!({"not": "an array"}).to_string()).expect("JSON has no NUL"); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, + NemoRelayStatus::Ok + ); + NemoRelayAsyncCallbackState::Complete +} + fn async_callback_user_data(kind: usize) -> *mut libc::c_void { Box::into_raw(Box::new(kind)).cast() } @@ -203,6 +219,71 @@ fn async_callback_wrappers_cover_all_middleware_shapes() { ); } +#[test] +fn async_conditional_and_stream_wrappers_validate_callback_results() { + let tool_rejection = wrap_async_tool_conditional_fn( + async_json_passthrough_callback, + async_callback_user_data(8), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(tool_rejection("tool".into(), json!({}))).unwrap(), + Some("blocked by async guardrail".into()) + ); + + let llm_rejection = wrap_async_llm_conditional_fn( + async_json_passthrough_callback, + async_callback_user_data(8), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(llm_rejection(make_request())).unwrap(), + Some("blocked by async guardrail".into()) + ); + + let tool_invalid = wrap_async_tool_conditional_fn( + async_json_passthrough_callback, + async_callback_user_data(9), + Some(free_async_callback_user_data), + ); + assert!( + resolve(tool_invalid("tool".into(), json!({}))) + .unwrap_err() + .to_string() + .contains("expected string or null") + ); + + let llm_invalid = wrap_async_llm_conditional_fn( + async_json_passthrough_callback, + async_callback_user_data(9), + Some(free_async_callback_user_data), + ); + assert!( + resolve(llm_invalid(make_request())) + .unwrap_err() + .to_string() + .contains("expected string or null") + ); + + let stream_intercept = wrap_async_llm_stream_execution_intercept_fn( + async_invalid_stream_callback, + std::ptr::null_mut(), + None, + ); + let next: nemo_relay::api::runtime::LlmStreamExecutionNextFn = Arc::new(|_request| { + Box::pin(async { + Ok(nemo_relay::api::runtime::LlmJsonStream::new( + tokio_stream::empty(), + )) + }) + }); + let result = resolve(stream_intercept("llm", make_request(), next)); + let Err(error) = result else { + panic!("a non-array async stream result must fail"); + }; + assert!(error.to_string().contains("must resolve to an array")); +} + #[test] fn async_execution_wrappers_continue_tool_and_llm_calls() { let tool_intercept = wrap_async_tool_execution_intercept_fn( From f8185db172627ec13c9310f880305c98d14bf443 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 18:29:30 -0400 Subject: [PATCH 44/52] fix: tighten async FFI parity and isolation Signed-off-by: Will Killian --- crates/ffi/build.rs | 14 ++++++++------ crates/ffi/src/callable.rs | 6 ++++++ crates/ffi/tests/integration/api_tests.rs | 10 ++++++++++ 3 files changed, 24 insertions(+), 6 deletions(-) diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index ae871818b..92f5240c0 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -112,12 +112,14 @@ fn parse_async_macro_invocations(source: &str) -> Vec { .collect::>(); remaining = &invocation[end + 2..]; - let Some(name) = arguments.first().copied() else { - continue; - }; - if !name.starts_with("nemo_relay_") || !name.ends_with("_async") { - continue; - } + let name = arguments + .first() + .copied() + .expect("async registration macro invocation is missing its export name"); + assert!( + name.starts_with("nemo_relay_") && name.ends_with("_async"), + "async registration macro exported unexpected name {name}; expected nemo_relay_*_async" + ); let callback_type = arguments .get(1) .unwrap_or_else(|| panic!("{name} macro invocation is missing its callback type")); diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index b7378936d..46a743b74 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -824,6 +824,9 @@ pub fn wrap_async_llm_conditional_fn( } /// Wrap a completion-based C LLM request sanitizer. +/// +/// The async invocation envelope includes `codec_kind` and `codec_id`, but not +/// the borrowed codec capability available to synchronous callbacks. pub fn wrap_async_llm_sanitize_request_fn( cb: NemoRelayAsyncJsonCb, user_data: *mut libc::c_void, @@ -855,6 +858,9 @@ pub fn wrap_async_llm_sanitize_request_fn( } /// Wrap a completion-based C LLM response sanitizer. +/// +/// The async invocation envelope includes `codec_kind` and `codec_id`, but not +/// the borrowed codec capability available to synchronous callbacks. pub fn wrap_async_llm_sanitize_response_fn( cb: NemoRelayAsyncJsonCb, user_data: *mut libc::c_void, diff --git a/crates/ffi/tests/integration/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index a457c4238..e13d80d24 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -771,6 +771,12 @@ fn scope_stack_propagation_and_thread_binding_validate_all_ffi_inputs() { NemoRelayStatus::NullPointer ); + let mut original_binding = ptr::null_mut(); + assert_eq!( + unsafe { nemo_relay_scope_stack_capture_thread(&mut original_binding) }, + NemoRelayStatus::Ok + ); + assert!(!original_binding.is_null()); let stack = unsafe { fresh_scope_stack() }; let mut binding = ptr::null_mut(); assert_eq!( @@ -782,6 +788,10 @@ fn scope_stack_propagation_and_thread_binding_validate_all_ffi_inputs() { unsafe { nemo_relay_scope_stack_restore_thread(binding) }, NemoRelayStatus::Ok ); + assert_eq!( + unsafe { nemo_relay_scope_stack_restore_thread(original_binding) }, + NemoRelayStatus::Ok + ); unsafe { nemo_relay_scope_stack_free(stack) }; } From 23daaf8219e88d4b415bb9e8d7de4b180901ed9d Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 19:18:25 -0400 Subject: [PATCH 45/52] test: cover async next callback failures Signed-off-by: Will Killian --- crates/ffi/nemo_relay.h | 3 ++ crates/ffi/src/callable.rs | 3 ++ .../ffi/tests/unit/callable_private_tests.rs | 29 +++++++++++++++++++ 3 files changed, 35 insertions(+) diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index fa52c428d..7999aec83 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -2585,6 +2585,9 @@ NemoRelayStatus nemo_relay_async_next_invoke(const struct NemoRelayAsyncNext *ne /** * Invoke the next execution layer and report its result through a callback. + * + * A non-`Ok` return means invocation was not scheduled and `callback` is + * never invoked; the caller owns any state it allocated for `user_data`. */ NemoRelayStatus nemo_relay_async_next_invoke_callback(const struct NemoRelayAsyncNext *next, const char *invocation_json, diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 46a743b74..9ec65b717 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -332,6 +332,9 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke( } /// Invoke the next execution layer and report its result through a callback. +/// +/// A non-`Ok` return means invocation was not scheduled and `callback` is +/// never invoked; the caller owns any state it allocated for `user_data`. #[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. #[unsafe(no_mangle)] pub unsafe extern "C" fn nemo_relay_async_next_invoke_callback( diff --git a/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index 13d0cc2d0..bea23fa94 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -298,4 +298,33 @@ fn async_next_callback_reports_tool_llm_and_stream_results() { assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); unsafe { nemo_relay_async_next_release(next_ref) }; } + + let next = Arc::new(NemoRelayAsyncNext { + inner: AsyncNextInner::Tool(Arc::new(|_value| { + Box::pin(async { Err(FlowError::Internal("next failed".into())) }) + })), + runtime: runtime.handle().clone(), + }); + let next_ref = Arc::into_raw(next); + let invocation = CString::new("{}").unwrap(); + let (sender, receiver) = tokio::sync::oneshot::channel::>(); + assert_eq!( + unsafe { + nemo_relay_async_next_invoke_callback( + next_ref, + invocation.as_ptr(), + send_next_result, + Box::into_raw(Box::new(sender)).cast(), + ) + }, + NemoRelayStatus::Ok + ); + assert!( + runtime + .block_on(receiver) + .unwrap() + .unwrap_err() + .contains("next failed") + ); + unsafe { nemo_relay_async_next_release(next_ref) }; } From 91dd5c576dda8e1bb167c2d98f2b7c543964e0ac Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 19:46:04 -0400 Subject: [PATCH 46/52] fix: scope reentrant flush guards to callbacks Signed-off-by: Will Killian --- crates/ffi/nemo_relay.h | 5 ++--- crates/ffi/src/api/llm_registry.rs | 8 +++++--- crates/ffi/src/callable.rs | 24 +++++++++++++++++++++++- 3 files changed, 30 insertions(+), 7 deletions(-) diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 7999aec83..b5ec74e9f 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -1291,9 +1291,8 @@ NemoRelayStatus nemo_relay_deregister_subscriber(const char *name); /** * Wait for subscriber callbacks queued before this call to finish. * - * Call this function outside native subscriber callbacks. A re-entrant call returns without - * waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can - * still run. + * A call from an asynchronous event-sanitizer callback returns without waiting to prevent a + * cycle with the serial dispatcher. */ NemoRelayStatus nemo_relay_flush_subscribers(void); diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index 218ca1138..d6463c3ac 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -434,12 +434,14 @@ pub unsafe extern "C" fn nemo_relay_deregister_subscriber(name: *const c_char) - /// Wait for subscriber callbacks queued before this call to finish. /// -/// Call this function outside native subscriber callbacks. A re-entrant call returns without -/// waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can -/// still run. +/// A call from an asynchronous event-sanitizer callback returns without waiting to prevent a +/// cycle with the serial dispatcher. #[unsafe(no_mangle)] pub extern "C" fn nemo_relay_flush_subscribers() -> NemoRelayStatus { clear_last_error(); + if crate::callable::event_sanitizer_callback_active() { + return NemoRelayStatus::Ok; + } match core_subscriber_api::flush_subscribers() { Ok(()) => NemoRelayStatus::Ok, Err(e) => status_from_error(&e), diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 9ec65b717..409484d0f 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -21,7 +21,7 @@ use std::future::Future; use std::pin::Pin; use std::ptr; use std::sync::Arc; -use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use libc::c_char; use nemo_relay::api::runtime::{ @@ -170,6 +170,27 @@ struct CompletionWait { receiver: tokio::sync::oneshot::Receiver>, } +static ACTIVE_EVENT_SANITIZER_CALLBACKS: AtomicUsize = AtomicUsize::new(0); + +struct ActiveEventSanitizerCallback; + +impl ActiveEventSanitizerCallback { + fn enter() -> Self { + ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_add(1, Ordering::AcqRel); + Self + } +} + +impl Drop for ActiveEventSanitizerCallback { + fn drop(&mut self) { + ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_sub(1, Ordering::AcqRel); + } +} + +pub(crate) fn event_sanitizer_callback_active() -> bool { + ACTIVE_EVENT_SANITIZER_CALLBACKS.load(Ordering::Acquire) != 0 +} + impl Drop for CompletionWait { fn drop(&mut self) { self.completion.cancelled.store(true, Ordering::Release); @@ -793,6 +814,7 @@ pub fn wrap_async_event_sanitize_fn( Arc::new(move |event: Arc, fields: EventSanitizeFields| { let user_data = user_data.clone(); Box::pin(async move { + let _active_callback = ActiveEventSanitizerCallback::enter(); let value = invoke_async_json( cb, user_data, From ca0293fe1189d24b430842d9c9a9d25add874f7b Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 19:51:32 -0400 Subject: [PATCH 47/52] fix: address async FFI review findings Signed-off-by: Will Killian --- crates/ffi/src/callable.rs | 4 +-- crates/ffi/tests/integration/api_tests.rs | 14 +++++++++ .../ffi/tests/unit/callable_private_tests.rs | 30 +++++++++++++++++++ go/nemo_relay/async_middleware_test.go | 23 ++++++++++++-- go/nemo_relay/callbacks.go | 8 ++--- go/nemo_relay/optimization_test.go | 2 +- 6 files changed, 72 insertions(+), 9 deletions(-) diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 409484d0f..6363bf126 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -310,8 +310,6 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke( let Some(invocation) = c_str_to_json(invocation_json) else { return NemoRelayStatus::InvalidJson; }; - unsafe { Arc::increment_strong_count(completion) }; - let completion = unsafe { Arc::from_raw(completion) }; let future: Pin> + Send>> = match &next.inner { AsyncNextInner::Tool(next) => { let next = next.clone(); @@ -338,6 +336,8 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke( Box::pin(async move { collect_async_stream_for_completion(next(request).await?).await }) } }; + unsafe { Arc::increment_strong_count(completion) }; + let completion = unsafe { Arc::from_raw(completion) }; next.runtime.spawn(async move { let result = future.await; if let Some(sender) = completion diff --git a/crates/ffi/tests/integration/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index e13d80d24..618968006 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -843,6 +843,20 @@ fn observability_component_helpers_serialize_defaults_and_validate_inputs() { NemoRelayStatus::InvalidJson ); assert!(rejected.is_null()); + + let wrong_shape = cstring(r#"{"version":"invalid"}"#); + let mut wrong_shape_out = ptr::null_mut(); + assert_eq!( + unsafe { + api::nemo_relay_observability_component_spec_json( + wrong_shape.as_ptr(), + false, + &mut wrong_shape_out, + ) + }, + NemoRelayStatus::InvalidJson + ); + assert!(wrong_shape_out.is_null()); } #[test] diff --git a/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index bea23fa94..48cd32a68 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -243,6 +243,36 @@ fn async_next_invocation_supports_tool_llm_and_stream_continuations() { nemo_relay_async_completion_release(completion_ref); } } + + let next = Arc::new(NemoRelayAsyncNext { + inner: AsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + runtime: runtime.handle().clone(), + }); + let next_ref = Arc::into_raw(next); + let (sender, _receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + let malformed_request = CString::new(r#"{"content":{}}"#).unwrap(); + assert_eq!( + unsafe { + nemo_relay_async_next_invoke(next_ref, malformed_request.as_ptr(), completion_ref) + }, + NemoRelayStatus::InvalidJson + ); + assert_eq!( + Arc::strong_count(&completion), + 2, + "rejected invocation must not retain the completion" + ); + unsafe { + nemo_relay_async_next_release(next_ref); + nemo_relay_async_completion_release(completion_ref); + } } #[test] diff --git a/go/nemo_relay/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go index 4261f0d08..e010964ea 100644 --- a/go/nemo_relay/async_middleware_test.go +++ b/go/nemo_relay/async_middleware_test.go @@ -8,6 +8,7 @@ import ( "encoding/json" "errors" "strings" + "sync" "testing" "time" ) @@ -317,9 +318,15 @@ func TestAsyncNextObservesOuterCancellationWithDetachedContext(t *testing.T) { const name = "go-async-detached-next" nextStarted := make(chan struct{}) releaseNext := make(chan struct{}) + nextDone := make(chan struct{}) + var releaseOnce sync.Once + release := func() { + releaseOnce.Do(func() { close(releaseNext) }) + } if err := RegisterToolExecutionInterceptAsync(name, 0, func(_ context.Context, invocation json.RawMessage, next AsyncNext) (any, error) { go func() { + defer close(nextDone) _, _ = next(context.Background(), invocation) }() <-nextStarted @@ -339,6 +346,14 @@ func TestAsyncNextObservesOuterCancellationWithDetachedContext(t *testing.T) { }) result <- err }() + defer func() { + release() + select { + case <-nextDone: + case <-time.After(time.Second): + t.Error("detached next continuation never settled during cleanup") + } + }() select { case err := <-result: @@ -346,9 +361,13 @@ func TestAsyncNextObservesOuterCancellationWithDetachedContext(t *testing.T) { t.Fatalf("execution error = %v, want intercept failure", err) } case <-time.After(time.Second): - close(releaseNext) t.Fatal("detached next context prevented intercept cleanup") } - close(releaseNext) + release() + select { + case <-nextDone: + case <-time.After(time.Second): + t.Fatal("detached next continuation never settled") + } }) } diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index 0cd23bbb3..b7b4626e2 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -722,7 +722,7 @@ func goAsyncMiddlewareTrampoline(userData unsafe.Pointer, invocationJSON *C.char invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) go func() { defer C.nemo_relay_async_completion_release(completion) - defer rejectAsyncCallbackPanic(completion) + defer rejectAsyncCallbackPanic(completion, "middleware") ctx, cancel := contextForCompletion(completion) defer cancel() value, err := fn(ctx, invocation) @@ -746,12 +746,12 @@ func goAsyncMiddlewareTrampoline(userData unsafe.Pointer, invocationJSON *C.char return asyncCallbackPending } -func rejectAsyncCallbackPanic(completion *C.NemoRelayAsyncCompletion) { +func rejectAsyncCallbackPanic(completion *C.NemoRelayAsyncCompletion, kind string) { recovered := recover() if recovered == nil { return } - message := C.CString(fmt.Sprintf("panic in async middleware callback: %v", recovered)) + message := C.CString(fmt.Sprintf("panic in async %s callback: %v", kind, recovered)) defer C.free(unsafe.Pointer(message)) C.nemo_relay_async_completion_reject(completion, message) } @@ -796,7 +796,7 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) go func() { defer C.nemo_relay_async_completion_release(completion) - defer rejectAsyncCallbackPanic(completion) + defer rejectAsyncCallbackPanic(completion, "execution intercept") ctx, cancel := contextForCompletion(completion) defer cancel() var nextMu sync.RWMutex diff --git a/go/nemo_relay/optimization_test.go b/go/nemo_relay/optimization_test.go index 491899ce5..a4fab1eaf 100644 --- a/go/nemo_relay/optimization_test.go +++ b/go/nemo_relay/optimization_test.go @@ -98,7 +98,7 @@ func TestLLMOptimizationContributionOmittedAppliedIsNonApplied(t *testing.T) { } } -func TestLLMOptimizationContributionRejectsMalformedAndUnknownWireShapes(t *testing.T) { +func TestLLMOptimizationContributionRejectsMalformedAndNonObjectWireShapes(t *testing.T) { var contribution LLMOptimizationContribution if err := json.Unmarshal([]byte(`not-json`), &contribution); err == nil { t.Fatal("expected malformed optimization contribution JSON to fail") From c1e22da74418c2d1384f4b989b2e9f994211c71e Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 20:14:08 -0400 Subject: [PATCH 48/52] fix: tighten async FFI review safeguards Signed-off-by: Will Killian --- crates/ffi/build.rs | 88 ++++++++++++++++++++++++- crates/ffi/nemo_relay.h | 5 +- crates/ffi/src/api/llm_registry.rs | 10 ++- crates/ffi/tests/unit/callable_tests.rs | 11 +++- go/nemo_relay/async_middleware_test.go | 1 + go/nemo_relay/callbacks.go | 3 + 6 files changed, 110 insertions(+), 8 deletions(-) diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index 92f5240c0..9ee7ea628 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -53,6 +53,11 @@ fn validate_async_registration_parity(crate_dir: &str) { println!("cargo:rerun-if-changed=cbindgen.toml"); println!("cargo:rerun-if-changed=src"); + let callable_path = format!("{crate_dir}/src/callable.rs"); + let callable = std::fs::read_to_string(&callable_path) + .unwrap_or_else(|error| panic!("read {callable_path}: {error}")); + validate_async_callback_abi(&callable); + let mut expected = Vec::new(); for source in REGISTRATION_SOURCES { let source_path = format!("{crate_dir}/{source}"); @@ -80,6 +85,87 @@ fn validate_async_registration_parity(crate_dir: &str) { ); } +fn normalize_whitespace(value: &str) -> String { + value.split_whitespace().collect::>().join(" ") +} + +fn rust_type_alias(source: &str, name: &str) -> String { + let prefix = format!("pub type {name}"); + let start = source + .find(&prefix) + .unwrap_or_else(|| panic!("src/callable.rs is missing {name}")); + let end = source[start..] + .find(';') + .map(|offset| start + offset + 1) + .unwrap_or_else(|| panic!("src/callable.rs has an unterminated {name} alias")); + normalize_whitespace(&source[start..end]) +} + +/// Keep the handwritten C typedef block tied to the Rust callback ABI that +/// cbindgen cannot derive through registration macros. +fn validate_async_callback_abi(callable: &str) { + let enum_start = callable + .find("pub enum NemoRelayAsyncCallbackState") + .expect("src/callable.rs is missing NemoRelayAsyncCallbackState"); + let enum_prefix = &callable[..enum_start]; + assert!( + enum_prefix + .rsplit_once("#[repr(u32)]") + .is_some_and(|(_, suffix)| suffix.len() < 256), + "NemoRelayAsyncCallbackState must retain its u32 representation" + ); + let enum_end = callable[enum_start..] + .find("\n}") + .map(|offset| enum_start + offset) + .expect("NemoRelayAsyncCallbackState is unterminated"); + let discriminants = callable[enum_start..enum_end] + .lines() + .map(str::trim) + .filter(|line| line.starts_with("Complete =") || line.starts_with("Pending =")) + .collect::>(); + assert_eq!( + discriminants, + ["Complete = 0,", "Pending = 1,"], + "NemoRelayAsyncCallbackState drifted from the C callback-state constants" + ); + + assert_eq!( + rust_type_alias(callable, "NemoRelayAsyncJsonCb"), + normalize_whitespace( + r#"pub type NemoRelayAsyncJsonCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, + ) -> NemoRelayAsyncCallbackState;"# + ), + "NemoRelayAsyncJsonCb drifted from ASYNC_REGISTRATIONS" + ); + assert_eq!( + rust_type_alias(callable, "NemoRelayAsyncInterceptCb"), + normalize_whitespace( + r#"pub type NemoRelayAsyncInterceptCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, + ) -> NemoRelayAsyncCallbackState;"# + ), + "NemoRelayAsyncInterceptCb drifted from ASYNC_REGISTRATIONS" + ); + for declaration in [ + "typedef uint32_t NemoRelayAsyncCallbackState;", + "NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0,", + "NEMO_RELAY_ASYNC_CALLBACK_STATE_PENDING = 1,", + "typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion);", + "typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncCompletion *completion);", + ] { + assert!( + ASYNC_REGISTRATIONS.contains(declaration), + "ASYNC_REGISTRATIONS is missing callback ABI declaration: {declaration}" + ); + } +} + fn parse_async_prototype(line: &str) -> Option { let line = line.strip_prefix("NemoRelayStatus ")?; let (name, parameters) = line.split_once('(')?; @@ -94,8 +180,6 @@ fn parse_async_macro_invocations(source: &str) -> Vec { const MACROS: &[(&str, bool)] = &[ ("global_async_registration!(", false), ("scope_async_registration!(", true), - ("global_async_event_registration!(", false), - ("scope_async_event_registration!(", true), ]; let mut prototypes = Vec::new(); diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index b5ec74e9f..843408d98 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -1291,8 +1291,9 @@ NemoRelayStatus nemo_relay_deregister_subscriber(const char *name); /** * Wait for subscriber callbacks queued before this call to finish. * - * A call from an asynchronous event-sanitizer callback returns without waiting to prevent a - * cycle with the serial dispatcher. + * A call made while an asynchronous C event sanitizer is pending returns `Internal` with a + * would-block error instead of waiting. Completion callbacks may resume on arbitrary threads, so + * callers that are not inside the sanitizer can retry after it settles. */ NemoRelayStatus nemo_relay_flush_subscribers(void); diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index d6463c3ac..777db2663 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -434,13 +434,17 @@ pub unsafe extern "C" fn nemo_relay_deregister_subscriber(name: *const c_char) - /// Wait for subscriber callbacks queued before this call to finish. /// -/// A call from an asynchronous event-sanitizer callback returns without waiting to prevent a -/// cycle with the serial dispatcher. +/// A call made while an asynchronous C event sanitizer is pending returns `Internal` with a +/// would-block error instead of waiting. Completion callbacks may resume on arbitrary threads, so +/// callers that are not inside the sanitizer can retry after it settles. #[unsafe(no_mangle)] pub extern "C" fn nemo_relay_flush_subscribers() -> NemoRelayStatus { clear_last_error(); if crate::callable::event_sanitizer_callback_active() { - return NemoRelayStatus::Ok; + crate::error::set_last_error( + "subscriber flush would block while an asynchronous event sanitizer is pending", + ); + return NemoRelayStatus::Internal; } match core_subscriber_api::flush_subscribers() { Ok(()) => NemoRelayStatus::Ok, diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index eb26c2333..09f8e7087 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -41,7 +41,14 @@ unsafe extern "C" fn async_json_passthrough_callback( "pending_marks": [], "optimization_contributions": [], }), - 7 => invocation["fields"].clone(), + 7 => { + assert_eq!( + crate::api::nemo_relay_flush_subscribers(), + NemoRelayStatus::Internal, + "flush inside an async event sanitizer must report would-block" + ); + invocation["fields"].clone() + } 8 => Json::String("blocked by async guardrail".into()), 9 => json!({"invalid": true}), _ => unreachable!("test callback kind must be known"), @@ -102,6 +109,8 @@ unsafe extern "C" fn async_next_callback( nemo_relay_async_next_release(next); nemo_relay_async_completion_release(completion); } + // A successful invoke retained a completion reference; this callback drops + // its own next and completion references before transferring Pending ownership. NemoRelayAsyncCallbackState::Pending } diff --git a/go/nemo_relay/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go index e010964ea..f3eb8879c 100644 --- a/go/nemo_relay/async_middleware_test.go +++ b/go/nemo_relay/async_middleware_test.go @@ -59,6 +59,7 @@ func TestAsyncMiddlewareGlobalRegistrationParity(t *testing.T) { if err := registration.register(name); err != nil { t.Fatalf("register: %v", err) } + t.Cleanup(func() { _ = registration.deregister(name) }) if err := registration.register(name); err == nil { t.Fatal("duplicate registration unexpectedly succeeded") } diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index b7b4626e2..2eeb20811 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -831,6 +831,9 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON } // Successful invocation transfers token ownership to the one-shot // result trampoline, even if this waiter is cancelled first. + // Keep the token registered when cancellation wins: unregistering + // before a late Rust callback would be a use-after-free. Runtime + // teardown that prevents delivery can therefore retain this token. select { case result := <-ch: return result.value, result.err From fcce7241144a7a17174e104d20118f4b669c38eb Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 20:17:59 -0400 Subject: [PATCH 49/52] perf(ffi): avoid buffering stream size checks Signed-off-by: Will Killian --- crates/ffi/src/callable.rs | 30 +++++++++++++------ .../ffi/tests/unit/callable_private_tests.rs | 16 ++++++++++ 2 files changed, 37 insertions(+), 9 deletions(-) diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 6363bf126..52d425c9d 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -101,9 +101,25 @@ enum AsyncNextInner { const ASYNC_STREAM_MAX_CHUNKS: usize = 4096; const ASYNC_STREAM_MAX_SERIALIZED_BYTES: usize = 16 * 1024 * 1024; +#[derive(Default)] +struct SerializedByteCounter { + bytes: usize, +} + +impl std::io::Write for SerializedByteCounter { + fn write(&mut self, buffer: &[u8]) -> std::io::Result { + self.bytes = self.bytes.saturating_add(buffer.len()); + Ok(buffer.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + async fn collect_async_stream_for_completion(mut stream: LlmJsonStream) -> Result { let mut chunks = Vec::new(); - let mut serialized_bytes = 0usize; + let mut serialized_bytes = SerializedByteCounter::default(); while let Some(chunk) = stream.next().await { let chunk = chunk?; if chunks.len() >= ASYNC_STREAM_MAX_CHUNKS { @@ -111,14 +127,10 @@ async fn collect_async_stream_for_completion(mut stream: LlmJsonStream) -> Resul "async stream continuation exceeded the {ASYNC_STREAM_MAX_CHUNKS}-chunk completion limit" ))); } - serialized_bytes = serialized_bytes.saturating_add( - serde_json::to_vec(&chunk) - .map_err(|error| { - FlowError::Internal(format!("failed to measure async stream chunk: {error}")) - })? - .len(), - ); - if serialized_bytes > ASYNC_STREAM_MAX_SERIALIZED_BYTES { + serde_json::to_writer(&mut serialized_bytes, &chunk).map_err(|error| { + FlowError::Internal(format!("failed to measure async stream chunk: {error}")) + })?; + if serialized_bytes.bytes > ASYNC_STREAM_MAX_SERIALIZED_BYTES { return Err(FlowError::Internal(format!( "async stream continuation exceeded the {ASYNC_STREAM_MAX_SERIALIZED_BYTES}-byte completion limit" ))); diff --git a/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index 48cd32a68..8bbdeee47 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -57,6 +57,22 @@ fn test_callable_private_helper_paths() { unsafe { nemo_relay_string_free_internal(raw) }; } +#[test] +fn serialized_byte_counter_accumulates_json_without_buffering() { + let values = [serde_json::json!({"chunk": 1}), serde_json::json!("two")]; + let expected: usize = values + .iter() + .map(|value| serde_json::to_vec(value).unwrap().len()) + .sum(); + let mut counter = SerializedByteCounter::default(); + + for value in &values { + serde_json::to_writer(&mut counter, value).unwrap(); + } + + assert_eq!(counter.bytes, expected); +} + #[test] fn async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlements() { let (sender, receiver) = tokio::sync::oneshot::channel(); From 0c172507b22f2f0745e11494b5ce611f4d10fcf5 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 22:02:53 -0400 Subject: [PATCH 50/52] fix(ffi): restore async completion JSON parsing Signed-off-by: Will Killian --- crates/ffi/src/callable.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 52d425c9d..3d53562a9 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -41,7 +41,7 @@ use nemo_relay::codec::request::AnnotatedLlmRequest as AnnotatedLLMRequest; use nemo_relay::codec::traits::LlmCodec; use nemo_relay::error::{FlowError, Result}; -use crate::convert::json_to_c_string; +use crate::convert::{c_str_to_json, json_to_c_string}; use crate::error::{NemoRelayStatus, clear_last_error, last_error_message, set_last_error}; use crate::types::{FfiEvent, FfiLLMRequest, FfiPluginContext}; From 4110c06cb5368e2a921a8c17c126b32181700bf0 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 22:08:35 -0400 Subject: [PATCH 51/52] fix(ffi): retain pending callback ownership safely Signed-off-by: Will Killian --- crates/ffi/build.rs | 4 +- crates/ffi/src/callable.rs | 40 ++++- .../tests/unit/api/coverage_sweeps_tests.rs | 8 +- .../ffi/tests/unit/callable_private_tests.rs | 138 +++++++++++++++++- crates/ffi/tests/unit/callable_tests.rs | 16 +- 5 files changed, 186 insertions(+), 20 deletions(-) diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index 9ee7ea628..5414aa443 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -136,7 +136,7 @@ fn validate_async_callback_abi(callable: &str) { user_data: *mut libc::c_void, invocation_json: *const c_char, completion: *const NemoRelayAsyncCompletion, - ) -> NemoRelayAsyncCallbackState;"# + ) -> u32;"# ), "NemoRelayAsyncJsonCb drifted from ASYNC_REGISTRATIONS" ); @@ -148,7 +148,7 @@ fn validate_async_callback_abi(callable: &str) { invocation_json: *const c_char, next: *const NemoRelayAsyncNext, completion: *const NemoRelayAsyncCompletion, - ) -> NemoRelayAsyncCallbackState;"# + ) -> u32;"# ), "NemoRelayAsyncInterceptCb drifted from ASYNC_REGISTRATIONS" ); diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 3d53562a9..c8d63ff9e 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -67,10 +67,23 @@ pub enum NemoRelayAsyncCallbackState { Pending = 1, } +impl TryFrom for NemoRelayAsyncCallbackState { + type Error = u32; + + fn try_from(value: u32) -> std::result::Result { + match value { + value if value == Self::Complete as u32 => Ok(Self::Complete), + value if value == Self::Pending as u32 => Ok(Self::Pending), + value => Err(value), + } + } +} + /// One-shot completion passed to asynchronous C callbacks. pub struct NemoRelayAsyncCompletion { sender: std::sync::Mutex>>>, cancelled: AtomicBool, + _callback_user_data: Option>, } /// Generic completion-based middleware callback. @@ -84,12 +97,13 @@ pub type NemoRelayAsyncJsonCb = unsafe extern "C" fn( user_data: *mut libc::c_void, invocation_json: *const c_char, completion: *const NemoRelayAsyncCompletion, -) -> NemoRelayAsyncCallbackState; +) -> u32; /// Runtime-owned asynchronous `next` continuation for execution intercepts. pub struct NemoRelayAsyncNext { inner: AsyncNextInner, runtime: tokio::runtime::Handle, + _callback_user_data: Option>, } enum AsyncNextInner { @@ -151,7 +165,7 @@ pub type NemoRelayAsyncInterceptCb = unsafe extern "C" fn( invocation_json: *const c_char, next: *const NemoRelayAsyncNext, completion: *const NemoRelayAsyncCompletion, -) -> NemoRelayAsyncCallbackState; +) -> u32; /// Result callback used by channel/future-style async `next` wrappers. /// @@ -218,11 +232,21 @@ async fn invoke_async_json( let completion = Arc::new(NemoRelayAsyncCompletion { sender: std::sync::Mutex::new(Some(sender)), cancelled: AtomicBool::new(false), + _callback_user_data: Some(user_data.clone()), }); let callback_ref = Arc::into_raw(completion.clone()); let invocation = json_to_c_string(&invocation); let state = unsafe { cb(user_data.ptr, invocation, callback_ref) }; unsafe { nemo_relay_string_free_internal(invocation) }; + let state = match NemoRelayAsyncCallbackState::try_from(state) { + Ok(state) => state, + Err(state) => { + unsafe { drop(Arc::from_raw(callback_ref)) }; + return Err(FlowError::Internal(format!( + "async C callback returned invalid state {state}" + ))); + } + }; if state == NemoRelayAsyncCallbackState::Complete { unsafe { drop(Arc::from_raw(callback_ref)) }; if completion @@ -260,16 +284,28 @@ async fn invoke_async_intercept( let completion = Arc::new(NemoRelayAsyncCompletion { sender: std::sync::Mutex::new(Some(sender)), cancelled: AtomicBool::new(false), + _callback_user_data: Some(user_data.clone()), }); let callback_ref = Arc::into_raw(completion.clone()); let next = Arc::new(NemoRelayAsyncNext { inner: next, runtime, + _callback_user_data: Some(user_data.clone()), }); let next_ref = Arc::into_raw(next); let invocation = json_to_c_string(&invocation); let state = unsafe { cb(user_data.ptr, invocation, next_ref, callback_ref) }; unsafe { nemo_relay_string_free_internal(invocation) }; + let state = match NemoRelayAsyncCallbackState::try_from(state) { + Ok(state) => state, + Err(state) => { + unsafe { drop(Arc::from_raw(callback_ref)) }; + unsafe { drop(Arc::from_raw(next_ref)) }; + return Err(FlowError::Internal(format!( + "async C intercept returned invalid state {state}" + ))); + } + }; if state == NemoRelayAsyncCallbackState::Complete { unsafe { drop(Arc::from_raw(callback_ref)) }; unsafe { drop(Arc::from_raw(next_ref)) }; diff --git a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs index 73343b59c..e2c97a071 100644 --- a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs +++ b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs @@ -17,8 +17,8 @@ unsafe extern "C" fn async_json_registration_callback( _user_data: *mut libc::c_void, _invocation_json: *const c_char, _completion: *const callable::NemoRelayAsyncCompletion, -) -> callable::NemoRelayAsyncCallbackState { - callable::NemoRelayAsyncCallbackState::Pending +) -> u32 { + callable::NemoRelayAsyncCallbackState::Pending as u32 } unsafe extern "C" fn async_intercept_registration_callback( @@ -26,8 +26,8 @@ unsafe extern "C" fn async_intercept_registration_callback( _invocation_json: *const c_char, _next: *const callable::NemoRelayAsyncNext, _completion: *const callable::NemoRelayAsyncCompletion, -) -> callable::NemoRelayAsyncCallbackState { - callable::NemoRelayAsyncCallbackState::Pending +) -> u32 { + callable::NemoRelayAsyncCallbackState::Pending as u32 } #[test] diff --git a/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index 8bbdeee47..c4c3c4e44 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -9,18 +9,60 @@ unsafe extern "C" fn complete_without_settling( _user_data: *mut libc::c_void, _invocation_json: *const c_char, _completion: *const NemoRelayAsyncCompletion, -) -> NemoRelayAsyncCallbackState { - NemoRelayAsyncCallbackState::Complete +) -> u32 { + NemoRelayAsyncCallbackState::Complete as u32 } unsafe extern "C" fn retain_pending_completion( user_data: *mut libc::c_void, _invocation_json: *const c_char, completion: *const NemoRelayAsyncCompletion, -) -> NemoRelayAsyncCallbackState { +) -> u32 { let slot = unsafe { &*user_data.cast::() }; slot.store(completion as usize, Ordering::Release); - NemoRelayAsyncCallbackState::Pending + NemoRelayAsyncCallbackState::Pending as u32 +} + +struct RetainedAsyncHandles { + completion: AtomicUsize, + next: AtomicUsize, + freed: Arc, +} + +unsafe extern "C" fn retain_pending_handles( + user_data: *mut libc::c_void, + _invocation_json: *const c_char, + next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, +) -> u32 { + let state = unsafe { &*user_data.cast::() }; + state + .completion + .store(completion as usize, Ordering::Release); + state.next.store(next as usize, Ordering::Release); + NemoRelayAsyncCallbackState::Pending as u32 +} + +unsafe extern "C" fn free_retained_async_handles(user_data: *mut libc::c_void) { + let state = unsafe { Box::from_raw(user_data.cast::()) }; + state.freed.store(true, Ordering::Release); +} + +unsafe extern "C" fn invalid_async_json_state( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _completion: *const NemoRelayAsyncCompletion, +) -> u32 { + 99 +} + +unsafe extern "C" fn invalid_async_intercept_state( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _next: *const NemoRelayAsyncNext, + _completion: *const NemoRelayAsyncCompletion, +) -> u32 { + 99 } unsafe extern "C" fn send_next_result( @@ -79,6 +121,7 @@ fn async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlements() { let completion = Arc::new(NemoRelayAsyncCompletion { sender: std::sync::Mutex::new(Some(sender)), cancelled: AtomicBool::new(false), + _callback_user_data: None, }); let completion_ref = Arc::into_raw(Arc::clone(&completion)); let invalid_json = CString::new("not-json").unwrap(); @@ -109,6 +152,7 @@ fn async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlements() { let completion = Arc::new(NemoRelayAsyncCompletion { sender: std::sync::Mutex::new(Some(sender)), cancelled: AtomicBool::new(false), + _callback_user_data: None, }); let completion_ref = Arc::into_raw(Arc::clone(&completion)); assert_eq!( @@ -189,6 +233,86 @@ fn async_callback_wrappers_reject_complete_callbacks_without_settlement() { ); } +#[test] +fn pending_async_handles_retain_callback_user_data_until_release() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let freed = Arc::new(AtomicBool::new(false)); + let state = Box::new(RetainedAsyncHandles { + completion: AtomicUsize::new(0), + next: AtomicUsize::new(0), + freed: Arc::clone(&freed), + }); + let state = Box::into_raw(state); + let user_data = Arc::new(UserData { + ptr: state.cast(), + free_fn: Some(free_retained_async_handles), + }); + + runtime.block_on(async { + let mut invocation = Box::pin(invoke_async_intercept( + retain_pending_handles, + user_data, + serde_json::json!({}), + AsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + )); + tokio::select! { + biased; + result = &mut invocation => panic!("pending callback unexpectedly settled: {result:?}"), + _ = tokio::task::yield_now() => {} + } + drop(invocation); + }); + + assert!(!freed.load(Ordering::Acquire)); + let completion = + unsafe { &*state }.completion.load(Ordering::Acquire) as *const NemoRelayAsyncCompletion; + let next = unsafe { &*state }.next.load(Ordering::Acquire) as *const NemoRelayAsyncNext; + assert!(!completion.is_null()); + assert!(!next.is_null()); + assert!(unsafe { nemo_relay_async_completion_is_cancelled(completion) }); + + unsafe { nemo_relay_async_completion_release(completion) }; + assert!(!freed.load(Ordering::Acquire)); + unsafe { nemo_relay_async_next_release(next) }; + assert!(freed.load(Ordering::Acquire)); +} + +#[test] +fn async_callbacks_reject_invalid_foreign_states() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let user_data = || { + Arc::new(UserData { + ptr: std::ptr::null_mut(), + free_fn: None, + }) + }; + + let error = runtime + .block_on(invoke_async_json( + invalid_async_json_state, + user_data(), + serde_json::json!({}), + )) + .unwrap_err(); + assert!(error.to_string().contains("invalid state 99")); + + let error = runtime + .block_on(invoke_async_intercept( + invalid_async_intercept_state, + user_data(), + serde_json::json!({}), + AsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + )) + .unwrap_err(); + assert!(error.to_string().contains("invalid state 99")); +} + #[test] fn async_next_invocation_supports_tool_llm_and_stream_continuations() { let runtime = tokio::runtime::Builder::new_current_thread() @@ -241,12 +365,14 @@ fn async_next_invocation_supports_tool_llm_and_stream_continuations() { let next = Arc::new(NemoRelayAsyncNext { inner, runtime: runtime.handle().clone(), + _callback_user_data: None, }); let next_ref = Arc::into_raw(next); let (sender, receiver) = tokio::sync::oneshot::channel(); let completion = Arc::new(NemoRelayAsyncCompletion { sender: std::sync::Mutex::new(Some(sender)), cancelled: AtomicBool::new(false), + _callback_user_data: None, }); let completion_ref = Arc::into_raw(Arc::clone(&completion)); assert_eq!( @@ -265,12 +391,14 @@ fn async_next_invocation_supports_tool_llm_and_stream_continuations() { Box::pin(async move { Ok(request.content) }) })), runtime: runtime.handle().clone(), + _callback_user_data: None, }); let next_ref = Arc::into_raw(next); let (sender, _receiver) = tokio::sync::oneshot::channel(); let completion = Arc::new(NemoRelayAsyncCompletion { sender: std::sync::Mutex::new(Some(sender)), cancelled: AtomicBool::new(false), + _callback_user_data: None, }); let completion_ref = Arc::into_raw(Arc::clone(&completion)); let malformed_request = CString::new(r#"{"content":{}}"#).unwrap(); @@ -326,6 +454,7 @@ fn async_next_callback_reports_tool_llm_and_stream_results() { let next = Arc::new(NemoRelayAsyncNext { inner, runtime: runtime.handle().clone(), + _callback_user_data: None, }); let next_ref = Arc::into_raw(next); let (sender, receiver) = @@ -350,6 +479,7 @@ fn async_next_callback_reports_tool_llm_and_stream_results() { Box::pin(async { Err(FlowError::Internal("next failed".into())) }) })), runtime: runtime.handle().clone(), + _callback_user_data: None, }); let next_ref = Arc::into_raw(next); let invocation = CString::new("{}").unwrap(); diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 09f8e7087..e2ee56e06 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -22,7 +22,7 @@ unsafe extern "C" fn async_json_passthrough_callback( user_data: *mut libc::c_void, invocation_json: *const c_char, completion: *const NemoRelayAsyncCompletion, -) -> NemoRelayAsyncCallbackState { +) -> u32 { let kind = unsafe { *(user_data.cast::()) }; let invocation: Json = serde_json::from_str( unsafe { CStr::from_ptr(invocation_json) } @@ -58,7 +58,7 @@ unsafe extern "C" fn async_json_passthrough_callback( unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, NemoRelayStatus::Ok ); - NemoRelayAsyncCallbackState::Complete + NemoRelayAsyncCallbackState::Complete as u32 } unsafe extern "C" fn async_invalid_stream_callback( @@ -66,13 +66,13 @@ unsafe extern "C" fn async_invalid_stream_callback( _invocation_json: *const c_char, _next: *const NemoRelayAsyncNext, completion: *const NemoRelayAsyncCompletion, -) -> NemoRelayAsyncCallbackState { +) -> u32 { let value = CString::new(json!({"not": "an array"}).to_string()).expect("JSON has no NUL"); assert_eq!( unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, NemoRelayStatus::Ok ); - NemoRelayAsyncCallbackState::Complete + NemoRelayAsyncCallbackState::Complete as u32 } fn async_callback_user_data(kind: usize) -> *mut libc::c_void { @@ -88,7 +88,7 @@ unsafe extern "C" fn async_next_callback( invocation_json: *const c_char, next: *const NemoRelayAsyncNext, completion: *const NemoRelayAsyncCompletion, -) -> NemoRelayAsyncCallbackState { +) -> u32 { let invocation: Json = serde_json::from_str( unsafe { CStr::from_ptr(invocation_json) } .to_str() @@ -111,7 +111,7 @@ unsafe extern "C" fn async_next_callback( } // A successful invoke retained a completion reference; this callback drops // its own next and completion references before transferring Pending ownership. - NemoRelayAsyncCallbackState::Pending + NemoRelayAsyncCallbackState::Pending as u32 } unsafe extern "C" fn async_tool_outcome_callback( @@ -119,7 +119,7 @@ unsafe extern "C" fn async_tool_outcome_callback( invocation_json: *const c_char, _next: *const NemoRelayAsyncNext, completion: *const NemoRelayAsyncCompletion, -) -> NemoRelayAsyncCallbackState { +) -> u32 { let invocation: Json = serde_json::from_str( unsafe { CStr::from_ptr(invocation_json) } .to_str() @@ -133,7 +133,7 @@ unsafe extern "C" fn async_tool_outcome_callback( unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, NemoRelayStatus::Ok ); - NemoRelayAsyncCallbackState::Complete + NemoRelayAsyncCallbackState::Complete as u32 } #[test] From c4ce06328899d9f905cea77090ba6a601297591a Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 28 Jul 2026 23:40:01 -0400 Subject: [PATCH 52/52] fix(ffi): stream async next incrementally Signed-off-by: Will Killian --- crates/ffi/build.rs | 59 ++- crates/ffi/nemo_relay.h | 129 ++++- crates/ffi/src/api/llm_registry.rs | 34 +- crates/ffi/src/api/mod.rs | 4 +- crates/ffi/src/api/scope_registry.rs | 15 +- crates/ffi/src/callable.rs | 494 ++++++++++++++---- .../tests/unit/api/coverage_sweeps_tests.rs | 53 +- .../ffi/tests/unit/callable_private_tests.rs | 325 ++++++++++-- crates/ffi/tests/unit/callable_tests.rs | 39 +- go/nemo_relay/README.md | 30 ++ go/nemo_relay/adaptive_runtime_test.go | 65 +-- go/nemo_relay/async_middleware_test.go | 180 ++++++- go/nemo_relay/callbacks.go | 271 +++++++++- go/nemo_relay/nemo_relay.go | 24 +- go/nemo_relay/optimization_test.go | 20 - 15 files changed, 1444 insertions(+), 298 deletions(-) diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index 5414aa443..49a69b616 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -152,12 +152,25 @@ fn validate_async_callback_abi(callable: &str) { ), "NemoRelayAsyncInterceptCb drifted from ASYNC_REGISTRATIONS" ); + assert_eq!( + rust_type_alias(callable, "NemoRelayAsyncStreamInterceptCb"), + normalize_whitespace( + r#"pub type NemoRelayAsyncStreamInterceptCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + next: *const NemoRelayAsyncNext, + stream: *const NemoRelayAsyncStream, + ) -> u32;"# + ), + "NemoRelayAsyncStreamInterceptCb drifted from ASYNC_REGISTRATIONS" + ); for declaration in [ "typedef uint32_t NemoRelayAsyncCallbackState;", "NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0,", "NEMO_RELAY_ASYNC_CALLBACK_STATE_PENDING = 1,", "typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion);", "typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncCompletion *completion);", + "typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncStreamInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncStream *stream);", ] { assert!( ASYNC_REGISTRATIONS.contains(declaration), @@ -230,7 +243,46 @@ fn parse_async_macro_invocations(source: &str) -> Vec { } const ASYNC_REGISTRATIONS: &str = r#" -/* Completion-based async middleware registrations generated from Rust macros. */ +/* + * Completion-based async middleware registrations generated from Rust macros. + * + * Callbacks can run on Relay runtime or publication threads. invocation_json + * and result-callback strings are borrowed only for the callback invocation; + * user_data must remain valid and thread-safe until free_fn runs. + * + * A callback returning COMPLETE must settle its completion, or finish/reject + * its stream, before returning. The runtime then releases the callback-owned + * handles. A callback returning PENDING owns its completion/stream and next + * references until it settles and releases each handle exactly once. While a + * handle reference remains valid, duplicate settlement returns + * NEMO_RELAY_STATUS_INVALID_ARG. After release, callers must not access the + * handle; doing so is undefined behavior. Relay introduces no + * implicit timeout; pending work must settle or observe cancellation through + * nemo_relay_async_completion_is_cancelled or + * nemo_relay_async_stream_is_cancelled. Each successful streaming next + * invocation returns a caller-owned invocation handle; cancel it to stop an + * idle continuation and release it exactly once after completion or + * cancellation. Cancellation waits for any active result callback, making its + * user_data unreachable before returning. Result callbacks must return false + * instead of cancelling their own invocation. + * + * invocation_json/result contracts: + * - event sanitizers: {"event":Event,"fields":EventSanitizeFields} + * -> EventSanitizeFields + * - tool sanitizers, conditional guardrails, and request intercepts: + * {"name":string,"value":JSON} -> JSON, string|null, or JSON respectively + * - tool execution intercepts: {"name":string,"value":JSON} + * -> ToolExecutionInterceptOutcome + * - LLM request/response sanitizers: + * {"request":LlmRequest,"context":LlmCodecIdentity} -> LlmRequest|null, or + * {"response":JSON,"context":LlmCodecIdentity} -> JSON|null + * - LLM conditional guardrails: {"request":LlmRequest} -> string|null + * - LLM request intercepts: + * {"name":string,"request":LlmRequest,"annotated":AnnotatedLlmRequest|null} + * -> LlmRequestInterceptOutcome + * - LLM execution and stream execution intercepts: + * {"name":string,"request":LlmRequest} -> JSON or incremental stream chunks + */ typedef uint32_t NemoRelayAsyncCallbackState; enum { NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0, @@ -238,6 +290,7 @@ enum { }; typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion); typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncCompletion *completion); +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncStreamInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncStream *stream); NemoRelayStatus nemo_relay_register_mark_sanitize_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_scope_sanitize_start_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_scope_sanitize_end_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); @@ -254,7 +307,7 @@ NemoRelayStatus nemo_relay_register_llm_sanitize_response_guardrail_async(const NemoRelayStatus nemo_relay_register_llm_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_llm_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_llm_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); -NemoRelayStatus nemo_relay_register_llm_stream_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_stream_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncStreamInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_tool_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_tool_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_tool_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); @@ -265,5 +318,5 @@ NemoRelayStatus nemo_relay_scope_register_llm_conditional_execution_guardrail_as NemoRelayStatus nemo_relay_scope_register_llm_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_tool_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_llm_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); -NemoRelayStatus nemo_relay_scope_register_llm_stream_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_stream_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncStreamInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); "#; diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 843408d98..4d674832e 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -245,6 +245,16 @@ typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; */ typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; +/** + * Callback-owned incremental output stream for async stream intercepts. + */ +typedef struct NemoRelayAsyncStream NemoRelayAsyncStream; + +/** + * Caller-owned handle for one asynchronous streaming `next` invocation. + */ +typedef struct NemoRelayAsyncStreamInvocation NemoRelayAsyncStreamInvocation; + typedef struct Option_NemoRelayCollectorCb Option_NemoRelayCollectorCb; typedef struct Option_NemoRelayFinalizerCb Option_NemoRelayFinalizerCb; @@ -467,6 +477,17 @@ typedef void (*NemoRelayAsyncNextResultCb)(void *user_data, const char *value_json, const char *error_message); +/** + * Incremental result callback used by streaming async `next` wrappers. + * + * `chunk_json` is non-null for a chunk. The final invocation sets `done` and + * may carry `error_message`. Return false to cancel the downstream stream. + */ +typedef bool (*NemoRelayAsyncNextStreamResultCb)(void *user_data, + const char *chunk_json, + const char *error_message, + bool done); + /** * Run the registered tool request intercept chain on the given arguments. * @@ -1291,9 +1312,9 @@ NemoRelayStatus nemo_relay_deregister_subscriber(const char *name); /** * Wait for subscriber callbacks queued before this call to finish. * - * A call made while an asynchronous C event sanitizer is pending returns `Internal` with a - * would-block error instead of waiting. Completion callbacks may resume on arbitrary threads, so - * callers that are not inside the sanitizer can retry after it settles. + * A call made while an asynchronous publication boundary is active may return + * before that boundary and later queued callbacks finish. Call this function + * again after the middleware settles to wait for the remaining work. */ NemoRelayStatus nemo_relay_flush_subscribers(void); @@ -2594,6 +2615,62 @@ NemoRelayStatus nemo_relay_async_next_invoke_callback(const struct NemoRelayAsyn NemoRelayAsyncNextResultCb callback, void *user_data); +/** + * Invoke a streaming continuation and report chunks incrementally. + * + * The callback runs on a Relay Tokio worker thread and receives one final + * invocation with `done=true`. Returning false from a chunk callback cancels + * and closes the downstream stream. On success, `out_invocation` receives one + * caller-owned reference. Cancel it to stop an idle continuation, and release + * it exactly once after the final callback or cancellation. Cancellation does + * not return while a result callback is active, so callback `user_data` is no + * longer reachable when it returns. Do not call cancellation from inside the + * result callback; return `false` instead. + */ +NemoRelayStatus nemo_relay_async_next_invoke_stream_callback(const struct NemoRelayAsyncNext *next, + const char *invocation_json, + NemoRelayAsyncNextStreamResultCb callback, + void *user_data, + const struct NemoRelayAsyncStreamInvocation **out_invocation); + +/** + * Cancel one asynchronous streaming `next` invocation and wait for any active + * result callback to return. + */ +NemoRelayStatus nemo_relay_async_stream_invocation_cancel(const struct NemoRelayAsyncStreamInvocation *invocation); + +/** + * Release one caller-owned asynchronous streaming invocation reference. + */ +void nemo_relay_async_stream_invocation_release(const struct NemoRelayAsyncStreamInvocation *invocation); + +/** + * Push one JSON chunk to an asynchronous stream-intercept output. + */ +NemoRelayStatus nemo_relay_async_stream_push_json(const struct NemoRelayAsyncStream *stream, + const char *chunk_json); + +/** + * Finish an asynchronous stream-intercept output. + */ +NemoRelayStatus nemo_relay_async_stream_finish(const struct NemoRelayAsyncStream *stream); + +/** + * Reject an asynchronous stream-intercept output. + */ +NemoRelayStatus nemo_relay_async_stream_reject(const struct NemoRelayAsyncStream *stream, + const char *message); + +/** + * Return whether the consumer cancelled an asynchronous stream output. + */ +bool nemo_relay_async_stream_is_cancelled(const struct NemoRelayAsyncStream *stream); + +/** + * Release a callback-owned asynchronous stream reference. + */ +void nemo_relay_async_stream_release(const struct NemoRelayAsyncStream *stream); + /** * Resolve an async C callback with owned JSON. */ @@ -3110,7 +3187,46 @@ char *nemo_relay_event_annotated_request(const struct FfiEvent *ptr); char *nemo_relay_event_annotated_response(const struct FfiEvent *ptr); -/* Completion-based async middleware registrations generated from Rust macros. */ +/* + * Completion-based async middleware registrations generated from Rust macros. + * + * Callbacks can run on Relay runtime or publication threads. invocation_json + * and result-callback strings are borrowed only for the callback invocation; + * user_data must remain valid and thread-safe until free_fn runs. + * + * A callback returning COMPLETE must settle its completion, or finish/reject + * its stream, before returning. The runtime then releases the callback-owned + * handles. A callback returning PENDING owns its completion/stream and next + * references until it settles and releases each handle exactly once. While a + * handle reference remains valid, duplicate settlement returns + * NEMO_RELAY_STATUS_INVALID_ARG. After release, callers must not access the + * handle; doing so is undefined behavior. Relay introduces no + * implicit timeout; pending work must settle or observe cancellation through + * nemo_relay_async_completion_is_cancelled or + * nemo_relay_async_stream_is_cancelled. Each successful streaming next + * invocation returns a caller-owned invocation handle; cancel it to stop an + * idle continuation and release it exactly once after completion or + * cancellation. Cancellation waits for any active result callback, making its + * user_data unreachable before returning. Result callbacks must return false + * instead of cancelling their own invocation. + * + * invocation_json/result contracts: + * - event sanitizers: {"event":Event,"fields":EventSanitizeFields} + * -> EventSanitizeFields + * - tool sanitizers, conditional guardrails, and request intercepts: + * {"name":string,"value":JSON} -> JSON, string|null, or JSON respectively + * - tool execution intercepts: {"name":string,"value":JSON} + * -> ToolExecutionInterceptOutcome + * - LLM request/response sanitizers: + * {"request":LlmRequest,"context":LlmCodecIdentity} -> LlmRequest|null, or + * {"response":JSON,"context":LlmCodecIdentity} -> JSON|null + * - LLM conditional guardrails: {"request":LlmRequest} -> string|null + * - LLM request intercepts: + * {"name":string,"request":LlmRequest,"annotated":AnnotatedLlmRequest|null} + * -> LlmRequestInterceptOutcome + * - LLM execution and stream execution intercepts: + * {"name":string,"request":LlmRequest} -> JSON or incremental stream chunks + */ typedef uint32_t NemoRelayAsyncCallbackState; enum { NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0, @@ -3118,6 +3234,7 @@ enum { }; typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion); typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncCompletion *completion); +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncStreamInterceptCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncNext *next, const struct NemoRelayAsyncStream *stream); NemoRelayStatus nemo_relay_register_mark_sanitize_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_scope_sanitize_start_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_scope_sanitize_end_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); @@ -3134,7 +3251,7 @@ NemoRelayStatus nemo_relay_register_llm_sanitize_response_guardrail_async(const NemoRelayStatus nemo_relay_register_llm_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_llm_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_register_llm_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); -NemoRelayStatus nemo_relay_register_llm_stream_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_stream_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncStreamInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_tool_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_tool_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_tool_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); @@ -3145,6 +3262,6 @@ NemoRelayStatus nemo_relay_scope_register_llm_conditional_execution_guardrail_as NemoRelayStatus nemo_relay_scope_register_llm_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_tool_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); NemoRelayStatus nemo_relay_scope_register_llm_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); -NemoRelayStatus nemo_relay_scope_register_llm_stream_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_stream_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncStreamInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); #endif /* NEMO_RELAY_H */ diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index 777db2663..2d9e8c9e2 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -2,16 +2,16 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayEventSubscriberCb, NemoRelayFreeFn, - NemoRelayLlmConditionalCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, - NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, c_char, - c_str_to_string, clear_last_error, core_registry_api, core_subscriber_api, status_from_error, - wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, - wrap_async_llm_request_intercept_fn, wrap_async_llm_sanitize_request_fn, - wrap_async_llm_sanitize_response_fn, wrap_async_llm_stream_execution_intercept_fn, - wrap_event_subscriber, wrap_llm_conditional_fn, wrap_llm_exec_intercept_fn, - wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, - wrap_llm_stream_exec_intercept_fn, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayAsyncStreamInterceptCb, + NemoRelayEventSubscriberCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, + NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, + NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, c_char, c_str_to_string, clear_last_error, + core_registry_api, core_subscriber_api, status_from_error, wrap_async_llm_conditional_fn, + wrap_async_llm_execution_intercept_fn, wrap_async_llm_request_intercept_fn, + wrap_async_llm_sanitize_request_fn, wrap_async_llm_sanitize_response_fn, + wrap_async_llm_stream_execution_intercept_fn, wrap_event_subscriber, wrap_llm_conditional_fn, + wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, + wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_intercept_fn, }; global_async_registration!( @@ -29,7 +29,7 @@ global_async_registration!( ); global_async_registration!( nemo_relay_register_llm_stream_execution_intercept_async, - NemoRelayAsyncInterceptCb, + NemoRelayAsyncStreamInterceptCb, core_registry_api::register_llm_stream_execution_intercept, wrap_async_llm_stream_execution_intercept_fn ); @@ -434,18 +434,12 @@ pub unsafe extern "C" fn nemo_relay_deregister_subscriber(name: *const c_char) - /// Wait for subscriber callbacks queued before this call to finish. /// -/// A call made while an asynchronous C event sanitizer is pending returns `Internal` with a -/// would-block error instead of waiting. Completion callbacks may resume on arbitrary threads, so -/// callers that are not inside the sanitizer can retry after it settles. +/// A call made while an asynchronous publication boundary is active may return +/// before that boundary and later queued callbacks finish. Call this function +/// again after the middleware settles to wait for the remaining work. #[unsafe(no_mangle)] pub extern "C" fn nemo_relay_flush_subscribers() -> NemoRelayStatus { clear_last_error(); - if crate::callable::event_sanitizer_callback_active() { - crate::error::set_last_error( - "subscriber flush would block while an asynchronous event sanitizer is pending", - ); - return NemoRelayStatus::Internal; - } match core_subscriber_api::flush_subscribers() { Ok(()) => NemoRelayStatus::Ok, Err(e) => status_from_error(&e), diff --git a/crates/ffi/src/api/mod.rs b/crates/ffi/src/api/mod.rs index 0b7005975..6629466cb 100644 --- a/crates/ffi/src/api/mod.rs +++ b/crates/ffi/src/api/mod.rs @@ -14,8 +14,8 @@ use std::sync::{Arc, OnceLock}; use std::time::Duration; use crate::callable::{ - NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayCodecDecodeFn, - NemoRelayCodecEncodeFn, NemoRelayCollectorCb, NemoRelayEventSanitizeCb, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayAsyncStreamInterceptCb, + NemoRelayCodecDecodeFn, NemoRelayCodecEncodeFn, NemoRelayCollectorCb, NemoRelayEventSanitizeCb, NemoRelayEventSubscriberCb, NemoRelayFinalizerCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, NemoRelayLlmExecCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayPluginRegisterCb, diff --git a/crates/ffi/src/api/scope_registry.rs b/crates/ffi/src/api/scope_registry.rs index 21d9767e9..2cbec1572 100644 --- a/crates/ffi/src/api/scope_registry.rs +++ b/crates/ffi/src/api/scope_registry.rs @@ -2,12 +2,13 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayEventSubscriberCb, NemoRelayFreeFn, - NemoRelayLlmConditionalCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, - NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, - NemoRelayToolConditionalCb, NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, c_char, - c_str_to_string, clear_last_error, core_registry_api, core_subscriber_api, set_last_error, - status_from_error, wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayAsyncStreamInterceptCb, + NemoRelayEventSubscriberCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, + NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, + NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, NemoRelayToolConditionalCb, + NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, c_char, c_str_to_string, + clear_last_error, core_registry_api, core_subscriber_api, set_last_error, status_from_error, + wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, wrap_async_llm_request_intercept_fn, wrap_async_llm_sanitize_request_fn, wrap_async_llm_sanitize_response_fn, wrap_async_llm_stream_execution_intercept_fn, wrap_async_tool_conditional_fn, wrap_async_tool_execution_intercept_fn, @@ -95,7 +96,7 @@ scope_async_registration!( ); scope_async_registration!( nemo_relay_scope_register_llm_stream_execution_intercept_async, - NemoRelayAsyncInterceptCb, + NemoRelayAsyncStreamInterceptCb, core_registry_api::scope_register_llm_stream_execution_intercept, wrap_async_llm_stream_execution_intercept_fn ); diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index c8d63ff9e..1e27a7a0a 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -21,7 +21,8 @@ use std::future::Future; use std::pin::Pin; use std::ptr; use std::sync::Arc; -use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::task::{Context, Poll}; use libc::c_char; use nemo_relay::api::runtime::{ @@ -32,7 +33,7 @@ use nemo_relay::api::runtime::{ ToolInterceptFn, ToolSanitizeFn, }; use serde_json::Value as Json; -use tokio_stream::StreamExt; +use tokio_stream::{Stream, StreamExt}; use nemo_relay::api::event::{Event, EventSanitizeFields}; use nemo_relay::api::llm::{LlmRequest, LlmRequestInterceptOutcome}; @@ -112,48 +113,6 @@ enum AsyncNextInner { LlmStream(LlmStreamExecutionNextFn), } -const ASYNC_STREAM_MAX_CHUNKS: usize = 4096; -const ASYNC_STREAM_MAX_SERIALIZED_BYTES: usize = 16 * 1024 * 1024; - -#[derive(Default)] -struct SerializedByteCounter { - bytes: usize, -} - -impl std::io::Write for SerializedByteCounter { - fn write(&mut self, buffer: &[u8]) -> std::io::Result { - self.bytes = self.bytes.saturating_add(buffer.len()); - Ok(buffer.len()) - } - - fn flush(&mut self) -> std::io::Result<()> { - Ok(()) - } -} - -async fn collect_async_stream_for_completion(mut stream: LlmJsonStream) -> Result { - let mut chunks = Vec::new(); - let mut serialized_bytes = SerializedByteCounter::default(); - while let Some(chunk) = stream.next().await { - let chunk = chunk?; - if chunks.len() >= ASYNC_STREAM_MAX_CHUNKS { - return Err(FlowError::Internal(format!( - "async stream continuation exceeded the {ASYNC_STREAM_MAX_CHUNKS}-chunk completion limit" - ))); - } - serde_json::to_writer(&mut serialized_bytes, &chunk).map_err(|error| { - FlowError::Internal(format!("failed to measure async stream chunk: {error}")) - })?; - if serialized_bytes.bytes > ASYNC_STREAM_MAX_SERIALIZED_BYTES { - return Err(FlowError::Internal(format!( - "async stream continuation exceeded the {ASYNC_STREAM_MAX_SERIALIZED_BYTES}-byte completion limit" - ))); - } - chunks.push(chunk); - } - Ok(Json::Array(chunks)) -} - /// Completion-based execution-intercept callback. /// /// A callback returning `Complete` must not release either `completion` or @@ -167,6 +126,37 @@ pub type NemoRelayAsyncInterceptCb = unsafe extern "C" fn( completion: *const NemoRelayAsyncCompletion, ) -> u32; +/// Callback-owned incremental output stream for async stream intercepts. +pub struct NemoRelayAsyncStream { + sender: std::sync::Mutex>>>, + cancelled: AtomicBool, + _callback_user_data: Option>, +} + +/// Caller-owned handle for one asynchronous streaming `next` invocation. +pub struct NemoRelayAsyncStreamInvocation { + state: Arc, + abort_handle: tokio::task::AbortHandle, +} + +struct AsyncStreamInvocationState { + cancelled: AtomicBool, + callback_gate: std::sync::Mutex<()>, +} + +/// Completion-based streaming execution-intercept callback. +/// +/// The callback emits replacement chunks with +/// [`nemo_relay_async_stream_push_json`] and completes with +/// [`nemo_relay_async_stream_finish`] or [`nemo_relay_async_stream_reject`]. +/// A pending callback must release `stream` and `next` exactly once. +pub type NemoRelayAsyncStreamInterceptCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + next: *const NemoRelayAsyncNext, + stream: *const NemoRelayAsyncStream, +) -> u32; + /// Result callback used by channel/future-style async `next` wrappers. /// /// Invoked on a Tokio runtime worker thread, not necessarily the thread that @@ -179,6 +169,17 @@ pub type NemoRelayAsyncNextResultCb = unsafe extern "C" fn( error_message: *const c_char, ); +/// Incremental result callback used by streaming async `next` wrappers. +/// +/// `chunk_json` is non-null for a chunk. The final invocation sets `done` and +/// may carry `error_message`. Return false to cancel the downstream stream. +pub type NemoRelayAsyncNextStreamResultCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + chunk_json: *const c_char, + error_message: *const c_char, + done: bool, +) -> bool; + struct SendUserData(*mut libc::c_void); // SAFETY: NemoRelayAsyncNextResultCb requires callers to keep user_data valid @@ -196,27 +197,6 @@ struct CompletionWait { receiver: tokio::sync::oneshot::Receiver>, } -static ACTIVE_EVENT_SANITIZER_CALLBACKS: AtomicUsize = AtomicUsize::new(0); - -struct ActiveEventSanitizerCallback; - -impl ActiveEventSanitizerCallback { - fn enter() -> Self { - ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_add(1, Ordering::AcqRel); - Self - } -} - -impl Drop for ActiveEventSanitizerCallback { - fn drop(&mut self) { - ACTIVE_EVENT_SANITIZER_CALLBACKS.fetch_sub(1, Ordering::AcqRel); - } -} - -pub(crate) fn event_sanitizer_callback_active() -> bool { - ACTIVE_EVENT_SANITIZER_CALLBACKS.load(Ordering::Acquire) != 0 -} - impl Drop for CompletionWait { fn drop(&mut self) { self.completion.cancelled.store(true, Ordering::Release); @@ -375,14 +355,7 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke( let next = next.clone(); Box::pin(async move { next(request).await }) } - AsyncNextInner::LlmStream(next) => { - let request = match serde_json::from_value(invocation) { - Ok(request) => request, - Err(_) => return NemoRelayStatus::InvalidJson, - }; - let next = next.clone(); - Box::pin(async move { collect_async_stream_for_completion(next(request).await?).await }) - } + AsyncNextInner::LlmStream(_) => return NemoRelayStatus::InvalidArg, }; unsafe { Arc::increment_strong_count(completion) }; let completion = unsafe { Arc::from_raw(completion) }; @@ -434,14 +407,7 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke_callback( let next = next.clone(); Box::pin(async move { next(request).await }) } - AsyncNextInner::LlmStream(next) => { - let request = match serde_json::from_value(invocation) { - Ok(request) => request, - Err(_) => return NemoRelayStatus::InvalidJson, - }; - let next = next.clone(); - Box::pin(async move { collect_async_stream_for_completion(next(request).await?).await }) - } + AsyncNextInner::LlmStream(_) => return NemoRelayStatus::InvalidArg, }; let user_data = SendUserData(user_data); next.runtime.spawn(async move { @@ -460,6 +426,277 @@ pub unsafe extern "C" fn nemo_relay_async_next_invoke_callback( NemoRelayStatus::Ok } +fn invoke_async_next_stream_callback( + invocation: &AsyncStreamInvocationState, + callback: NemoRelayAsyncNextStreamResultCb, + user_data: *mut libc::c_void, + chunk_json: *const c_char, + error_message: *const c_char, + done: bool, +) -> bool { + let _callback_guard = invocation + .callback_gate + .lock() + .unwrap_or_else(|error| error.into_inner()); + if invocation.cancelled.load(Ordering::Acquire) { + return false; + } + unsafe { callback(user_data, chunk_json, error_message, done) } +} + +/// Invoke a streaming continuation and report chunks incrementally. +/// +/// The callback runs on a Relay Tokio worker thread and receives one final +/// invocation with `done=true`. Returning false from a chunk callback cancels +/// and closes the downstream stream. On success, `out_invocation` receives one +/// caller-owned reference. Cancel it to stop an idle continuation, and release +/// it exactly once after the final callback or cancellation. Cancellation does +/// not return while a result callback is active, so callback `user_data` is no +/// longer reachable when it returns. Do not call cancellation from inside the +/// result callback; return `false` instead. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_next_invoke_stream_callback( + next: *const NemoRelayAsyncNext, + invocation_json: *const c_char, + callback: NemoRelayAsyncNextStreamResultCb, + user_data: *mut libc::c_void, + out_invocation: *mut *const NemoRelayAsyncStreamInvocation, +) -> NemoRelayStatus { + if out_invocation.is_null() { + return NemoRelayStatus::NullPointer; + } + unsafe { *out_invocation = ptr::null() }; + let Some(next) = (unsafe { next.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + let Some(invocation) = c_str_to_json(invocation_json) else { + return NemoRelayStatus::InvalidJson; + }; + let AsyncNextInner::LlmStream(next_fn) = &next.inner else { + return NemoRelayStatus::InvalidArg; + }; + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(_) => return NemoRelayStatus::InvalidJson, + }; + let next_fn = next_fn.clone(); + let user_data = SendUserData(user_data); + let invocation_state = Arc::new(AsyncStreamInvocationState { + cancelled: AtomicBool::new(false), + callback_gate: std::sync::Mutex::new(()), + }); + let task_invocation_state = invocation_state.clone(); + let task = next.runtime.spawn(async move { + match next_fn(request).await { + Ok(mut stream) => { + while let Some(result) = stream.next().await { + match result { + Ok(chunk) => { + let chunk = json_to_c_string(&chunk); + let keep_going = invoke_async_next_stream_callback( + &task_invocation_state, + callback, + user_data.as_ptr(), + chunk, + ptr::null(), + false, + ); + unsafe { nemo_relay_string_free_internal(chunk) }; + if !keep_going { + let _ = stream.close().await; + return; + } + } + Err(error) => { + let error = CString::new(error.to_string()).unwrap_or_default(); + invoke_async_next_stream_callback( + &task_invocation_state, + callback, + user_data.as_ptr(), + ptr::null(), + error.as_ptr(), + true, + ); + return; + } + } + } + if let Err(error) = stream.close().await { + let error = CString::new(error.to_string()).unwrap_or_default(); + invoke_async_next_stream_callback( + &task_invocation_state, + callback, + user_data.as_ptr(), + ptr::null(), + error.as_ptr(), + true, + ); + } else { + invoke_async_next_stream_callback( + &task_invocation_state, + callback, + user_data.as_ptr(), + ptr::null(), + ptr::null(), + true, + ); + } + } + Err(error) => { + let error = CString::new(error.to_string()).unwrap_or_default(); + invoke_async_next_stream_callback( + &task_invocation_state, + callback, + user_data.as_ptr(), + ptr::null(), + error.as_ptr(), + true, + ); + } + } + }); + let invocation = Arc::new(NemoRelayAsyncStreamInvocation { + state: invocation_state, + abort_handle: task.abort_handle(), + }); + unsafe { *out_invocation = Arc::into_raw(invocation) }; + NemoRelayStatus::Ok +} + +/// Cancel one asynchronous streaming `next` invocation and wait for any active +/// result callback to return. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_stream_invocation_cancel( + invocation: *const NemoRelayAsyncStreamInvocation, +) -> NemoRelayStatus { + let Some(invocation) = (unsafe { invocation.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if !invocation.state.cancelled.swap(true, Ordering::AcqRel) { + invocation.abort_handle.abort(); + } + let _callback_guard = invocation + .state + .callback_gate + .lock() + .unwrap_or_else(|error| error.into_inner()); + NemoRelayStatus::Ok +} + +/// Release one caller-owned asynchronous streaming invocation reference. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_stream_invocation_release( + invocation: *const NemoRelayAsyncStreamInvocation, +) { + if !invocation.is_null() { + unsafe { drop(Arc::from_raw(invocation)) }; + } +} + +/// Push one JSON chunk to an asynchronous stream-intercept output. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_stream_push_json( + stream: *const NemoRelayAsyncStream, + chunk_json: *const c_char, +) -> NemoRelayStatus { + let Some(stream) = (unsafe { stream.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if stream.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let Some(chunk) = c_str_to_json(chunk_json) else { + return NemoRelayStatus::InvalidJson; + }; + let sender = stream + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()); + match sender.as_ref() { + Some(sender) if sender.send(Ok(chunk)).is_ok() => NemoRelayStatus::Ok, + _ => NemoRelayStatus::InvalidArg, + } +} + +/// Finish an asynchronous stream-intercept output. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_stream_finish( + stream: *const NemoRelayAsyncStream, +) -> NemoRelayStatus { + let Some(stream) = (unsafe { stream.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if stream.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + match stream + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + { + Some(_) => NemoRelayStatus::Ok, + None => NemoRelayStatus::InvalidArg, + } +} + +/// Reject an asynchronous stream-intercept output. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_stream_reject( + stream: *const NemoRelayAsyncStream, + message: *const c_char, +) -> NemoRelayStatus { + let Some(stream) = (unsafe { stream.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if stream.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let message = if message.is_null() { + "async C stream callback rejected".to_string() + } else { + unsafe { CStr::from_ptr(message) } + .to_string_lossy() + .into_owned() + }; + let sender = stream + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take(); + match sender { + Some(sender) => { + let _ = sender.send(Err(FlowError::Internal(message))); + NemoRelayStatus::Ok + } + None => NemoRelayStatus::InvalidArg, + } +} + +/// Return whether the consumer cancelled an asynchronous stream output. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_stream_is_cancelled( + stream: *const NemoRelayAsyncStream, +) -> bool { + unsafe { stream.as_ref() }.is_none_or(|stream| stream.cancelled.load(Ordering::Acquire)) +} + +/// Release a callback-owned asynchronous stream reference. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_stream_release(stream: *const NemoRelayAsyncStream) { + if !stream.is_null() { + unsafe { drop(Arc::from_raw(stream)) }; + } +} + /// Resolve an async C callback with owned JSON. #[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. #[unsafe(no_mangle)] @@ -862,7 +1099,6 @@ pub fn wrap_async_event_sanitize_fn( Arc::new(move |event: Arc, fields: EventSanitizeFields| { let user_data = user_data.clone(); Box::pin(async move { - let _active_callback = ActiveEventSanitizerCallback::enter(); let value = invoke_async_json( cb, user_data, @@ -1026,15 +1262,33 @@ pub fn wrap_async_llm_execution_intercept_fn( ) } -/// Wrap a completion-based C LLM stream execution intercept. -/// -/// The completion ABI resolves one JSON value, so a stream intercept must -/// resolve to an array of chunks. Relay replays that array as a stream after -/// completion; incremental chunk delivery is not available through this ABI. -/// When the callback invokes `next`, Relay rejects more than 4096 chunks or -/// 16 MiB of serialized chunk data while collecting that continuation. +struct AsyncCallbackOutputStream { + receiver: tokio::sync::mpsc::UnboundedReceiver>, + state: Arc, +} + +impl Stream for AsyncCallbackOutputStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.receiver).poll_recv(cx) + } +} + +impl Drop for AsyncCallbackOutputStream { + fn drop(&mut self) { + self.state.cancelled.store(true, Ordering::Release); + self.state + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take(); + } +} + +/// Wrap an incremental completion-based C LLM stream execution intercept. pub fn wrap_async_llm_stream_execution_intercept_fn( - cb: NemoRelayAsyncInterceptCb, + cb: NemoRelayAsyncStreamInterceptCb, user_data: *mut libc::c_void, free_fn: NemoRelayFreeFn, ) -> LlmStreamExecutionFn { @@ -1044,19 +1298,55 @@ pub fn wrap_async_llm_stream_execution_intercept_fn( let user_data = user_data.clone(); let invocation = serde_json::json!({"name": name, "request": request}); Box::pin(async move { - let value = invoke_async_intercept( - cb, - user_data, - invocation, - AsyncNextInner::LlmStream(next), - ) - .await?; - let chunks = value.as_array().cloned().ok_or_else(|| { - FlowError::Internal("async stream intercept must resolve to an array".into()) + let runtime = tokio::runtime::Handle::try_current().map_err(|error| { + FlowError::Internal(format!( + "async C stream intercept requires a Tokio runtime: {error}" + )) })?; - Ok(LlmJsonStream::new(tokio_stream::iter( - chunks.into_iter().map(Ok), - ))) + let (sender, receiver) = tokio::sync::mpsc::unbounded_channel(); + let stream = Arc::new(NemoRelayAsyncStream { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + _callback_user_data: Some(user_data.clone()), + }); + let stream_ref = Arc::into_raw(stream.clone()); + let next = Arc::new(NemoRelayAsyncNext { + inner: AsyncNextInner::LlmStream(next), + runtime, + _callback_user_data: Some(user_data.clone()), + }); + let next_ref = Arc::into_raw(next); + let invocation = json_to_c_string(&invocation); + let state = unsafe { cb(user_data.ptr, invocation, next_ref, stream_ref) }; + unsafe { nemo_relay_string_free_internal(invocation) }; + let state = match NemoRelayAsyncCallbackState::try_from(state) { + Ok(state) => state, + Err(state) => { + unsafe { drop(Arc::from_raw(stream_ref)) }; + unsafe { drop(Arc::from_raw(next_ref)) }; + return Err(FlowError::Internal(format!( + "async C stream intercept returned invalid state {state}" + ))); + } + }; + if state == NemoRelayAsyncCallbackState::Complete { + unsafe { drop(Arc::from_raw(stream_ref)) }; + unsafe { drop(Arc::from_raw(next_ref)) }; + if stream + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .is_some() + { + return Err(FlowError::Internal( + "async C stream intercept returned Complete without finishing".into(), + )); + } + } + Ok(LlmJsonStream::new(AsyncCallbackOutputStream { + receiver, + state: stream, + })) }) }, ) diff --git a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs index e2c97a071..977a92a7b 100644 --- a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs +++ b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs @@ -30,6 +30,15 @@ unsafe extern "C" fn async_intercept_registration_callback( callable::NemoRelayAsyncCallbackState::Pending as u32 } +unsafe extern "C" fn async_stream_intercept_registration_callback( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _next: *const callable::NemoRelayAsyncNext, + _stream: *const callable::NemoRelayAsyncStream, +) -> u32 { + callable::NemoRelayAsyncCallbackState::Pending as u32 +} + #[test] fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { let _lock = TEST_MUTEX.lock().unwrap_or_else(|e| e.into_inner()); @@ -90,6 +99,24 @@ fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { assert_eq!(unsafe { $deregister(name.as_ptr()) }, NemoRelayStatus::Ok); }}; } + macro_rules! global_stream_intercept { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + name.as_ptr(), + 0, + async_stream_intercept_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(unsafe { $deregister(name.as_ptr()) }, NemoRelayStatus::Ok); + }}; + } global_json!( nemo_relay_register_mark_sanitize_guardrail_async, @@ -143,7 +170,7 @@ fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { nemo_relay_register_llm_execution_intercept_async, nemo_relay_deregister_llm_execution_intercept ); - global_intercept!( + global_stream_intercept!( nemo_relay_register_llm_stream_execution_intercept_async, nemo_relay_deregister_llm_stream_execution_intercept ); @@ -226,6 +253,28 @@ fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { ); }}; } + macro_rules! scope_stream_intercept { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + scope_uuid.as_ptr(), + name.as_ptr(), + 0, + async_stream_intercept_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { $deregister(scope_uuid.as_ptr(), name.as_ptr()) }, + NemoRelayStatus::Ok + ); + }}; + } scope_json!( nemo_relay_scope_register_mark_sanitize_guardrail_async, @@ -279,7 +328,7 @@ fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { nemo_relay_scope_register_llm_execution_intercept_async, nemo_relay_scope_deregister_llm_execution_intercept ); - scope_intercept!( + scope_stream_intercept!( nemo_relay_scope_register_llm_stream_execution_intercept_async, nemo_relay_scope_deregister_llm_stream_execution_intercept ); diff --git a/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index c4c3c4e44..adddd6f4a 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -4,6 +4,7 @@ //! Unit tests for callable private in the NeMo Relay FFI crate. use super::*; +use std::sync::atomic::AtomicUsize; unsafe extern "C" fn complete_without_settling( _user_data: *mut libc::c_void, @@ -86,6 +87,42 @@ unsafe extern "C" fn send_next_result( let _ = sender.send(result); } +unsafe extern "C" fn send_next_stream_result( + user_data: *mut libc::c_void, + chunk_json: *const c_char, + error_message: *const c_char, + done: bool, +) -> bool { + let state = unsafe { + &*user_data + .cast::, String>>>() + }; + let result = if !error_message.is_null() { + Err(unsafe { CStr::from_ptr(error_message) } + .to_string_lossy() + .into_owned()) + } else if done { + Ok(None) + } else { + serde_json::from_str(unsafe { CStr::from_ptr(chunk_json) }.to_str().unwrap()) + .map(Some) + .map_err(|error| error.to_string()) + }; + let keep_going = state.send(result).is_ok(); + if done || !keep_going { + unsafe { + drop( + Box::from_raw( + user_data.cast::, String>, + >>(), + ), + ) + }; + } + keep_going +} + #[test] fn test_callable_private_helper_paths() { clear_last_error(); @@ -99,22 +136,6 @@ fn test_callable_private_helper_paths() { unsafe { nemo_relay_string_free_internal(raw) }; } -#[test] -fn serialized_byte_counter_accumulates_json_without_buffering() { - let values = [serde_json::json!({"chunk": 1}), serde_json::json!("two")]; - let expected: usize = values - .iter() - .map(|value| serde_json::to_vec(value).unwrap().len()) - .sum(); - let mut counter = SerializedByteCounter::default(); - - for value in &values { - serde_json::to_writer(&mut counter, value).unwrap(); - } - - assert_eq!(counter.bytes, expected); -} - #[test] fn async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlements() { let (sender, receiver) = tokio::sync::oneshot::channel(); @@ -314,7 +335,7 @@ fn async_callbacks_reject_invalid_foreign_states() { } #[test] -fn async_next_invocation_supports_tool_llm_and_stream_continuations() { +fn async_next_invocation_supports_tool_and_llm_continuations() { let runtime = tokio::runtime::Builder::new_current_thread() .enable_all() .build() @@ -340,25 +361,6 @@ fn async_next_invocation_supports_tool_llm_and_stream_continuations() { .unwrap(), serde_json::json!({"llm": true}), ), - ( - AsyncNextInner::LlmStream(Arc::new(|_request| { - Box::pin(async { - Ok(LlmJsonStream::new(tokio_stream::iter(vec![ - Ok(serde_json::json!({"chunk": 1})), - Ok(serde_json::json!({"chunk": 2})), - ]))) - }) - })), - CString::new( - serde_json::to_string(&LlmRequest { - headers: serde_json::Map::new(), - content: serde_json::json!({"stream": true}), - }) - .unwrap(), - ) - .unwrap(), - serde_json::json!([{"chunk": 1}, {"chunk": 2}]), - ), ]; for (inner, invocation, expected) in cases { @@ -420,7 +422,7 @@ fn async_next_invocation_supports_tool_llm_and_stream_continuations() { } #[test] -fn async_next_callback_reports_tool_llm_and_stream_results() { +fn async_next_callback_reports_tool_and_llm_results() { let runtime = tokio::runtime::Builder::new_current_thread() .enable_all() .build() @@ -438,17 +440,6 @@ fn async_next_callback_reports_tool_llm_and_stream_results() { CString::new(r#"{"headers":{},"content":{"llm":true}}"#).unwrap(), serde_json::json!({"llm": true}), ), - ( - AsyncNextInner::LlmStream(Arc::new(|_request| { - Box::pin(async { - Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( - serde_json::json!({"stream": true}), - )]))) - }) - })), - CString::new(r#"{"headers":{},"content":{}}"#).unwrap(), - serde_json::json!([{ "stream": true }]), - ), ]; for (inner, invocation, expected) in cases { let next = Arc::new(NemoRelayAsyncNext { @@ -504,3 +495,241 @@ fn async_next_callback_reports_tool_llm_and_stream_results() { ); unsafe { nemo_relay_async_next_release(next_ref) }; } + +#[test] +fn async_next_stream_callback_reports_chunks_incrementally() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let next = Arc::new(NemoRelayAsyncNext { + inner: AsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![ + Ok(serde_json::json!({"chunk": 1})), + Ok(serde_json::json!({"chunk": 2})), + ]))) + }) + })), + runtime: runtime.handle().clone(), + _callback_user_data: None, + }); + let next_ref = Arc::into_raw(next); + let invocation = CString::new(r#"{"headers":{},"content":{}}"#).unwrap(); + let (sender, mut receiver) = + tokio::sync::mpsc::unbounded_channel::, String>>(); + let mut stream_invocation = std::ptr::null(); + assert_eq!( + unsafe { + nemo_relay_async_next_invoke_stream_callback( + next_ref, + invocation.as_ptr(), + send_next_stream_result, + Box::into_raw(Box::new(sender)).cast(), + &raw mut stream_invocation, + ) + }, + NemoRelayStatus::Ok + ); + let values = runtime.block_on(async move { + let mut values = Vec::new(); + while let Some(result) = receiver.recv().await { + match result.unwrap() { + Some(value) => values.push(value), + None => break, + } + } + values + }); + assert_eq!( + values, + vec![ + serde_json::json!({"chunk": 1}), + serde_json::json!({"chunk": 2}) + ] + ); + unsafe { nemo_relay_async_stream_invocation_release(stream_invocation) }; + unsafe { nemo_relay_async_next_release(next_ref) }; +} + +#[test] +fn async_next_stream_invocation_cancellation_aborts_idle_continuation() { + struct DropSignal(Arc); + impl Drop for DropSignal { + fn drop(&mut self) { + self.0.store(true, Ordering::Release); + } + } + + unsafe extern "C" fn record_unexpected_callback( + user_data: *mut libc::c_void, + _chunk_json: *const c_char, + _error_message: *const c_char, + _done: bool, + ) -> bool { + unsafe { &*user_data.cast::() }.store(true, Ordering::Release); + false + } + + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let started_tx = Arc::new(std::sync::Mutex::new(Some(started_tx))); + let dropped = Arc::new(AtomicBool::new(false)); + let next = Arc::new(NemoRelayAsyncNext { + inner: AsyncNextInner::LlmStream(Arc::new({ + let started_tx = started_tx.clone(); + let dropped = dropped.clone(); + move |_request| { + let started_tx = started_tx.clone(); + let guard = DropSignal(dropped.clone()); + Box::pin(async move { + let _guard = guard; + if let Some(started_tx) = started_tx + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + { + let _ = started_tx.send(()); + } + std::future::pending::>().await + }) + } + })), + runtime: runtime.handle().clone(), + _callback_user_data: None, + }); + let next_ref = Arc::into_raw(next); + let invocation = CString::new(r#"{"headers":{},"content":{}}"#).unwrap(); + let callback_called = AtomicBool::new(false); + let mut stream_invocation = std::ptr::null(); + assert_eq!( + unsafe { + nemo_relay_async_next_invoke_stream_callback( + next_ref, + invocation.as_ptr(), + record_unexpected_callback, + std::ptr::from_ref(&callback_called).cast_mut().cast(), + &raw mut stream_invocation, + ) + }, + NemoRelayStatus::Ok + ); + runtime.block_on(started_rx).unwrap(); + assert_eq!( + unsafe { nemo_relay_async_stream_invocation_cancel(stream_invocation) }, + NemoRelayStatus::Ok + ); + runtime.block_on(async { + tokio::time::timeout(std::time::Duration::from_secs(1), async { + while !dropped.load(Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await + .expect("idle continuation was not aborted"); + }); + assert!(!callback_called.load(Ordering::Acquire)); + unsafe { + nemo_relay_async_stream_invocation_release(stream_invocation); + nemo_relay_async_next_release(next_ref); + } +} + +#[test] +fn async_next_stream_cancellation_waits_for_active_callback() { + struct BlockingCallback { + entered: std::sync::mpsc::Sender<()>, + release: std::sync::Mutex>, + } + + unsafe extern "C" fn block_in_callback( + user_data: *mut libc::c_void, + _chunk_json: *const c_char, + _error_message: *const c_char, + _done: bool, + ) -> bool { + let state = unsafe { &*user_data.cast::() }; + let _ = state.entered.send(()); + let _ = state + .release + .lock() + .unwrap_or_else(|error| error.into_inner()) + .recv(); + true + } + + let runtime = tokio::runtime::Builder::new_multi_thread() + .worker_threads(1) + .enable_all() + .build() + .unwrap(); + let next = Arc::new(NemoRelayAsyncNext { + inner: AsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( + serde_json::json!({"chunk": 1}), + )]))) + }) + })), + runtime: runtime.handle().clone(), + _callback_user_data: None, + }); + let next_ref = Arc::into_raw(next); + let invocation_json = CString::new(r#"{"headers":{},"content":{}}"#).unwrap(); + let (entered_tx, entered_rx) = std::sync::mpsc::channel(); + let (release_tx, release_rx) = std::sync::mpsc::channel(); + let callback_state = Box::into_raw(Box::new(BlockingCallback { + entered: entered_tx, + release: std::sync::Mutex::new(release_rx), + })); + let mut stream_invocation = std::ptr::null(); + assert_eq!( + unsafe { + nemo_relay_async_next_invoke_stream_callback( + next_ref, + invocation_json.as_ptr(), + block_in_callback, + callback_state.cast(), + &raw mut stream_invocation, + ) + }, + NemoRelayStatus::Ok + ); + entered_rx + .recv_timeout(std::time::Duration::from_secs(1)) + .expect("stream callback did not start"); + + let (cancel_done_tx, cancel_done_rx) = std::sync::mpsc::channel(); + let invocation_address = stream_invocation as usize; + let cancel_thread = std::thread::spawn(move || { + let status = unsafe { + nemo_relay_async_stream_invocation_cancel( + invocation_address as *const NemoRelayAsyncStreamInvocation, + ) + }; + let _ = cancel_done_tx.send(status); + }); + assert!( + cancel_done_rx + .recv_timeout(std::time::Duration::from_millis(50)) + .is_err(), + "cancellation returned while callback user_data was still active" + ); + release_tx.send(()).unwrap(); + assert_eq!( + cancel_done_rx + .recv_timeout(std::time::Duration::from_secs(1)) + .expect("cancellation did not finish after callback returned"), + NemoRelayStatus::Ok + ); + cancel_thread.join().unwrap(); + + unsafe { + drop(Box::from_raw(callback_state)); + nemo_relay_async_stream_invocation_release(stream_invocation); + nemo_relay_async_next_release(next_ref); + } +} diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index e2ee56e06..7ec6f756a 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -44,8 +44,8 @@ unsafe extern "C" fn async_json_passthrough_callback( 7 => { assert_eq!( crate::api::nemo_relay_flush_subscribers(), - NemoRelayStatus::Internal, - "flush inside an async event sanitizer must report would-block" + NemoRelayStatus::Ok, + "flush inside event publication must return without waiting" ); invocation["fields"].clone() } @@ -61,15 +61,30 @@ unsafe extern "C" fn async_json_passthrough_callback( NemoRelayAsyncCallbackState::Complete as u32 } -unsafe extern "C" fn async_invalid_stream_callback( +unsafe extern "C" fn async_unfinished_stream_callback( _user_data: *mut libc::c_void, _invocation_json: *const c_char, _next: *const NemoRelayAsyncNext, - completion: *const NemoRelayAsyncCompletion, + _stream: *const NemoRelayAsyncStream, ) -> u32 { - let value = CString::new(json!({"not": "an array"}).to_string()).expect("JSON has no NUL"); + NemoRelayAsyncCallbackState::Complete as u32 +} + +unsafe extern "C" fn async_immediate_stream_callback( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _next: *const NemoRelayAsyncNext, + stream: *const NemoRelayAsyncStream, +) -> u32 { + for chunk in [json!({"chunk": 1}), json!({"chunk": 2})] { + let chunk = CString::new(chunk.to_string()).expect("JSON has no NUL"); + assert_eq!( + unsafe { nemo_relay_async_stream_push_json(stream, chunk.as_ptr()) }, + NemoRelayStatus::Ok + ); + } assert_eq!( - unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, + unsafe { nemo_relay_async_stream_finish(stream) }, NemoRelayStatus::Ok ); NemoRelayAsyncCallbackState::Complete as u32 @@ -275,7 +290,7 @@ fn async_conditional_and_stream_wrappers_validate_callback_results() { ); let stream_intercept = wrap_async_llm_stream_execution_intercept_fn( - async_invalid_stream_callback, + async_unfinished_stream_callback, std::ptr::null_mut(), None, ); @@ -288,9 +303,9 @@ fn async_conditional_and_stream_wrappers_validate_callback_results() { }); let result = resolve(stream_intercept("llm", make_request(), next)); let Err(error) = result else { - panic!("a non-array async stream result must fail"); + panic!("an unfinished immediate stream callback must fail"); }; - assert!(error.to_string().contains("must resolve to an array")); + assert!(error.to_string().contains("without finishing")); } #[test] @@ -319,9 +334,9 @@ fn async_execution_wrappers_continue_tool_and_llm_calls() { } #[test] -fn async_stream_execution_wrapper_collects_the_continued_stream() { +fn async_stream_execution_wrapper_delivers_chunks_incrementally() { let intercept = wrap_async_llm_stream_execution_intercept_fn( - async_next_callback, + async_immediate_stream_callback, std::ptr::null_mut(), None, ); @@ -339,7 +354,7 @@ fn async_stream_execution_wrapper_collects_the_continued_stream() { let mut stream = resolve(intercept("llm", make_request(), next)).unwrap(); assert_eq!( resolve(async { stream.next().await.unwrap().unwrap() }), - json!({"model": "test-model", "chunk": 1}) + json!({"chunk": 1}) ); assert_eq!( resolve(async { stream.next().await.unwrap().unwrap() }), diff --git a/go/nemo_relay/README.md b/go/nemo_relay/README.md index 7df389c52..4e98fc458 100644 --- a/go/nemo_relay/README.md +++ b/go/nemo_relay/README.md @@ -61,6 +61,36 @@ The Go package provides the following capabilities: - **Local source-first workflow**: Build the FFI library locally, then test or consume the Go module from the checkout. +## Async Middleware + +Every asynchronous registration has a global and scope-local Go wrapper. +`AsyncMiddlewareFunc` receives one of these JSON envelopes and returns the +corresponding value: + +| Middleware | Invocation | Result | +| --- | --- | --- | +| Event sanitizer | `{"event": Event, "fields": EventSanitizeFields}` | `EventSanitizeFields` | +| Tool sanitizer or request intercept | `{"name": string, "value": JSON}` | JSON | +| Tool conditional guardrail | `{"name": string, "value": JSON}` | `string` or `null` | +| LLM request sanitizer | `{"request": LlmRequest, "context": LlmCodecIdentity}` | `LlmRequest` or `null` | +| LLM response sanitizer | `{"response": JSON, "context": LlmCodecIdentity}` | JSON or `null` | +| LLM conditional guardrail | `{"request": LlmRequest}` | `string` or `null` | +| LLM request intercept | `{"name": string, "request": LlmRequest, "annotated": AnnotatedLlmRequest \| null}` | `LlmRequestInterceptOutcome` | + +Execution intercepts also receive a `next` function. Tool execution uses the +`{"name": string, "value": JSON}` envelope and returns a +`ToolExecutionInterceptOutcome`; LLM execution uses +`{"name": string, "request": LlmRequest}` and returns JSON. +`AsyncStreamExecutionInterceptFunc` uses the LLM envelope and produces an +`AsyncStreamItem` channel. Its `next` helper streams downstream chunks +incrementally instead of collecting the response. + +Callbacks run in goroutines and may be entered from Relay runtime or publication +threads. Callback values must therefore be safe for concurrent use. The +callback context is cancelled when Relay abandons the native invocation, and +stream producers must stop when that context is cancelled. Relay does not add +an implicit middleware timeout. + ## Installation Build the FFI library from a repository checkout before using the Go binding: diff --git a/go/nemo_relay/adaptive_runtime_test.go b/go/nemo_relay/adaptive_runtime_test.go index 1051e053a..d63dbcda5 100644 --- a/go/nemo_relay/adaptive_runtime_test.go +++ b/go/nemo_relay/adaptive_runtime_test.go @@ -44,7 +44,7 @@ func TestValidateAdaptiveConfigAndOwnedRuntime(t *testing.T) { if err != nil { t.Fatalf(newAdaptiveRuntimeFailedMsg, err) } - defer func() { _ = runtime.Shutdown() }() + defer runtime.Shutdown() if err := runtime.Register(); err != nil { t.Fatalf("Register failed: %v", err) } @@ -159,7 +159,7 @@ func TestAdaptiveRuntimeBindScopeRejectsNilScope(t *testing.T) { if err != nil { t.Fatalf(newAdaptiveRuntimeFailedMsg, err) } - defer func() { _ = runtime.Shutdown() }() + defer runtime.Shutdown() if runtime.BindScope(nil) == nil { t.Fatal("expected BindScope to reject nil scope") @@ -172,36 +172,12 @@ func TestSetLatencySensitivityRejectsInvalidValue(t *testing.T) { } } -func TestAdaptiveRuntimeRejectsNilHandles(t *testing.T) { - var nilRuntime *AdaptiveRuntime - if err := nilRuntime.Register(); err == nil { - t.Fatal("expected nil Register to fail") - } - if err := nilRuntime.Deregister(); err == nil { - t.Fatal("expected nil Deregister to fail") - } - if err := nilRuntime.Shutdown(); err == nil { - t.Fatal("expected nil Shutdown to fail") - } - if err := nilRuntime.WaitForIdle(); err == nil { - t.Fatal("expected nil WaitForIdle to fail") - } - if _, err := nilRuntime.Report(); err == nil { - t.Fatal("expected nil Report to fail") - } - if err := nilRuntime.BindScope(nil); err == nil { - t.Fatal("expected nil BindScope to fail") - } - if _, err := nilRuntime.BuildCacheRequestFacts(CacheRequestFactsInput{}); err == nil { - t.Fatal("expected nil BuildCacheRequestFacts to fail") - } -} - func TestAdaptiveRuntimeLifecycleRejectsUseAfterShutdown(t *testing.T) { - runtime, err := NewAdaptiveRuntime(NewAdaptiveConfig()) + runtime, err := NewAdaptiveRuntime(testAdaptiveRuntimeConfig("openai")) if err != nil { t.Fatalf(newAdaptiveRuntimeFailedMsg, err) } + if err := runtime.Shutdown(); err != nil { t.Fatalf("Shutdown failed: %v", err) } @@ -237,39 +213,6 @@ func assertAdaptiveRuntimeClosed(t *testing.T, runtime *AdaptiveRuntime) { } } -func TestAdaptiveRuntimeHelpersRejectInvalidInputs(t *testing.T) { - if _, err := BuildCacheTelemetryEvent(CacheTelemetryEventInput{ - Provider: "unsupported", - RequestID: "018f13f0-7c1a-7a80-8000-000000000001", - }); err == nil || !strings.Contains(err.Error(), "provider") { - t.Fatalf("expected unsupported provider rejection, got %v", err) - } - if _, err := BuildCacheTelemetryEvent(CacheTelemetryEventInput{ - Provider: "openai", - RequestID: "not-a-uuid", - }); err == nil || !strings.Contains(err.Error(), "request_id") { - t.Fatalf("expected invalid request ID rejection, got %v", err) - } - - runtime, err := NewAdaptiveRuntime(testAdaptiveRuntimeConfig("openai")) - if err != nil { - t.Fatalf(newAdaptiveRuntimeFailedMsg, err) - } - defer runtime.Shutdown() - if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{ - Provider: "openai", - RequestID: "not-a-uuid", - AnnotatedRequest: json.RawMessage(`{}`), - }); err == nil || !strings.Contains(err.Error(), "request_id") { - t.Fatalf("expected invalid request ID rejection, got %v", err) - } - if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{ - AnnotatedRequest: json.RawMessage(`not-json`), - }); err == nil { - t.Fatal("expected malformed annotated request JSON to fail before the FFI call") - } -} - func TestAdaptiveRuntimePublicHelpersPropagateJSONMarshalFailures(t *testing.T) { oldMarshal := jsonMarshal t.Cleanup(func() { jsonMarshal = oldMarshal }) diff --git a/go/nemo_relay/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go index f3eb8879c..8e3086054 100644 --- a/go/nemo_relay/async_middleware_test.go +++ b/go/nemo_relay/async_middleware_test.go @@ -7,6 +7,7 @@ import ( "context" "encoding/json" "errors" + "io" "strings" "sync" "testing" @@ -21,6 +22,12 @@ func asyncExecutionNoop(context.Context, json.RawMessage, AsyncNext) (any, error return nil, nil } +func asyncStreamExecutionNoop(context.Context, json.RawMessage, AsyncStreamNext) (<-chan AsyncStreamItem, error) { + ch := make(chan AsyncStreamItem) + close(ch) + return ch, nil +} + func TestAsyncMiddlewareGlobalRegistrationParity(t *testing.T) { registrations := []struct { name string @@ -50,7 +57,9 @@ func TestAsyncMiddlewareGlobalRegistrationParity(t *testing.T) { }, DeregisterLlmConditionalExecutionGuardrail}, {"llm-request", func(name string) error { return RegisterLlmRequestInterceptAsync(name, 0, false, asyncMiddlewareNoop) }, DeregisterLlmRequestIntercept}, {"llm-execution", func(name string) error { return RegisterLlmExecutionInterceptAsync(name, 0, asyncExecutionNoop) }, DeregisterLlmExecutionIntercept}, - {"llm-stream-execution", func(name string) error { return RegisterLlmStreamExecutionInterceptAsync(name, 0, asyncExecutionNoop) }, DeregisterLlmStreamExecutionIntercept}, + {"llm-stream-execution", func(name string) error { + return RegisterLlmStreamExecutionInterceptAsync(name, 0, asyncStreamExecutionNoop) + }, DeregisterLlmStreamExecutionIntercept}, } for _, registration := range registrations { @@ -131,7 +140,7 @@ func TestAsyncMiddlewareScopeLocalRegistrationParity(t *testing.T) { return ScopeRegisterLlmExecutionInterceptAsync(scopeUUID, name, 0, asyncExecutionNoop) }, func(name string) error { return ScopeDeregisterLlmExecutionIntercept(scopeUUID, name) }}, {"llm-stream-execution", func(name string) error { - return ScopeRegisterLlmStreamExecutionInterceptAsync(scopeUUID, name, 0, asyncExecutionNoop) + return ScopeRegisterLlmStreamExecutionInterceptAsync(scopeUUID, name, 0, asyncStreamExecutionNoop) }, func(name string) error { return ScopeDeregisterLlmStreamExecutionIntercept(scopeUUID, name) }}, } @@ -226,6 +235,173 @@ func TestAsyncToolMiddlewareCompletionAndNext(t *testing.T) { }) } +func TestAsyncLlmStreamExecutionInterceptEmitsChunksIncrementally(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + const name = "go-async-llm-stream-execution" + err := RegisterLlmStreamExecutionInterceptAsync(name, 0, + func(ctx context.Context, _ json.RawMessage, _ AsyncStreamNext) (<-chan AsyncStreamItem, error) { + chunks := make(chan AsyncStreamItem) + go func() { + defer close(chunks) + for _, chunk := range []json.RawMessage{json.RawMessage(`{"chunk":1}`), json.RawMessage(`{"chunk":2}`)} { + select { + case chunks <- AsyncStreamItem{Chunk: chunk}: + case <-ctx.Done(): + return + } + } + }() + return chunks, nil + }, + ) + if err != nil { + t.Fatalf("register stream execution intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterLlmStreamExecutionIntercept(name) }) + + stream, err := LlmStreamCallExecute( + "go-async-stream", + makeRequest(), + func(json.RawMessage) (json.RawMessage, error) { + return json.Marshal("data: {\"chunk\":1}\n\ndata: {\"chunk\":2}\n\ndata: [DONE]\n\n") + }, + nil, + nil, + ) + if err != nil { + t.Fatalf("execute stream: %v", err) + } + defer stream.Close() + + var chunks []json.RawMessage + for { + chunk, err := stream.Next() + if err == io.EOF { + break + } + if err != nil { + t.Fatalf("read stream: %v", err) + } + chunks = append(chunks, append(json.RawMessage(nil), chunk...)) + } + if len(chunks) != 2 { + t.Fatalf("chunks = %q, want two incremental chunks", chunks) + } + }) +} + +func TestAsyncLlmStreamExecutionNextPreservesTerminalError(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + const name = "go-async-llm-stream-next-error" + err := RegisterLlmStreamExecutionInterceptAsync(name, 0, + func(ctx context.Context, invocation json.RawMessage, next AsyncStreamNext) (<-chan AsyncStreamItem, error) { + var payload struct { + Request json.RawMessage `json:"request"` + } + if err := json.Unmarshal(invocation, &payload); err != nil { + return nil, err + } + return next(ctx, payload.Request) + }, + ) + if err != nil { + t.Fatalf("register stream execution intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterLlmStreamExecutionIntercept(name) }) + + stream, err := LlmStreamCallExecute( + name, + makeRequest(), + func(json.RawMessage) (json.RawMessage, error) { + return nil, errors.New("downstream stream failed") + }, + nil, + nil, + ) + if err != nil { + t.Fatalf("execute stream: %v", err) + } + defer stream.Close() + + _, err = stream.Next() + if err == nil || !strings.Contains(err.Error(), "downstream stream failed") { + t.Fatalf("stream error = %v, want downstream terminal error", err) + } + }) +} + +func TestAsyncLlmStreamExecutionNextCancellationStopsIdleDownstream(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + const outerName = "go-async-llm-stream-next-cancel" + const innerName = "go-async-llm-stream-idle-downstream" + + err := RegisterLlmStreamExecutionInterceptAsync(innerName, 10, + func(ctx context.Context, _ json.RawMessage, _ AsyncStreamNext) (<-chan AsyncStreamItem, error) { + ch := make(chan AsyncStreamItem) + go func() { + <-ctx.Done() + close(ch) + }() + return ch, nil + }, + ) + if err != nil { + t.Fatalf("register idle downstream intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterLlmStreamExecutionIntercept(innerName) }) + + err = RegisterLlmStreamExecutionInterceptAsync(outerName, 0, + func(ctx context.Context, invocation json.RawMessage, next AsyncStreamNext) (<-chan AsyncStreamItem, error) { + var payload struct { + Request json.RawMessage `json:"request"` + } + if err := json.Unmarshal(invocation, &payload); err != nil { + return nil, err + } + nextCtx, cancelNext := context.WithCancel(ctx) + downstream, err := next(nextCtx, payload.Request) + if err != nil { + cancelNext() + return nil, err + } + cancelNext() + select { + case _, ok := <-downstream: + if ok { + return nil, errors.New("cancelled downstream produced an item") + } + case <-time.After(2 * time.Second): + return nil, errors.New("cancelled idle downstream did not close") + } + ch := make(chan AsyncStreamItem) + close(ch) + return ch, nil + }, + ) + if err != nil { + t.Fatalf("register outer stream intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterLlmStreamExecutionIntercept(outerName) }) + + stream, err := LlmStreamCallExecute( + outerName, + makeRequest(), + func(json.RawMessage) (json.RawMessage, error) { + return json.RawMessage(`{"unused":true}`), nil + }, + nil, + nil, + ) + if err != nil { + t.Fatalf("execute stream: %v", err) + } + defer stream.Close() + if _, err := stream.Next(); err != io.EOF { + t.Fatalf("stream result = %v, want EOF after cancellation", err) + } + }) +} + func TestAsyncToolMiddlewarePropagatesCallbackAndNextErrors(t *testing.T) { runTestWithScopeStack(t, func(t *testing.T) { const conditionalName = "go-async-tool-conditional-error" diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index 2eeb20811..edbff52cf 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -52,14 +52,26 @@ typedef char* (*NemoRelayEventSanitizeFn)(void* user_data, const FfiEvent* event typedef struct FfiPluginContext FfiPluginContext; typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; +typedef struct NemoRelayAsyncStream NemoRelayAsyncStream; +typedef struct NemoRelayAsyncStreamInvocation NemoRelayAsyncStreamInvocation; typedef void (*NemoRelayAsyncNextResultCb)(void*, const char*, const char*); +typedef bool (*NemoRelayAsyncNextStreamResultCb)(void*, const char*, const char*, bool); extern int32_t nemo_relay_async_completion_resolve_json(const NemoRelayAsyncCompletion*, const char*); extern int32_t nemo_relay_async_completion_reject(const NemoRelayAsyncCompletion*, const char*); extern bool nemo_relay_async_completion_is_cancelled(const NemoRelayAsyncCompletion*); extern void nemo_relay_async_completion_release(const NemoRelayAsyncCompletion*); extern int32_t nemo_relay_async_next_invoke_callback(const NemoRelayAsyncNext*, const char*, NemoRelayAsyncNextResultCb, void*); +extern int32_t nemo_relay_async_next_invoke_stream_callback(const NemoRelayAsyncNext*, const char*, NemoRelayAsyncNextStreamResultCb, void*, const NemoRelayAsyncStreamInvocation**); +extern int32_t nemo_relay_async_stream_invocation_cancel(const NemoRelayAsyncStreamInvocation*); +extern void nemo_relay_async_stream_invocation_release(const NemoRelayAsyncStreamInvocation*); extern void nemo_relay_async_next_release(const NemoRelayAsyncNext*); +extern int32_t nemo_relay_async_stream_push_json(const NemoRelayAsyncStream*, const char*); +extern int32_t nemo_relay_async_stream_finish(const NemoRelayAsyncStream*); +extern int32_t nemo_relay_async_stream_reject(const NemoRelayAsyncStream*, const char*); +extern bool nemo_relay_async_stream_is_cancelled(const NemoRelayAsyncStream*); +extern void nemo_relay_async_stream_release(const NemoRelayAsyncStream*); extern void goAsyncNextResultTrampoline(void*, char*, char*); +extern bool goAsyncNextStreamResultTrampoline(void*, char*, char*, bool); // Middleware chain next function types typedef char* (*NemoRelayToolExecNextFn)(const char* args_json, void* next_ctx); @@ -190,15 +202,37 @@ type ToolSanitizeFunc func(name string, args json.RawMessage) json.RawMessage type ToolConditionalFunc func(name string, args json.RawMessage) *string // AsyncMiddlewareFunc is the common completion-based middleware callback. -// The JSON envelope identifies the middleware family and invocation fields. +// +// Relay invokes the callback from a goroutine and cancels ctx if the native +// invocation is abandoned. There is no implicit timeout. The invocation and +// result JSON contracts are documented in the Async Middleware section of the +// package README. type AsyncMiddlewareFunc func(ctx context.Context, invocation json.RawMessage) (any, error) -// AsyncNext invokes the remaining execution chain and returns its eventual result. +// AsyncNext invokes the remaining execution chain and returns its eventual +// result. It is valid only while its enclosing intercept callback is running. type AsyncNext func(ctx context.Context, invocation json.RawMessage) (json.RawMessage, error) -// AsyncExecutionInterceptFunc is an asynchronous execution intercept with an awaitable next helper. +// AsyncExecutionInterceptFunc is an asynchronous execution intercept with an +// awaitable next helper. type AsyncExecutionInterceptFunc func(ctx context.Context, invocation json.RawMessage, next AsyncNext) (any, error) +// AsyncStreamItem is one chunk or terminal error from an asynchronous stream. +// A channel producer must stop when its callback context is cancelled. +type AsyncStreamItem struct { + Chunk json.RawMessage + Err error +} + +// AsyncStreamNext invokes the remaining streaming execution chain without +// collecting it into a single result. It is valid only while its enclosing +// intercept callback is running. +type AsyncStreamNext func(ctx context.Context, invocation json.RawMessage) (<-chan AsyncStreamItem, error) + +// AsyncStreamExecutionInterceptFunc produces chunks incrementally. Returning +// the channel from next preserves the downstream stream without buffering it. +type AsyncStreamExecutionInterceptFunc func(ctx context.Context, invocation json.RawMessage, next AsyncStreamNext) (<-chan AsyncStreamItem, error) + const asyncCallbackPending = C.uint32_t(1) const asyncCancellationPollInterval = 10 * time.Millisecond @@ -864,6 +898,237 @@ func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON return asyncCallbackPending } +type asyncNextStreamState struct { + ch chan AsyncStreamItem + ctx context.Context + cancel context.CancelFunc + mu sync.Mutex + done chan struct{} + closed bool +} + +func (state *asyncNextStreamState) deliver(item AsyncStreamItem) bool { + state.mu.Lock() + defer state.mu.Unlock() + if state.closed { + return false + } + select { + case state.ch <- item: + return true + case <-state.ctx.Done(): + return false + } +} + +func (state *asyncNextStreamState) finish(item *AsyncStreamItem) { + state.mu.Lock() + defer state.mu.Unlock() + if state.closed { + return + } + state.closed = true + if item != nil { + select { + case state.ch <- *item: + case <-state.ctx.Done(): + } + } + close(state.ch) + close(state.done) + state.cancel() +} + +//export goAsyncNextStreamResultTrampoline +func goAsyncNextStreamResultTrampoline(userData unsafe.Pointer, chunkJSON *C.char, errorMessage *C.char, done C.bool) C.bool { + state, ok := lookupClosure(userData).(*asyncNextStreamState) + if !ok { + unregisterClosure(userData) + return C.bool(false) + } + if bool(done) { + var terminal *AsyncStreamItem + if errorMessage != nil { + item := AsyncStreamItem{Err: errors.New(C.GoString(errorMessage))} + terminal = &item + } + state.finish(terminal) + unregisterClosure(userData) + return C.bool(true) + } + chunk := append(json.RawMessage(nil), []byte(C.GoString(chunkJSON))...) + if state.deliver(AsyncStreamItem{Chunk: chunk}) { + return C.bool(true) + } + state.finish(nil) + unregisterClosure(userData) + return C.bool(false) +} + +func contextForAsyncStream(stream *C.NemoRelayAsyncStream) (context.Context, func()) { + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + finished := make(chan struct{}) + go func() { + defer close(finished) + ticker := time.NewTicker(asyncCancellationPollInterval) + defer ticker.Stop() + for { + select { + case <-done: + return + case <-ticker.C: + if bool(C.nemo_relay_async_stream_is_cancelled(stream)) { + cancel() + return + } + } + } + }() + var once sync.Once + return ctx, func() { + once.Do(func() { + close(done) + <-finished + cancel() + }) + } +} + +func rejectAsyncStreamPanic(stream *C.NemoRelayAsyncStream) { + recovered := recover() + if recovered == nil { + return + } + message := C.CString(fmt.Sprintf("panic in async stream execution intercept: %v", recovered)) + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_stream_reject(stream, message) +} + +//export goAsyncStreamExecutionInterceptTrampoline +func goAsyncStreamExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON *C.char, next *C.NemoRelayAsyncNext, stream *C.NemoRelayAsyncStream) C.uint32_t { + fn, ok := lookupClosure(userData).(AsyncStreamExecutionInterceptFunc) + if !ok { + message := C.CString("nemo_relay: async stream execution intercept callback is not registered") + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_stream_reject(stream, message) + C.nemo_relay_async_stream_release(stream) + C.nemo_relay_async_next_release(next) + return asyncCallbackPending + } + invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) + go func() { + defer C.nemo_relay_async_stream_release(stream) + defer rejectAsyncStreamPanic(stream) + ctx, cancel := contextForAsyncStream(stream) + defer cancel() + var nextMu sync.RWMutex + nextOpen := true + defer func() { + cancel() + nextMu.Lock() + nextOpen = false + nextMu.Unlock() + C.nemo_relay_async_next_release(next) + }() + nextFn := func(nextCtx context.Context, payload json.RawMessage) (<-chan AsyncStreamItem, error) { + nextMu.RLock() + defer nextMu.RUnlock() + if !nextOpen { + return nil, context.Canceled + } + combinedCtx, combinedCancel := context.WithCancel(ctx) + go func() { + select { + case <-nextCtx.Done(): + combinedCancel() + case <-combinedCtx.Done(): + } + }() + // Keep one terminal slot so a downstream error can be delivered + // before cancellation closes the stream. + ch := make(chan AsyncStreamItem, 1) + state := &asyncNextStreamState{ + ch: ch, + ctx: combinedCtx, + cancel: combinedCancel, + done: make(chan struct{}), + } + token := registerClosure(state) + cPayload := C.CString(string(payload)) + var invocation *C.NemoRelayAsyncStreamInvocation + status := C.nemo_relay_async_next_invoke_stream_callback( + next, cPayload, + (C.NemoRelayAsyncNextStreamResultCb)(C.goAsyncNextStreamResultTrampoline), token, + &invocation, + ) + C.free(unsafe.Pointer(cPayload)) + if err := checkStatus(status); err != nil { + state.finish(nil) + unregisterClosure(token) + return nil, err + } + go func() { + select { + case <-state.done: + case <-combinedCtx.Done(): + select { + case <-state.done: + default: + C.nemo_relay_async_stream_invocation_cancel(invocation) + state.finish(nil) + unregisterClosure(token) + } + } + C.nemo_relay_async_stream_invocation_release(invocation) + }() + return ch, nil + } + output, err := fn(ctx, invocation, nextFn) + if err != nil { + message := C.CString(err.Error()) + C.nemo_relay_async_stream_reject(stream, message) + C.free(unsafe.Pointer(message)) + return + } + if output == nil { + message := C.CString("async stream execution intercept returned a nil channel") + C.nemo_relay_async_stream_reject(stream, message) + C.free(unsafe.Pointer(message)) + return + } + for { + select { + case <-ctx.Done(): + return + case item, ok := <-output: + if !ok { + C.nemo_relay_async_stream_finish(stream) + return + } + if item.Err != nil { + message := C.CString(item.Err.Error()) + C.nemo_relay_async_stream_reject(stream, message) + C.free(unsafe.Pointer(message)) + return + } + chunk := C.CString(string(item.Chunk)) + status := C.nemo_relay_async_stream_push_json(stream, chunk) + C.free(unsafe.Pointer(chunk)) + if err := checkStatus(status); err != nil { + if !bool(C.nemo_relay_async_stream_is_cancelled(stream)) { + message := C.CString(err.Error()) + C.nemo_relay_async_stream_reject(stream, message) + C.free(unsafe.Pointer(message)) + } + return + } + } + } + }() + return asyncCallbackPending +} + //export goToolConditionalTrampoline func goToolConditionalTrampoline(userData unsafe.Pointer, name *C.char, argsJSON *C.char) *C.char { fn := lookupClosure(userData).(ToolConditionalFunc) diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index bb931c9a5..eb5150f9c 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -46,8 +46,10 @@ typedef struct NemoRelayLlmSanitizeResponseContext { uint32_t codec_kind; const typedef void (*NemoRelayFreeFn)(void* user_data); typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; +typedef struct NemoRelayAsyncStream NemoRelayAsyncStream; typedef uint32_t (*NemoRelayAsyncJsonCb)(void*, const char*, const NemoRelayAsyncCompletion*); typedef uint32_t (*NemoRelayAsyncInterceptCb)(void*, const char*, const NemoRelayAsyncNext*, const NemoRelayAsyncCompletion*); +typedef uint32_t (*NemoRelayAsyncStreamInterceptCb)(void*, const char*, const NemoRelayAsyncNext*, const NemoRelayAsyncStream*); // Core API extern int32_t nemo_relay_get_handle(FfiScopeHandle** out); @@ -175,7 +177,7 @@ extern int32_t nemo_relay_register_llm_execution_intercept(const char* name, int extern int32_t nemo_relay_register_llm_execution_intercept_async(const char*, int32_t, NemoRelayAsyncInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_execution_intercept(const char* name); extern int32_t nemo_relay_register_llm_stream_execution_intercept(const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); -extern int32_t nemo_relay_register_llm_stream_execution_intercept_async(const char*, int32_t, NemoRelayAsyncInterceptCb, void*, NemoRelayFreeFn); +extern int32_t nemo_relay_register_llm_stream_execution_intercept_async(const char*, int32_t, NemoRelayAsyncStreamInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_stream_execution_intercept(const char* name); // Subscribers @@ -241,7 +243,7 @@ extern int32_t nemo_relay_scope_register_llm_execution_intercept(const char* sco extern int32_t nemo_relay_scope_register_llm_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_execution_intercept(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_stream_execution_intercept(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); -extern int32_t nemo_relay_scope_register_llm_stream_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_stream_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncStreamInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_stream_execution_intercept(const char* scope_uuid, const char* name); // Scope-local subscribers @@ -301,6 +303,7 @@ extern void nemo_relay_otel_subscriber_free(void*); extern char* goToolSanitizeTrampoline(void*, const char*, const char*); extern uint32_t goAsyncMiddlewareTrampoline(void*, const char*, const NemoRelayAsyncCompletion*); extern uint32_t goAsyncExecutionInterceptTrampoline(void*, const char*, const NemoRelayAsyncNext*, const NemoRelayAsyncCompletion*); +extern uint32_t goAsyncStreamExecutionInterceptTrampoline(void*, const char*, const NemoRelayAsyncNext*, const NemoRelayAsyncStream*); extern char* goEventSanitizeTrampoline(void*, const FfiEvent*, const char*); extern char* goToolConditionalTrampoline(void*, const char*, const char*); extern char* goToolExecTrampoline(void*, const char*); @@ -1638,9 +1641,9 @@ func RegisterLlmStreamExecutionIntercept(name string, priority int32, execFn LLM } // RegisterLlmStreamExecutionInterceptAsync registers an asynchronous streaming LLM intercept. -func RegisterLlmStreamExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { +func RegisterLlmStreamExecutionInterceptAsync(name string, priority int32, fn AsyncStreamExecutionInterceptFunc) error { return withGlobalAsyncMiddleware(name, priority, fn, func(name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { - return C.nemo_relay_register_llm_stream_execution_intercept_async(name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + return C.nemo_relay_register_llm_stream_execution_intercept_async(name, priority, C.NemoRelayAsyncStreamInterceptCb(C.goAsyncStreamExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) }) } @@ -1686,9 +1689,10 @@ func DeregisterSubscriber(name string) error { // FlushSubscribers waits for subscriber callbacks queued before this call to // finish. Native event-producing APIs enqueue subscriber work and return // without waiting for observer callbacks. Call this function outside native -// subscriber callbacks. A re-entrant call returns without waiting to avoid -// blocking the dispatcher, so callbacks later in the same dispatch snapshot -// can still run. +// subscriber callbacks. A call made while an asynchronous publication +// boundary is active may return before that boundary and later queued +// callbacks finish. Call FlushSubscribers again after the middleware settles +// to wait for the remaining work. func FlushSubscribers() error { return checkStatus(C.nemo_relay_flush_subscribers()) } @@ -2308,7 +2312,7 @@ func (s *OpenTelemetrySubscriber) Close() { // --------------------------------------------------------------------------- type asyncMiddlewareCallback interface { - AsyncMiddlewareFunc | AsyncExecutionInterceptFunc + AsyncMiddlewareFunc | AsyncExecutionInterceptFunc | AsyncStreamExecutionInterceptFunc } func withGlobalAsyncMiddleware[T asyncMiddlewareCallback](name string, priority int32, fn T, call func(*C.char, C.int32_t, unsafe.Pointer) C.int32_t) error { @@ -2621,9 +2625,9 @@ func ScopeRegisterLlmExecutionInterceptAsync(scopeUUID, name string, priority in } // ScopeRegisterLlmStreamExecutionInterceptAsync registers an asynchronous scope-local streaming LLM intercept. -func ScopeRegisterLlmStreamExecutionInterceptAsync(scopeUUID, name string, priority int32, fn AsyncExecutionInterceptFunc) error { +func ScopeRegisterLlmStreamExecutionInterceptAsync(scopeUUID, name string, priority int32, fn AsyncStreamExecutionInterceptFunc) error { return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { - return C.nemo_relay_scope_register_llm_stream_execution_intercept_async(scope, name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + return C.nemo_relay_scope_register_llm_stream_execution_intercept_async(scope, name, priority, C.NemoRelayAsyncStreamInterceptCb(C.goAsyncStreamExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) }) } diff --git a/go/nemo_relay/optimization_test.go b/go/nemo_relay/optimization_test.go index a4fab1eaf..60712a8da 100644 --- a/go/nemo_relay/optimization_test.go +++ b/go/nemo_relay/optimization_test.go @@ -98,26 +98,6 @@ func TestLLMOptimizationContributionOmittedAppliedIsNonApplied(t *testing.T) { } } -func TestLLMOptimizationContributionRejectsMalformedAndNonObjectWireShapes(t *testing.T) { - var contribution LLMOptimizationContribution - if err := json.Unmarshal([]byte(`not-json`), &contribution); err == nil { - t.Fatal("expected malformed optimization contribution JSON to fail") - } - if err := json.Unmarshal([]byte(`[]`), &contribution); err == nil { - t.Fatal("expected non-object optimization contribution JSON to fail") - } - - contribution = LLMOptimizationContribution{ - Producer: "test", - Kind: "custom", - PayloadSchema: &LLMOptimizationDataSchema{Name: "test", Version: "v1"}, - Payload: json.RawMessage(`not-json`), - } - if _, err := json.Marshal(contribution); err == nil { - t.Fatal("expected malformed payload JSON to fail") - } -} - func TestLLMRequestInterceptOptimizationContributionsRoundTrip(t *testing.T) { fixture, contribution := optimizationContributionFixture(t) const interceptName = "go_optimization_fixture"