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..3a95f06ac 100644 --- a/crates/adaptive/src/adaptive_hints_intercept.rs +++ b/crates/adaptive/src/adaptive_hints_intercept.rs @@ -174,14 +174,14 @@ 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 this = this.clone(); 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); @@ -196,9 +196,9 @@ impl AdaptiveHintsIntercept { inject_agent_hints(&mut request, &mut annotated, &hints); } - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + let outcome = + nemo_relay::api::llm::LlmRequestInterceptOutcome::new(request, annotated); + Box::pin(async move { Ok(outcome) }) }, ) } 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/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/llm.rs b/crates/core/src/api/llm.rs index c34829241..cfd107173 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_reserved_sanitized_event, 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,34 @@ 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())? + }; + 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( handle: &LlmHandle, marks: Vec, subscribers: &[EventSubscriberFn], @@ -517,28 +528,78 @@ 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( +pub(crate) async fn emit_reserved_optimization_marks( handle: &LlmHandle, subscribers: &[EventSubscriberFn], - mut sanitize: impl FnMut(Event) -> Option, - mut enqueue: impl FnMut(&Event, &[EventSubscriberFn]) -> bool, ) { + 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 +/// 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: 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 +615,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 +633,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 @@ -615,11 +711,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`, @@ -641,7 +739,78 @@ 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() + .map_err(|error| FlowError::Internal(error.to_string()))?; + 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(); + 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) } @@ -651,6 +820,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 @@ -672,27 +910,112 @@ struct LlmCallEndBehavior { /// 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. +/// 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 /// 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() + .map_err(|error| FlowError::Internal(error.to_string()))?; + 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 = params.response; + let fallback_data = params.data; + 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(), + ) + }; + 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 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; + 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 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(payload.data) + .metadata_opt(end_metadata) + .annotated_response_opt(payload.annotated_response) + .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]>, @@ -725,55 +1048,18 @@ 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()), - &entries, - ); - 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); - 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); + emit_optimization_marks(handle, &subscribers).await; + let payload = build_llm_end_payload( + handle, + response, + data, + annotated_response, + response_codec, + &entries, + behavior, + ) + .await; let event = { let context = global_context(); let state = context @@ -783,17 +1069,18 @@ 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(), ) }; - 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 + if let Some(error) = payload.decode_error && behavior.response_codec_errors_fatal { Err(error) @@ -842,7 +1129,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 +1155,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 +1193,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 +1281,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 +1311,7 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { codec, &optimization_recorder, ) + .await }) .await?; @@ -1043,12 +1337,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 +1382,8 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { attach_estimated_cost: true, }, Some(&lifecycle_subscribers), - )?; + ) + .await?; Ok(response) } Err(error) => { @@ -1098,7 +1394,8 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { end_metadata, response_codec, Some(&lifecycle_subscribers), - ); + ) + .await; Err(error) } } @@ -1186,7 +1483,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 +1513,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu codec, &optimization_recorder, ) + .await }) .await?; @@ -1239,12 +1539,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 +1590,8 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu end_metadata, response_codec, Some(&lifecycle_subscribers), - ); + ) + .await; Err(error) } } @@ -1318,7 +1620,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 +1638,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 +1664,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 +1687,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..d52577f07 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( + Arc, + 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 eb3dbb950..b83e3d170 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, @@ -39,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}; @@ -564,7 +568,7 @@ impl NemoRelayContextState { )) } - fn emit_guardrail_scope_start( + async fn emit_guardrail_scope_start( name: &str, parent_uuid: Option, metadata: Option, @@ -591,13 +595,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], @@ -616,7 +620,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); } } @@ -633,13 +637,36 @@ 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 { - let fields = (entry.payload)(&event, event.sanitize_fields()); - event.apply_sanitize_fields(fields); + let fields = event.sanitize_fields(); + let callback = Arc::clone(&entry.payload); + let context = Arc::new(event); + let callback_context = Arc::clone(&context); + let outcome = AssertUnwindSafe(async move { callback(callback_context, fields).await }) + .catch_unwind() + .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", + 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 } @@ -672,14 +699,36 @@ 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); + 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 } @@ -712,14 +761,36 @@ 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); + 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 } @@ -769,7 +840,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], @@ -787,8 +858,22 @@ impl NemoRelayContextState { "target_name": name, }), subscribers, - ); - let result = (entry.payload)(name, args); + ) + .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, @@ -804,7 +889,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)); } @@ -847,14 +932,27 @@ 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)?; + 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; } @@ -964,14 +1062,45 @@ 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() { + 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", + sanitizer = entry.name.as_str(), + preserved_value = "last_valid_request"; + "LLM request sanitizer failed; preserving the last valid request: {error}" + ); + 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); + } + } + } } value } @@ -1003,14 +1132,45 @@ 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() { + 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", + sanitizer = entry.name.as_str(), + preserved_value = "last_valid_response"; + "LLM response sanitizer failed; preserving the last valid response: {error}" + ); + 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); + } + } + } } value } @@ -1059,7 +1219,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], @@ -1075,8 +1235,20 @@ impl NemoRelayContextState { "kind": "llm_conditional_execution", }), subscribers, - ); - let result = (entry.payload)(request); + ) + .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, @@ -1092,7 +1264,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)); } @@ -1139,7 +1311,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, @@ -1154,11 +1326,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, @@ -1172,7 +1345,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, request_value, annotated_value)?; + 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", @@ -1280,3 +1468,7 @@ impl Default for NemoRelayContextState { Self::new() } } + +#[cfg(test)] +#[path = "../../../tests/unit/runtime_state_tests.rs"] +mod tests; diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index ffb4ef3b0..fc77a9bcc 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -4,11 +4,21 @@ //! 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; +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; + 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}; @@ -24,21 +34,69 @@ mod native { enum DispatcherMessage { Deliver { event: Box, + transform: Option, + sanitizers: Vec>, subscribers: Vec, scope_stack: ScopeStackHandle, }, Flush { done: Sender<()>, }, + Barrier { + publications: 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) }; } + tokio::task_local! { + 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)); + Self + } + } + + impl Drop for DispatchGuard { + fn drop(&mut self) { + IN_DISPATCHER.with(|flag| flag.set(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() { @@ -46,38 +104,101 @@ 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(), }; - 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( + 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, + }; + 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, + 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) + } + + /// 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 (publication_tx, publication_rx) = mpsc::channel(); + sender + .send(DispatcherMessage::Barrier { + publications: publication_rx, + }) + .ok() + .map(|_| AsyncPublication { + sender: publication_tx, + }) } 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 { @@ -98,10 +219,63 @@ mod native { Ok(()) } + pub(super) fn in_dispatcher_callback() -> bool { + IN_DISPATCHER.with(Cell::get) || ASYNC_PUBLICATION_MESSAGES.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 { + 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 + } + } + fn dispatcher_sender() -> std::result::Result, String> { 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() @@ -120,25 +294,47 @@ 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(()); } } + DispatcherMessage::Barrier { publications } => { + if let Ok(publications) = publications.recv() { + for publication in publications { + handle_message(publication); + } + } + } message => handle_message(message), } } } - 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), + message @ DispatcherMessage::Barrier { .. } => { + pending.push_back(message); + break; + } message => handle_message(message), } } @@ -149,23 +345,38 @@ mod native { match message { DispatcherMessage::Deliver { event, + transform, + sanitizers, subscribers, scope_stack, - } => deliver_event(event, subscribers, scope_stack), + } => deliver_event(event, transform, sanitizers, subscribers, scope_stack), DispatcherMessage::Flush { done } => { let _ = done.send(()); } + DispatcherMessage::Barrier { publications } => { + if let Ok(publications) = publications.recv() { + for publication in publications { + handle_message(publication); + } + } + } } } fn deliver_event( event: Box, + transform: Option, + 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 _dispatch_guard = DispatchGuard::enter(); + let Some(event) = sanitize_event_snapshot(*event, transform, sanitizers) else { + restore_thread_scope_stack(previous_scope_stack); + return; + }; for subscriber in subscribers { if catch_unwind(AssertUnwindSafe(|| subscriber(&event))).is_err() { log::error!( @@ -175,9 +386,159 @@ 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() { + 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 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) + } + } + } + + #[cfg(test)] + mod tests { + use super::*; + + #[test] + 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(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(); + } + } +} + +#[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. @@ -185,7 +546,71 @@ 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) +} + +/// 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( + 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 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, +/// 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( + publication: Option, + future: F, +) -> F::Output { + native::with_async_publication_context(publication, future).await +} + /// Wait for all queued subscriber callbacks submitted before this call. 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/core/src/api/scope.rs b/crates/core/src/api/scope.rs index 60a1aa53d..98afdfbb7 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; @@ -51,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"))))] @@ -216,14 +227,16 @@ 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); - 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_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(); @@ -241,12 +254,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) } @@ -273,11 +290,18 @@ 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(); - let (scope, event, subscribers) = { - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let (scope, event, subscribers, emission_scope_stack) = { + 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() { @@ -302,13 +326,21 @@ 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); + // 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); - 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,21 +367,25 @@ 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); 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 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().expect("scope stack lock poisoned"); + 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(); @@ -368,10 +404,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..5491b3182 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; @@ -42,21 +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).await) +} + +/// Snapshot the event sanitizers visible to an event without invoking them. +/// +/// 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>> { let entries = { - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + 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(_) => { @@ -88,9 +122,7 @@ pub(crate) fn sanitize_event_with_scope_stack( } } }; - Some(NemoRelayContextState::event_sanitize_snapshot_chain( - event, &entries, - )) + Some(entries) } pub(crate) fn ensure_runtime_owner() -> Result<()> { @@ -196,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>, @@ -228,7 +260,9 @@ 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); @@ -246,7 +280,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/subscriber.rs b/crates/core/src/api/subscriber.rs index 02ce9a322..337544167 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 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/api/tool.rs b/crates/core/src/api/tool.rs index 6fbe6cd70..14b974651 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}; @@ -81,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"))))] @@ -180,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. @@ -196,21 +218,106 @@ 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 /// 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() + .map_err(|error| FlowError::Internal(error.to_string()))?; + 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 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(); + 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(); + 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 { + let sanitizers = snapshot_event_sanitizers(&mark, &scope_stack).unwrap_or_default(); + 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()?; @@ -230,25 +337,13 @@ 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, &entries, - ); + ) + .await; let (handle, event, marks) = { let context = global_context(); let state = context @@ -286,23 +381,21 @@ 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) { + if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, &subscribers); } for mark in marks { - NemoRelayContextState::emit_event(&mark, &subscribers); + if let Some(mark) = sanitize_event(mark).await { + NemoRelayContextState::emit_event(&mark, &subscribers); + } } Ok((handle, subscribers)) } /// 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. @@ -316,20 +409,82 @@ 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 /// 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() + .map_err(|error| FlowError::Internal(error.to_string()))?; + 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(); + 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(()) } -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 +513,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 +552,19 @@ fn tool_call_end_with_pending_marks( mark.category_profile, )) }) - .filter_map(sanitize_event) .collect::>(); - if let Some(event) = sanitize_event(event) { + if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, subscribers); } for mark in marks { - NemoRelayContextState::emit_event(&mark, subscribers); + if let Some(mark) = sanitize_event(mark).await { + 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 +577,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 +650,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 +685,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 +697,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 +728,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 +759,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 +772,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 +796,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 +820,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/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index f1217c746..0f51a25c4 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,744 @@ 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>>>, + 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) }; + 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( + 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(NativeAsyncResult::Json(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(NativeAsyncResult::Json) + .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.map(NativeAsyncResult::Json) }) + } + 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 { next(request).await.map(NativeAsyncResult::LlmStream) }) + } + }; + 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? + .into_json()?; + 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? + .into_json()? + { + 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? + .into_json()? + { + 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? + .into_json()?; + 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? + .into_json()?; + 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? + .into_json()?, + ) + .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? + .into_json()?, + ) + .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? + .into_json()?, + ) + .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(); + 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() + }) + }) +} + +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?; + 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(), + )), + } + }) + }) +} + +unsafe extern "C" fn native_plugin_context_register_async_middleware( + ctx: *mut NemoRelayNativePluginContext, + kind: u32, + 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 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 + .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 +2513,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 +2569,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 +2581,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 +2624,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 +2742,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 +2757,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 +2901,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 +2940,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..1d0bb2ecf 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: Arc, _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,25 @@ 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 { + 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_name; + "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..a1c02179b 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -34,19 +34,21 @@ 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; 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, } } @@ -172,7 +178,7 @@ impl LlmStreamWrapper { &self.scope_stack } - fn finish(&mut self) { + fn finish(&mut self, background_thread: bool) { if self.ended { return; } @@ -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, background_thread); } 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, @@ -216,82 +236,133 @@ 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 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 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() - }); - } - 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)) + .flatten(); + let interruption = (interrupted + && !has_authoritative_final_usage(annotated_response.as_ref())) + .then_some("stream_interrupted"); + handle + .optimization_recorder + .close_for_finalization(interruption); + emit_reserved_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, } - Err(_) => None, + }; + if let Some(event) = event_snapshot + && let Some(event) = sanitize_event_with_scope_stack(event, &scope_stack).await + { + let _ = subscriber_dispatcher::dispatch_reserved_sanitized_event( + event, + Vec::new(), + &subscribers, + scope_stack.clone(), + ); } }; - if let Some(event) = event_snapshot - && let Some(event) = sanitize_event_with_scope_stack(event, &self.scope_stack) - { - NemoRelayContextState::emit_event(&event, &self.subscribers); + 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 + // 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; + } + match tokio::runtime::Handle::try_current() { + Ok(handle) => Some(handle.spawn(finalize)), + Err(_) => { + std::thread::spawn(move || { + if let Ok(runtime) = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + { + runtime.block_on(finalize); + } + }); + None + } } } @@ -317,10 +388,15 @@ impl LlmStreamWrapper { Err(_) => 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); + 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, + &self.subscribers, + self.scope_stack.clone(), + ); } } } @@ -328,11 +404,37 @@ 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); + return match this.terminal_result.take() { + Some(result) => Poll::Ready(Some(result)), + None => Poll::Ready(None), + }; } // Poll the inner stream @@ -346,19 +448,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, } @@ -373,7 +477,12 @@ 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}")) + })?; + } this.close_result = Some(result.clone()); this.close_result .as_ref() @@ -641,7 +750,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 350cf1414..46c05c265 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -3,18 +3,28 @@ use std::ffi::c_void; use std::ptr; +use std::sync::atomic::{AtomicBool, Ordering}; use nemo_relay_plugin::{ CategoryProfile, ConfigDiagnostic, DiagnosticLevel, Event, EventCategory, EventSanitizeFields, - Json, LlmJsonStream, LlmRequest, LlmRequestInterceptOutcome, NemoRelayNativeHostApiV1, - 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}; 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" @@ -55,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", @@ -130,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", @@ -165,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, @@ -280,6 +297,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_v3 = unsafe { &*(host as *const NemoRelayNativeHostApiV3) }; + let mut plugin = NemoRelayNativePluginV1::default(); + 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_v3)).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 +614,372 @@ 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 as u32, + 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, +) -> u32 { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete as u32; + }; + 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 as u32 +} + +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, +) -> u32 { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + 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()) + .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 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", + ) + }; + } else { + unsafe { + (host.async_completion_resolve_json)(completion, result); + (host.v1.string_free)(result); + } + } + NemoRelayNativeAsyncCallbackState::Complete as u32 +} + +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, +) -> u32 { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + 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()) + .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 as u32; + }; + if pending { + ASYNC_PENDING_ENTERED.store(true, Ordering::Release); + 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 as u32; + } + 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 as u32 +} + +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, +) -> u32 { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + 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 as u32; + } + 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") + }; + 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 { (host.async_next_release)(next) }; + return NemoRelayNativeAsyncCallbackState::Complete as u32; + } + 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 as u32 + } else { + unsafe { + reject_async_completion( + host, + completion, + "failed to invoke async tool execution next", + ) + }; + unsafe { (host.async_next_release)(next) }; + NemoRelayNativeAsyncCallbackState::Complete as u32 + } +} + +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, @@ -598,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 @@ -638,6 +1046,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..5b14fe82f 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,32 @@ 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) - ); + let scope_events = success_events + .iter() + .filter(|event| event.kind() == "scope") + .collect::>(); assert_eq!( - success_events.last().unwrap().category().unwrap().as_str(), - "llm" + 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.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..7e8763850 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, @@ -1626,6 +1645,7 @@ fn test_scope_local_guardrail_lifecycle() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!( call_count.load(Ordering::SeqCst), 1, @@ -1700,8 +1720,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 +1738,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 +1754,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 +1777,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 +1908,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 +1920,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 +1959,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 +2010,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 +2024,7 @@ fn test_sanitize_guardrails_pipe_data() { args.as_object_mut() .unwrap() .insert("field_a".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -2018,7 +2039,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 +2082,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 +2096,7 @@ fn test_response_sanitize_guardrails_pipe() { .as_object_mut() .unwrap() .insert("sanitized".into(), json!(true)); - result + ready(result) }), ) .unwrap(); @@ -2127,8 +2148,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 +2166,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 +2194,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 +2212,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 +2238,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 +2248,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 +2267,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 +2301,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 +2329,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 +2357,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 +2385,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 +2393,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 +2406,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 +2419,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 +2442,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 +2454,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 +2466,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 +2479,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 +2490,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 +2502,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 +2536,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 +2548,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 +2613,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 +2625,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 +2637,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 +2652,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 +2665,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 +2677,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 +2734,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 +2746,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 +2809,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 +2880,7 @@ async fn test_full_pipeline_integration() { args.as_object_mut() .unwrap() .insert("intercepted".into(), json!(true)); - Ok(args) + ready(args) }), ) .unwrap(); @@ -2864,7 +2892,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 +2904,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 +2931,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 +2989,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 +3020,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 +3057,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 +3070,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 +3084,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 +3097,7 @@ fn test_deregister_removes_from_chain() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!(call_count.load(Ordering::SeqCst), 1); // Deregister @@ -3070,6 +3112,7 @@ fn test_deregister_removes_from_chain() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!( call_count.load(Ordering::SeqCst), 1, @@ -3091,7 +3134,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 +3185,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 +3279,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 +3293,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 +3316,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 +3332,7 @@ fn test_llm_request_intercept_pending_marks_preserve_order_and_break_chain() { content: json!({"prompt": "hello"}), }, ) + .await .unwrap(); assert_eq!( @@ -3322,7 +3368,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 +3377,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 +3497,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 +3514,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 +3540,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 +3748,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 +3803,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 +3849,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 +3943,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 +3955,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 +4070,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 +4164,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 +4217,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 +4246,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 +4266,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 +4313,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..8e9782def 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,144 @@ 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 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") + }; + // 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(); + 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"); + cleanup.mark_plugin_configuration_active(); + + 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::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 + .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); +} + #[tokio::test] async fn native_validation_diagnostics_prevent_initialization() { let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; @@ -752,7 +892,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 +933,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 +1358,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 +1379,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 +1535,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..999369412 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; @@ -27,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; @@ -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 @@ -1961,3 +1972,164 @@ 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(); + + 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(); + 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/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 e83fc7d6d..5bdfd7184 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -3,14 +3,21 @@ //! Integration tests for native subscriber dispatch behavior. +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, +}; 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 nemo_relay::error::FlowError; +use serde_json::json; static TEST_MUTEX: Mutex<()> = Mutex::new(()); @@ -126,3 +133,130 @@ fn dispatcher_continues_after_subscriber_panic() { assert_eq!(observed.lock().unwrap().as_slice(), ["after-panic"]); } + +#[test] +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 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.clone())), + ) + .unwrap(); + register_mark_sanitize_guardrail( + "fail-open-mark-sanitizer", + 10, + Arc::new(|_, _| { + Box::pin(async { + Err(FlowError::Internal( + "intentional event-sanitizer failure".to_string(), + )) + }) + }), + ) + .unwrap(); + + 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[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(); +} + +#[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); + Box::pin(async move { Ok(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_publishes_the_snapshot_when_an_async_sanitizer_panics() { + let _lock = TEST_MUTEX.lock().unwrap(); + flush_subscribers().unwrap(); + reset_global(); + setup_isolated_thread(); + + 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.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(); + register_mark_sanitize_guardrail( + "panic-mark-sanitizer", + 10, + Arc::new(|_, _| Box::pin(async { panic!("intentional event-sanitizer panic") })), + ) + .unwrap(); + + emit_mark("panic-fallback"); + flush_subscribers().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})) + ); + drop(events); + deregister_mark_sanitize_guardrail("successful-mark-sanitizer").unwrap(); + 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..63320c7a8 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)] // The process-wide test mutex intentionally serializes runtime state. async fn installed_callbacks_apply_surface_specific_fallbacks() { struct RuntimeCleanup { registrations: Option, @@ -1411,62 +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); - assert_eq!(sanitized.data(), None); - assert_eq!(sanitized.metadata(), None); - } - - let entries = state.tool_sanitize_request_entries(&[]); - assert_eq!( - NemoRelayContextState::tool_sanitize_request_snapshot_chain( - "tool", - tool_request.clone(), - &entries, + ( + state.collect_event_subscribers(&[]), + NemoRelayContextState::event_sanitize_entries(&state.mark_sanitize_guardrails, &[]), + NemoRelayContextState::event_sanitize_entries( + &state.scope_sanitize_start_guardrails, + &[], ), - tool_request - ); - let entries = state.tool_sanitize_response_entries(&[]); - assert_eq!( - NemoRelayContextState::tool_sanitize_response_snapshot_chain( - "tool", - tool_response.clone(), - &entries, + NemoRelayContextState::event_sanitize_entries( + &state.scope_sanitize_end_guardrails, + &[], ), - tool_response - ); - let entries = state.llm_sanitize_request_entries(&[]); - assert!( - 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" - ); - let entries = state.llm_sanitize_response_entries(&[]); - assert!( - 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" - ); + 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); + + 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 65c6e5b17..8e90747da 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( @@ -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() @@ -760,7 +777,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 +808,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 +1542,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..934e0c3ee 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -257,10 +257,166 @@ 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}), + ), + ]; + + 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); + } + } + + 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] +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/runtime_state_tests.rs b/crates/core/tests/unit/runtime_state_tests.rs new file mode 100644 index 000000000..09f01f279 --- /dev/null +++ b/crates/core/tests/unit/runtime_state_tests.rs @@ -0,0 +1,173 @@ +// 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 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") })); + 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 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(), + 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 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") })); + 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!(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") })); + let error = NemoRelayContextState::llm_conditional_execution_snapshot_chain( + &request, + &[RegistryRecord::new( + "llm-conditional-panic", + 0, + llm_conditional, + )], + &[], + None, + None, + ) + .await + .unwrap_err(); + 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") })); + 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!(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") })); + 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!(matches!( + error, + FlowError::Internal(ref message) if message.contains("llm-intercept-panic") + )); +} diff --git a/crates/core/tests/unit/shared_tests.rs b/crates/core/tests/unit/shared_tests.rs index 9d17ce188..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}; @@ -168,21 +168,27 @@ 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(); + 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| { - 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))) + Arc::new(move |_name, mut request, annotated| { + let callback_observed_without_codec = Arc::clone(&callback_observed_without_codec); + Box::pin(async move { + *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()); + Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + }) }), ) .unwrap(); @@ -196,7 +202,9 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { }, None, ) + .await .unwrap(); + assert_eq!(*observed_without_codec.lock().unwrap(), Some(true)); assert_eq!( request_without_codec.headers.get("x-no-codec"), Some(&json!(true)) @@ -210,15 +218,25 @@ 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| { - 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))) + Arc::new(move |_name, mut request, annotated| { + let callback_observed_with_codec = Arc::clone(&callback_observed_with_codec); + Box::pin(async move { + *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))) + }) }), ) .unwrap(); @@ -233,8 +251,13 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { }, Some(codec), ) + .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)) @@ -259,8 +282,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 +293,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 +307,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 +326,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 +341,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 +352,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 +396,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 +424,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 +462,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 +505,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/build.rs b/crates/ffi/build.rs index 013b53194..49a69b616 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -5,14 +5,318 @@ 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"); + 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() - { - bindings.write_to_file(format!("{crate_dir}/nemo_relay.h")); + .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)] +struct AsyncPrototype { + name: String, + parameters: Vec, +} + +/// 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", + ]; + + 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}"); + let contents = std::fs::read_to_string(&source_path) + .unwrap_or_else(|error| panic!("read {source_path}: {error}")); + expected.extend(parse_async_macro_invocations(&contents)); } + 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)); + for duplicates in declared.windows(2) { + assert_ne!( + duplicates[0].name, duplicates[1].name, + "ASYNC_REGISTRATIONS contains duplicate declaration for {}", + duplicates[0].name + ); + } + assert_eq!( + declared, expected, + "ASYNC_REGISTRATIONS must exactly match the macro-generated Rust FFI exports" + ); +} + +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, + ) -> u32;"# + ), + "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, + ) -> u32;"# + ), + "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), + "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('(')?; + let parameters = parameters.strip_suffix(");")?; + Some(AsyncPrototype { + name: name.to_owned(), + parameters: parameters.split(", ").map(str::to_owned).collect(), + }) +} + +fn parse_async_macro_invocations(source: &str) -> Vec { + const MACROS: &[(&str, bool)] = &[ + ("global_async_registration!(", false), + ("scope_async_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 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")); + 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, + }); + } + } + prototypes +} + +const ASYNC_REGISTRATIONS: &str = r#" +/* + * 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, + 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); +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, 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); +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, NemoRelayAsyncStreamInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +"#; diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 70c480b24..4d674832e 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -235,6 +235,26 @@ 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; + +/** + * 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; @@ -252,6 +272,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 +385,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 +433,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); @@ -434,6 +465,29 @@ typedef char *(*NemoRelayToolExecInterceptCb)(void *user_data, */ typedef char *(*NemoRelayToolExecCb)(void *user_data, const char *args_json); +/** + * 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); + +/** + * 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. * @@ -1258,9 +1312,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. + * 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); @@ -2535,6 +2589,110 @@ 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. + * + * 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, + const struct NemoRelayAsyncCompletion *completion); + +/** + * 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, + 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. + */ +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. @@ -3028,4 +3186,82 @@ 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. + * + * 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, + 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); +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, 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); +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, NemoRelayAsyncStreamInterceptCb 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..45524edf7 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,25 @@ unsafe fn register_global( .unwrap_or_else(|error| status_from_error(&error)) } +global_async_registration!( + nemo_relay_register_mark_sanitize_guardrail_async, + NemoRelayAsyncJsonCb, + core_registry_api::register_mark_sanitize_guardrail, + wrap_async_event_sanitize_fn +); +global_async_registration!( + nemo_relay_register_scope_sanitize_start_guardrail_async, + NemoRelayAsyncJsonCb, + core_registry_api::register_scope_sanitize_start_guardrail, + wrap_async_event_sanitize_fn +); +global_async_registration!( + nemo_relay_register_scope_sanitize_end_guardrail_async, + NemoRelayAsyncJsonCb, + core_registry_api::register_scope_sanitize_end_guardrail, + wrap_async_event_sanitize_fn +); + unsafe fn deregister_global(name: *const c_char, surface: Surface) -> NemoRelayStatus { clear_last_error(); let name = match c_str_to_string(name) { @@ -102,6 +122,25 @@ unsafe fn register_scope( .unwrap_or_else(|error| status_from_error(&error)) } +scope_async_registration!( + nemo_relay_scope_register_mark_sanitize_guardrail_async, + NemoRelayAsyncJsonCb, + core_registry_api::scope_register_mark_sanitize_guardrail, + wrap_async_event_sanitize_fn +); +scope_async_registration!( + nemo_relay_scope_register_scope_sanitize_start_guardrail_async, + NemoRelayAsyncJsonCb, + core_registry_api::scope_register_scope_sanitize_start_guardrail, + wrap_async_event_sanitize_fn +); +scope_async_registration!( + nemo_relay_scope_register_scope_sanitize_end_guardrail_async, + NemoRelayAsyncJsonCb, + core_registry_api::scope_register_scope_sanitize_end_guardrail, + wrap_async_event_sanitize_fn +); + 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..2d9e8c9e2 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -2,14 +2,57 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ + 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_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, + 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!( + nemo_relay_register_llm_sanitize_request_guardrail_async, + NemoRelayAsyncJsonCb, + core_registry_api::register_llm_sanitize_request_guardrail, + wrap_async_llm_sanitize_request_fn +); + +global_async_registration!( + nemo_relay_register_llm_execution_intercept_async, + NemoRelayAsyncInterceptCb, + core_registry_api::register_llm_execution_intercept, + wrap_async_llm_execution_intercept_fn +); +global_async_registration!( + nemo_relay_register_llm_stream_execution_intercept_async, + NemoRelayAsyncStreamInterceptCb, + core_registry_api::register_llm_stream_execution_intercept, + wrap_async_llm_stream_execution_intercept_fn +); +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 +); +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 +); +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 +); + // --------------------------------------------------------------------------- // LLM guardrail registrations // --------------------------------------------------------------------------- @@ -391,9 +434,9 @@ 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 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(); diff --git a/crates/ffi/src/api/mod.rs b/crates/ffi/src/api/mod.rs index 03103f58a..6629466cb 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::{ + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayAsyncStreamInterceptCb, 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, @@ -65,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; @@ -99,6 +168,17 @@ 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(), + )); + } + tokio_runtime().block_on(future) +} + // --------------------------------------------------------------------------- // Standalone middleware chains // --------------------------------------------------------------------------- @@ -141,7 +221,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 +259,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 +313,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 +476,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/api/scope_registry.rs b/crates/ffi/src/api/scope_registry.rs index 50efd644b..2cbec1572 100644 --- a/crates/ffi/src/api/scope_registry.rs +++ b/crates/ffi/src/api/scope_registry.rs @@ -2,15 +2,20 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ + 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_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, + 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 +31,76 @@ fn parse_scope_uuid(scope_uuid: *const c_char) -> Result; +/// 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, +} + +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. +/// +/// `invocation_json` is borrowed for the duration of the call. The completion +/// 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, + completion: *const NemoRelayAsyncCompletion, +) -> u32; + +/// Runtime-owned asynchronous `next` continuation for execution intercepts. +pub struct NemoRelayAsyncNext { + inner: AsyncNextInner, + runtime: tokio::runtime::Handle, + _callback_user_data: Option>, +} + +enum AsyncNextInner { + Tool(ToolExecutionNextFn), + Llm(LlmExecutionNextFn), + LlmStream(LlmStreamExecutionNextFn), +} + +/// 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, + next: *const NemoRelayAsyncNext, + 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 +/// 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, +); + +/// 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 +// 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>, +} + +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), + _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 + .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), + _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)) }; + 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. +/// +/// 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( + 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; + }; + 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(_) => return NemoRelayStatus::InvalidJson, + }; + let next = next.clone(); + Box::pin(async move { next(request).await }) + } + AsyncNextInner::LlmStream(_) => return NemoRelayStatus::InvalidArg, + }; + 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 + .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. +/// +/// 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( + 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(_) => return NemoRelayStatus::InvalidArg, + }; + 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_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_ptr(), ptr::null(), error.as_ptr()) }; + } + } + }); + 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)] +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. @@ -76,7 +804,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 +900,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; @@ -310,6 +1045,313 @@ 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: Arc, 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. +/// +/// 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, + 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 = 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}), + ) + .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. +/// +/// 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, + 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 = 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}), + ) + .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), + )) + }, + ) +} + +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: NemoRelayAsyncStreamInterceptCb, + 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 runtime = tokio::runtime::Handle::try_current().map_err(|error| { + FlowError::Internal(format!( + "async C stream intercept requires a Tokio runtime: {error}" + )) + })?; + 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, + })) + }) + }, + ) +} + // --------------------------------------------------------------------------- // Wrapper functions: C callback -> core trait objects // --------------------------------------------------------------------------- @@ -321,14 +1363,18 @@ 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 { + 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 sanitize callback returned null"); + unsafe { nemo_relay_string_free_internal(result_ptr) }; + result + }) }) } @@ -339,22 +1385,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 +1414,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 + }) }) } @@ -388,12 +1440,13 @@ 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) }; - 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 }) }) } @@ -444,12 +1497,14 @@ 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) }; 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}" @@ -515,13 +1570,14 @@ 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) }; 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 }) }, ) @@ -590,6 +1646,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)) }; @@ -597,8 +1654,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)) }) @@ -617,64 +1675,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 +1749,104 @@ 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 Err(error); + } + }; + 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 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)); + } + 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 Err(error); } }; 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 { 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; + unsafe { nemo_relay_string_free_internal(response_json) }; + return match last_error_message() { + Some(message) => Err(FlowError::Internal(message)), + None => Ok(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 = 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); + } } - } - result + result.map(Some) + }) }) } @@ -783,6 +1869,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, @@ -790,20 +1892,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 + }) }) } @@ -818,13 +1923,14 @@ 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) }; 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 }) }) } @@ -843,12 +1949,14 @@ 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) }; 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)); @@ -918,14 +2026,20 @@ 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: Arc, fields: EventSanitizeFields| { + let ud = ud.clone(); + Box::pin(async move { + 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)).map_err(|error| { + FlowError::Internal(format!("invalid event sanitizer result: {error}")) + }); + unsafe { nemo_relay_string_free_internal(result_ptr) }; + result + }) }) } @@ -1025,7 +2139,11 @@ 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_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}"))) } fn ptr_to_opt_string(ptr: *mut c_char) -> Option { @@ -1046,6 +2164,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/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index 140f089ce..618968006 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -661,6 +661,204 @@ 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 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!( + 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 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!( + 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 + ); + assert_eq!( + unsafe { nemo_relay_scope_stack_restore_thread(original_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()); + + 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] fn llm_request_accessors_round_trip() { let headers = cstring(r#"{"x-trace":"1"}"#); diff --git a/crates/ffi/tests/integration/callable_extra_tests.rs b/crates/ffi/tests/integration/callable_extra_tests.rs index b176f4988..2a284f07e 100644 --- a/crates/ffi/tests/integration/callable_extra_tests.rs +++ b/crates/ffi/tests/integration/callable_extra_tests.rs @@ -8,6 +8,8 @@ use std::ptr; use tokio_stream::StreamExt; +use super::test_support::resolve; + unsafe extern "C" fn tool_conditional_error_cb( _user_data: *mut libc::c_void, _name: *const c_char, @@ -172,7 +174,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 +245,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 +256,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 +264,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 +367,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 +388,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/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/api/coverage_sweeps_tests.rs b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs index b778b9a2c..977a92a7b 100644 --- a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs +++ b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs @@ -13,6 +13,329 @@ 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, +) -> u32 { + callable::NemoRelayAsyncCallbackState::Pending as u32 +} + +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, +) -> u32 { + 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()); + 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); + }}; + } + 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, + 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_stream_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 + ); + }}; + } + 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, + 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_stream_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) }; + unsafe { nemo_relay_scope_stack_free(stack) }; +} + 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 185cadd79..79cc9c296 100644 --- a/crates/ffi/tests/unit/api/registry_tests.rs +++ b/crates/ffi/tests/unit/api/registry_tests.rs @@ -243,6 +243,9 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { nemo_relay_deregister_mark_sanitize_guardrail(invalid_guard.as_ptr()), NemoRelayStatus::Ok ); + // 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); let mut owner = ptr::null_mut(); @@ -343,6 +346,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"); @@ -404,13 +409,14 @@ 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() .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/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index a2684dd2f..adddd6f4a 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -4,6 +4,124 @@ //! 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, + _invocation_json: *const c_char, + _completion: *const NemoRelayAsyncCompletion, +) -> 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, +) -> u32 { + let slot = unsafe { &*user_data.cast::() }; + slot.store(completion as usize, Ordering::Release); + 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( + 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); +} + +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() { @@ -17,3 +135,601 @@ 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), + _callback_user_data: None, + }); + 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), + _callback_user_data: None, + }); + 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 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 = + 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 + ); + 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 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_and_llm_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}), + ), + ]; + + for (inner, invocation, expected) in cases { + 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!( + 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); + } + } + + let next = Arc::new(NemoRelayAsyncNext { + inner: AsyncNextInner::Llm(Arc::new(|request| { + 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(); + 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] +fn async_next_callback_reports_tool_and_llm_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}), + ), + ]; + for (inner, invocation, expected) in cases { + 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::>(); + 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) }; + } + + let next = Arc::new(NemoRelayAsyncNext { + inner: AsyncNextInner::Tool(Arc::new(|_value| { + 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(); + 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) }; +} + +#[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 c58098b4a..7ec6f756a 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -11,11 +11,358 @@ 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); } +unsafe extern "C" fn async_json_passthrough_callback( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, +) -> u32 { + 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 => { + assert_eq!( + crate::api::nemo_relay_flush_subscribers(), + NemoRelayStatus::Ok, + "flush inside event publication must return without waiting" + ); + 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"); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, + NemoRelayStatus::Ok + ); + NemoRelayAsyncCallbackState::Complete as u32 +} + +unsafe extern "C" fn async_unfinished_stream_callback( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _next: *const NemoRelayAsyncNext, + _stream: *const NemoRelayAsyncStream, +) -> u32 { + 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_stream_finish(stream) }, + NemoRelayStatus::Ok + ); + NemoRelayAsyncCallbackState::Complete as u32 +} + +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, +) -> u32 { + 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); + } + // A successful invoke retained a completion reference; this callback drops + // its own next and completion references before transferring Pending ownership. + NemoRelayAsyncCallbackState::Pending as u32 +} + +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, +) -> u32 { + 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 as u32 +} + +#[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(Arc::new(event), fields.clone())).unwrap(), + fields + ); +} + +#[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_unfinished_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!("an unfinished immediate stream callback must fail"); + }; + assert!(error.to_string().contains("without finishing")); +} + +#[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_delivers_chunks_incrementally() { + let intercept = wrap_async_llm_stream_execution_intercept_fn( + async_immediate_stream_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!({"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; @@ -328,7 +675,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 +685,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,87 +758,96 @@ 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})) ); 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!( - malformed_response( - json!({"secret": "must be omitted"}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ), - 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); - 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 request sanitizer wrapper"); + assert!( + request_error + .to_string() + .contains("runtime codec ID contains an embedded NUL") ); assert!( last_error_message() @@ -501,12 +857,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 preserved"}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), + )) + .expect_err("an embedded runtime codec ID must fail the response sanitizer wrapper"); + assert!( + response_error + .to_string() + .contains("runtime codec ID contains an embedded NUL") ); assert!( last_error_message() @@ -542,7 +901,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 +1015,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(Arc::new(event.clone()), original_fields.clone())).unwrap(); assert_eq!(sanitized.data, Some(json!({"safe": true}))); assert_eq!( sanitized @@ -666,14 +1030,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!( - invalid(&event, original_fields.clone()), - 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!( - null(&event, original_fields.clone()), - 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/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/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/api/mod.rs b/crates/node/src/api/mod.rs index 24bc87b80..6f10cb2f6 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,11 @@ 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_event_sanitizer( + env, func, + )?); + Ok(callable::wrap_js_event_sanitize_promise_fn(callback)) } type NodeLlmCodec = ( @@ -1857,15 +1749,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 +2666,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 +2677,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 +2723,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 +2784,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 +2827,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])* @@ -2995,23 +2918,24 @@ 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, 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) } @@ -3028,23 +2952,24 @@ 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, 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 +2992,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 +3033,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) } @@ -3237,17 +3171,34 @@ 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. +/// 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`; callers that -/// need JS callback side effects should await an event-loop tick after this returns. -#[napi] -pub fn flush_subscribers() -> Result<()> { - core_subscriber_api::flush_subscribers().map_err(to_napi_err) +/// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this +/// 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. +#[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(), + ) } // --------------------------------------------------------------------------- @@ -3258,8 +3209,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( @@ -3268,7 +3221,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<()> { @@ -3326,8 +3279,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])* @@ -3395,7 +3354,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) } @@ -3434,9 +3397,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])* @@ -3531,7 +3504,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, @@ -3539,7 +3513,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<()> { @@ -3549,9 +3523,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) } @@ -3575,7 +3549,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, @@ -3583,7 +3558,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<()> { @@ -3593,9 +3568,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) } @@ -3633,7 +3608,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) } @@ -3670,19 +3649,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) } @@ -3858,7 +3841,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 }, @@ -3898,6 +3885,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..9d08f1ba0 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -207,6 +207,354 @@ 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(); + let publication = nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); + Box::pin(async move { + 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()); + }) + }) + }) +} + +/// 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(); + 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!( + "failed to serialize JS LLM sanitize request: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let context = js_llm_sanitize_request_context(&context); + 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 { + 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(); + let publication = nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); + Box::pin(async move { + let context = js_llm_sanitize_response_context(&context); + 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)) + }) + }) +} + +/// 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 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(request) + .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 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 { + 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: Arc, 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 +597,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 +624,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 +651,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 +713,66 @@ 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).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, + "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 +786,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, 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 +854,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 +1021,32 @@ 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).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, + 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 + }) }) } @@ -778,80 +1158,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: &Event, fields: CoreEventSanitizeFields| { - 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 CoreEventSanitizeFields::default(); - } - }; - 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| { - 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() - } - } - }) -} - // --------------------------------------------------------------------------- // Codec wrappers // --------------------------------------------------------------------------- diff --git a/crates/node/src/callback_factory.rs b/crates/node/src/callback_factory.rs index 891373dbf..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,16 +101,36 @@ 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, publication) { + if (error != null) { + let message = 'unknown error'; + try { + message = String(error?.message ?? error); + } catch {} + reject(message); + return; + } + 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) { - reject(error); + let message = 'unknown error'; + try { + message = String(error?.message ?? error); + } catch {} + reject(message); return; } - Promise.resolve().then(() => ( - next === undefined ? fn(arg0) : fn(arg0, next) - )).then((value) => jsonValue(value === undefined ? null : value)).then(resolve, reject); + callPromise(fn, arg0, spread, next, resolve, reject, true); }; }, + + eventSanitizerCallbackActive() { + return eventSanitizerContext.getStore()?.active === true; + }, }; })()"#; @@ -122,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 bb5435207..40f780bc0 100644 --- a/crates/node/src/promise_call.rs +++ b/crates/node/src/promise_call.rs @@ -53,7 +53,9 @@ enum PrimaryArg { struct CallArgs { arg0: PrimaryArg, + spread: bool, next: Option, + publication: bool, completion: CallCompletion, } @@ -76,19 +78,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 +157,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() })?; @@ -196,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)?, @@ -208,7 +208,19 @@ 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 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) })?; @@ -222,7 +234,24 @@ 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, false) + .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, false) + .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 @@ -232,14 +261,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), None).await + 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, 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), 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 @@ -249,8 +299,13 @@ impl PromiseAwareFn { args: Json, next: JsonStreamNextFn, ) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), 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. @@ -260,7 +315,13 @@ impl PromiseAwareFn { } } - async fn call_inner(&self, arg0: PrimaryArg, next: Option) -> FlowResult { + async fn call_inner( + &self, + arg0: PrimaryArg, + spread: bool, + next: Option, + publication: bool, + ) -> FlowResult { let (sender, receiver) = tokio::sync::oneshot::channel(); let tsfn = self .tsfn @@ -272,7 +333,9 @@ impl PromiseAwareFn { let status = tsfn.call( Ok(CallArgs { arg0, + spread, next, + publication, completion: CallCompletion::new(sender), }), napi::threadsafe_function::ThreadsafeFunctionCallMode::NonBlocking, 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 812a2a0ef..2380f2cbc 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) { @@ -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'); @@ -107,13 +107,139 @@ 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('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('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 = { scalar: () => 'invalid', emptyObject: () => ({}), array: () => [], - promise: () => Promise.resolve({ data: { changed: true } }), + promise: () => Promise.resolve([]), }; try { for (const [kind, sanitizer] of Object.entries(invalidResults)) { @@ -129,14 +255,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/); + assertSanitizerFieldsPreserved(events.at(-1), { kept: kind }); + assert.match(lib.getLastCallbackError(), /invalid JS event sanitizer result/); } } finally { lib.deregisterSubscriber('node-event-sanitize-invalid-sub'); @@ -151,7 +277,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'); @@ -163,12 +289,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)) { @@ -184,7 +310,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); @@ -193,7 +319,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 +327,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) => ({ @@ -215,12 +341,12 @@ 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', ); - assertSanitizerFieldsCleared(start); + assertSanitizerFieldsPreserved(start, { kept: true }); assert.match(lib.getLastCallbackError() ?? '', /background sanitizer boom/i); } finally { lib.deregisterScopeSanitizeStartGuardrail('node-background-throw-seed'); @@ -243,7 +369,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 +397,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(); @@ -289,7 +415,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, { @@ -312,9 +438,9 @@ 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)); + 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 3115b2136..c2bad301d 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'); @@ -51,7 +52,7 @@ function rejectWith(value) { } async function flushSubscriberCallbacks() { - flushSubscribers(); + await flushSubscribers(); for (let i = 0; i < 10; i += 1) { await new Promise((resolve) => setImmediate(resolve)); } @@ -291,7 +292,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 +390,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', ); @@ -560,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 = ` @@ -637,6 +679,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)); @@ -718,7 +787,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 +797,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 +814,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.equal(getLastCallbackError(), 'internal error: unknown error'); deregisterLlmSanitizeRequestGuardrail('node_llm_san_req_throw'); const result = await llmCallExecute( @@ -840,7 +918,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 +928,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 +944,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 +971,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 +1120,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'); @@ -1127,7 +1248,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); }); @@ -1146,14 +1267,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( @@ -1166,7 +1287,7 @@ describe('LLM intercepts', () => { null, null, ), - /unknown error/i, + /internal error: 42/i, ); }); @@ -1340,11 +1461,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 a4cb41aed..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,13 +30,6 @@ function rejectWithPrimitive(value) { return Promise.reject(value); } -async function flushSubscriberCallbacks() { - flushSubscribers(); - for (let i = 0; i < 10; i += 1) { - await new Promise((resolve) => setImmediate(resolve)); - } -} - // =========================================================================== // Scope operations // =========================================================================== @@ -104,7 +98,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 +202,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 +254,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', @@ -273,14 +267,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'); @@ -355,21 +349,18 @@ 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'); } }); - 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 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'); @@ -384,7 +375,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'); @@ -407,7 +398,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/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 125475561..4576d8253 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,21 +48,6 @@ function sparseArray() { return values; } -async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { - 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()) { - if (Date.now() >= deadline) { - throw new Error('timed out waiting for subscriber callbacks'); - } - await new Promise((resolve) => setImmediate(resolve)); - } -} - // =========================================================================== // Tool lifecycle // =========================================================================== @@ -608,6 +594,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 { @@ -696,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'); @@ -782,6 +829,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'); @@ -934,7 +1016,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( @@ -948,7 +1030,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 414fb13f6..8852f8543 100644 --- a/crates/pii-redaction/src/builtin.rs +++ b/crates/pii-redaction/src/builtin.rs @@ -461,12 +461,16 @@ 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), - }, - ) + let backend = Arc::new(backend); + Arc::new(move |_name: String, payload: Json| { + let backend = Arc::clone(&backend); + 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 { @@ -485,133 +489,145 @@ 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| { - if scope_categories.is_some_and(|(sanitize_llm, sanitize_tool)| { - matches!(event, Event::Scope(_)) + let backend = Arc::clone(&backend); + Box::pin(async move { + if scope_categories.is_some_and(|(sanitize_llm, sanitize_tool)| { + matches!(event.as_ref(), 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.as_ref(), 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) + }) }) } pub(super) fn llm_sanitize_request_callback( backend: CompiledBuiltinBackend, ) -> LlmSanitizeRequestFn { + let backend = Arc::new(backend); 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 = Arc::clone(&backend); + 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) + }) }) } pub(super) fn llm_sanitize_response_callback( backend: CompiledBuiltinBackend, ) -> LlmSanitizeResponseFn { + let backend = Arc::new(backend); 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 = Arc::clone(&backend); + 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..2e80f1172 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, + Arc::new(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, + Arc::new(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, + Arc::new(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, + Arc::new(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, + Arc::new(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, + Arc::new(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,16 @@ 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(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); + let sanitized = redact(Arc::new(event), fields).await.unwrap(); assert_eq!( sanitized.data.unwrap(), json!({ @@ -1016,8 +1067,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 +1119,7 @@ fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations None, )); let sanitized = callback( - &event, + Arc::new(event), EventSanitizeFields { data: Some(json!({"already": "sanitized by the response callback"})), category_profile: Some( @@ -1079,7 +1130,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 +1169,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 +1194,8 @@ fn preserved_custom_marks_remain_eligible_for_a_later_email_profile() { .unwrap(), ); - let sanitized = email(&event, trajectory(&event, fields)); + 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!( @@ -1365,7 +1419,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 +1488,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 +1507,7 @@ fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() None, )); let sanitized = callback( - &event, + Arc::new(event), EventSanitizeFields { data: Some(json!({"email": "person@example.com"})), category_profile: Some( @@ -1459,7 +1517,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 +1528,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 +1553,15 @@ fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() .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()), metadata: Some(json!({"owner": "person@example.com"})), }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.data.unwrap()["content"], @@ -1511,8 +1573,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 +1609,15 @@ 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()), 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 +1626,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 +1645,7 @@ fn event_sanitizer_discards_category_profile_when_sanitization_fails() { None, )); let sanitized = callback( - &event, + Arc::new(event), EventSanitizeFields { data: None, category_profile: Some(CategoryProfile { @@ -1593,7 +1657,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..ecafb951f 100644 --- a/crates/plugin/README.md +++ b/crates/plugin/README.md @@ -42,8 +42,13 @@ 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. 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. ## Installation diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index 727e09aa2..1e08c05f7 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,172 @@ 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, +} + +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)] +pub enum NemoRelayNativeAsyncCallbackState { + /// The callback settled its completion before returning. + Complete = 0, + /// The callback retained its completion for later settlement. + 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 { + _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`] 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`]. +/// +/// 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. + /// + /// `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: u32, + 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 +2405,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 as u32, + 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..a56b2fdc6 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -16,22 +16,55 @@ 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, - 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}; +#[test] +fn async_abi_discriminants_reject_unknown_values() { + 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), + Ok(NemoRelayNativeAsyncCallbackState::Pending) + ); + assert!(NemoRelayNativeAsyncCallbackState::try_from(2).is_err()); +} + struct TestString(Vec); struct RegisteredSubscriber { @@ -300,8 +333,8 @@ static LLM_REQUEST_INTERCEPT_REGISTRATION: Mutex(), test_host().struct_size @@ -328,6 +361,12 @@ fn native_abi_v2_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 +387,12 @@ fn native_abi_v2_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 +402,22 @@ fn native_abi_v2_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), diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index f5c2ccd51..53f48f33e 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1313,12 +1313,45 @@ 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 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)); + } + 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 +1362,41 @@ 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(); + 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(); + 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 +1409,46 @@ 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 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, + 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 +1458,39 @@ 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(); + 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(); + 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 + }) } // --------------------------------------------------------------------------- @@ -1394,9 +1524,9 @@ 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. +/// 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/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 57fffd5d5..279040388 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -32,14 +32,16 @@ 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; use serde_json::Value as Json; 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; @@ -53,6 +55,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 +98,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,14 +131,54 @@ 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)) + .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( + 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: 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(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 { + 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> { @@ -408,25 +468,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 { - 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(); + 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 { + 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 = 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)) } - }; - 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() + .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?; + Python::attach(|py| { + py_to_json(result.bind(py)) + .map_err(|e| FlowError::Internal(format!("tool py_to_json failed: {e}"))) }) }) }) @@ -434,45 +499,50 @@ pub fn wrap_py_tool_fn(py_fn: Py) -> ToolSanitizeFn { /// 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 +885,45 @@ 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| { - 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; + 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 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) } - }; - 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 + .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| { + 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 +931,31 @@ 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); + 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_with_locals(py, result, task_locals.as_ref()) + })) + .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 +967,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 +1109,37 @@ 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| { - 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; + 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 = 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)) } - }; - 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 + .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| { + 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 +1183,68 @@ 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); + let task_locals = capture_python_task_locals(); + Arc::new(move |event: Arc, fields: EventSanitizeFields| { + let py_fn = py_fn.clone(); + let task_locals = capture_python_task_locals().or_else(|| task_locals.clone()); + Box::pin(async move { + let result = Python::attach( + |py| -> FlowResult, PyValueFuture>> { + let py_event = match event.as_ref() { + 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) => { + return Err(FlowError::Internal(error.to_string())); + } + }; + let fields_json = match serde_json::to_value(&fields) { + Ok(value) => value, + Err(error) => { + return Err(FlowError::Internal(error.to_string())); + } + }; + let py_fields = match json_to_py(py, &fields_json) { + Ok(value) => value, + Err(error) => { + return Err(FlowError::Internal(error.to_string())); + } + }; + 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.unbind(), task_locals.as_ref()) + }, + ); + 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..7eab65355 100644 --- a/crates/python/tests/coverage/coverage_tests.rs +++ b/crates/python/tests/coverage/coverage_tests.rs @@ -661,51 +661,72 @@ 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})); + let error = runtime + .block_on(tool_fail("demo".to_string(), json!({"x": 1}))) + .unwrap_err(); + assert!( + error.to_string().contains("tool boom"), + "unexpected tool error: {error}" + ); 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( + let error = runtime + .block_on(llm_sanitize( request.clone(), nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ), - None + )) + .unwrap_err(); + assert!( + 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()); 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") @@ -713,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!( - tool_req("demo", json!({"x": 1})) - .unwrap_err() - .to_string() - .contains("Python tool callable failed") + 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(); - assert_eq!( - llm_resp( + let error = runtime + .block_on(llm_resp( json!({"ok": true}), nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ), - None + )) + .unwrap_err(); + assert!( + error.to_string().contains("resp boom"), + "unexpected LLM response sanitizer error: {error}" ); 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..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", @@ -472,18 +485,51 @@ 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 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 { @@ -492,7 +538,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,20 +550,38 @@ 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") ); 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 fe815cc1b..d8ece09e7 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,77 @@ fn load_module<'py>(py: Python<'py>, code: &str) -> Bound<'py, PyModule> { PyModule::from_code(py, &code, &file_name, &module_name).unwrap() } +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" + )) + .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(); + 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 { LlmRequest { headers: serde_json::Map::new(), @@ -153,7 +224,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 +240,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 +255,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 +276,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 { @@ -651,20 +739,28 @@ 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(); Python::attach(|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") @@ -683,23 +779,122 @@ 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(), + )(Arc::new(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(), + 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(), + )(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(), + )(Arc::new(event), fields.clone())) + .unwrap_err(); + assert!( + invalid + .to_string() + .contains("invalid event sanitizer result") + ); + }); +} + +#[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") +"#, ); - assert_eq!(raised, EventSanitizeFields::default()); + 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()); - let invalid = wrap_py_event_sanitize_fn(module.getattr("invalid").unwrap().unbind())( - &event, - fields.clone(), + 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(); + }); + }); +} + +#[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() +"#, ); - assert_eq!(invalid, EventSanitizeFields::default()); + ( + 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/docs/about-nemo-relay/concepts/middleware.mdx b/docs/about-nemo-relay/concepts/middleware.mdx index 0afcd2d56..f74a3bb89 100644 --- a/docs/about-nemo-relay/concepts/middleware.mdx +++ b/docs/about-nemo-relay/concepts/middleware.mdx @@ -19,6 +19,23 @@ 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, 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 +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 +142,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 @@ -298,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/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/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 8ded00de5..9ebeb5eae 100644 --- a/docs/reference/event-sanitizers.mdx +++ b/docs/reference/event-sanitizers.mdx @@ -49,10 +49,26 @@ 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()`. +## Async Delivery and Ordering + +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. + +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 Where you register a sanitizer determines how long it stays active. @@ -119,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) + }) }), )?; @@ -174,17 +192,28 @@ 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. +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 +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..4fc5529e0 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,75 @@ 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. +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. +### 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 | + +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: + +```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, 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 The registration names remain unchanged for global, plugin-context, and @@ -49,9 +107,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)) + }) }), )?; ``` @@ -115,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. @@ -150,15 +211,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 +228,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 +283,16 @@ 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. 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 @@ -243,7 +309,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/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/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go new file mode 100644 index 000000000..8e3086054 --- /dev/null +++ b/go/nemo_relay/async_middleware_test.go @@ -0,0 +1,550 @@ +// 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" + "io" + "strings" + "sync" + "testing" + "time" +) + +func asyncMiddlewareNoop(context.Context, json.RawMessage) (any, error) { + return nil, nil +} + +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 + 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, asyncStreamExecutionNoop) + }, 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) + } + t.Cleanup(func() { _ = registration.deregister(name) }) + 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) + } + }) + } +} + +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, asyncStreamExecutionNoop) + }, 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.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) + } + }) +} + +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 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" + 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) + } + }) +} + +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" + 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 + 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 + }() + defer func() { + release() + select { + case <-nextDone: + case <-time.After(time.Second): + t.Error("detached next continuation never settled during cleanup") + } + }() + + 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): + t.Fatal("detached next context prevented intercept cleanup") + } + 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 8096c7447..edbff52cf 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -50,6 +50,28 @@ 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 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); @@ -87,10 +109,13 @@ typedef NemoRelayCodecEncodeCb NemoRelayCodecEncodeFn; import "C" import ( + "context" "encoding/json" "errors" + "fmt" "sync" "sync/atomic" + "time" "unsafe" ) @@ -102,6 +127,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)))) @@ -119,10 +145,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()) @@ -130,27 +152,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) + } } // --------------------------------------------------------------------------- @@ -167,6 +201,125 @@ 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. +// +// 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. 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. +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 + +type completionCancellationWatch struct { + completion *C.NemoRelayAsyncCompletion + cancel context.CancelFunc + probeMu *sync.Mutex +} + +var ( + completionCancellationMu sync.Mutex + completionCancellationWatches = make(map[uint64]completionCancellationWatch) + completionCancellationNextID atomic.Uint64 + completionCancellationRunning bool +) + +func startCompletionCancellationMonitor() { + 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() + _, 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() + } + } + watch.probeMu.Unlock() + } + } + }() +} + +func contextForCompletion(completion *C.NemoRelayAsyncCompletion) (context.Context, func()) { + ctx, cancel := context.WithCancel(context.Background()) + 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() + }) + } +} + // 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 +743,392 @@ 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) + defer rejectAsyncCallbackPanic(completion, "middleware") + 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 +} + +func rejectAsyncCallbackPanic(completion *C.NemoRelayAsyncCompletion, kind string) { + recovered := recover() + if recovered == nil { + return + } + 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) +} + +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 { + unregisterClosure(userData) + 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) + defer rejectAsyncCallbackPanic(completion, "execution intercept") + ctx, cancel := contextForCompletion(completion) + defer cancel() + 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() + C.nemo_relay_async_next_release(next) + }() + outerCtx := ctx + nextFn := func(nextCtx 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 { + // 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. + // 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 + case <-nextCtx.Done(): + return nil, nextCtx.Err() + case <-outerCtx.Done(): + return nil, outerCtx.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 +} + +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/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()) } }) } 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++ diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index 978e40228..eb5150f9c 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -44,6 +44,12 @@ 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 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); @@ -121,46 +127,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, NemoRelayAsyncStreamInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_stream_execution_intercept(const char* name); // Subscribers @@ -170,46 +187,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, 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 @@ -267,6 +301,9 @@ 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 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*); @@ -1198,11 +1235,31 @@ func registerEventSanitizer(name string, priority int32, fn EventSanitizeFunc, k return checkStatus(status) } +func registerAsyncEventSanitizer(name string, priority int32, fn AsyncMiddlewareFunc, kind int) error { + 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. 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,13 @@ func RegisterToolSanitizeRequestGuardrail(name string, priority int32, fn ToolSa )) } +// RegisterToolSanitizeRequestGuardrailAsync registers an asynchronous tool request sanitizer. +func RegisterToolSanitizeRequestGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) 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_tool_sanitize_request_guardrail_async(name, 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 +1351,13 @@ func RegisterToolSanitizeResponseGuardrail(name string, priority int32, fn ToolS )) } +// RegisterToolSanitizeResponseGuardrailAsync registers an asynchronous tool response sanitizer. +func RegisterToolSanitizeResponseGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) 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_tool_sanitize_response_guardrail_async(name, 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 +1385,13 @@ func RegisterToolConditionalExecutionGuardrail(name string, priority int32, fn T )) } +// RegisterToolConditionalExecutionGuardrailAsync registers an asynchronous tool guardrail. +func RegisterToolConditionalExecutionGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) 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_tool_conditional_execution_guardrail_async(name, 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 +1418,13 @@ 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 { + 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 // intercept by name. func DeregisterToolRequestIntercept(name string) error { @@ -1354,6 +1449,13 @@ func RegisterToolExecutionIntercept(name string, priority int32, execFn ToolExec )) } +// RegisterToolExecutionInterceptAsync registers an asynchronous tool execution intercept. +func RegisterToolExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) 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_tool_execution_intercept_async(name, 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 +1482,13 @@ func RegisterLlmSanitizeRequestGuardrail(name string, priority int32, fn LLMRequ )) } +// RegisterLlmSanitizeRequestGuardrailAsync registers an asynchronous LLM request sanitizer. +func RegisterLlmSanitizeRequestGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) 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_sanitize_request_guardrail_async(name, 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 +1511,13 @@ func RegisterLlmSanitizeResponseGuardrail(name string, priority int32, fn LLMRes )) } +// RegisterLlmSanitizeResponseGuardrailAsync registers an asynchronous LLM response sanitizer. +func RegisterLlmSanitizeResponseGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) 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_sanitize_response_guardrail_async(name, 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 +1544,13 @@ func RegisterLlmConditionalExecutionGuardrail(name string, priority int32, fn LL )) } +// RegisterLlmConditionalExecutionGuardrailAsync registers an asynchronous LLM guardrail. +func RegisterLlmConditionalExecutionGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) 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_conditional_execution_guardrail_async(name, 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 +1577,13 @@ 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 { + 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 // intercept by name. func DeregisterLlmRequestIntercept(name string) error { @@ -1478,6 +1608,13 @@ func RegisterLlmExecutionIntercept(name string, priority int32, execFn LLMExecut )) } +// RegisterLlmExecutionInterceptAsync registers an asynchronous LLM execution intercept. +func RegisterLlmExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) 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_execution_intercept_async(name, 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 +1640,13 @@ func RegisterLlmStreamExecutionIntercept(name string, priority int32, execFn LLM )) } +// RegisterLlmStreamExecutionInterceptAsync registers an asynchronous streaming LLM intercept. +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.NemoRelayAsyncStreamInterceptCb(C.goAsyncStreamExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + // DeregisterLlmStreamExecutionIntercept removes a previously registered LLM // stream execution intercept by name. func DeregisterLlmStreamExecutionIntercept(name string) error { @@ -1545,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()) } @@ -2166,6 +2311,86 @@ func (s *OpenTelemetrySubscriber) Close() { // Scope-local guardrail/intercept registration (Tool) // --------------------------------------------------------------------------- +type asyncMiddlewareCallback interface { + 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 { + 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)) +} + +// 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 +2589,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 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.NemoRelayAsyncStreamInterceptCb(C.goAsyncStreamExecutionInterceptTrampoline), 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/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}`, 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/_event_sanitizer_context.py b/python/nemo_relay/_event_sanitizer_context.py new file mode 100644 index 000000000..7991ebf4a --- /dev/null +++ b/python/nemo_relay/_event_sanitizer_context.py @@ -0,0 +1,43 @@ +# 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) + + +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) + try: + result = callback(*args) + finally: + _ACTIVE.reset(token) + if inspect.isawaitable(result): + return _await_result(result) + return result 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/nemo_relay/subscribers.py b/python/nemo_relay/subscribers.py index b54b42e2e..aeabaf557 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,13 @@ 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 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 return _native_flush() 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..c0dc3c78e 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 @@ -56,7 +57,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,10 +70,63 @@ 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 +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} + + +@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 diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index 9a3662be4..22bc1aa18 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 @@ -167,6 +168,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") @@ -177,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 @@ -275,7 +312,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 +321,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 +335,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 +347,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 +361,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 +383,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 +405,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 +441,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 +453,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") @@ -658,6 +701,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) -> dict: + del context + await asyncio.sleep(0) + assert asyncio.get_running_loop() is originating_loop + return {"sanitized": response["raw"]} + + async def stream_func(request) -> AsyncIterator[dict]: + 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() diff --git a/python/tests/test_tools.py b/python/tests/test_tools.py index 3619902f3..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") @@ -361,7 +404,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 +417,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 +528,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 +540,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")