diff --git a/crates/adaptive/src/acg_component.rs b/crates/adaptive/src/acg_component.rs index 443e376ff..c71ec811d 100644 --- a/crates/adaptive/src/acg_component.rs +++ b/crates/adaptive/src/acg_component.rs @@ -582,26 +582,32 @@ pub(crate) fn create_acg_llm_request_intercept( provider: String, plugin: Arc, ) -> LlmRequestInterceptFn { - Arc::new(move |_name: &str, request: LlmRequest, annotated| { - let input_content = request.content.clone(); - let translated = - translate_request(&request, &agent_id, &provider, plugin.as_ref(), &hot_cache) - .unwrap_or(request); - if annotated.is_some() && translated.content != input_content { - let translated_annotated = build_semantic_request_view(&translated) - .map_err(|error| nemo_relay::error::FlowError::Internal(error.to_string()))? - .annotated_request; - return Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - LlmRequest { - headers: translated.headers, - content: input_content, - }, - Some(translated_annotated), - )); - } - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - translated, annotated, - )) + Arc::new(move |_name: String, request: LlmRequest, annotated| { + let hot_cache = hot_cache.clone(); + let agent_id = agent_id.clone(); + let provider = provider.clone(); + let plugin = plugin.clone(); + Box::pin(async move { + let input_content = request.content.clone(); + let translated = + translate_request(&request, &agent_id, &provider, plugin.as_ref(), &hot_cache) + .unwrap_or(request); + if annotated.is_some() && translated.content != input_content { + let translated_annotated = build_semantic_request_view(&translated) + .map_err(|error| nemo_relay::error::FlowError::Internal(error.to_string()))? + .annotated_request; + return Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + LlmRequest { + headers: translated.headers, + content: input_content, + }, + Some(translated_annotated), + )); + } + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + translated, annotated, + )) + }) }) } diff --git a/crates/adaptive/src/adaptive_hints_intercept.rs b/crates/adaptive/src/adaptive_hints_intercept.rs index c9f245505..b4b649264 100644 --- a/crates/adaptive/src/adaptive_hints_intercept.rs +++ b/crates/adaptive/src/adaptive_hints_intercept.rs @@ -174,31 +174,34 @@ impl AdaptiveHintsIntercept { pub fn into_request_fn(self) -> LlmRequestInterceptFn { let this = Arc::new(self); Arc::new( - move |_name: &str, + move |_name: String, mut request: LlmRequest, mut annotated: Option| { - let scope_path = extract_scope_path(); - let manual_ls = read_manual_latency_sensitivity(); - let scope_depth = scope_path.len(); - let call_index = this.call_counter.fetch_add(1, Ordering::Relaxed); - - let effective_agent_id = this.effective_agent_id(); - let cached_hints = - this.load_hints(&scope_path, &effective_agent_id, call_index, scope_depth); - let final_hints = apply_manual_latency_override( - cached_hints, - manual_ls, - &effective_agent_id, - scope_depth, - ); - - if let Some(hints) = final_hints { - inject_agent_hints(&mut request, &mut annotated, &hints); - } - - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + let this = this.clone(); + Box::pin(async move { + let scope_path = extract_scope_path(); + let manual_ls = read_manual_latency_sensitivity(); + let scope_depth = scope_path.len(); + let call_index = this.call_counter.fetch_add(1, Ordering::Relaxed); + + let effective_agent_id = this.effective_agent_id(); + let cached_hints = + this.load_hints(&scope_path, &effective_agent_id, call_index, scope_depth); + let final_hints = apply_manual_latency_override( + cached_hints, + manual_ls, + &effective_agent_id, + scope_depth, + ); + + if let Some(hints) = final_hints { + inject_agent_hints(&mut request, &mut annotated, &hints); + } + + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }, ) } diff --git a/crates/adaptive/src/lib.rs b/crates/adaptive/src/lib.rs index ebe78d534..74ec930b7 100644 --- a/crates/adaptive/src/lib.rs +++ b/crates/adaptive/src/lib.rs @@ -14,6 +14,10 @@ pub(crate) static TEST_GLOBAL_CONTEXT_MUTEX: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); +#[cfg(test)] +#[path = "../tests/support/mod.rs"] +pub(crate) mod test_support; + pub mod acg; pub mod acg_component; pub mod acg_learner; diff --git a/crates/adaptive/tests/integration/runtime_integration_tests.rs b/crates/adaptive/tests/integration/runtime_integration_tests.rs index dd65079ea..f0e64fb1a 100644 --- a/crates/adaptive/tests/integration/runtime_integration_tests.rs +++ b/crates/adaptive/tests/integration/runtime_integration_tests.rs @@ -605,6 +605,7 @@ async fn test_adaptive_plugin_registers_and_passes_calls_through() { content: json!({"messages": []}), }, ) + .await .unwrap(); assert_eq!(request.request.content["messages"], json!([])); @@ -739,9 +740,11 @@ impl Plugin for HeaderPlugin { false, Arc::new(|_name, mut request, annotated| { request.headers.insert("x-plugin".into(), json!("set")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + Box::pin(async move { + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }), )?; ctx.register_tool_request_intercept( @@ -752,7 +755,7 @@ impl Plugin for HeaderPlugin { if let Json::Object(ref mut map) = args { map.insert("x-tool-plugin".into(), json!(true)); } - Ok(args) + Box::pin(async move { Ok(args) }) }), )?; ctx.register_llm_execution_intercept( @@ -823,6 +826,7 @@ async fn test_top_level_plugin_registers_request_and_execution_intercepts() { content: json!({"messages": []}), }, ) + .await .unwrap(); assert_eq!(request.request.headers.get("x-plugin"), Some(&json!("set"))); diff --git a/crates/adaptive/tests/support/mod.rs b/crates/adaptive/tests/support/mod.rs new file mode 100644 index 000000000..c9f4775d0 --- /dev/null +++ b/crates/adaptive/tests/support/mod.rs @@ -0,0 +1,12 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::future::Future; + +pub(crate) fn block_on(future: F) -> F::Output { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .expect("test runtime should build") + .block_on(future) +} diff --git a/crates/adaptive/tests/unit/acg_component_tests.rs b/crates/adaptive/tests/unit/acg_component_tests.rs index 51477bb3b..69004ee72 100644 --- a/crates/adaptive/tests/unit/acg_component_tests.rs +++ b/crates/adaptive/tests/unit/acg_component_tests.rs @@ -1087,11 +1087,11 @@ fn acg_component_request_intercept_passes_original_request_and_annotation_when_t plugin, ); - let outcome = intercept( - "anthropic", + let outcome = crate::test_support::block_on(intercept( + "anthropic".to_string(), invalid_request.clone(), Some(annotated.clone()), - ) + )) .expect("request intercept should pass through"); let translated = outcome.request; let returned_annotated = outcome.annotated_request; @@ -1335,7 +1335,12 @@ fn acg_component_request_intercept_rewrites_annotation_without_mutating_provider plugin, ); - let outcome = intercept("anthropic", request, Some(original_annotation.clone())).unwrap(); + let outcome = crate::test_support::block_on(intercept( + "anthropic".to_string(), + request, + Some(original_annotation.clone()), + )) + .unwrap(); assert_eq!(outcome.request.content, original_content); let annotation = outcome diff --git a/crates/adaptive/tests/unit/adaptive_hints_intercept_tests.rs b/crates/adaptive/tests/unit/adaptive_hints_intercept_tests.rs index e3ef485b5..19296c842 100644 --- a/crates/adaptive/tests/unit/adaptive_hints_intercept_tests.rs +++ b/crates/adaptive/tests/unit/adaptive_hints_intercept_tests.rs @@ -4,6 +4,7 @@ //! Unit tests for adaptive hints intercept in the NeMo Relay adaptive crate. use super::*; + use std::sync::{Mutex, OnceLock}; use crate::trie::data_models::{LlmCallPrediction, PredictionMetrics}; @@ -196,14 +197,14 @@ fn test_adaptive_hints_intercept_injects_prediction_hints_and_manual_override() stream: None, extra: serde_json::Map::new(), }; - let outcome = req_fn( - "model", + let outcome = crate::test_support::block_on(req_fn( + "model".to_string(), LlmRequest { headers: serde_json::Map::new(), content: serde_json::json!({}), }, Some(annotated.clone()), - ) + )) .unwrap(); let request = outcome.request; let returned_annotated = outcome.annotated_request; @@ -266,14 +267,14 @@ fn test_adaptive_hints_intercept_uses_defaults_and_ignores_poisoned_cache() { })); let req_fn = AdaptiveHintsIntercept::new(hot_cache, "fallback-agent".to_string()).into_request_fn(); - let outcome = req_fn( - "model", + let outcome = crate::test_support::block_on(req_fn( + "model".to_string(), LlmRequest { headers: serde_json::Map::new(), content: serde_json::json!({}), }, None, - ) + )) .unwrap(); let request = outcome.request; let annotated = outcome.annotated_request; @@ -305,14 +306,14 @@ fn test_adaptive_hints_intercept_uses_defaults_and_ignores_poisoned_cache() { }); let poisoned_req_fn = AdaptiveHintsIntercept::new(poisoned_cache, "fallback-agent".to_string()).into_request_fn(); - let poisoned_outcome = poisoned_req_fn( - "model", + let poisoned_outcome = crate::test_support::block_on(poisoned_req_fn( + "model".to_string(), LlmRequest { headers: serde_json::Map::new(), content: serde_json::json!({"existing": true}), }, None, - ) + )) .unwrap(); let poisoned_request = poisoned_outcome.request; assert!( diff --git a/crates/adaptive/tests/unit/plugin_component_tests.rs b/crates/adaptive/tests/unit/plugin_component_tests.rs index 044504d19..5a49feb25 100644 --- a/crates/adaptive/tests/unit/plugin_component_tests.rs +++ b/crates/adaptive/tests/unit/plugin_component_tests.rs @@ -364,6 +364,7 @@ async fn adaptive_plugin_registers_runtime_and_rolls_back_registration() { content: json!({}), }, ) + .await .unwrap(); assert!(request.request.headers.is_empty()); diff --git a/crates/adaptive/tests/unit/runtime_features_tests.rs b/crates/adaptive/tests/unit/runtime_features_tests.rs index 5a53c7734..efdb60541 100644 --- a/crates/adaptive/tests/unit/runtime_features_tests.rs +++ b/crates/adaptive/tests/unit/runtime_features_tests.rs @@ -139,9 +139,11 @@ fn assert_llm_request_intercept_registered(name: &str) { i32::MAX, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + Box::pin(async move { + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }), ), name, @@ -154,9 +156,11 @@ fn assert_llm_request_intercept_absent(name: &str) { i32::MAX, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + Box::pin(async move { + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }), ) .unwrap(); @@ -565,6 +569,7 @@ async fn adaptive_hints_feature_registers_request_intercept() { content: json!({}), }, ) + .await .unwrap(); assert!(request.request.headers.contains_key(AGENT_HINTS_HEADER_KEY)); @@ -730,9 +735,11 @@ async fn registration_context_registers_all_supported_callback_types() { 5, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( - request, annotated, - )) + Box::pin(async move { + Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + request, annotated, + )) + }) }), ) .unwrap(); diff --git a/crates/adaptive/tests/unit/runtime_tests.rs b/crates/adaptive/tests/unit/runtime_tests.rs index f14d83fb0..cfad601bd 100644 --- a/crates/adaptive/tests/unit/runtime_tests.rs +++ b/crates/adaptive/tests/unit/runtime_tests.rs @@ -637,6 +637,7 @@ async fn adaptive_runtime_bind_scope_requires_registration_and_passes_through_wi }; let translated = llm_request_intercepts("anthropic", request.clone()) + .await .expect("request intercept chain should pass through when no hot-cache state exists"); assert_eq!(translated.request.content, request.content); diff --git a/crates/cli/src/sessions/mod.rs b/crates/cli/src/sessions/mod.rs index 732f9c31e..5f5510966 100644 --- a/crates/cli/src/sessions/mod.rs +++ b/crates/cli/src/sessions/mod.rs @@ -582,6 +582,10 @@ impl SessionManager { .map_err(CliError::from) }) .await?; + // Manual lifecycle events publish on the serial dispatcher. This + // test-only seam returns after the matching end event is observable so + // a subsequent synthetic provider call cannot overtake it. + nemo_relay::api::subscriber::flush_subscribers().map_err(CliError::from)?; let mut sessions = self.inner.lock().await; if let Some(session) = sessions.get_mut(&session_id) { session.record_completed_llm_response(response_for_hints, owner_subagent_id); @@ -784,9 +788,9 @@ impl Session { NormalizedEvent::SubagentStarted(event) => self.start_subagent(event).await, NormalizedEvent::SubagentEnded(event) => self.end_subagent(event).await, NormalizedEvent::LlmHint(event) => self.add_llm_hint(event), - NormalizedEvent::LlmStarted(event) => self.start_hook_llm(event), - NormalizedEvent::LlmEnded(event) => self.end_hook_llm(event), - NormalizedEvent::ToolStarted(event) => self.start_tool(event), + NormalizedEvent::LlmStarted(event) => self.start_hook_llm(event).await, + NormalizedEvent::LlmEnded(event) => self.end_hook_llm(event).await, + NormalizedEvent::ToolStarted(event) => self.start_tool(event).await, NormalizedEvent::ToolEnded(event) => self.end_tool(event).await, NormalizedEvent::PromptSubmitted(event) => self.start_turn(event).await, NormalizedEvent::Compaction(event) => self.mark("compaction", event), @@ -1142,8 +1146,8 @@ impl Session { if self.turn_scope.is_none() { return Ok(Vec::new()); } - self.close_active_llms(reason)?; - self.close_active_tools(reason)?; + self.close_active_llms(reason).await?; + self.close_active_tools(reason).await?; let closed_subagents = self.close_active_subagents(reason).await?; let output = self.last_turn_llm_output.take().unwrap_or(output); self.clear_correlation_state(); @@ -1182,7 +1186,7 @@ impl Session { } // Ends all active hook-observed LLM calls before closing their containing scopes. - fn close_active_llms(&mut self, reason: &str) -> Result<(), CliError> { + async fn close_active_llms(&mut self, reason: &str) -> Result<(), CliError> { let active_llms: Vec<_> = self.llms.drain().map(|(_, handle)| handle).collect(); for handle in active_llms { llm_call_end( @@ -1198,7 +1202,7 @@ impl Session { // Ends all active tool calls with a synthetic close result before ending their containing scopes. // Draining first avoids holding mutable map state while the runtime emits lifecycle events. - fn close_active_tools(&mut self, reason: &str) -> Result<(), CliError> { + async fn close_active_tools(&mut self, reason: &str) -> Result<(), CliError> { let active_tools: Vec<_> = self .tools .drain() @@ -1428,7 +1432,7 @@ impl Session { // ignored so repeated pre hooks do not create parallel handles for one provider call. Aliased // child-session LLMs carry their subagent owner in metadata and are resolved by // `hook_llm_owner`. - fn start_hook_llm(&mut self, event: LlmEvent) -> Result<(), CliError> { + async fn start_hook_llm(&mut self, event: LlmEvent) -> Result<(), CliError> { self.ensure_turn_started(event.metadata.clone())?; if self.llms.contains_key(&event.api_call_id) { return Ok(()); @@ -1454,7 +1458,7 @@ impl Session { // Ends a hook-observed LLM call, synthesizing a start if only the post hook arrives. The same // alias metadata recovery used by `start_hook_llm` keeps post-only aliased child LLMs under the // subagent instead of falling back to the root agent. - fn end_hook_llm(&mut self, event: LlmEvent) -> Result<(), CliError> { + async fn end_hook_llm(&mut self, event: LlmEvent) -> Result<(), CliError> { self.ensure_turn_started(event.metadata.clone())?; let (parent, metadata) = self.hook_llm_owner(event.metadata); let handle = match self.llms.remove(&event.api_call_id) { @@ -1511,7 +1515,7 @@ impl Session { // Starts a tool call under an explicit subagent when available, otherwise under the turn // scope. Duplicate tool IDs are ignored so repeated pre-tool hooks do not create parallel // handles for one agent tool invocation. - fn start_tool(&mut self, event: ToolEvent) -> Result<(), CliError> { + async fn start_tool(&mut self, event: ToolEvent) -> Result<(), CliError> { self.ensure_turn_started(event.metadata.clone())?; if self.tools.contains_key(&event.tool_call_id) { return Ok(()); @@ -1529,7 +1533,7 @@ impl Session { let active_tool_arguments = arguments.clone(); let active_tool_name = event.tool_name.clone(); let active_tool_owner_subagent_id = owner.subagent_id.clone(); - tool_conditional_execution(event.tool_name.as_str(), &arguments)?; + tool_conditional_execution(event.tool_name.as_str(), &arguments).await?; let metadata = tool_correlation_metadata( self.event_identity_metadata(event.metadata), owner.status, diff --git a/crates/cli/tests/coverage/shared/server_tests.rs b/crates/cli/tests/coverage/shared/server_tests.rs index c624182d8..bf9ee30e6 100644 --- a/crates/cli/tests/coverage/shared/server_tests.rs +++ b/crates/cli/tests/coverage/shared/server_tests.rs @@ -2336,7 +2336,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(); @@ -2562,21 +2564,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(); @@ -2619,10 +2624,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(); @@ -2685,17 +2692,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(); @@ -2743,29 +2753,33 @@ async fn gateway_request_codecs_apply_buffered_and_streaming_edits_on_all_genera 1, false, Arc::new(move |_name, mut request, annotated| { - let Some(marker) = request - .headers - .get("x-codec-matrix") - .and_then(Value::as_str) - .map(str::to_string) - else { - return Ok(LlmRequestInterceptOutcome::new(request, annotated)); - }; - let mut annotated = annotated.expect("generation route must expose an annotation"); - captured_annotations.lock().unwrap().push(json!({ - "marker": marker, - "annotation": annotated, - })); - let nemo_relay::codec::request::Message::User { content, .. } = - &mut annotated.messages[0] - else { - panic!("expected the first request item to be a portable user message"); - }; - *content = nemo_relay::codec::request::MessageContent::Text(format!("edited-{marker}")); - request - .headers - .insert("x-codec-edited".into(), json!(marker)); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + let captured_annotations = captured_annotations.clone(); + Box::pin(async move { + let Some(marker) = request + .headers + .get("x-codec-matrix") + .and_then(Value::as_str) + .map(str::to_string) + else { + return Ok(LlmRequestInterceptOutcome::new(request, annotated)); + }; + let mut annotated = annotated.expect("generation route must expose an annotation"); + captured_annotations.lock().unwrap().push(json!({ + "marker": marker, + "annotation": annotated, + })); + let nemo_relay::codec::request::Message::User { content, .. } = + &mut annotated.messages[0] + else { + panic!("expected the first request item to be a portable user message"); + }; + *content = + nemo_relay::codec::request::MessageContent::Text(format!("edited-{marker}")); + request + .headers + .insert("x-codec-edited".into(), json!(marker)); + Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + }) }), ) .unwrap(); diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index c34829241..06ac5119b 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -1,6 +1,7 @@ // SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 +use std::future::Future; use std::sync::Arc; use chrono::{DateTime, TimeDelta, Utc}; @@ -19,6 +20,9 @@ use crate::api::optimization::{ use crate::api::runtime::LlmCodecIdentity; use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::global_context; +use crate::api::runtime::subscriber_dispatcher::{ + dispatch_sanitized_event, dispatch_transformed_event, +}; use crate::api::runtime::{ EventSubscriberFn, LlmCollectorFn, LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn, @@ -30,7 +34,7 @@ use crate::api::scope::{EmitMarkEventParams, ScopeHandle}; use crate::api::shared::{ ensure_runtime_owner, inject_dynamo_session_ids, metadata_with_otel_status, resolve_parent_uuid, run_request_intercepts_with_codec_and_recorder, - sanitize_event_with_scope_stack, snapshot_event_subscribers, + sanitize_event_with_scope_stack, snapshot_event_sanitizers, snapshot_event_subscribers, }; use crate::codec::request::{AnnotatedLlmRequest, Message}; use crate::codec::response::{AnnotatedLlmResponse, attach_estimated_cost_for_provider}; @@ -400,28 +404,7 @@ fn limit_annotated_request_history_to_current_user_turn( ) } -fn emit_llm_start( - handle: &LlmHandle, - request: &LlmRequest, - annotated_request: Option>, - request_codec: Option>, -) -> Result<()> { - ensure_runtime_owner()?; - let subscribers = { - let scope_stack = handle.captured_scope_stack(); - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); - snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())? - }; - emit_llm_start_with_subscribers( - handle, - request, - annotated_request, - request_codec, - &subscribers, - ) -} - -fn emit_llm_start_with_subscribers( +async fn emit_llm_start_with_subscribers( handle: &LlmHandle, request: &LlmRequest, annotated_request: Option>, @@ -446,7 +429,8 @@ fn emit_llm_start_with_subscribers( observable_request.clone(), LlmSanitizeRequestContext::for_request_codec(request_codec.clone()), &entries, - ); + ) + .await; let request_changed = sanitized_request .as_ref() .is_some_and(|sanitized_request| sanitized_request != &observable_request); @@ -480,7 +464,7 @@ fn emit_llm_start_with_subscribers( .map_err(|error| FlowError::Internal(error.to_string()))?; state.build_llm_start_event(handle, input, annotated_request) }; - if let Some(event) = sanitize_event_with_scope_stack(event, scope_stack) { + if let Some(event) = sanitize_event_with_scope_stack(event, scope_stack).await { NemoRelayContextState::emit_event(&event, subscribers); } Ok(()) @@ -495,7 +479,33 @@ fn remove_observability_credential_headers(mut request: LlmRequest) -> LlmReques request } -fn emit_pending_request_marks( +/// Synchronous test seam retained for lifecycle unit tests. Public manual +/// lifecycle emission is synchronous too, but its work is queued; this helper +/// exercises the managed start-event transformation directly. +#[cfg(test)] +fn emit_llm_start( + handle: &LlmHandle, + request: &LlmRequest, + annotated_request: Option>, + request_codec: Option>, +) -> Result<()> { + let subscribers = { + let scope_stack = handle.captured_scope_stack(); + let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())? + }; + tokio::runtime::Runtime::new() + .map_err(|error| FlowError::Internal(error.to_string()))? + .block_on(emit_llm_start_with_subscribers( + handle, + request, + annotated_request, + request_codec, + &subscribers, + )) +} + +async fn emit_pending_request_marks( handle: &LlmHandle, marks: Vec, subscribers: &[EventSubscriberFn], @@ -517,28 +527,58 @@ fn emit_pending_request_marks( mark.category, mark.category_profile, )); - if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) { + if let Some(event) = + sanitize_event_with_scope_stack(event, handle.captured_scope_stack()).await + { NemoRelayContextState::emit_event(&event, subscribers); } } Ok(()) } -pub(crate) fn emit_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) { - emit_optimization_marks_with( +pub(crate) async fn emit_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) { + emit_optimization_marks_with_async( handle, subscribers, |event| sanitize_event_with_scope_stack(event, handle.captured_scope_stack()), |event, subscribers| NemoRelayContextState::try_emit_event(event, subscribers), - ); + ) + .await; } -fn emit_optimization_marks_with( +/// Queue optimization marks from a synchronous lifecycle API. +/// +/// The public manual lifecycle APIs must not await middleware. Capture each +/// event's sanitizer chain now and enqueue the immutable snapshots ahead of +/// the corresponding end event, preserving publication order. +fn enqueue_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) { + let contributions = handle.optimization_recorder.unemitted_with_timestamps(); + if contributions.is_empty() || ensure_runtime_owner().is_err() { + return; + } + let scope_stack = handle.captured_scope_stack().clone(); + for (contribution, recorded_at) in contributions { + let event = optimization_mark_event(handle, &contribution, recorded_at); + let Some(sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) else { + break; + }; + if dispatch_sanitized_event(event, sanitizers, subscribers, scope_stack.clone()) { + handle.optimization_recorder.mark_emitted(1); + } else { + break; + } + } +} + +async fn emit_optimization_marks_with_async( handle: &LlmHandle, subscribers: &[EventSubscriberFn], - mut sanitize: impl FnMut(Event) -> Option, + mut sanitize: F, mut enqueue: impl FnMut(&Event, &[EventSubscriberFn]) -> bool, -) { +) where + F: FnMut(Event) -> Fut, + Fut: Future>, +{ let contributions = handle.optimization_recorder.unemitted_with_timestamps(); if contributions.is_empty() { return; @@ -554,30 +594,8 @@ fn emit_optimization_marks_with( return; } for (contribution, recorded_at) in contributions { - let offset = contribution.sequence.unwrap_or(0).saturating_add(2); - let offset = i64::try_from(offset).unwrap_or(i64::MAX); - let request_ordered_timestamp = handle.started_at + TimeDelta::microseconds(offset); - let timestamp = recorded_at.max(request_ordered_timestamp); - let data = serde_json::to_value(&contribution).unwrap_or(Json::Null); - let event = Event::Mark(MarkEvent::new( - BaseEvent::builder() - .name("nemo_relay.llm.optimization") - .parent_uuid(handle.uuid) - .timestamp(timestamp) - .data(data) - .data_schema(DataSchema { - name: "nemo.relay.llm_optimization_contribution".to_string(), - version: "1".to_string(), - }) - .build(), - Some(EventCategory::custom()), - Some( - CategoryProfile::builder() - .subtype("nemo_relay.llm.optimization") - .build(), - ), - )); - let Some(event) = sanitize(event) else { + let event = optimization_mark_event(handle, &contribution, recorded_at); + let Some(event) = sanitize(event).await else { // Sanitizers currently rewrite fields rather than intentionally // dropping events. `None` means the sanitizer context was // unavailable, so preserve this ordered suffix for a later retry. @@ -594,6 +612,63 @@ fn emit_optimization_marks_with( } } +fn optimization_mark_event( + handle: &LlmHandle, + contribution: &crate::codec::optimization::LlmOptimizationContribution, + recorded_at: DateTime, +) -> Event { + let offset = contribution.sequence.unwrap_or(0).saturating_add(2); + let offset = i64::try_from(offset).unwrap_or(i64::MAX); + let request_ordered_timestamp = handle.started_at + TimeDelta::microseconds(offset); + Event::Mark(MarkEvent::new( + BaseEvent::builder() + .name("nemo_relay.llm.optimization") + .parent_uuid(handle.uuid) + .timestamp(recorded_at.max(request_ordered_timestamp)) + .data(serde_json::to_value(contribution).unwrap_or(Json::Null)) + .data_schema(DataSchema { + name: "nemo.relay.llm_optimization_contribution".to_string(), + version: "1".to_string(), + }) + .build(), + Some(EventCategory::custom()), + Some( + CategoryProfile::builder() + .subtype("nemo_relay.llm.optimization") + .build(), + ), + )) +} + +/// Synchronous test seam for optimization-mark accounting. Production paths +/// always use [`emit_optimization_marks_with_async`]; unit tests use this seam +/// to isolate cursor behavior from asynchronous event publication. +#[cfg(test)] +fn emit_optimization_marks_with( + handle: &LlmHandle, + subscribers: &[EventSubscriberFn], + mut sanitize: F, + mut enqueue: impl FnMut(&Event, &[EventSubscriberFn]) -> bool, +) where + F: FnMut(Event) -> Option, +{ + let contributions = handle.optimization_recorder.unemitted_with_timestamps(); + if contributions.is_empty() || ensure_runtime_owner().is_err() { + return; + } + for (contribution, recorded_at) in contributions { + let event = optimization_mark_event(handle, &contribution, recorded_at); + let Some(event) = sanitize(event) else { + break; + }; + if enqueue(&event, subscribers) { + handle.optimization_recorder.mark_emitted(1); + } else { + break; + } + } +} + /// Start a manual LLM lifecycle span. /// /// This emits an LLM-start event after applying sanitize-request guardrails to @@ -641,7 +716,77 @@ pub fn llm_call(params: LlmCallParams<'_>) -> Result { .timestamp_opt(params.timestamp) .build(); let handle = create_llm_handle(handle_params)?; - emit_llm_start(&handle, params.request, params.annotated_request, None)?; + let scope_stack = handle.captured_scope_stack().clone(); + let (entries, subscribers, agent_is_fresh) = { + let mut scope_guard = scope_stack.write().expect("scope stack lock poisoned"); + let scope_locals = scope_guard.collect_scope_local_registries(|registries| { + ®istries.llm_sanitize_request_guardrails + }); + let subscribers = + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + let entries = state.llm_sanitize_request_entries(&scope_locals); + drop(state); + let agent_is_fresh = scope_guard.take_agent_freshness(handle.parent_uuid); + (entries, subscribers, agent_is_fresh) + }; + // Middleware and event publication only observe a credential-free copy. + // Keep `params.request` untouched: it remains the caller/provider request. + let request = remove_observability_credential_headers(params.request.clone()); + let annotated_request = params.annotated_request; + let event = { + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + state.build_llm_start_event(&handle, None, None) + }; + let queued_handle = handle.clone(); + if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { + dispatch_transformed_event( + event, + Box::new(move |event| { + Box::pin(async move { + let mut sanitized_request = + NemoRelayContextState::llm_sanitize_request_snapshot_chain( + request.clone(), + LlmSanitizeRequestContext::default(), + &entries, + ) + .await; + let request_changed = sanitized_request + .as_ref() + .is_some_and(|sanitized| sanitized != &request); + let mut annotation = if sanitized_request.is_none() || request_changed { + None + } else { + annotated_request + }; + if !agent_is_fresh && let Some(sanitized_request) = sanitized_request.as_mut() { + project_llm_request_to_current_user_turn( + sanitized_request, + &mut annotation, + None, + ); + } + let input = sanitized_request + .as_ref() + .and_then(|request| serde_json::to_value(request).ok()); + let context = global_context(); + match context.read() { + Ok(state) => state.build_llm_start_event(&queued_handle, input, annotation), + Err(_) => event, + } + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); + } Ok(handle) } @@ -682,17 +827,137 @@ struct LlmCallEndBehavior { /// Sanitize-response guardrails affect only the emitted end-event payload, not /// the caller-owned `response` value. pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> { - llm_call_end_with_behavior( - params, - LlmCallEndBehavior { - response_codec_errors_fatal: true, - attach_estimated_cost: false, - }, - None, - ) + ensure_runtime_owner()?; + let scope_stack = params.handle.captured_scope_stack().clone(); + let (entries, subscribers) = { + let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_locals = scope_guard.collect_scope_local_registries(|registries| { + ®istries.llm_sanitize_response_guardrails + }); + let subscribers = + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + ( + state.llm_sanitize_response_entries(&scope_locals), + subscribers, + ) + }; + let response = if params.response.is_null() { + params.data.unwrap_or(params.response) + } else { + params.response + }; + let response_was_null_without_fallback = response.is_null(); + let handle = params.handle.clone(); + let metadata = params.metadata; + let timestamp = params.timestamp; + let annotated_response = params.annotated_response; + let response_codec = params.response_codec; + handle.optimization_recorder.close_for_finalization(None); + enqueue_optimization_marks(&handle, &subscribers); + let event = { + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + state.build_llm_end_event( + EndLlmHandleParams::builder() + .handle(&handle) + .data(Json::Null) + .metadata_opt(metadata.clone()) + .annotated_response_opt(annotated_response.clone()) + .timestamp_opt(timestamp) + .build(), + ) + }; + if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { + dispatch_transformed_event( + event, + Box::new(move |event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( + response.clone(), + LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), + &entries, + ) + .await; + let changed = sanitized + .as_ref() + .is_some_and(|sanitized| sanitized != &response); + let data = match sanitized { + Some(response) + if response_was_null_without_fallback && response.is_null() => + { + None + } + response => response, + }; + let annotation_omitted = data.as_ref().is_none_or(Json::is_null); + let (mut annotation, decode_error) = if annotation_omitted { + (None, None) + } else { + resolve_llm_end_annotation( + (!changed).then_some(annotated_response).flatten(), + response_codec, + data.as_ref(), + &LlmCallEndBehavior { + response_codec_errors_fatal: false, + attach_estimated_cost: false, + }, + &handle.name, + ) + }; + if let Some(error) = decode_error { + log::error!( + target: "nemo_relay.runtime", + event = "manual_llm_response_codec_failed"; + "Manual LLM response annotation failed during queued publication: {error}" + ); + } + let pricing = crate::codec::response::active_pricing_resolver(); + let summary = finalize_optimization_summary( + &handle.optimization_recorder, + annotation.as_mut(), + handle.model_name.as_deref(), + &pricing, + ); + if !annotation_omitted + && annotation.is_none() + && let Some(summary) = summary + { + annotation = Some(AnnotatedLlmResponse { + optimization_summary: Some(summary), + ..AnnotatedLlmResponse::default() + }); + } + let context = global_context(); + let Ok(state) = context.read() else { + return event; + }; + let end_metadata = metadata_with_otel_status(metadata, "OK", None); + state.build_llm_end_event( + EndLlmHandleParams::builder() + .handle(&handle) + .data_opt(data) + .metadata_opt(end_metadata) + .annotated_response_opt(annotation.map(Arc::new)) + .timestamp_opt(timestamp) + .build(), + ) + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); + } + Ok(()) } -fn llm_call_end_with_behavior( +async fn llm_call_end_with_behavior( params: LlmCallEndParams<'_>, behavior: LlmCallEndBehavior, lifecycle_subscribers: Option<&[EventSubscriberFn]>, @@ -735,7 +1000,8 @@ fn llm_call_end_with_behavior( response.clone(), LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), &entries, - ); + ) + .await; let response_changed = sanitized_response .as_ref() .is_some_and(|sanitized_response| sanitized_response != &response); @@ -756,7 +1022,7 @@ fn llm_call_end_with_behavior( ) }; handle.optimization_recorder.close_for_finalization(None); - emit_optimization_marks(handle, &subscribers); + emit_optimization_marks(handle, &subscribers).await; let pricing = crate::codec::response::active_pricing_resolver(); let summary = finalize_optimization_summary( &handle.optimization_recorder, @@ -790,7 +1056,8 @@ fn llm_call_end_with_behavior( .build(), ) }; - if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) { + if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()).await + { NemoRelayContextState::emit_event(&event, &subscribers); } if let Some(error) = decode_error @@ -842,7 +1109,7 @@ fn resolve_llm_end_annotation( } } -fn emit_llm_end_without_output( +async fn emit_llm_end_without_output( handle: &LlmHandle, metadata: Option, response_codec: Option>, @@ -868,17 +1135,20 @@ fn emit_llm_end_without_output( (entries, subscribers) }; let had_fallback_data = handle.data.is_some(); - let data = handle.data.clone().and_then(|data| { + let data = if let Some(data) = handle.data.clone() { NemoRelayContextState::llm_sanitize_response_snapshot_chain( data, LlmSanitizeResponseContext::for_response_codec(response_codec), &entries, ) - }); + .await + } else { + None + }; let annotation_omitted = (had_fallback_data && data.is_none()) || data.as_ref().is_some_and(Json::is_null); handle.optimization_recorder.close_for_finalization(None); - emit_optimization_marks(handle, &subscribers); + emit_optimization_marks(handle, &subscribers).await; let pricing = crate::codec::response::active_pricing_resolver(); let annotated_response = (!annotation_omitted) .then(|| { @@ -903,7 +1173,8 @@ fn emit_llm_end_without_output( .map_err(|error| FlowError::Internal(error.to_string()))?; state.end_llm_handle(handle, data, metadata, annotated_response) }; - if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) { + if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()).await + { NemoRelayContextState::emit_event(&event, &subscribers); } Ok(()) @@ -990,7 +1261,9 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { &subscribers, parent_uuid, guardrail_metadata, - )? { + ) + .await? + { let mut rejection_data = json!({}); if let Some(object) = rejection_data.as_object_mut() { object.insert("rejected".into(), json!(true)); @@ -1018,6 +1291,7 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { codec, &optimization_recorder, ) + .await }) .await?; @@ -1043,12 +1317,13 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { annotated_request.clone(), request_codec.clone(), &lifecycle_subscribers, - )?; - emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers)?; + ) + .await?; + emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers).await?; handle .optimization_recorder .record_all(optimization_contributions); - emit_optimization_marks(&handle, &lifecycle_subscribers); + emit_optimization_marks(&handle, &lifecycle_subscribers).await; let execution_name = name.clone(); let event_uuid = handle.uuid; @@ -1087,7 +1362,8 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { attach_estimated_cost: true, }, Some(&lifecycle_subscribers), - )?; + ) + .await?; Ok(response) } Err(error) => { @@ -1098,7 +1374,8 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { end_metadata, response_codec, Some(&lifecycle_subscribers), - ); + ) + .await; Err(error) } } @@ -1186,7 +1463,9 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu &subscribers, parent_uuid, guardrail_metadata, - )? { + ) + .await? + { let mut rejection_data = json!({}); if let Some(object) = rejection_data.as_object_mut() { object.insert("rejected".into(), json!(true)); @@ -1214,6 +1493,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu codec, &optimization_recorder, ) + .await }) .await?; @@ -1239,12 +1519,13 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu annotated_request, request_codec.clone(), &lifecycle_subscribers, - )?; - emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers)?; + ) + .await?; + emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers).await?; handle .optimization_recorder .record_all(optimization_contributions); - emit_optimization_marks(&handle, &lifecycle_subscribers); + emit_optimization_marks(&handle, &lifecycle_subscribers).await; let execution_name = name.clone(); let event_uuid = handle.uuid; @@ -1289,7 +1570,8 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu end_metadata, response_codec, Some(&lifecycle_subscribers), - ); + ) + .await; Err(error) } } @@ -1318,7 +1600,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu /// /// This helper does not emit the returned marks because it does not own an LLM /// lifecycle. Callers must attach them to the lifecycle they own. -pub fn llm_request_intercepts( +pub async fn llm_request_intercepts( name: &str, request: LlmRequest, ) -> Result { @@ -1336,7 +1618,8 @@ pub fn llm_request_intercepts( }; let mut outcome = NemoRelayContextState::llm_request_intercepts_snapshot_chain( name, request, None, &entries, false, - )?; + ) + .await?; inject_dynamo_session_ids(&mut outcome.request); Ok(outcome) } @@ -1361,7 +1644,7 @@ pub fn llm_request_intercepts( /// This helper is useful for preflight checks when the caller needs the /// rejection result without starting an LLM span. Guardrail scopes are still /// emitted for the conditional checks themselves. -pub fn llm_conditional_execution(request: &LlmRequest) -> Result<()> { +pub async fn llm_conditional_execution(request: &LlmRequest) -> Result<()> { ensure_runtime_owner()?; let (entries, subscribers, parent_uuid) = { let scope_stack = current_scope_stack(); @@ -1384,7 +1667,9 @@ pub fn llm_conditional_execution(request: &LlmRequest) -> Result<()> { &subscribers, parent_uuid, None, - )? { + ) + .await? + { return Err(FlowError::GuardrailRejected(error)); } Ok(()) diff --git a/crates/core/src/api/runtime/callbacks.rs b/crates/core/src/api/runtime/callbacks.rs index e40558da1..b07a82fde 100644 --- a/crates/core/src/api/runtime/callbacks.rs +++ b/crates/core/src/api/runtime/callbacks.rs @@ -27,8 +27,14 @@ use crate::json::Json; /// /// The callback receives the current event as immutable context and the fields /// it may replace. Later callbacks observe fields returned by earlier entries. -pub type EventSanitizeFn = - Arc EventSanitizeFields + Send + Sync>; +pub type EventSanitizeFn = Arc< + dyn Fn( + Event, + EventSanitizeFields, + ) -> Pin> + Send>> + + Send + + Sync, +>; /// Sanitize a tool request payload before the runtime records it. /// @@ -42,7 +48,8 @@ pub type EventSanitizeFn = /// /// # Returns /// Sanitized JSON payload for the emitted event. -pub type ToolSanitizeFn = Arc Json + Send + Sync>; +pub type ToolSanitizeFn = + Arc Pin> + Send>> + Send + Sync>; /// Decide whether a tool call is allowed to continue. /// /// The callback receives the tool name and the current argument payload. It can @@ -64,7 +71,11 @@ pub type ToolSanitizeFn = Arc Json + Send + Sync>; /// # Errors /// The callback can return any [`FlowError`](crate::error::FlowError) to abort /// guardrail evaluation. -pub type ToolConditionalFn = Arc Result> + Send + Sync>; +pub type ToolConditionalFn = Arc< + dyn Fn(String, Json) -> Pin>> + Send>> + + Send + + Sync, +>; /// Rewrite tool arguments before execution. /// /// Tool request intercepts run in priority order and can transform the JSON @@ -80,7 +91,8 @@ pub type ToolConditionalFn = Arc Result> + /// # Errors /// The callback can return any [`FlowError`](crate::error::FlowError) to abort /// the request-intercept chain. -pub type ToolInterceptFn = Arc Result + Send + Sync>; +pub type ToolInterceptFn = + Arc Pin> + Send>> + Send + Sync>; /// Continuation type invoked by tool execution intercepts. /// /// Execution intercepts receive this callable as their `next` continuation and @@ -308,8 +320,14 @@ impl LlmSanitizeResponseContext { /// /// The context is always supplied and distinguishes no codec, built-in codecs, /// runtime-registered codecs, and opaque active codecs. -pub type LlmSanitizeRequestFn = - Arc Option + Send + Sync>; +pub type LlmSanitizeRequestFn = Arc< + dyn Fn( + LlmRequest, + LlmSanitizeRequestContext, + ) -> Pin>> + Send>> + + Send + + Sync, +>; /// Sanitize an LLM response before the runtime records it. /// /// These callbacks rewrite the JSON response payload captured on LLM-end @@ -325,8 +343,14 @@ pub type LlmSanitizeRequestFn = /// /// The context is always supplied and distinguishes no codec, built-in codecs, /// runtime-registered codecs, and opaque active codecs. -pub type LlmSanitizeResponseFn = - Arc Option + Send + Sync>; +pub type LlmSanitizeResponseFn = Arc< + dyn Fn( + Json, + LlmSanitizeResponseContext, + ) -> Pin>> + Send>> + + Send + + Sync, +>; /// Decide whether an LLM call is allowed to continue. /// /// The callback receives the current [`LlmRequest`] and can allow execution, @@ -346,7 +370,11 @@ pub type LlmSanitizeResponseFn = /// # Errors /// The callback can return any [`FlowError`](crate::error::FlowError) to abort /// guardrail evaluation. -pub type LlmConditionalFn = Arc Result> + Send + Sync>; +pub type LlmConditionalFn = Arc< + dyn Fn(LlmRequest) -> Pin>> + Send>> + + Send + + Sync, +>; /// Rewrite or annotate an LLM request before execution. /// /// Request intercepts can transform the wire request, attach or replace a @@ -368,7 +396,11 @@ pub type LlmConditionalFn = Arc Result> + /// The callback can return any [`FlowError`](crate::error::FlowError) to abort /// the request-intercept chain. pub type LlmRequestInterceptFn = Arc< - dyn Fn(&str, LlmRequest, Option) -> Result + dyn Fn( + String, + LlmRequest, + Option, + ) -> Pin> + Send>> + Send + Sync, >; diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index eb3dbb950..e2098b9a7 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -564,7 +564,7 @@ impl NemoRelayContextState { )) } - fn emit_guardrail_scope_start( + async fn emit_guardrail_scope_start( name: &str, parent_uuid: Option, metadata: Option, @@ -591,13 +591,13 @@ impl NemoRelayContextState { EventCategory::from(handle.scope_type), None, )); - if let Some(event) = sanitize_event(event) { + if let Some(event) = sanitize_event(event).await { Self::emit_event(&event, subscribers); } handle } - fn emit_guardrail_scope_end( + async fn emit_guardrail_scope_end( handle: &ScopeHandle, output: Json, subscribers: &[EventSubscriberFn], @@ -616,7 +616,7 @@ impl NemoRelayContextState { EventCategory::from(handle.scope_type), None, )); - if let Some(event) = sanitize_event(event) { + if let Some(event) = sanitize_event(event).await { Self::emit_event(&event, subscribers); } } @@ -633,13 +633,22 @@ 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(); + match (entry.payload)(event.clone(), fields).await { + Ok(fields) => event.apply_sanitize_fields(fields), + Err(error) => log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_failed", + sanitizer = entry.name.as_str(), + event_name = event.name(); + "Event sanitizer failed; preserving the last valid event snapshot: {error}" + ), + } } event } @@ -672,14 +681,23 @@ impl NemoRelayContextState { /// /// # Returns /// The sanitized JSON payload after every provided guardrail has run. - pub(crate) fn tool_sanitize_request_snapshot_chain( + pub(crate) async fn tool_sanitize_request_snapshot_chain( name: &str, args: Json, entries: &[Guardrail], ) -> Json { let mut value = args; for entry in entries { - value = (entry.payload)(name, value); + match (entry.payload)(name.to_string(), value.clone()).await { + Ok(next) => value = next, + Err(error) => log::error!( + target: "nemo_relay.runtime", + event = "tool_request_sanitizer_failed", + sanitizer = entry.name.as_str(), + tool_name = name; + "Tool request sanitizer failed; preserving the last valid payload: {error}" + ), + } } value } @@ -712,14 +730,23 @@ impl NemoRelayContextState { /// /// # Returns /// The sanitized JSON payload after every provided guardrail has run. - pub(crate) fn tool_sanitize_response_snapshot_chain( + pub(crate) async fn tool_sanitize_response_snapshot_chain( name: &str, result: Json, entries: &[Guardrail], ) -> Json { let mut value = result; for entry in entries { - value = (entry.payload)(name, value); + match (entry.payload)(name.to_string(), value.clone()).await { + Ok(next) => value = next, + Err(error) => log::error!( + target: "nemo_relay.runtime", + event = "tool_response_sanitizer_failed", + sanitizer = entry.name.as_str(), + tool_name = name; + "Tool response sanitizer failed; preserving the last valid payload: {error}" + ), + } } value } @@ -769,7 +796,7 @@ impl NemoRelayContextState { /// # Errors /// Propagates any error returned by a guardrail callback after emitting the /// corresponding guardrail scope end event. - pub(crate) fn tool_conditional_execution_snapshot_chain( + pub(crate) async fn tool_conditional_execution_snapshot_chain( name: &str, args: &Json, entries: &[Guardrail], @@ -787,8 +814,9 @@ impl NemoRelayContextState { "target_name": name, }), subscribers, - ); - let result = (entry.payload)(name, args); + ) + .await; + let result = (entry.payload)(name.to_string(), args.clone()).await; let output = match &result { Ok(Some(reason)) => json!({ "allowed": false, @@ -804,7 +832,7 @@ impl NemoRelayContextState { "error": error.to_string(), }), }; - Self::emit_guardrail_scope_end(&handle, output, subscribers); + Self::emit_guardrail_scope_end(&handle, output, subscribers).await; if let Some(error) = result? { return Ok(Some(error)); } @@ -847,14 +875,14 @@ impl NemoRelayContextState { /// # Notes /// If an intercept entry has `break_chain` enabled, later intercepts are /// skipped after that entry runs. - pub(crate) fn tool_request_intercepts_snapshot_chain( + pub(crate) async fn tool_request_intercepts_snapshot_chain( name: &str, args: Json, entries: &[Intercept], ) -> crate::error::Result { let mut value = args; for entry in entries { - value = (entry.payload.callable)(name, value)?; + value = (entry.payload.callable)(name.to_string(), value).await?; if entry.payload.break_chain { break; } @@ -964,14 +992,28 @@ impl NemoRelayContextState { /// /// # Returns /// The sanitized [`LlmRequest`] after every provided guardrail has run. - pub(crate) fn llm_sanitize_request_snapshot_chain( + pub(crate) async fn llm_sanitize_request_snapshot_chain( request: LlmRequest, context: LlmSanitizeRequestContext, entries: &[Guardrail], ) -> Option { let mut value = Some(request); for entry in entries { - value = value.and_then(|value| (entry.payload)(value, context.clone())); + if let Some(current) = value.take() { + match (entry.payload)(current.clone(), context.clone()).await { + Ok(next) => value = next, + Err(error) => { + log::error!( + target: "nemo_relay.runtime", + event = "llm_request_sanitizer_failed", + sanitizer = entry.name.as_str(), + preserved_value = "unsanitized_request"; + "LLM request sanitizer failed; preserving the last valid unsanitized request: {error}" + ); + value = Some(current); + } + } + } } value } @@ -1003,14 +1045,28 @@ impl NemoRelayContextState { /// /// # Returns /// The sanitized response payload after every provided guardrail has run. - pub(crate) fn llm_sanitize_response_snapshot_chain( + pub(crate) async fn llm_sanitize_response_snapshot_chain( response: Json, context: LlmSanitizeResponseContext, entries: &[Guardrail], ) -> Option { let mut value = Some(response); for entry in entries { - value = value.and_then(|value| (entry.payload)(value, context.clone())); + if let Some(current) = value.take() { + match (entry.payload)(current.clone(), context.clone()).await { + Ok(next) => value = next, + Err(error) => { + log::error!( + target: "nemo_relay.runtime", + event = "llm_response_sanitizer_failed", + sanitizer = entry.name.as_str(), + preserved_value = "unsanitized_response"; + "LLM response sanitizer failed; preserving the last valid unsanitized response: {error}" + ); + value = Some(current); + } + } + } } value } @@ -1059,7 +1115,7 @@ impl NemoRelayContextState { /// # Errors /// Propagates any error returned by a guardrail callback after emitting the /// corresponding guardrail scope end event. - pub(crate) fn llm_conditional_execution_snapshot_chain( + pub(crate) async fn llm_conditional_execution_snapshot_chain( request: &LlmRequest, entries: &[Guardrail], subscribers: &[EventSubscriberFn], @@ -1075,8 +1131,9 @@ impl NemoRelayContextState { "kind": "llm_conditional_execution", }), subscribers, - ); - let result = (entry.payload)(request); + ) + .await; + let result = (entry.payload)(request.clone()).await; let output = match &result { Ok(Some(reason)) => json!({ "allowed": false, @@ -1092,7 +1149,7 @@ impl NemoRelayContextState { "error": error.to_string(), }), }; - Self::emit_guardrail_scope_end(&handle, output, subscribers); + Self::emit_guardrail_scope_end(&handle, output, subscribers).await; if let Some(error) = result? { return Ok(Some(error)); } @@ -1139,7 +1196,7 @@ impl NemoRelayContextState { /// # Notes /// If an intercept entry has `break_chain` enabled, later intercepts are /// skipped after that entry runs. - pub(crate) fn llm_request_intercepts_snapshot_chain( + pub(crate) async fn llm_request_intercepts_snapshot_chain( name: &str, request: LlmRequest, annotated: Option, @@ -1154,11 +1211,12 @@ impl NemoRelayContextState { codec_active, None, ) + .await } /// Run a request-intercept snapshot while ingesting optimization evidence /// directly into the managed call's bounded accumulator. - pub(crate) fn llm_request_intercepts_snapshot_chain_with_recorder( + pub(crate) async fn llm_request_intercepts_snapshot_chain_with_recorder( name: &str, request: LlmRequest, annotated: Option, @@ -1172,7 +1230,8 @@ impl NemoRelayContextState { let mut optimization_contributions = Vec::new(); for entry in entries { let input_content = request_value.content.clone(); - let outcome = (entry.payload.callable)(name, request_value, annotated_value)?; + let outcome = + (entry.payload.callable)(name.to_string(), request_value, annotated_value).await?; if codec_active && outcome.request.content != input_content { return Err(crate::error::FlowError::InvalidArgument(format!( "LLM request intercept '{}' changed request.content while a request codec is active; modify annotated_request instead", diff --git a/crates/core/src/api/runtime/subscriber_dispatcher.rs b/crates/core/src/api/runtime/subscriber_dispatcher.rs index ffb4ef3b0..51df9f716 100644 --- a/crates/core/src/api/runtime/subscriber_dispatcher.rs +++ b/crates/core/src/api/runtime/subscriber_dispatcher.rs @@ -4,8 +4,17 @@ //! 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; @@ -24,17 +33,25 @@ mod native { enum DispatcherMessage { Deliver { event: Box, + transform: Option, + sanitizers: Vec>, subscribers: Vec, scope_stack: ScopeStackHandle, }, Flush { done: Sender<()>, }, + Barrier { + done: Receiver<()>, + }, } static DISPATCHER: OnceLock, String>> = OnceLock::new(); + static SANITIZER_RUNTIME: OnceLock> = + OnceLock::new(); static DISPATCHER_FAILURE_LOGGED: AtomicBool = AtomicBool::new(false); + static SANITIZER_RUNTIME_FAILURE_LOGGED: AtomicBool = AtomicBool::new(false); thread_local! { static IN_DISPATCHER: Cell = const { Cell::new(false) }; @@ -46,6 +63,8 @@ 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(), }; @@ -76,6 +95,51 @@ mod native { } } + pub(super) fn dispatch_sanitized_event( + event: Event, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, + ) -> bool { + let message = DispatcherMessage::Deliver { + event: Box::new(event), + transform: None, + sanitizers, + subscribers: subscribers.to_vec(), + scope_stack, + }; + send_dispatch_message(message) + } + + pub(super) fn dispatch_transformed_event( + event: Event, + transform: EventTransformFn, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, + ) -> bool { + let message = DispatcherMessage::Deliver { + event: Box::new(event), + transform: Some(transform), + sanitizers, + subscribers: subscribers.to_vec(), + scope_stack, + }; + send_dispatch_message(message) + } + + /// Insert a FIFO barrier for work that will enqueue a publication from an + /// async task. A later flush waits for the task to signal completion, then + /// drains the event it queued before acknowledging the flush. + pub(super) fn register_async_publication() -> Option> { + let sender = dispatcher_sender().ok()?; + let (done_tx, done_rx) = mpsc::channel(); + sender + .send(DispatcherMessage::Barrier { done: done_rx }) + .ok() + .map(|_| done_tx) + } + pub(super) fn flush_subscribers() -> Result<()> { if IN_DISPATCHER.with(Cell::get) { return Ok(()); @@ -102,6 +166,30 @@ mod native { DISPATCHER.get_or_init(start_dispatcher).clone() } + fn send_dispatch_message(message: DispatcherMessage) -> bool { + match dispatcher_sender() { + Ok(sender) if sender.send(message).is_ok() => true, + Ok(_) => { + log::warn!( + target: "nemo_relay.runtime", + event = "subscriber_event_dropped", + reason = "dispatcher_disconnected"; + "Subscriber event was dropped because the dispatcher stopped" + ); + false + } + Err(error) if !DISPATCHER_FAILURE_LOGGED.swap(true, Ordering::AcqRel) => { + log::error!( + target: "nemo_relay.runtime", + event = "subscriber_dispatcher_failed"; + "Subscriber dispatcher failed to start: {error}" + ); + false + } + Err(_) => false, + } + } + fn start_dispatcher() -> std::result::Result, String> { let (tx, rx) = mpsc::channel::(); let sender = std::thread::Builder::new() @@ -129,6 +217,9 @@ mod native { let _ = pending.send(()); } } + DispatcherMessage::Barrier { done } => { + let _ = done.recv(); + } message => handle_message(message), } } @@ -139,6 +230,9 @@ mod native { while let Ok(message) = rx.try_recv() { match message { DispatcherMessage::Flush { done } => pending_flushes.push(done), + DispatcherMessage::Barrier { done } => { + let _ = done.recv(); + } message => handle_message(message), } } @@ -149,23 +243,35 @@ 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 { done } => { + let _ = done.recv(); + } } } 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 Some(event) = sanitize_event_snapshot(*event, transform, sanitizers) else { + IN_DISPATCHER.with(|flag| flag.set(false)); + restore_thread_scope_stack(previous_scope_stack); + return; + }; for subscriber in subscribers { if catch_unwind(AssertUnwindSafe(|| subscriber(&event))).is_err() { log::error!( @@ -178,6 +284,75 @@ mod native { IN_DISPATCHER.with(|flag| flag.set(false)); restore_thread_scope_stack(previous_scope_stack); } + + /// Apply a transform and sanitizers on the dispatcher thread. A transform + /// failure drops the event because it may be responsible for inserting the + /// sanitized payload. A sanitizer failure retains the transformed snapshot + /// and continues publication (fail open). + fn sanitize_event_snapshot( + event: Event, + transform: Option, + sanitizers: Vec>, + ) -> Option { + let runtime = match SANITIZER_RUNTIME.get_or_init(|| { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| error.to_string()) + }) { + Ok(runtime) => runtime, + Err(error) => { + if !SANITIZER_RUNTIME_FAILURE_LOGGED.swap(true, Ordering::AcqRel) { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_runtime_failed"; + "Event sanitizer runtime failed; dropping events: {error}" + ); + } + return None; + } + }; + let transformed = match catch_unwind(AssertUnwindSafe(|| { + runtime.block_on(async move { + match transform { + Some(transform) => transform(event).await, + None => event, + } + }) + })) { + Ok(event) => event, + Err(_) => { + log::error!( + target: "nemo_relay.runtime", + event = "event_transform_panicked"; + "Event transform panicked; dropping the event" + ); + return None; + } + }; + if sanitizers.is_empty() { + return Some(transformed); + } + let original = transformed.clone(); + Some( + match catch_unwind(AssertUnwindSafe(|| { + runtime.block_on(NemoRelayContextState::event_sanitize_snapshot_chain( + transformed, + &sanitizers, + )) + })) { + Ok(event) => event, + Err(_) => { + log::error!( + target: "nemo_relay.runtime", + event = "event_sanitizer_panicked"; + "Event sanitizer panicked; publishing the transformed event snapshot" + ); + original + } + }, + ) + } } /// Queue an event for subscriber delivery. @@ -185,6 +360,37 @@ 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) +} + +/// Queue a snapshot for a middleware-specific asynchronous transformation, +/// followed by event sanitization and subscriber delivery. +pub(crate) fn dispatch_transformed_event( + event: Event, + transform: EventTransformFn, + sanitizers: Vec>, + subscribers: &[EventSubscriberFn], + scope_stack: ScopeStackHandle, +) -> bool { + native::dispatch_transformed_event(event, transform, sanitizers, subscribers, scope_stack) +} + +/// Register a FIFO barrier for async work that will queue a subscriber event. +/// +/// Dropping the returned sender releases the barrier, so error paths cannot +/// leave the dispatcher blocked. +pub(crate) fn register_async_publication() -> Option> { + native::register_async_publication() +} + /// Wait for all queued subscriber callbacks submitted before this call. pub fn flush_subscribers() -> Result<()> { native::flush_subscribers() diff --git a/crates/core/src/api/scope.rs b/crates/core/src/api/scope.rs index 60a1aa53d..0164f6102 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; @@ -221,7 +222,7 @@ pub fn get_handle() -> Result { pub fn push_scope(params: PushScopeParams<'_>) -> Result { ensure_runtime_owner()?; let parent_uuid = resolve_parent_uuid(params.parent); - let (handle, event, subscribers) = { + let (handle, event, subscribers, emission_scope_stack) = { let scope_stack = current_scope_stack(); let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); let scope_subscribers = scope_guard.collect_scope_local_subscribers(); @@ -241,12 +242,16 @@ pub fn push_scope(params: PushScopeParams<'_>) -> Result { .build(); let handle = state.create_scope_handle(handle_params); let event = state.build_scope_start_event(&handle, params.input); - (handle, event, subscribers) + (handle, event, subscribers, scope_stack.clone()) }; - let event = sanitize_event(event); task_scope_push(handle.clone()); - if let Some(event) = event { - NemoRelayContextState::emit_event(&event, &subscribers); + if let Some(sanitizers) = snapshot_event_sanitizers(&event, &emission_scope_stack) { + let _ = subscriber_dispatcher::dispatch_sanitized_event( + event, + sanitizers, + &subscribers, + emission_scope_stack, + ); } Ok(handle) } @@ -276,7 +281,7 @@ pub fn push_scope(params: PushScopeParams<'_>) -> Result { pub fn pop_scope(params: PopScopeParams<'_>) -> Result<()> { ensure_runtime_owner()?; let scope_stack = current_scope_stack(); - let (scope, event, subscribers) = { + let (scope, event, subscribers, emission_scope_stack) = { let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); let top = scope_guard.top(); if top.uuid != *params.handle_uuid { @@ -302,13 +307,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(()) } @@ -341,15 +354,29 @@ 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| { + log::error!( + target: "nemo_relay.runtime", + event = "mark_event_scope_stack_unavailable"; + "Mark event was dropped because the scope stack lock is poisoned: {error}" + ); + FlowError::Internal(error.to_string()) + })?; let subscribers = snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; scope_guard.mark_agent_fresh(parent_uuid); subscribers } else { - let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_guard = scope_stack.read().map_err(|error| { + log::error!( + target: "nemo_relay.runtime", + event = "mark_event_scope_stack_unavailable"; + "Mark event was dropped because the scope stack lock is poisoned: {error}" + ); + FlowError::Internal(error.to_string()) + })?; snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())? }; let context = global_context(); @@ -368,10 +395,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..b97c147bb 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>, @@ -246,7 +278,8 @@ fn run_request_intercepts_with_codec_inner( &entries, codec.is_some(), recorder, - )?; + ) + .await?; let mut request = outcome.request; inject_dynamo_session_ids(&mut request); let pending_marks = outcome.pending_marks; diff --git a/crates/core/src/api/tool.rs b/crates/core/src/api/tool.rs index 6fbe6cd70..7d6d10f71 100644 --- a/crates/core/src/api/tool.rs +++ b/crates/core/src/api/tool.rs @@ -7,12 +7,15 @@ use crate::api::event::{BaseEvent, Event, MarkEvent, PendingMarkSpec}; use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::current_scope_stack; use crate::api::runtime::global_context; +use crate::api::runtime::subscriber_dispatcher::{ + dispatch_sanitized_event, dispatch_transformed_event, +}; use crate::api::runtime::{EventSubscriberFn, ToolExecutionNextFn, with_active_event_uuid}; use crate::api::scope::event; use crate::api::scope::{EmitMarkEventParams, ScopeHandle}; use crate::api::shared::{ ensure_runtime_owner, metadata_with_otel_status, resolve_parent_uuid, sanitize_event, - snapshot_event_subscribers, + snapshot_event_sanitizers, snapshot_event_subscribers, }; use crate::api::skill_load; use crate::error::{FlowError, Result}; @@ -206,11 +209,107 @@ pub struct ToolCallEndParams<'a> { /// Sanitize-request guardrails affect only the emitted start-event payload, not /// the caller-owned `args` value. pub fn tool_call(params: ToolCallParams<'_>) -> Result { - let (handle, _) = tool_call_with_subscriber_snapshot(params)?; + ensure_runtime_owner()?; + let scope_stack = current_scope_stack(); + let (entries, subscribers) = { + let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_locals = scope_guard.collect_scope_local_registries(|registries| { + ®istries.tool_sanitize_request_guardrails + }); + let subscribers = + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + ( + state.tool_sanitize_request_entries(&scope_locals), + subscribers, + ) + }; + let handled_skill_loads = params + .metadata + .as_ref() + .and_then(Json::as_object) + .and_then(|metadata| metadata.get(skill_load::HANDLED_METADATA_KEY)) + .and_then(Json::as_bool) + .is_some_and(|handled| handled); + let skill_loads = if handled_skill_loads { + Vec::new() + } else if let Some(skill_loads) = skill_load::precomputed(params.metadata.as_ref()) { + skill_loads + } else { + skill_load::detect(params.name, ¶ms.args) + }; + let raw_args = params.args; + let (handle, event, marks) = { + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + let handle = state.create_tool_handle( + CreateToolHandleParams::builder() + .name(params.name) + .parent_uuid_opt(resolve_parent_uuid(params.parent)) + .attributes(params.attributes) + .data_opt(params.data) + .metadata_opt(params.metadata) + .tool_call_id_opt(params.tool_call_id) + .timestamp_opt(params.timestamp) + .build(), + ); + let event = state.build_tool_start_event(&handle, None); + let marks = skill_loads + .into_iter() + .map(|skill_load| { + state.create_event(MarkEvent::new( + BaseEvent::builder() + .name("skill.load") + .parent_uuid(handle.uuid) + .timestamp(handle.started_at) + .data(json!({"skill_name": skill_load.name})) + .metadata(json!({ + "skill_load_source": <&str>::from(skill_load.source), + "tool_name": handle.name, + })) + .build(), + None, + None, + )) + }) + .collect::>(); + (handle, event, marks) + }; + let tool_name = handle.name.clone(); + if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { + dispatch_transformed_event( + event, + Box::new(move |mut event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::tool_sanitize_request_snapshot_chain( + &tool_name, raw_args, &entries, + ) + .await; + let mut fields = event.sanitize_fields(); + fields.data = Some(sanitized); + event.apply_sanitize_fields(fields); + event + }) + }), + event_sanitizers, + &subscribers, + scope_stack.clone(), + ); + } + for mark in marks { + if let Some(sanitizers) = snapshot_event_sanitizers(&mark, &scope_stack) { + dispatch_sanitized_event(mark, sanitizers, &subscribers, scope_stack.clone()); + } + } Ok(handle) } -fn tool_call_with_subscriber_snapshot( +async fn tool_call_with_subscriber_snapshot( params: ToolCallParams<'_>, ) -> Result<(ToolHandle, Vec)> { ensure_runtime_owner()?; @@ -248,7 +347,8 @@ fn tool_call_with_subscriber_snapshot( params.name, params.args, &entries, - ); + ) + .await; let (handle, event, marks) = { let context = global_context(); let state = context @@ -286,14 +386,16 @@ fn tool_call_with_subscriber_snapshot( .collect::>(); (handle, event, marks) }; - let marks = marks - .into_iter() - .filter_map(sanitize_event) - .collect::>(); - if let Some(event) = sanitize_event(event) { + let mut sanitized_marks = Vec::with_capacity(marks.len()); + for mark in marks { + if let Some(mark) = sanitize_event(mark).await { + sanitized_marks.push(mark); + } + } + if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, &subscribers); } - for mark in marks { + for mark in sanitized_marks { NemoRelayContextState::emit_event(&mark, &subscribers); } Ok((handle, subscribers)) @@ -326,10 +428,69 @@ fn tool_call_with_subscriber_snapshot( /// Sanitize-response guardrails affect only the emitted end-event payload, not /// the caller-owned `result` value. pub fn tool_call_end(params: ToolCallEndParams<'_>) -> Result<()> { - tool_call_end_with_pending_marks(params, Vec::new(), None) + ensure_runtime_owner()?; + let scope_stack = current_scope_stack(); + let (entries, subscribers) = { + let scope_guard = scope_stack.read().expect("scope stack lock poisoned"); + let scope_locals = scope_guard.collect_scope_local_registries(|registries| { + ®istries.tool_sanitize_response_guardrails + }); + let subscribers = + snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?; + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + ( + state.tool_sanitize_response_entries(&scope_locals), + subscribers, + ) + }; + let result = params.result; + let fallback = params.data; + let event = { + let context = global_context(); + let state = context + .read() + .map_err(|error| FlowError::Internal(error.to_string()))?; + state.build_tool_end_event( + EndToolHandleParams::builder() + .handle(params.handle) + .data(Json::Null) + .metadata_opt(params.metadata) + .timestamp_opt(params.timestamp) + .build(), + ) + }; + let tool_name = params.handle.name.clone(); + if let Some(event_sanitizers) = snapshot_event_sanitizers(&event, &scope_stack) { + dispatch_transformed_event( + event, + Box::new(move |mut event| { + Box::pin(async move { + let sanitized = NemoRelayContextState::tool_sanitize_response_snapshot_chain( + &tool_name, result, &entries, + ) + .await; + let mut fields = event.sanitize_fields(); + fields.data = if sanitized.is_null() { + fallback + } else { + Some(sanitized) + }; + event.apply_sanitize_fields(fields); + event + }) + }), + event_sanitizers, + &subscribers, + scope_stack, + ); + } + Ok(()) } -fn tool_call_end_with_pending_marks( +async fn tool_call_end_with_pending_marks( params: ToolCallEndParams<'_>, pending_marks: Vec, lifecycle_subscribers: Option<&[EventSubscriberFn]>, @@ -358,7 +519,8 @@ fn tool_call_end_with_pending_marks( ¶ms.handle.name, params.result, &entries, - ); + ) + .await; let data = if sanitized_result.is_null() { params.data } else { @@ -396,18 +558,23 @@ fn tool_call_end_with_pending_marks( mark.category_profile, )) }) - .filter_map(sanitize_event) .collect::>(); - if let Some(event) = sanitize_event(event) { + let mut sanitized_marks = Vec::with_capacity(marks.len()); + for mark in marks { + if let Some(mark) = sanitize_event(mark).await { + sanitized_marks.push(mark); + } + } + if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, subscribers); } - for mark in marks { + for mark in sanitized_marks { NemoRelayContextState::emit_event(&mark, subscribers); } Ok(()) } -fn emit_tool_end_without_output( +async fn emit_tool_end_without_output( handle: &ToolHandle, metadata: Option, lifecycle_subscribers: &[EventSubscriberFn], @@ -420,7 +587,7 @@ fn emit_tool_end_without_output( .map_err(|error| FlowError::Internal(error.to_string()))?; state.end_tool_handle(handle, handle.data.clone(), metadata) }; - if let Some(event) = sanitize_event(event) { + if let Some(event) = sanitize_event(event).await { NemoRelayContextState::emit_event(&event, lifecycle_subscribers); } Ok(()) @@ -493,7 +660,9 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { &subscribers, parent_uuid, guardrail_metadata, - )? { + ) + .await? + { let mut rejection_data = json!({}); if let Some(object) = rejection_data.as_object_mut() { object.insert("rejected".into(), json!(true)); @@ -526,7 +695,8 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { &name, args, &intercept_entries, - )?; + ) + .await?; let (handle, lifecycle_subscribers) = tool_call_with_subscriber_snapshot( ToolCallParams::builder() @@ -537,7 +707,8 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { .data_opt(data.clone()) .metadata_opt(metadata.clone()) .build(), - )?; + ) + .await?; let execution = { let scope_stack = current_scope_stack(); @@ -567,13 +738,15 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { .build(), pending_marks, Some(&lifecycle_subscribers), - )?; + ) + .await?; Ok(result) } Err(error) => { let end_metadata = metadata_with_otel_status(metadata, "ERROR", Some(error.to_string())); - let _ = emit_tool_end_without_output(&handle, end_metadata, &lifecycle_subscribers); + let _ = + emit_tool_end_without_output(&handle, end_metadata, &lifecycle_subscribers).await; Err(error) } } @@ -596,7 +769,7 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result { /// /// # Notes /// Conditional guardrails and execution intercepts are not run by this helper. -pub fn tool_request_intercepts(name: &str, args: Json) -> Result { +pub async fn tool_request_intercepts(name: &str, args: Json) -> Result { ensure_runtime_owner()?; let entries = { let scope_stack = current_scope_stack(); @@ -609,7 +782,7 @@ pub fn tool_request_intercepts(name: &str, args: Json) -> Result { .map_err(|error| FlowError::Internal(error.to_string()))?; state.tool_request_intercept_entries(&scope_locals) }; - NemoRelayContextState::tool_request_intercepts_snapshot_chain(name, args, &entries) + NemoRelayContextState::tool_request_intercepts_snapshot_chain(name, args, &entries).await } /// Run only the tool conditional-execution guardrail chain. @@ -633,7 +806,7 @@ pub fn tool_request_intercepts(name: &str, args: Json) -> Result { /// This helper is useful for preflight checks when the caller needs the /// rejection result without starting a tool span. Guardrail scopes are still /// emitted for the conditional checks themselves. -pub fn tool_conditional_execution(name: &str, args: &Json) -> Result<()> { +pub async fn tool_conditional_execution(name: &str, args: &Json) -> Result<()> { ensure_runtime_owner()?; let (entries, subscribers, parent_uuid) = { let scope_stack = current_scope_stack(); @@ -657,7 +830,9 @@ pub fn tool_conditional_execution(name: &str, args: &Json) -> Result<()> { &subscribers, parent_uuid, None, - )? { + ) + .await? + { return Err(FlowError::GuardrailRejected(error)); } Ok(()) diff --git a/crates/core/src/logging/rotation.rs b/crates/core/src/logging/rotation.rs index 3314377c1..fa55780bb 100644 --- a/crates/core/src/logging/rotation.rs +++ b/crates/core/src/logging/rotation.rs @@ -146,3 +146,7 @@ pub(crate) fn rotated_log_path(base_path: &Path, index: usize) -> PathBuf { } base_path.with_file_name(file_name) } + +#[cfg(test)] +#[path = "../../tests/coverage/logging_rotation_tests.rs"] +mod tests; diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index f1217c746..7481d86ee 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -14,24 +14,27 @@ use std::panic::{AssertUnwindSafe, catch_unwind}; use std::path::{Path, PathBuf}; use std::pin::Pin; use std::ptr; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, OnceLock}; use std::task::{Context, Poll}; use chrono::{DateTime, Utc}; use libloading::{Library, Symbol}; use nemo_relay_plugin::{ - NEMO_RELAY_NATIVE_ABI_VERSION, NemoRelayNativeEventSanitizeCb, - NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, - NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, - NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, - NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, - NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, - NemoRelayNativeLlmStreamV1, NemoRelayNativePluginContext, NemoRelayNativePluginEntry, - NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, - NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, NemoRelayNativeString, - NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, NemoRelayNativeToolJsonCb, - NemoRelayNativeWithScopeStackCb, NemoRelayStatus, + NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, + NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncCompletion, + NemoRelayNativeAsyncMiddlewareCb, NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, + NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, + NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeLlmCodecKind, + NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmRequestCodec, + NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, + NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, + NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, + NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativePluginContext, + NemoRelayNativePluginEntry, NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, + NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, + NemoRelayNativeString, NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, + NemoRelayNativeToolJsonCb, NemoRelayNativeWithScopeStackCb, NemoRelayStatus, }; use semver::{Version, VersionReq}; use serde_json::{Map, Value as Json}; @@ -374,7 +377,14 @@ fn load_one_native_plugin( library_path.display() )) })?; - let status = entry(native_host_api(), &mut plugin); + let mut status = entry(native_host_api(), &mut plugin); + // SDKs compiled against ABI v2 correctly reject a v3 table. Retry + // their entry point with the frozen v2 prefix instead of making a + // runtime upgrade a breaking change for installed native plugins. + if status == NemoRelayStatus::InvalidArg { + drop_native_plugin_descriptor(&mut plugin); + status = entry(native_host_api_legacy(), &mut plugin); + } if status != NemoRelayStatus::Ok { drop_native_plugin_descriptor(&mut plugin); return Err(PluginError::RegistrationFailed(format!( @@ -784,10 +794,19 @@ unsafe extern "C" fn native_llm_response_codec_decode( } fn native_host_api() -> *const NemoRelayNativeHostApiV1 { + static HOST_API: OnceLock = OnceLock::new(); + &HOST_API.get_or_init(build_native_host_api_v3).v1 as *const NemoRelayNativeHostApiV1 +} + +fn native_host_api_legacy() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); + HOST_API.get_or_init(build_native_host_api_legacy) as *const _ +} + +fn build_native_host_api_legacy() -> NemoRelayNativeHostApiV1 { static RELAY_VERSION: &[u8] = concat!(env!("CARGO_PKG_VERSION"), "\0").as_bytes(); - HOST_API.get_or_init(|| NemoRelayNativeHostApiV1 { - abi_version: NEMO_RELAY_NATIVE_ABI_VERSION, + NemoRelayNativeHostApiV1 { + abi_version: NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, struct_size: std::mem::size_of::(), relay_version: RELAY_VERSION.as_ptr().cast(), string_new: native_string_new, @@ -841,7 +860,23 @@ fn native_host_api() -> *const NemoRelayNativeHostApiV1 { native_plugin_context_register_scope_sanitize_start_guardrail, plugin_context_register_scope_sanitize_end_guardrail: native_plugin_context_register_scope_sanitize_end_guardrail, - }) as *const _ + } +} + +fn build_native_host_api_v3() -> NemoRelayNativeHostApiV3 { + let mut v1 = build_native_host_api_legacy(); + v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION; + v1.struct_size = std::mem::size_of::(); + NemoRelayNativeHostApiV3 { + v1, + async_completion_resolve_json: native_async_completion_resolve_json, + async_completion_reject: native_async_completion_reject, + async_completion_is_cancelled: native_async_completion_is_cancelled, + async_completion_release: native_async_completion_release, + async_next_invoke: native_async_next_invoke, + async_next_release: native_async_next_release, + plugin_context_register_async_middleware: native_plugin_context_register_async_middleware, + } } fn read_native_string(value: *const NemoRelayNativeString) -> crate::plugin::Result { @@ -1305,6 +1340,680 @@ fn make_user_data( }) } +/// One-shot state retained by a v3 native async callback. +struct NativeAsyncCompletion { + sender: Mutex>>>, + cancelled: AtomicBool, + // A pending native callback can continue running after its completion + // wakes the awaiting task. Keep the callback's dynamic-library instance + // alive until native code explicitly releases this handle. + _callback_user_data: Option>, +} + +struct NativeAsyncWait { + completion: Arc, + receiver: tokio::sync::oneshot::Receiver>, +} + +impl Drop for NativeAsyncWait { + fn drop(&mut self) { + self.completion.cancelled.store(true, Ordering::Release); + } +} + +enum NativeAsyncNextInner { + Tool(ToolExecutionNextFn), + Llm(LlmExecutionNextFn), + LlmStream(LlmStreamExecutionNextFn), +} + +struct NativeAsyncNext { + inner: NativeAsyncNextInner, + runtime: tokio::runtime::Handle, + // The native callback owns this handle independently of its completion. + // Retaining the library here prevents an unload while it still uses `next`. + _callback_user_data: Option>, +} + +async fn invoke_native_async_callback( + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: Arc, + invocation: Json, + next: Option, +) -> FlowResult { + let runtime = if next.is_some() { + Some(tokio::runtime::Handle::try_current().map_err(|error| { + FlowError::Internal(format!( + "native async intercept requires a Tokio runtime: {error}" + )) + })?) + } else { + None + }; + let invocation = native_string_from_json(&invocation) + .ok_or_else(|| FlowError::Internal("failed to allocate native async invocation".into()))? + as usize; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + _callback_user_data: Some(user_data.clone()), + }); + let completion_ref = Arc::into_raw(completion.clone()) as usize; + let next_ref = match (next, runtime) { + (Some(inner), Some(runtime)) => Some(Arc::into_raw(Arc::new(NativeAsyncNext { + inner, + runtime, + _callback_user_data: Some(user_data.clone()), + })) as usize), + (None, None) => None, + _ => unreachable!("runtime is present exactly for native async intercepts"), + }; + let state = match catch_unwind(AssertUnwindSafe(|| unsafe { + cb( + user_data.ptr, + invocation as *const NemoRelayNativeString, + next_ref + .map(|next| next as *const NemoRelayNativeAsyncNext) + .unwrap_or(ptr::null()), + completion_ref as *const NemoRelayNativeAsyncCompletion, + ) + })) { + Ok(state) => state, + Err(_) => { + unsafe { + drop(Arc::from_raw( + completion_ref as *const NativeAsyncCompletion, + )); + if let Some(next_ref) = next_ref { + drop(Arc::from_raw(next_ref as *const NativeAsyncNext)); + } + native_string_free(invocation as *mut NemoRelayNativeString); + } + return Err(FlowError::Internal("native async callback panicked".into())); + } + }; + unsafe { native_string_free(invocation as *mut NemoRelayNativeString) }; + if state == NemoRelayNativeAsyncCallbackState::Complete { + unsafe { + drop(Arc::from_raw( + completion_ref as *const NativeAsyncCompletion, + )); + if let Some(next_ref) = next_ref { + drop(Arc::from_raw(next_ref as *const NativeAsyncNext)); + } + } + if completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .is_some() + { + return Err(FlowError::Internal( + "native async callback returned Complete without settling".into(), + )); + } + } + let mut wait = NativeAsyncWait { + completion, + receiver, + }; + (&mut wait.receiver) + .await + .map_err(|_| FlowError::Internal("native async callback dropped without settling".into()))? +} + +unsafe extern "C" fn native_async_completion_resolve_json( + completion: *const NemoRelayNativeAsyncCompletion, + value_json: *const NemoRelayNativeString, +) -> NemoRelayStatus { + let Some(completion) = (unsafe { (completion as *const NativeAsyncCompletion).as_ref() }) + else { + return NemoRelayStatus::NullPointer; + }; + if completion.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let value = match parse_json_arg(value_json, "native async completion result") { + Ok(value) => value, + Err(status) => return status, + }; + let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + else { + return NemoRelayStatus::InvalidArg; + }; + let _ = sender.send(Ok(value)); + NemoRelayStatus::Ok +} + +unsafe extern "C" fn native_async_completion_reject( + completion: *const NemoRelayNativeAsyncCompletion, + message: *const NemoRelayNativeString, +) -> NemoRelayStatus { + let Some(completion) = (unsafe { (completion as *const NativeAsyncCompletion).as_ref() }) + else { + return NemoRelayStatus::NullPointer; + }; + if completion.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let message = if message.is_null() { + "native async callback rejected".to_string() + } else { + match read_native_string(message) { + Ok(message) => message, + Err(error) => { + set_native_last_error(error.to_string()); + return NemoRelayStatus::InvalidArg; + } + } + }; + let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + else { + return NemoRelayStatus::InvalidArg; + }; + let _ = sender.send(Err(FlowError::Internal(message))); + NemoRelayStatus::Ok +} + +unsafe extern "C" fn native_async_completion_is_cancelled( + completion: *const NemoRelayNativeAsyncCompletion, +) -> bool { + unsafe { (completion as *const NativeAsyncCompletion).as_ref() } + .is_none_or(|completion| completion.cancelled.load(Ordering::Acquire)) +} + +unsafe extern "C" fn native_async_completion_release( + completion: *const NemoRelayNativeAsyncCompletion, +) { + if !completion.is_null() { + unsafe { drop(Arc::from_raw(completion as *const NativeAsyncCompletion)) }; + } +} + +unsafe extern "C" fn native_async_next_release(next: *const NemoRelayNativeAsyncNext) { + if !next.is_null() { + unsafe { drop(Arc::from_raw(next as *const NativeAsyncNext)) }; + } +} + +/// Invokes the runtime continuation without blocking the calling native thread. +unsafe extern "C" fn native_async_next_invoke( + next: *const NemoRelayNativeAsyncNext, + invocation_json: *const NemoRelayNativeString, + completion: *const NemoRelayNativeAsyncCompletion, +) -> NemoRelayStatus { + let Some(next) = (unsafe { (next as *const NativeAsyncNext).as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if completion.is_null() { + return NemoRelayStatus::NullPointer; + } + let invocation = match parse_json_arg(invocation_json, "native async next invocation") { + Ok(value) => value, + Err(status) => return status, + }; + unsafe { Arc::increment_strong_count(completion as *const NativeAsyncCompletion) }; + let completion = unsafe { Arc::from_raw(completion as *const NativeAsyncCompletion) }; + let future: Pin> + Send>> = match &next.inner { + NativeAsyncNextInner::Tool(next) => { + let next = next.clone(); + Box::pin(async move { + serde_json::to_value(ToolExecutionInterceptOutcome::new(next(invocation).await?)) + .map_err(|error| { + FlowError::Internal(format!( + "failed to serialize native async tool outcome: {error}" + )) + }) + }) + } + NativeAsyncNextInner::Llm(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(error) => { + let _ = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); + return NemoRelayStatus::InvalidArg; + } + }; + let next = next.clone(); + Box::pin(async move { next(request).await }) + } + NativeAsyncNextInner::LlmStream(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(error) => { + let _ = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); + return NemoRelayStatus::InvalidArg; + } + }; + let next = next.clone(); + Box::pin(async move { + let mut stream = next(request).await?; + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk?); + } + Ok(Json::Array(chunks)) + }) + } + }; + next.runtime.spawn(async move { + let result = future.await; + if let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + { + let _ = sender.send(result); + } + }); + NemoRelayStatus::Ok +} + +fn wrap_native_async_tool_json( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> ToolSanitizeFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, value| { + let user_data = user_data.clone(); + Box::pin(async move { + let value = invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "value": value}), + None, + ) + .await?; + Ok(value) + }) + }) +} + +fn wrap_native_async_tool_conditional( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> ToolConditionalFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, value| { + let user_data = user_data.clone(); + Box::pin(async move { + match invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "value": value}), + None, + ) + .await? + { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => Err(FlowError::Internal(format!( + "native async tool conditional callback returned {other}; expected string or null" + ))), + } + }) + }) +} + +fn wrap_native_async_llm_conditional( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmConditionalFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |request| { + let user_data = user_data.clone(); + Box::pin(async move { + match invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"request": request}), + None, + ) + .await? + { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => Err(FlowError::Internal(format!( + "native async LLM conditional callback returned {other}; expected string or null" + ))), + } + }) + }) +} + +fn wrap_native_async_llm_sanitize_request( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmSanitizeRequestFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |request, context| { + let user_data = user_data.clone(); + let codec = format!("{:?}", context.codec()); + Box::pin(async move { + let value = invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"request": request, "context": {"codec": codec}}), + None, + ) + .await?; + if value.is_null() { + Ok(None) + } else { + serde_json::from_value(value) + .map(Some) + .map_err(|error| FlowError::Internal(error.to_string())) + } + }) + }) +} + +fn wrap_native_async_llm_sanitize_response( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmSanitizeResponseFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |response, context| { + let user_data = user_data.clone(); + let codec = format!("{:?}", context.codec()); + Box::pin(async move { + let value = invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"response": response, "context": {"codec": codec}}), + None, + ) + .await?; + Ok((!value.is_null()).then_some(value)) + }) + }) +} + +fn wrap_native_async_llm_request_intercept( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmRequestInterceptFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, request, annotated| { + let user_data = user_data.clone(); + Box::pin(async move { + serde_json::from_value( + invoke_native_async_callback( + cb, + user_data, + serde_json::json!({ + "name": name, + "request": request, + "annotated": annotated, + }), + None, + ) + .await?, + ) + .map_err(|error| { + FlowError::Internal(format!( + "invalid native async LLM intercept outcome: {error}" + )) + }) + }) + }) +} + +fn wrap_native_async_event_sanitize( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> EventSanitizeFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |event, fields| { + let user_data = user_data.clone(); + Box::pin(async move { + serde_json::from_value( + invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"event": event, "fields": fields}), + None, + ) + .await?, + ) + .map_err(|error| { + FlowError::Internal(format!("invalid native async event fields: {error}")) + }) + }) + }) +} + +fn wrap_native_async_tool_execution( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> ToolExecutionFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, args, next| { + let user_data = user_data.clone(); + let invocation = serde_json::json!({"name": name, "value": args}); + Box::pin(async move { + serde_json::from_value( + invoke_native_async_callback( + cb, + user_data, + invocation, + Some(NativeAsyncNextInner::Tool(next)), + ) + .await?, + ) + .map_err(|error| { + FlowError::Internal(format!("invalid native async tool outcome: {error}")) + }) + }) + }) +} + +fn wrap_native_async_llm_execution( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmExecutionFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, request, next| { + let user_data = user_data.clone(); + Box::pin(invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "request": request}), + Some(NativeAsyncNextInner::Llm(next)), + )) + }) +} + +fn wrap_native_async_llm_stream_execution( + instance: Arc, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> LlmStreamExecutionFn { + let user_data = make_user_data(instance, user_data, free_fn); + Arc::new(move |name, request, next| { + let user_data = user_data.clone(); + let name = name.to_owned(); + Box::pin(async move { + let value = invoke_native_async_callback( + cb, + user_data, + serde_json::json!({"name": name, "request": request}), + Some(NativeAsyncNextInner::LlmStream(next)), + ) + .await?; + let chunks = value.as_array().cloned().ok_or_else(|| { + FlowError::Internal( + "native async LLM stream intercept must resolve to an array".into(), + ) + })?; + Ok(LlmJsonStream::new(tokio_stream::iter( + chunks.into_iter().map(Ok), + ))) + }) + }) +} + +unsafe extern "C" fn native_plugin_context_register_async_middleware( + ctx: *mut NemoRelayNativePluginContext, + kind: NemoRelayNativeAsyncMiddlewareKind, + name: *const NemoRelayNativeString, + priority: i32, + break_chain: bool, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> NemoRelayStatus { + clear_native_last_error(); + let host_ctx = match host_ctx_mut(ctx) { + Ok(ctx) => ctx, + Err(status) => return status, + }; + let instance = host_ctx.instance.clone(); + let name = match read_name(name) { + Ok(name) => name, + Err(status) => return status, + }; + let context = unsafe { &mut *host_ctx.ctx }; + let registration = match kind { + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest => context + .register_tool_sanitize_request_guardrail( + &name, + priority, + wrap_native_async_tool_json(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeResponse => context + .register_tool_sanitize_response_guardrail( + &name, + priority, + wrap_native_async_tool_json(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ToolConditionalExecution => context + .register_tool_conditional_execution_guardrail( + &name, + priority, + wrap_native_async_tool_conditional(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ToolRequestIntercept => context + .register_tool_request_intercept( + &name, + priority, + break_chain, + wrap_native_async_tool_json(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ToolExecutionIntercept => context + .register_tool_execution_intercept( + &name, + priority, + wrap_native_async_tool_execution(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeRequest => context + .register_llm_sanitize_request_guardrail( + &name, + priority, + wrap_native_async_llm_sanitize_request(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeResponse => context + .register_llm_sanitize_response_guardrail( + &name, + priority, + wrap_native_async_llm_sanitize_response(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmConditionalExecution => context + .register_llm_conditional_execution_guardrail( + &name, + priority, + wrap_native_async_llm_conditional(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept => { + if let Err(error) = validate_annotated_request_consumer_compatibility( + &instance.relay_compat, + &instance.plugin_kind, + ) { + return status_from_plugin_error(error); + } + context.register_llm_request_intercept( + &name, + priority, + break_chain, + wrap_native_async_llm_request_intercept(instance, cb, user_data, free_fn), + ) + } + NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept => context + .register_llm_execution_intercept( + &name, + priority, + wrap_native_async_llm_execution(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept => context + .register_llm_stream_execution_intercept( + &name, + priority, + wrap_native_async_llm_stream_execution(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::MarkSanitize => context + .register_mark_sanitize_guardrail( + &name, + priority, + wrap_native_async_event_sanitize(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeStart => context + .register_scope_sanitize_start_guardrail( + &name, + priority, + wrap_native_async_event_sanitize(instance, cb, user_data, free_fn), + ), + NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeEnd => context + .register_scope_sanitize_end_guardrail( + &name, + priority, + wrap_native_async_event_sanitize(instance, cb, user_data, free_fn), + ), + }; + match registration { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_plugin_error(error), + } +} + fn host_ctx_mut<'a>( ctx: *mut NemoRelayNativePluginContext, ) -> Result<&'a mut NativeHostPluginContext, NemoRelayStatus> { @@ -1740,7 +2449,8 @@ fn wrap_event_sanitize_fn( ) -> EventSanitizeFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |event, fields| { - call_event_sanitize_callback(cb, user_data.ptr, event, &fields).unwrap_or_default() + let user_data = user_data.clone(); + Box::pin(async move { call_event_sanitize_callback(cb, user_data.ptr, &event, &fields) }) }) } @@ -1795,7 +2505,8 @@ fn wrap_tool_json_fn( ) -> ToolSanitizeFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |name, payload| { - call_tool_json_callback(cb, user_data.ptr, name, &payload).unwrap_or(Json::Null) + let user_data = user_data.clone(); + Box::pin(async move { call_tool_json_callback(cb, user_data.ptr, &name, &payload) }) }) } @@ -1806,7 +2517,10 @@ fn wrap_tool_intercept_fn( free_fn: NemoRelayNativeFreeFn, ) -> ToolInterceptFn { let user_data = make_user_data(instance, user_data, free_fn); - Arc::new(move |name, payload| call_tool_json_callback(cb, user_data.ptr, name, &payload)) + Arc::new(move |name, payload| { + let user_data = user_data.clone(); + Box::pin(async move { call_tool_json_callback(cb, user_data.ptr, &name, &payload) }) + }) } fn call_tool_json_callback( @@ -1846,32 +2560,35 @@ fn wrap_tool_conditional_fn( ) -> ToolConditionalFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |name, args| { - clear_native_last_error(); - let name_string = native_string_from_str(name) - .ok_or_else(|| FlowError::Internal("failed to allocate native name".into()))?; - let args_string = native_string_from_json(args) - .ok_or_else(|| FlowError::Internal("failed to allocate native args".into()))?; - let mut out = ptr::null_mut(); - let status = unsafe { cb(user_data.ptr, name_string, args_string, &mut out) }; - unsafe { - native_string_free(name_string); - native_string_free(args_string); - } - if status != NemoRelayStatus::Ok { - if !out.is_null() { - unsafe { native_string_free(out) }; + let user_data = user_data.clone(); + Box::pin(async move { + clear_native_last_error(); + let name_string = native_string_from_str(&name) + .ok_or_else(|| FlowError::Internal("failed to allocate native name".into()))?; + let args_string = native_string_from_json(&args) + .ok_or_else(|| FlowError::Internal("failed to allocate native args".into()))?; + let mut out = ptr::null_mut(); + let status = unsafe { cb(user_data.ptr, name_string, args_string, &mut out) }; + unsafe { + native_string_free(name_string); + native_string_free(args_string); } - return Err(flow_error_from_status( - status, - "native tool conditional failed", - )); - } - if out.is_null() { - Ok(None) - } else { - let reason = take_native_string(out)?; - Ok(Some(reason)) - } + if status != NemoRelayStatus::Ok { + if !out.is_null() { + unsafe { native_string_free(out) }; + } + return Err(flow_error_from_status( + status, + "native tool conditional failed", + )); + } + if out.is_null() { + Ok(None) + } else { + let reason = take_native_string(out)?; + Ok(Some(reason)) + } + }) }) } @@ -1961,9 +2678,10 @@ fn wrap_llm_sanitize_request_fn( ) -> LlmSanitizeRequestFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |request, context| { - call_llm_sanitize_request_callback(cb, user_data.ptr, &request, context) - .ok() - .flatten() + let user_data = user_data.clone(); + Box::pin( + async move { call_llm_sanitize_request_callback(cb, user_data.ptr, &request, context) }, + ) }) } @@ -1975,9 +2693,10 @@ fn wrap_llm_sanitize_response_fn( ) -> LlmSanitizeResponseFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |payload, context| { - call_llm_sanitize_response_callback(cb, user_data.ptr, &payload, context) - .ok() - .flatten() + let user_data = user_data.clone(); + Box::pin(async move { + call_llm_sanitize_response_callback(cb, user_data.ptr, &payload, context) + }) }) } @@ -2118,30 +2837,34 @@ fn wrap_llm_conditional_fn( ) -> LlmConditionalFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |request| { - clear_native_last_error(); - let request_json = serde_json::to_value(request).map_err(|err| { - FlowError::Internal(format!("failed to serialize LLM request: {err}")) - })?; - let request_string = native_string_from_json(&request_json) - .ok_or_else(|| FlowError::Internal("failed to allocate native LLM request".into()))?; - let mut out = ptr::null_mut(); - let status = unsafe { cb(user_data.ptr, request_string, &mut out) }; - unsafe { native_string_free(request_string) }; - if status != NemoRelayStatus::Ok { - if !out.is_null() { - unsafe { native_string_free(out) }; + let user_data = user_data.clone(); + Box::pin(async move { + clear_native_last_error(); + let request_json = serde_json::to_value(request).map_err(|err| { + FlowError::Internal(format!("failed to serialize LLM request: {err}")) + })?; + let request_string = native_string_from_json(&request_json).ok_or_else(|| { + FlowError::Internal("failed to allocate native LLM request".into()) + })?; + let mut out = ptr::null_mut(); + let status = unsafe { cb(user_data.ptr, request_string, &mut out) }; + unsafe { native_string_free(request_string) }; + if status != NemoRelayStatus::Ok { + if !out.is_null() { + unsafe { native_string_free(out) }; + } + return Err(flow_error_from_status( + status, + "native LLM conditional failed", + )); } - return Err(flow_error_from_status( - status, - "native LLM conditional failed", - )); - } - if out.is_null() { - Ok(None) - } else { - let reason = take_native_string(out)?; - Ok(Some(reason)) - } + if out.is_null() { + Ok(None) + } else { + let reason = take_native_string(out)?; + Ok(Some(reason)) + } + }) }) } @@ -2153,58 +2876,62 @@ fn wrap_llm_request_intercept_fn( ) -> LlmRequestInterceptFn { let user_data = make_user_data(instance, user_data, free_fn); Arc::new(move |name, request, annotated| { - clear_native_last_error(); - let name_string = native_string_from_str(name) - .ok_or_else(|| FlowError::Internal("failed to allocate native name".into()))?; - let request_json = serde_json::to_value(&request).map_err(|err| { - FlowError::Internal(format!("failed to serialize LLM request: {err}")) - })?; - let request_string = native_string_from_json(&request_json) - .ok_or_else(|| FlowError::Internal("failed to allocate native LLM request".into()))?; - let annotated_string = match &annotated { - Some(annotated) => { - let value = serde_json::to_value(annotated).map_err(|err| { - FlowError::Internal(format!("failed to serialize annotated request: {err}")) - })?; - native_string_from_json(&value).ok_or_else(|| { - FlowError::Internal("failed to allocate annotated request".into()) - })? + let user_data = user_data.clone(); + Box::pin(async move { + clear_native_last_error(); + let name_string = native_string_from_str(&name) + .ok_or_else(|| FlowError::Internal("failed to allocate native name".into()))?; + let request_json = serde_json::to_value(&request).map_err(|err| { + FlowError::Internal(format!("failed to serialize LLM request: {err}")) + })?; + let request_string = native_string_from_json(&request_json).ok_or_else(|| { + FlowError::Internal("failed to allocate native LLM request".into()) + })?; + let annotated_string = match &annotated { + Some(annotated) => { + let value = serde_json::to_value(annotated).map_err(|err| { + FlowError::Internal(format!("failed to serialize annotated request: {err}")) + })?; + native_string_from_json(&value).ok_or_else(|| { + FlowError::Internal("failed to allocate annotated request".into()) + })? + } + None => ptr::null_mut(), + }; + let mut out_outcome = ptr::null_mut(); + let status = unsafe { + cb( + user_data.ptr, + name_string, + request_string, + annotated_string, + &mut out_outcome, + ) + }; + unsafe { + native_string_free(name_string); + native_string_free(request_string); + native_string_free(annotated_string); } - None => ptr::null_mut(), - }; - let mut out_outcome = ptr::null_mut(); - let status = unsafe { - cb( - user_data.ptr, - name_string, - request_string, - annotated_string, - &mut out_outcome, - ) - }; - unsafe { - native_string_free(name_string); - native_string_free(request_string); - native_string_free(annotated_string); - } - if status != NemoRelayStatus::Ok { + if status != NemoRelayStatus::Ok { + unsafe { + native_string_free(out_outcome); + } + return Err(flow_error_from_status( + status, + "native LLM request intercept failed", + )); + } + let outcome_json = json_from_native_string( + out_outcome, + "native LLM request intercept returned null outcome", + ); unsafe { native_string_free(out_outcome); } - return Err(flow_error_from_status( - status, - "native LLM request intercept failed", - )); - } - let outcome_json = json_from_native_string( - out_outcome, - "native LLM request intercept returned null outcome", - ); - unsafe { - native_string_free(out_outcome); - } - serde_json::from_value::(outcome_json?).map_err(|err| { - FlowError::Internal(format!("invalid LLM request intercept outcome JSON: {err}")) + serde_json::from_value::(outcome_json?).map_err(|err| { + FlowError::Internal(format!("invalid LLM request intercept outcome JSON: {err}")) + }) }) }) } diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index 79e2e70ce..5262dd784 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -59,9 +59,9 @@ use tower::service_fn; use crate::api::event::{Event, EventSanitizeFields}; use crate::api::llm::{LLM_REQUEST_INTERCEPT_OUTCOME_SCHEMA, LlmRequest}; use crate::api::runtime::{ - LlmCodecIdentity, LlmExecutionNextFn, LlmJsonStream, LlmSanitizeRequestContext, - LlmSanitizeResponseContext, LlmStreamExecutionNextFn, ToolExecutionNextFn, current_scope_stack, - with_scope_stack, + EventSanitizeFn, LlmCodecIdentity, LlmExecutionNextFn, LlmJsonStream, + LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn, + ToolExecutionNextFn, current_scope_stack, with_scope_stack, }; use crate::api::scope::{ EmitMarkEventParams, PopScopeParams, PushScopeParams, ScopeAttributes, ScopeHandle, ScopeType, @@ -1119,14 +1119,16 @@ impl WorkerPluginInstance { ) -> crate::plugin::Result<()> { let instance = Arc::new(self.clone_for_callback()); let callback_name = name.to_owned(); - let callback = Arc::new(move |event: &Event, _fields: EventSanitizeFields| { - instance - .invoke_event_sanitize(&callback_name, surface, event) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - EventSanitizeFields::default() + let callback: EventSanitizeFn = + Arc::new(move |event: Event, _fields: EventSanitizeFields| { + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_event_sanitize(&callback_name, surface, &event) + .await }) - }); + }); match surface { RegistrationSurface::MarkSanitizeGuardrail => { ctx.register_mark_sanitize_guardrail(name, priority, callback) @@ -1157,18 +1159,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |tool_name, value| { - instance - .invoke_tool_json( - &callback_name, - surface, - tool_name, - value.clone(), - None, - ) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - value - }) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_tool_json(&callback_name, surface, &tool_name, value, None) + .await + }) }), ), RegistrationSurface::ToolSanitizeResponseGuardrail => ctx @@ -1176,18 +1173,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |tool_name, value| { - instance - .invoke_tool_json( - &callback_name, - surface, - tool_name, - value.clone(), - None, - ) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - value - }) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_tool_json(&callback_name, surface, &tool_name, value, None) + .await + }) }), ), RegistrationSurface::ToolConditionalExecutionGuardrail => ctx @@ -1195,7 +1187,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |tool_name, value| { - instance.invoke_tool_guardrail(&callback_name, tool_name, value.clone()) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_tool_guardrail(&callback_name, &tool_name, value) + .await + }) }), ), RegistrationSurface::ToolRequestIntercept => ctx.register_tool_request_intercept( @@ -1203,7 +1201,13 @@ impl WorkerPluginInstance { priority, registration.break_chain, Arc::new(move |tool_name, value| { - instance.invoke_tool_json(&callback_name, surface, tool_name, value, None) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_tool_json(&callback_name, surface, &tool_name, value, None) + .await + }) }), ), RegistrationSurface::ToolExecutionIntercept => ctx.register_tool_execution_intercept( @@ -1244,12 +1248,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |request, context| { - instance - .invoke_llm_sanitize_request(&callback_name, request.clone(), context) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - None - }) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_llm_sanitize_request(&callback_name, request, context) + .await + }) }), ), RegistrationSurface::LlmSanitizeResponseGuardrail => ctx @@ -1257,12 +1262,13 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |value, context| { - instance - .invoke_llm_sanitize_response(&callback_name, value.clone(), context) - .unwrap_or_else(|_| { - instance.log_callback_fallback(&callback_name, surface); - None - }) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_llm_sanitize_response(&callback_name, value, context) + .await + }) }), ), RegistrationSurface::LlmConditionalExecutionGuardrail => ctx @@ -1270,7 +1276,11 @@ impl WorkerPluginInstance { name, priority, Arc::new(move |request| { - instance.invoke_llm_guardrail(&callback_name, request.clone()) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance.invoke_llm_guardrail(&callback_name, request).await + }) }), ), RegistrationSurface::LlmRequestIntercept => ctx.register_llm_request_intercept( @@ -1278,12 +1288,18 @@ impl WorkerPluginInstance { priority, registration.break_chain, Arc::new(move |model_name, request, annotated| { - instance.invoke_llm_request_intercept( - &callback_name, - model_name, - request, - annotated, - ) + let instance = instance.clone(); + let callback_name = callback_name.clone(); + Box::pin(async move { + instance + .invoke_llm_request_intercept( + &callback_name, + &model_name, + request, + annotated, + ) + .await + }) }), ), RegistrationSurface::LlmExecutionIntercept => ctx.register_llm_execution_intercept( @@ -1476,7 +1492,7 @@ impl WorkerPluginCallback { } } - fn invoke_event_sanitize( + async fn invoke_event_sanitize( &self, registration_name: &str, surface: RegistrationSurface, @@ -1488,7 +1504,7 @@ impl WorkerPluginCallback { None, Some(invoke_request_payload_event(event)), ); - let value = json_from_invoke_response(self.invoke_blocking(request)?)?; + let value = json_from_invoke_response(self.invoke_async(request).await?)?; serde_json::from_value(value).map_err(|err| { FlowError::Internal(format!( "worker returned invalid event sanitize fields: {err}" @@ -1496,7 +1512,7 @@ impl WorkerPluginCallback { }) } - fn invoke_tool_json( + async fn invoke_tool_json( &self, registration_name: &str, surface: RegistrationSurface, @@ -1510,10 +1526,10 @@ impl WorkerPluginCallback { continuation_id, Some(invoke_request_payload_tool(tool_name, value)), ); - json_from_invoke_response(self.invoke_blocking(request)?) + json_from_invoke_response(self.invoke_async(request).await?) } - fn invoke_tool_guardrail( + async fn invoke_tool_guardrail( &self, registration_name: &str, tool_name: &str, @@ -1525,7 +1541,7 @@ impl WorkerPluginCallback { None, Some(invoke_request_payload_tool(tool_name, value)), ); - guardrail_from_invoke_response(self.invoke_blocking(request)?) + guardrail_from_invoke_response(self.invoke_async(request).await?) } async fn invoke_tool_execution( @@ -1568,7 +1584,7 @@ impl WorkerPluginCallback { } } - fn invoke_llm_sanitize_request( + async fn invoke_llm_sanitize_request( &self, registration_name: &str, request: LlmRequest, @@ -1606,7 +1622,7 @@ impl WorkerPluginCallback { context.codec_capability_id = Some(capability_id.clone()); capability_id }); - let response = self.invoke_blocking(invoke); + let response = self.invoke_async(invoke).await; if let Some(capability_id) = capability_id { self.host_state.remove_codec(&capability_id); } @@ -1618,7 +1634,7 @@ impl WorkerPluginCallback { }) } - fn invoke_llm_sanitize_response( + async fn invoke_llm_sanitize_response( &self, registration_name: &str, response: Json, @@ -1656,14 +1672,14 @@ impl WorkerPluginCallback { context.codec_capability_id = Some(capability_id.clone()); capability_id }); - let response = self.invoke_blocking(invoke); + let response = self.invoke_async(invoke).await; if let Some(capability_id) = capability_id { self.host_state.remove_codec(&capability_id); } optional_json_from_invoke_response(response?) } - fn invoke_llm_guardrail( + async fn invoke_llm_guardrail( &self, registration_name: &str, request: LlmRequest, @@ -1674,10 +1690,10 @@ impl WorkerPluginCallback { None, Some(invoke_request_payload_llm("", Some(request), None, None)), ); - guardrail_from_invoke_response(self.invoke_blocking(invoke)?) + guardrail_from_invoke_response(self.invoke_async(invoke).await?) } - fn invoke_llm_request_intercept( + async fn invoke_llm_request_intercept( &self, registration_name: &str, model_name: &str, @@ -1695,7 +1711,7 @@ impl WorkerPluginCallback { None, )), ); - let response = self.invoke_blocking(invoke)?; + let response = self.invoke_async(invoke).await?; match response.result { Some(invoke_response_result::Result::LlmRequest(result)) => { let outcome = required_envelope(result.outcome, "llm request intercept outcome")?; @@ -1853,8 +1869,22 @@ impl WorkerPluginCallback { } async fn invoke_async(&self, request: InvokeRequest) -> FlowResult { - self.invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) - .await + let callback_name = request.registration_name.clone(); + let surface = request.surface; + let result = self + .invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) + .await; + if let Err(error) = &result { + log::warn!( + target: "nemo_relay.worker", + event = "worker_callback_failed", + plugin_id = self.plugin_kind.as_str(), + callback = callback_name.as_str(), + surface; + "Worker plugin callback failed: {error}" + ); + } + result } async fn invoke_async_with_timeout( diff --git a/crates/core/src/stream.rs b/crates/core/src/stream.rs index 3fee8dd3e..756f066a7 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -39,14 +39,16 @@ use crate::api::optimization::finalize_optimization_summary; use crate::api::runtime::LlmSanitizeResponseContext; use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::global_context; +use crate::api::runtime::subscriber_dispatcher; use crate::api::runtime::{ EventSubscriberFn, LlmJsonStream, LlmStreamInner, ScopeStackHandle, current_scope_stack, }; -use crate::api::shared::metadata_with_otel_status; -use crate::api::shared::sanitize_event_with_scope_stack; +use crate::api::shared::{ + metadata_with_otel_status, sanitize_event_with_scope_stack, snapshot_event_sanitizers, +}; use crate::codec::response::{AnnotatedLlmResponse, attach_estimated_cost_for_provider}; use crate::codec::traits::LlmResponseCodec; -use crate::error::Result; +use crate::error::{FlowError, Result}; use crate::json::Json; use serde_json::Map; @@ -78,6 +80,8 @@ pub struct LlmStreamWrapper { chunk_index: u64, ended: bool, close_result: Option>, + finalization: Option>, + terminal_result: Option>, } impl LlmStreamWrapper { @@ -157,6 +161,8 @@ impl LlmStreamWrapper { chunk_index: 0, ended: false, close_result: None, + finalization: None, + terminal_result: None, } } @@ -182,7 +188,12 @@ impl LlmStreamWrapper { "ERROR", Some("stream dropped before clean completion".to_string()), ); - self.emit_end_event(metadata, true); + // Drop cannot await the async finalizer. Close the recorder before + // spawning it so late optimization evidence is rejected immediately. + self.handle + .optimization_recorder + .close_for_finalization(Some("stream_interrupted")); + self.finalization = self.emit_end_event(metadata, true, true); } fn finish_with_status( @@ -197,14 +208,23 @@ impl LlmStreamWrapper { self.ended = true; let metadata = metadata_with_otel_status(self.metadata.clone(), status_code, status_message); - self.emit_end_event(metadata, interrupted); + self.finalization = self.emit_end_event(metadata, interrupted, false); } /// Emit the LLM END event with aggregated response data. /// /// Calls the finalizer to produce the aggregated response, runs sanitize /// response guardrails, and emits the END event. - fn emit_end_event(&mut self, metadata: Option, interrupted: bool) { + fn emit_end_event( + &mut self, + metadata: Option, + interrupted: bool, + background_thread: bool, + ) -> Option> { + // The finalizer below runs on the caller's Tokio runtime. Register a + // dispatcher barrier before spawning it so a synchronous subscriber + // flush after this stream is dropped cannot overtake the END event. + let publication_barrier = subscriber_dispatcher::register_async_publication(); let aggregated = match self.finalizer.take() { Some(finalizer) => finalizer(), None => Json::Null, @@ -230,68 +250,108 @@ impl LlmStreamWrapper { Err(_) => None, } }; - let Some(entries) = snapshot else { - return; - }; - let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( - response, - self.sanitize_context.clone(), - &entries, - ); - let data = match sanitized { - Some(response) if response_was_null_without_fallback && response.is_null() => None, - response => response, - }; - let annotation_omitted = data.as_ref().is_none_or(Json::is_null); - let mut annotated_response: Option = (!annotation_omitted) - .then(|| { - data.as_ref().and_then(|response| { - self.response_codec.as_ref().and_then(|codec| { - let mut decoded = codec.decode_response(response).ok()?; - attach_estimated_cost_for_provider(&mut decoded, Some(&self.handle.name)); - Some(decoded) + let entries = snapshot?; + let handle = self.handle.clone(); + let scope_stack = self.scope_stack.clone(); + let subscribers = self.subscribers.clone(); + let response_codec = self.response_codec.clone(); + let sanitize_context = self.sanitize_context.clone(); + let finalize = async move { + let sanitized = NemoRelayContextState::llm_sanitize_response_snapshot_chain( + response, + sanitize_context, + &entries, + ) + .await; + let data = match sanitized { + Some(response) if response_was_null_without_fallback && response.is_null() => None, + response => response, + }; + let annotation_omitted = data.as_ref().is_none_or(Json::is_null); + let mut annotated_response: Option = (!annotation_omitted) + .then(|| { + data.as_ref().and_then(|response| { + response_codec.as_ref().and_then(|codec| { + let mut decoded = codec.decode_response(response).ok()?; + attach_estimated_cost_for_provider(&mut decoded, Some(&handle.name)); + Some(decoded) + }) }) }) - }) - .flatten(); - let interruption = (interrupted - && !has_authoritative_final_usage(annotated_response.as_ref())) - .then_some("stream_interrupted"); - self.handle - .optimization_recorder - .close_for_finalization(interruption); - emit_optimization_marks(&self.handle, &self.subscribers); - let pricing = crate::codec::response::active_pricing_resolver(); - let summary = finalize_optimization_summary( - &self.handle.optimization_recorder, - annotated_response.as_mut(), - self.handle.model_name.as_deref(), - &pricing, - ); - if !annotation_omitted - && annotated_response.is_none() - && let Some(summary) = summary - { - annotated_response = Some(AnnotatedLlmResponse { - optimization_summary: Some(summary), - ..AnnotatedLlmResponse::default() + .flatten(); + let interruption = (interrupted + && !has_authoritative_final_usage(annotated_response.as_ref())) + .then_some("stream_interrupted"); + handle + .optimization_recorder + .close_for_finalization(interruption); + emit_optimization_marks(&handle, &subscribers).await; + let pricing = crate::codec::response::active_pricing_resolver(); + let summary = finalize_optimization_summary( + &handle.optimization_recorder, + annotated_response.as_mut(), + handle.model_name.as_deref(), + &pricing, + ); + if !annotation_omitted + && annotated_response.is_none() + && let Some(summary) = summary + { + annotated_response = Some(AnnotatedLlmResponse { + optimization_summary: Some(summary), + ..AnnotatedLlmResponse::default() + }); + } + let annotated_response = annotated_response.map(Arc::new); + let event_snapshot = { + let ctx = global_context(); + let state = ctx.read(); + match state { + Ok(state) => { + Some(state.end_llm_handle(&handle, data, metadata, annotated_response)) + } + Err(_) => None, + } + }; + if let Some(event) = event_snapshot + && let Some(event) = sanitize_event_with_scope_stack(event, &scope_stack).await + { + let _ = subscriber_dispatcher::dispatch_sanitized_event( + event, + Vec::new(), + &subscribers, + scope_stack.clone(), + ); + } + if let Some(done) = publication_barrier { + let _ = done.send(()); + } + }; + if background_thread { + // `Drop` can run while the current-thread Tokio executor is + // synchronously flushing subscribers. Use a dedicated runtime so + // the FIFO publication barrier can still be released. + std::thread::spawn(move || { + if let Ok(runtime) = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + { + runtime.block_on(finalize); + } }); + return None; } - let annotated_response = annotated_response.map(Arc::new); - let event_snapshot = { - let ctx = global_context(); - let state = ctx.read(); - match state { - Ok(state) => { - Some(state.end_llm_handle(&self.handle, data, metadata, annotated_response)) + match tokio::runtime::Handle::try_current() { + Ok(handle) => Some(handle.spawn(finalize)), + Err(_) => { + if let Ok(runtime) = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + { + runtime.block_on(finalize); } - Err(_) => None, + None } - }; - if let Some(event) = event_snapshot - && let Some(event) = sanitize_event_with_scope_stack(event, &self.scope_stack) - { - NemoRelayContextState::emit_event(&event, &self.subscribers); } } @@ -318,9 +378,14 @@ impl LlmStreamWrapper { } }; if let Some(event) = event_snapshot - && let Some(event) = sanitize_event_with_scope_stack(event, &self.scope_stack) + && let Some(sanitizers) = snapshot_event_sanitizers(&event, &self.scope_stack) { - NemoRelayContextState::emit_event(&event, &self.subscribers); + let _ = subscriber_dispatcher::dispatch_sanitized_event( + event, + sanitizers, + &self.subscribers, + self.scope_stack.clone(), + ); } } } @@ -328,8 +393,31 @@ impl LlmStreamWrapper { impl Stream for LlmStreamWrapper { type Item = Result; - fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - let this = self.get_mut(); + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.as_mut().get_mut(); + + // The END event runs async because response and event sanitizers may + // await. Do not expose stream termination until that work has queued + // the event: callers commonly flush subscribers immediately after + // exhausting a stream, and that flush must include its END event. + if let Some(finalization) = this.finalization.as_mut() { + return match Pin::new(finalization).poll(cx) { + Poll::Pending => Poll::Pending, + Poll::Ready(Ok(())) => { + this.finalization = None; + match this.terminal_result.take() { + Some(result) => Poll::Ready(Some(result)), + None => Poll::Ready(None), + } + } + Poll::Ready(Err(error)) => { + this.finalization = None; + Poll::Ready(Some(Err(FlowError::Internal(format!( + "stream finalization task failed: {error}" + ))))) + } + }; + } if this.ended { return Poll::Ready(None); @@ -346,19 +434,21 @@ impl Stream for LlmStreamWrapper { Ok(()) => Poll::Ready(Some(Ok(raw_chunk))), Err(e) => { let message = e.to_string(); + this.terminal_result = Some(Err(e)); this.finish_with_status("ERROR", Some(message), true); - Poll::Ready(Some(Err(e))) + self.poll_next(cx) } } } Poll::Ready(Some(Err(e))) => { let message = e.to_string(); + this.terminal_result = Some(Err(e)); this.finish_with_status("ERROR", Some(message), true); - Poll::Ready(Some(Err(e))) + self.poll_next(cx) } Poll::Ready(None) => { this.finish_with_status("OK", None, false); - Poll::Ready(None) + self.poll_next(cx) } Poll::Pending => Poll::Pending, } @@ -374,6 +464,11 @@ impl LlmStreamInner for LlmStreamWrapper { } let result = this.inner.close().await; this.finish(); + if let Some(finalization) = this.finalization.take() { + finalization.await.map_err(|error| { + FlowError::Internal(format!("stream finalization task failed: {error}")) + })?; + } this.close_result = Some(result.clone()); this.close_result .as_ref() diff --git a/crates/core/tests/coverage/logging_rotation_tests.rs b/crates/core/tests/coverage/logging_rotation_tests.rs new file mode 100644 index 000000000..842fa7927 --- /dev/null +++ b/crates/core/tests/coverage/logging_rotation_tests.rs @@ -0,0 +1,34 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use super::*; + +#[test] +fn rotating_writer_rotates_retains_and_reports_closed_file_errors() { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("nested").join("relay.log"); + let mut writer = SizeRotatingFileWriter::new(path.clone(), 4, 2).unwrap(); + assert_eq!(writer.write(b"abcd").unwrap(), 4); + writer.flush().unwrap(); + assert_eq!(writer.write(b"e").unwrap(), 1); + writer.flush().unwrap(); + + assert_eq!(std::fs::read(rotated_log_path(&path, 1)).unwrap(), b"abcd"); + assert_eq!(std::fs::read(&path).unwrap(), b"e"); + + writer.file = None; + assert!(writer.write(b"x").is_err()); + assert!(writer.flush().is_err()); +} + +#[test] +fn rotation_helpers_handle_empty_and_missing_files() { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("missing.log"); + rotate_files(&path, 2).unwrap(); + assert_eq!( + rotated_log_path(&path, 2), + directory.path().join("missing.2.log") + ); + create_parent_directory(std::path::Path::new("plain.log")).unwrap(); +} diff --git a/crates/core/tests/coverage/logging_sink_tests.rs b/crates/core/tests/coverage/logging_sink_tests.rs index 044ce2b47..88f142655 100644 --- a/crates/core/tests/coverage/logging_sink_tests.rs +++ b/crates/core/tests/coverage/logging_sink_tests.rs @@ -2,10 +2,16 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - DROP_REPORT_INTERVAL_MILLIS, DropNoticeRateLimiter, dropped_record_error_handler, - log_level_filter, now_millis, spdlog_level, stderr_error_handler, + DROP_REPORT_INTERVAL_MILLIS, DropNoticeRateLimiter, build_logger, dropped_record_error_handler, + log_level_filter, now_millis, reserved_sink_paths, resolve_log_path, rotated_log_path, + spdlog_level, stderr_error_handler, }; use crate::logging::LogLevel; +use crate::logging::{ + FileLogRotationConfig, FileLogSinkConfig, LogFormat, LogSinkConfig, LoggingConfig, + MAX_FILE_SINK_QUEUE_ENTRIES, +}; +use std::path::PathBuf; #[test] fn drop_notice_rate_limiter_reports_immediately_then_once_per_interval() { @@ -33,3 +39,60 @@ fn sink_helpers_cover_boundary_levels_time_and_emergency_handlers() { "expected test error", ))); } + +#[test] +fn logger_builder_rejects_duplicate_conflicting_and_invalid_file_sinks() { + let directory = tempfile::tempdir().unwrap(); + let log_path = directory.path().join("relay.log"); + let file_sink = |path: PathBuf, rotation| { + LogSinkConfig::File(FileLogSinkConfig { + path, + level: LogLevel::Info, + format: LogFormat::Jsonl, + queue_capacity: 8, + rotation, + }) + }; + + assert!(resolve_log_path(std::path::Path::new("")).is_err()); + let rotation = FileLogRotationConfig::new(32, 1).unwrap(); + assert_eq!(reserved_sink_paths(&log_path, Some(rotation)).len(), 2); + + let duplicate = LoggingConfig { + sinks: vec![ + file_sink(log_path.clone(), None), + file_sink(log_path.clone(), None), + ], + ..LoggingConfig::default() + }; + let error = match build_logger(&duplicate, "root".into()) { + Ok(_) => panic!("duplicate file sinks must be rejected"), + Err(error) => error, + }; + assert!(error.to_string().contains("duplicate logging sink path")); + + let conflict = LoggingConfig { + sinks: vec![ + file_sink(log_path.clone(), Some(rotation)), + file_sink(rotated_log_path(&log_path, 1), None), + ], + ..LoggingConfig::default() + }; + let error = match build_logger(&conflict, "root".into()) { + Ok(_) => panic!("active and rotated file paths must not overlap"), + Err(error) => error, + }; + assert!(error.to_string().contains("conflicts")); + + let mut invalid_capacity = LoggingConfig { + sinks: vec![file_sink(log_path, None)], + ..LoggingConfig::default() + }; + let LogSinkConfig::File(file_sink) = &mut invalid_capacity.sinks[0]; + file_sink.queue_capacity = MAX_FILE_SINK_QUEUE_ENTRIES + 1; + let error = match build_logger(&invalid_capacity, "root".into()) { + Ok(_) => panic!("oversized async queues must be rejected"), + Err(error) => error, + }; + assert!(error.to_string().contains("queue_capacity")); +} diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index 350cf1414..acade0d0c 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -7,7 +7,9 @@ use std::ptr; use nemo_relay_plugin::{ CategoryProfile, ConfigDiagnostic, DiagnosticLevel, Event, EventCategory, EventSanitizeFields, Json, LlmJsonStream, LlmRequest, LlmRequestInterceptOutcome, NemoRelayNativeHostApiV1, - NemoRelayNativePluginContext, NemoRelayNativePluginV1, NemoRelayNativeString, NemoRelayStatus, + NemoRelayNativeHostApiV3, NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncMiddlewareCb, + NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativePluginContext, NemoRelayNativePluginV1, + NemoRelayNativeString, NemoRelayStatus, NemoRelayNativeToolNextFn, NativePlugin, PendingMarkSpec, PluginContext, PluginRuntime, ScopeCategory, ScopeType, ToolExecutionInterceptOutcome, }; @@ -280,6 +282,33 @@ fn mark_json(mut value: Json, key: &str) -> Json { nemo_relay_plugin::nemo_relay_plugin!(nemo_relay_fixture_native_plugin, || FixtureNativePlugin); +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_fixture_async_entry( + host: *const NemoRelayNativeHostApiV1, + out: *mut NemoRelayNativePluginV1, +) -> NemoRelayStatus { + if host.is_null() || out.is_null() { + return NemoRelayStatus::NullPointer; + } + let host_v1 = unsafe { &*host }; + if host_v1.abi_version < 3 + || host_v1.struct_size < std::mem::size_of::() + { + return NemoRelayStatus::InvalidArg; + } + let host_v2 = unsafe { &*(host as *const NemoRelayNativeHostApiV3) }; + let mut plugin = NemoRelayNativePluginV1::default(); + plugin.plugin_kind = unsafe { raw_host_string(&host_v2.v1, "fixture_async") }; + if plugin.plugin_kind.is_null() { + return NemoRelayStatus::Internal; + } + plugin.user_data = Box::into_raw(Box::new(*host_v2)).cast(); + plugin.register = Some(raw_register_async_tool_request); + plugin.drop = Some(raw_drop_async_host); + unsafe { *out = plugin }; + NemoRelayStatus::Ok +} + #[unsafe(no_mangle)] pub unsafe extern "C" fn nemo_relay_fixture_observability_collision( host: *const NemoRelayNativeHostApiV1, @@ -570,6 +599,260 @@ unsafe extern "C" fn raw_register_event_sanitize_errors( status } +unsafe extern "C" fn raw_register_async_tool_request( + user_data: *mut c_void, + _plugin_config_json: *const NemoRelayNativeString, + ctx: *mut NemoRelayNativePluginContext, +) -> NemoRelayStatus { + if user_data.is_null() { + return NemoRelayStatus::NullPointer; + } + let host = unsafe { &*(user_data as *const NemoRelayNativeHostApiV3) }; + let registrations: [ + (NemoRelayNativeAsyncMiddlewareKind, &str, NemoRelayNativeAsyncMiddlewareCb); + 14 + ] = [ + (NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest, "fixture_async_tool_sanitize_request", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeResponse, "fixture_async_tool_sanitize_response", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::ToolConditionalExecution, "fixture_async_tool_conditional", raw_async_allow_callback), + (NemoRelayNativeAsyncMiddlewareKind::ToolRequestIntercept, "fixture_async_request", raw_async_tool_request_callback), + (NemoRelayNativeAsyncMiddlewareKind::ToolExecutionIntercept, "fixture_async_execution", raw_async_tool_execution_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeRequest, "fixture_async_llm_sanitize_request", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeResponse, "fixture_async_llm_sanitize_response", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmConditionalExecution, "fixture_async_llm_conditional", raw_async_allow_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept, "fixture_async_llm_request", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept, "fixture_async_llm_execution", raw_async_tool_execution_callback), + (NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept, "fixture_async_llm_stream", raw_async_tool_execution_callback), + (NemoRelayNativeAsyncMiddlewareKind::MarkSanitize, "fixture_async_mark", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeStart, "fixture_async_scope_start", raw_async_passthrough_callback), + (NemoRelayNativeAsyncMiddlewareKind::ScopeSanitizeEnd, "fixture_async_scope_end", raw_async_passthrough_callback), + ]; + for (kind, registration_name, callback) in registrations { + let name = unsafe { raw_host_string(&host.v1, registration_name) }; + if name.is_null() { + return NemoRelayStatus::Internal; + } + let status = unsafe { + (host.plugin_context_register_async_middleware)( + ctx, kind, name, 0, false, callback, user_data, None, + ) + }; + unsafe { (host.v1.string_free)(name) }; + if status != NemoRelayStatus::Ok { + return status; + } + } + NemoRelayStatus::Ok +} + +unsafe extern "C" fn raw_async_allow_callback( + user_data: *mut c_void, + _invocation_json: *const NemoRelayNativeString, + _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> NemoRelayNativeAsyncCallbackState { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let result = unsafe { raw_host_string(&host.v1, "null") }; + if result.is_null() { + unsafe { reject_async_completion(host, completion, "failed to allocate async allow result") }; + } else { + unsafe { + (host.async_completion_resolve_json)(completion, result); + (host.v1.string_free)(result); + } + } + NemoRelayNativeAsyncCallbackState::Complete +} + +unsafe extern "C" fn raw_async_passthrough_callback( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> NemoRelayNativeAsyncCallbackState { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let result = unsafe { raw_host_string_value(&host.v1, invocation_json) } + .and_then(|value| serde_json::from_str::(&value).ok()) + .and_then(|invocation| { + invocation.get("annotated").map(|annotated| { + json!({ + "request": invocation["request"], + "annotated_request": annotated, + "pending_marks": [], + "optimization_contributions": [], + }) + }).or_else(|| { + ["value", "request", "response", "fields"] + .into_iter() + .find_map(|key| invocation.get(key).cloned()) + }) + }) + .and_then(|value| serde_json::to_string(&value).ok()); + let Some(result) = result else { + unsafe { reject_async_completion(host, completion, "invalid async passthrough invocation") }; + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let result = unsafe { raw_host_string(&host.v1, &result) }; + if result.is_null() { + unsafe { reject_async_completion(host, completion, "failed to allocate async passthrough result") }; + } else { + unsafe { + (host.async_completion_resolve_json)(completion, result); + (host.v1.string_free)(result); + } + } + NemoRelayNativeAsyncCallbackState::Complete +} + +unsafe extern "C" fn raw_async_tool_request_callback( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + _next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> NemoRelayNativeAsyncCallbackState { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let invocation = unsafe { raw_host_string_value(&host.v1, invocation_json) } + .and_then(|json| serde_json::from_str::(&json).ok()) + .and_then(|mut invocation| { + let pending = invocation["name"].as_str() == Some("async-pending"); + let duplicate = invocation["name"].as_str() == Some("async-double"); + invocation + .get_mut("value") + .and_then(Json::as_object_mut) + .map(|value| { + value.insert("native_async".into(), json!(true)); + (Json::Object(value.clone()), pending, duplicate) + }) + }) + .and_then(|(value, pending, duplicate)| { + serde_json::to_string(&value) + .ok() + .map(|value| (value, pending, duplicate)) + }); + let Some((result, pending, duplicate)) = invocation else { + unsafe { reject_async_completion(host, completion, "invalid async tool request invocation") }; + return NemoRelayNativeAsyncCallbackState::Complete; + }; + if pending { + let host = *host; + let completion = completion as usize; + std::thread::spawn(move || { + std::thread::sleep(std::time::Duration::from_millis(10)); + let result = unsafe { raw_host_string(&host.v1, &result) }; + if !result.is_null() { + unsafe { + (host.async_completion_resolve_json)( + completion as *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + result, + ); + (host.v1.string_free)(result); + (host.async_completion_release)( + completion as *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + ); + } + } else { + unsafe { + reject_async_completion( + &host, + completion as *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + "failed to allocate async tool request result", + ); + (host.async_completion_release)( + completion as *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + ); + } + } + }); + return NemoRelayNativeAsyncCallbackState::Pending; + } + let result = unsafe { raw_host_string(&host.v1, &result) }; + if !result.is_null() { + unsafe { + (host.async_completion_resolve_json)(completion, result); + if duplicate { + let _ = (host.async_completion_resolve_json)(completion, result); + } + (host.v1.string_free)(result); + } + } else { + unsafe { reject_async_completion(host, completion, "failed to allocate async tool request result") }; + } + NemoRelayNativeAsyncCallbackState::Complete +} + +unsafe extern "C" fn raw_async_tool_execution_callback( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const nemo_relay_plugin::NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> NemoRelayNativeAsyncCallbackState { + let Some(host) = (unsafe { (user_data as *const NemoRelayNativeHostApiV3).as_ref() }) else { + return NemoRelayNativeAsyncCallbackState::Complete; + }; + if next.is_null() || completion.is_null() { + unsafe { reject_async_completion(host, completion, "async tool execution requires next and completion") }; + return NemoRelayNativeAsyncCallbackState::Complete; + } + let value = unsafe { raw_host_string_value(&host.v1, invocation_json) } + .and_then(|json| serde_json::from_str::(&json).ok()) + .and_then(|mut invocation| { + if let Some(value) = invocation.get_mut("value").and_then(Json::as_object_mut) { + value.insert("native_async_execution".into(), json!(true)); + Some(Json::Object(value.clone())) + } else { + invocation.get("request").cloned() + } + }) + .and_then(|value| serde_json::to_string(&value).ok()); + let Some(value) = value else { + unsafe { reject_async_completion(host, completion, "invalid async tool execution invocation") }; + return NemoRelayNativeAsyncCallbackState::Complete; + }; + let value = unsafe { raw_host_string(&host.v1, &value) }; + if value.is_null() { + unsafe { reject_async_completion(host, completion, "failed to allocate async tool execution invocation") }; + return NemoRelayNativeAsyncCallbackState::Complete; + } + let status = unsafe { (host.async_next_invoke)(next, value, completion) }; + unsafe { + (host.v1.string_free)(value); + } + if status == NemoRelayStatus::Ok { + unsafe { + (host.async_next_release)(next); + (host.async_completion_release)(completion); + } + NemoRelayNativeAsyncCallbackState::Pending + } else { + unsafe { reject_async_completion(host, completion, "failed to invoke async tool execution next") }; + NemoRelayNativeAsyncCallbackState::Complete + } +} + +unsafe fn reject_async_completion( + host: &NemoRelayNativeHostApiV3, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, + message: &str, +) { + if completion.is_null() { + return; + } + let message = unsafe { raw_host_string(&host.v1, message) }; + if message.is_null() { + return; + } + unsafe { + let _ = (host.async_completion_reject)(completion, message); + (host.v1.string_free)(message); + } +} + unsafe extern "C" fn raw_tool_outcome_callback( user_data: *mut c_void, name: *const NemoRelayNativeString, @@ -638,6 +921,12 @@ unsafe extern "C" fn raw_drop_host(user_data: *mut c_void) { } } +unsafe extern "C" fn raw_drop_async_host(user_data: *mut c_void) { + if !user_data.is_null() { + drop(unsafe { Box::from_raw(user_data as *mut NemoRelayNativeHostApiV3) }); + } +} + unsafe fn raw_host_from_user_data<'a>( user_data: *mut c_void, ) -> Option<&'a NemoRelayNativeHostApiV1> { diff --git a/crates/core/tests/integration/api_surface_tests.rs b/crates/core/tests/integration/api_surface_tests.rs index 7a6045ac0..61f56251f 100644 --- a/crates/core/tests/integration/api_surface_tests.rs +++ b/crates/core/tests/integration/api_surface_tests.rs @@ -7,6 +7,9 @@ use std::sync::{Arc, Mutex}; +mod test_support; +use test_support::ready; + use chrono::{DateTime, TimeDelta, Utc}; use futures::StreamExt; use nemo_relay::api::event::{CategoryProfile, Event, ScopeCategory}; @@ -65,6 +68,7 @@ use nemo_relay::api::tool::{ tool_call, tool_call_end, tool_call_execute, tool_conditional_execution, tool_request_intercepts, }; +use nemo_relay::codec::optimization::LlmOptimizationContribution; use nemo_relay::error::{FlowError, Result}; use nemo_relay::json::Json; use serde_json::{Map, json}; @@ -90,7 +94,7 @@ fn event_sanitizers_rewrite_only_observability_fields_in_priority_order() { Arc::new(|event, mut fields| { assert_eq!(event.data().unwrap()["order"], json!(["early"])); fields.data = Some(json!({"order": ["early", "late"]})); - fields + ready(fields) }), ) .unwrap(); @@ -101,7 +105,7 @@ fn event_sanitizers_rewrite_only_observability_fields_in_priority_order() { fields.data = Some(json!({"order": ["early"]})); fields.metadata = Some(json!({"redacted": true})); fields.category_profile = Some(CategoryProfile::builder().subtype("sanitized").build()); - fields + ready(fields) }), ) .unwrap(); @@ -151,7 +155,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 20, Arc::new(|_, mut fields| { fields.data = Some(json!({"mark": "global"})); - fields + ready(fields) }), ) .unwrap(); @@ -160,7 +164,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 20, Arc::new(|_, mut fields| { fields.metadata = Some(json!({"scope_end": true})); - fields + ready(fields) }), ) .unwrap(); @@ -178,7 +182,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 10, Arc::new(|_, mut fields| { fields.data = Some(json!({"mark": "local"})); - fields + ready(fields) }), ) .unwrap(); @@ -188,7 +192,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 10, Arc::new(|_, mut fields| { fields.metadata = Some(json!({"scope_start": true})); - fields + ready(fields) }), ) .unwrap(); @@ -198,7 +202,7 @@ fn mark_and_scope_local_sanitizers_cover_marks_and_tool_scopes() { 10, Arc::new(|_, mut fields| { fields.data = Some(json!({"scope_end": "local"})); - fields + ready(fields) }), ) .unwrap(); @@ -512,7 +516,7 @@ fn skill_load_detection_uses_original_arguments_before_observability_sanitizatio register_tool_sanitize_request_guardrail( "strip-skill-path", 1, - Arc::new(|_name, _args| json!({"path": "[redacted]"})), + Arc::new(|_name, _args| ready(json!({"path": "[redacted]"}))), ) .unwrap(); let events = capture_events("sanitized-skill-load-api-events"); @@ -575,7 +579,7 @@ async fn managed_skill_load_marks_survive_failures_repeat_per_call_and_skip_bloc register_tool_conditional_execution_guardrail( "block-skill-load", 1, - Arc::new(|_name, _args| Ok(Some("blocked before start".into()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("blocked before start".into())) })), ) .unwrap(); let blocked = tool_call_execute( @@ -703,6 +707,56 @@ fn test_manual_lifecycle_timestamp_overrides() { deregister_subscriber("timestamp-api-events").unwrap(); } +#[test] +fn test_manual_llm_end_queues_optimization_marks_before_end_event() { + let _lock = TEST_MUTEX.lock().unwrap(); + reset_global(); + setup_isolated_thread(); + + let events = capture_events("manual-optimization-events"); + let request = make_llm_request(json!({"messages": []})); + let handle = llm_call( + LlmCallParams::builder() + .name("manual-optimized-llm") + .request(&request) + .build(), + ) + .unwrap(); + assert!( + handle + .optimization_recorder + .record(LlmOptimizationContribution::new( + "test.manual", + "test_manual_kind", + )) + ); + + llm_call_end( + nemo_relay::api::llm::LlmCallEndParams::builder() + .handle(&handle) + .response(json!({"ok": true})) + .build(), + ) + .unwrap(); + + let names = captured_events_snapshot(&events) + .into_iter() + .filter(|event| { + event.name() == "manual-optimized-llm" || event.name() == "nemo_relay.llm.optimization" + }) + .map(|event| event.name().to_owned()) + .collect::>(); + assert_eq!( + names, + [ + "manual-optimized-llm", + "nemo_relay.llm.optimization", + "manual-optimized-llm", + ] + ); + deregister_subscriber("manual-optimization-events").unwrap(); +} + #[test] fn test_manual_lifecycle_default_end_timestamps_follow_explicit_starts() { let _lock = TEST_MUTEX.lock().unwrap(); @@ -834,10 +888,19 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { reset_global(); setup_isolated_thread(); - register_mark_sanitize_guardrail("mark-sanitize", 1, Arc::new(|_, fields| fields)).unwrap(); + register_mark_sanitize_guardrail( + "mark-sanitize", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); expect_already_exists( - register_mark_sanitize_guardrail("mark-sanitize", 1, Arc::new(|_, fields| fields)) - .unwrap_err(), + register_mark_sanitize_guardrail( + "mark-sanitize", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap_err(), "mark-sanitize", ); assert!(deregister_mark_sanitize_guardrail("mark-sanitize").unwrap()); @@ -846,28 +909,32 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_scope_sanitize_start_guardrail( "scope-start-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); assert!(deregister_scope_sanitize_start_guardrail("scope-start-sanitize").unwrap()); assert!(!deregister_scope_sanitize_start_guardrail("scope-start-sanitize").unwrap()); - register_scope_sanitize_end_guardrail("scope-end-sanitize", 1, Arc::new(|_, fields| fields)) - .unwrap(); + register_scope_sanitize_end_guardrail( + "scope-end-sanitize", + 1, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) + .unwrap(); assert!(deregister_scope_sanitize_end_guardrail("scope-end-sanitize").unwrap()); assert!(!deregister_scope_sanitize_end_guardrail("scope-end-sanitize").unwrap()); register_tool_sanitize_request_guardrail( "tool-sanitize-request", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); expect_already_exists( register_tool_sanitize_request_guardrail( "tool-sanitize-request", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap_err(), "tool-sanitize-request", @@ -878,7 +945,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_tool_sanitize_response_guardrail( "tool-sanitize-response", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); assert!(deregister_tool_sanitize_response_guardrail("tool-sanitize-response").unwrap()); @@ -886,13 +953,18 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_tool_conditional_execution_guardrail( "tool-conditional", 1, - Arc::new(|_name, _args| Ok(None)), + Arc::new(|_name, _args| Box::pin(async { Ok(None) })), ) .unwrap(); assert!(deregister_tool_conditional_execution_guardrail("tool-conditional").unwrap()); - register_tool_request_intercept("tool-request", 1, false, Arc::new(|_name, args| Ok(args))) - .unwrap(); + register_tool_request_intercept( + "tool-request", + 1, + false, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ) + .unwrap(); assert!(deregister_tool_request_intercept("tool-request").unwrap()); register_tool_execution_intercept( @@ -906,7 +978,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_llm_sanitize_request_guardrail( "llm-sanitize-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); assert!(deregister_llm_sanitize_request_guardrail("llm-sanitize-request").unwrap()); @@ -914,7 +986,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_llm_sanitize_response_guardrail( "llm-sanitize-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); assert!(deregister_llm_sanitize_response_guardrail("llm-sanitize-response").unwrap()); @@ -922,7 +994,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { register_llm_conditional_execution_guardrail( "llm-conditional", 1, - Arc::new(|_request| Ok(None)), + Arc::new(|_request| Box::pin(async { Ok(None) })), ) .unwrap(); assert!(deregister_llm_conditional_execution_guardrail("llm-conditional").unwrap()); @@ -932,7 +1004,7 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { 1, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -1023,7 +1095,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "mark-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); expect_already_exists( @@ -1031,7 +1103,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "mark-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap_err(), "mark-sanitize", @@ -1043,7 +1115,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "scope-start-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); assert!( @@ -1059,7 +1131,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "scope-end-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); assert!( @@ -1073,7 +1145,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "tool-sanitize-request", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); expect_already_exists( @@ -1081,7 +1153,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "tool-sanitize-request", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap_err(), "tool-sanitize-request", @@ -1095,7 +1167,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "tool-sanitize-response", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); assert!( @@ -1107,7 +1179,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "tool-conditional", 1, - Arc::new(|_name, _args| Ok(None)), + Arc::new(|_name, _args| Box::pin(async { Ok(None) })), ) .unwrap(); assert!( @@ -1120,7 +1192,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss "tool-request", 1, false, - Arc::new(|_name, args| Ok(args)), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); assert!(scope_deregister_tool_request_intercept(&scope.uuid, "tool-request").unwrap()); @@ -1138,7 +1210,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "llm-sanitize-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); assert!( @@ -1150,7 +1222,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "llm-sanitize-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); assert!( @@ -1162,7 +1234,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "llm-conditional", 1, - Arc::new(|_request| Ok(None)), + Arc::new(|_request| Box::pin(async { Ok(None) })), ) .unwrap(); assert!( @@ -1176,7 +1248,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss 1, false, Arc::new(|_name, request, annotated| { - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -1229,7 +1301,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "missing-mark-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap_err(), "scope", @@ -1239,7 +1311,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "missing-scope-start-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap_err(), "scope", @@ -1249,7 +1321,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "missing-scope-end-sanitize", 1, - Arc::new(|_, fields| fields), + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap_err(), "scope", @@ -1259,7 +1331,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss &scope.uuid, "missing-tool-sanitize", 1, - Arc::new(|_name, args| args), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap_err(), "scope", @@ -1270,7 +1342,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss "missing-tool-request", 1, false, - Arc::new(|_name, args| Ok(args)), + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap_err(), "scope", @@ -1307,7 +1379,7 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { args.as_object_mut() .unwrap() .insert("sanitized_request".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -1319,7 +1391,7 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { .as_object_mut() .unwrap() .insert("sanitized_response".into(), json!(true)); - result + ready(result) }), ) .unwrap(); @@ -1331,7 +1403,7 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { assert_eq!(event.input().unwrap()["sanitized_request"], true); fields.metadata = Some(json!({"generic_start": true})); } - fields + ready(fields) }), ) .unwrap(); @@ -1343,7 +1415,7 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { assert_eq!(event.output().unwrap()["sanitized_response"], true); fields.metadata = Some(json!({"generic_end": true})); } - fields + ready(fields) }), ) .unwrap(); @@ -1401,12 +1473,14 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { args.as_object_mut() .unwrap() .insert("intercepted".into(), json!(true)); - Ok(args) + ready(args) }), ) .unwrap(); assert_eq!( - tool_request_intercepts("tool-api", json!({"value": 2})).unwrap()["intercepted"], + tool_request_intercepts("tool-api", json!({"value": 2})) + .await + .unwrap()["intercepted"], json!(true) ); deregister_tool_request_intercept("tool-request").unwrap(); @@ -1414,11 +1488,11 @@ async fn test_tool_api_emits_sanitized_events_and_covers_error_paths() { register_tool_conditional_execution_guardrail( "tool-reject", 1, - Arc::new(|_name, _args| Ok(Some("tool denied".into()))), + Arc::new(|_name, _args| Box::pin(async { Ok(Some("tool denied".into())) })), ) .unwrap(); assert!(matches!( - tool_conditional_execution("tool-api", &json!({"value": 3})), + tool_conditional_execution("tool-api", &json!({"value": 3})).await, Err(FlowError::GuardrailRejected(reason)) if reason == "tool denied" )); assert!(matches!( @@ -1492,7 +1566,7 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { 1, Arc::new(|mut request, _context| { request.headers.insert("x-sanitized".into(), json!(true)); - Some(request) + ready(Some(request)) }), ) .unwrap(); @@ -1504,7 +1578,7 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { .as_object_mut() .unwrap() .insert("sanitized_response".into(), json!(true)); - Some(response) + ready(Some(response)) }), ) .unwrap(); @@ -1557,7 +1631,7 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { false, Arc::new(|_name, mut request, annotated| { request.headers.insert("x-intercepted".into(), json!(true)); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( request, annotated, )) }), @@ -1567,6 +1641,7 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { "llm-api", make_llm_request(json!({"messages": [{"role": "user", "content": "hello"}]})), ) + .await .unwrap(); assert_eq!( intercepted.request.headers.get("x-intercepted"), @@ -1577,11 +1652,11 @@ async fn test_llm_api_emits_sanitized_events_and_covers_error_paths() { register_llm_conditional_execution_guardrail( "llm-reject", 1, - Arc::new(|_request| Ok(Some("llm denied".into()))), + Arc::new(|_request| Box::pin(async { Ok(Some("llm denied".into())) })), ) .unwrap(); assert!(matches!( - llm_conditional_execution(&make_llm_request(json!({"messages": []}))), + llm_conditional_execution(&make_llm_request(json!({"messages": []}))).await, Err(FlowError::GuardrailRejected(reason)) if reason == "llm denied" )); assert!(matches!( @@ -1661,7 +1736,7 @@ async fn test_llm_stream_chunk_marks_track_successful_chunks() { Arc::new(|event, mut fields| { assert_eq!(event.name(), "llm.chunk"); fields.metadata = Some(json!({"sanitized": true})); - fields + ready(fields) }), ) .unwrap(); @@ -1703,6 +1778,7 @@ async fn test_llm_stream_chunk_marks_track_successful_chunks() { yielded.push(item.unwrap()); } assert_eq!(yielded, raw_chunks); + stream.close().await.unwrap(); let captured = captured_events_snapshot(&events); assert_eq!(captured.len(), 4); @@ -1775,6 +1851,7 @@ async fn test_llm_stream_chunk_mark_survives_collector_failure() { Err(FlowError::Internal(message)) if message == "collector failed" )); assert!(stream.next().await.is_none()); + stream.close().await.unwrap(); let captured = captured_events_snapshot(&events); assert_eq!(captured.len(), 3); @@ -1841,31 +1918,34 @@ async fn test_llm_stream_api_covers_success_rejection_and_execution_error_paths( chunks, vec![json!({"messages": [{"role": "user", "content": "hello"}]})] ); + stream.close().await.unwrap(); let success_events = captured_events_snapshot(&events); - assert_eq!(success_events[0].kind(), "scope"); - assert_eq!( - success_events[0].scope_category(), - Some(ScopeCategory::Start) - ); - assert_eq!(success_events[0].category().unwrap().as_str(), "llm"); - assert_eq!(success_events.last().unwrap().kind(), "scope"); - assert_eq!( - success_events.last().unwrap().scope_category(), - Some(ScopeCategory::End) - ); - assert_eq!( - success_events.last().unwrap().category().unwrap().as_str(), - "llm" - ); + let success_start = success_events + .iter() + .find(|event| { + event.kind() == "scope" && event.scope_category() == Some(ScopeCategory::Start) + }) + .expect("stream start event"); + let success_end = success_events + .iter() + .rev() + .find(|event| event.kind() == "scope" && event.scope_category() == Some(ScopeCategory::End)) + .expect("stream end event"); + assert_eq!(success_start.kind(), "scope"); + assert_eq!(success_start.scope_category(), Some(ScopeCategory::Start)); + assert_eq!(success_start.category().unwrap().as_str(), "llm"); + assert_eq!(success_end.kind(), "scope"); + assert_eq!(success_end.scope_category(), Some(ScopeCategory::End)); + assert_eq!(success_end.category().unwrap().as_str(), "llm"); assert_eq!( - success_events.last().unwrap().output().unwrap(), + success_end.output().unwrap(), &json!([{"messages": [{"role": "user", "content": "hello"}]}]) ); register_llm_conditional_execution_guardrail( "llm-stream-reject", 1, - Arc::new(|_request| Ok(Some("stream denied".into()))), + Arc::new(|_request| Box::pin(async { Ok(Some("stream denied".into())) })), ) .unwrap(); let reject_collector: Box Result<()> + Send> = Box::new(|_chunk| Ok(())); diff --git a/crates/core/tests/integration/middleware_tests.rs b/crates/core/tests/integration/middleware_tests.rs index 41093cdcb..ae3df1732 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, @@ -149,8 +152,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(); @@ -164,7 +167,7 @@ fn test_sanitize_guardrail_priority_ordering() { 1, Arc::new(move |_name, args| { o1.lock().unwrap().push(1); - args + ready(args) }), ) .unwrap(); @@ -176,7 +179,7 @@ fn test_sanitize_guardrail_priority_ordering() { 3, Arc::new(move |_name, args| { o3.lock().unwrap().push(3); - args + ready(args) }), ) .unwrap(); @@ -188,7 +191,7 @@ fn test_sanitize_guardrail_priority_ordering() { 2, Arc::new(move |_name, args| { o2.lock().unwrap().push(2); - args + ready(args) }), ) .unwrap(); @@ -201,6 +204,7 @@ fn test_sanitize_guardrail_priority_ordering() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); let recorded = order.lock().unwrap(); assert_eq!( @@ -217,8 +221,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(); @@ -232,7 +236,7 @@ fn test_request_intercept_priority_ordering() { false, Arc::new(move |_name, args| { o1.lock().unwrap().push(1); - Ok(args) + ready(args) }), ) .unwrap(); @@ -244,7 +248,7 @@ fn test_request_intercept_priority_ordering() { false, Arc::new(move |_name, args| { o3.lock().unwrap().push(3); - Ok(args) + ready(args) }), ) .unwrap(); @@ -256,13 +260,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!( @@ -278,8 +284,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(); @@ -293,7 +299,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(); @@ -305,13 +311,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"]); @@ -326,14 +332,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!( @@ -354,8 +360,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(); @@ -370,7 +376,7 @@ fn test_break_chain_stops_subsequent_intercepts() { args.as_object_mut() .unwrap() .insert("breaker_ran".into(), json!(true)); - Ok(args) + ready(args) }), ) .unwrap(); @@ -385,12 +391,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); @@ -410,8 +416,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(); @@ -425,7 +431,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(); @@ -437,12 +443,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), @@ -696,7 +702,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(); @@ -1332,7 +1338,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(); @@ -1366,8 +1372,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) })); @@ -1405,12 +1415,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(); @@ -1490,13 +1504,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(); @@ -1536,9 +1554,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) } }), ) @@ -1578,8 +1596,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"); @@ -1594,7 +1612,7 @@ fn test_scope_local_guardrail_lifecycle() { 1, Arc::new(move |_name, args| { cc.fetch_add(1, Ordering::SeqCst); - args + ready(args) }), ) .unwrap(); @@ -1607,6 +1625,7 @@ fn test_scope_local_guardrail_lifecycle() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!( call_count.load(Ordering::SeqCst), 1, @@ -1703,8 +1722,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"); @@ -1721,7 +1740,7 @@ fn test_scope_local_and_global_guardrail_merge_priority() { args.as_object_mut() .unwrap() .insert("global".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -1737,7 +1756,7 @@ fn test_scope_local_and_global_guardrail_merge_priority() { args.as_object_mut() .unwrap() .insert("local".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -1760,6 +1779,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(); @@ -1890,7 +1910,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(); @@ -1902,7 +1922,7 @@ async fn test_conditional_rejection_prevents_intercepts() { false, Arc::new(move |_name, args| { ic.store(true, Ordering::SeqCst); - Ok(args) + ready(args) }), ) .unwrap(); @@ -1941,7 +1961,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(); @@ -1992,8 +2012,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(); @@ -2006,7 +2026,7 @@ fn test_sanitize_guardrails_pipe_data() { args.as_object_mut() .unwrap() .insert("field_a".into(), json!(true)); - args + ready(args) }), ) .unwrap(); @@ -2021,7 +2041,7 @@ fn test_sanitize_guardrails_pipe_data() { args.as_object_mut() .unwrap() .insert("field_b".into(), json!(has_a)); - args + ready(args) }), ) .unwrap(); @@ -2064,8 +2084,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(); @@ -2078,7 +2098,7 @@ fn test_response_sanitize_guardrails_pipe() { .as_object_mut() .unwrap() .insert("sanitized".into(), json!(true)); - result + ready(result) }), ) .unwrap(); @@ -2130,8 +2150,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(); @@ -2148,7 +2168,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}"); @@ -2176,8 +2196,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(); @@ -2194,7 +2214,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()); @@ -2220,8 +2240,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(); @@ -2230,7 +2250,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(); } @@ -2249,7 +2269,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); @@ -2283,8 +2303,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(); @@ -2311,23 +2331,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"]); @@ -2335,8 +2359,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(); @@ -2363,7 +2387,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, )) }), @@ -2371,7 +2395,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, )) }), @@ -2384,6 +2408,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"]); @@ -2396,6 +2421,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"]); @@ -2418,7 +2444,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(); @@ -2430,7 +2456,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(); @@ -2442,7 +2468,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(); @@ -2455,7 +2481,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(); @@ -2466,7 +2492,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(); @@ -2478,7 +2504,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(); @@ -2512,7 +2538,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(); @@ -2524,7 +2550,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(); @@ -2589,7 +2615,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(); @@ -2601,7 +2627,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(); @@ -2613,7 +2639,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, )) }), @@ -2628,7 +2654,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, )) }), @@ -2641,7 +2667,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(); @@ -2653,7 +2679,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(); @@ -2710,7 +2736,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(); @@ -2722,7 +2748,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(); @@ -2785,6 +2811,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, &[ @@ -2855,7 +2882,7 @@ async fn test_full_pipeline_integration() { args.as_object_mut() .unwrap() .insert("intercepted".into(), json!(true)); - Ok(args) + ready(args) }), ) .unwrap(); @@ -2867,7 +2894,7 @@ async fn test_full_pipeline_integration() { 1, Arc::new(move |_name, args| { o2.lock().unwrap().push("sanitize_request".into()); - args + ready(args) }), ) .unwrap(); @@ -2879,7 +2906,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(); @@ -2906,7 +2933,7 @@ async fn test_full_pipeline_integration() { 1, Arc::new(move |_name, result| { o5.lock().unwrap().push("sanitize_response".into()); - result + ready(result) }), ) .unwrap(); @@ -2964,15 +2991,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() { @@ -2987,19 +3022,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()); @@ -3019,8 +3059,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(); @@ -3032,8 +3072,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(); @@ -3046,7 +3086,7 @@ fn test_deregister_removes_from_chain() { 1, Arc::new(move |_name, args| { cc.fetch_add(1, Ordering::SeqCst); - args + ready(args) }), ) .unwrap(); @@ -3059,6 +3099,7 @@ fn test_deregister_removes_from_chain() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!(call_count.load(Ordering::SeqCst), 1); // Deregister @@ -3073,6 +3114,7 @@ fn test_deregister_removes_from_chain() { .build(), ) .unwrap(); + flush_subscribers().unwrap(); assert_eq!( call_count.load(Ordering::SeqCst), 1, @@ -3094,7 +3136,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(); @@ -3145,12 +3187,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(); @@ -3239,9 +3281,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, )) }), @@ -3253,15 +3295,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(); @@ -3276,8 +3318,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(); @@ -3290,6 +3334,7 @@ fn test_llm_request_intercept_pending_marks_preserve_order_and_break_chain() { content: json!({"prompt": "hello"}), }, ) + .await .unwrap(); assert_eq!( @@ -3325,7 +3370,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(); @@ -3334,24 +3379,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(); @@ -3452,7 +3499,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(); @@ -3469,7 +3516,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(); @@ -3495,8 +3542,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(); @@ -3701,7 +3750,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(); @@ -3756,6 +3805,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); @@ -3801,8 +3851,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(); @@ -3893,8 +3945,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(); @@ -3903,7 +3957,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(); @@ -4018,7 +4072,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, )) }), @@ -4112,7 +4166,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, )) }), @@ -4165,6 +4219,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), @@ -4193,19 +4248,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(); @@ -4213,11 +4268,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) => { @@ -4260,12 +4315,14 @@ async fn test_empty_chain_passthrough() { } /// Standalone intercept chain with no registrations returns input unchanged. -#[test] -fn test_empty_request_intercept_chain() { +#[tokio::test] +async fn test_empty_request_intercept_chain() { let _lock = TEST_MUTEX.lock().unwrap(); reset_global(); setup_isolated_thread(); - let result = tool_request_intercepts("tool", json!({"key": "val"})).unwrap(); + let result = tool_request_intercepts("tool", json!({"key": "val"})) + .await + .unwrap(); assert_eq!(result["key"], "val"); } diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index ee8a860be..c669f7c46 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -239,6 +239,7 @@ async fn sdk_cdylib_registers_tool_request_intercept() { .expect("outer scope should push"); let outer_uuid = outer.uuid; let rewritten = tool_request_intercepts("demo_tool", json!({ "input": "value" })) + .await .expect("native request intercept should run"); let tool_result = tool_call_execute( ToolCallExecuteParams::builder() @@ -446,6 +447,7 @@ async fn sdk_cdylib_registers_tool_request_intercept() { .expect("thread outer scope should push"); let thread_outer_uuid = thread_outer.uuid; let rewritten = tool_request_intercepts("demo_tool", json!({ "input": "thread" })) + .await .expect("native request intercept should run with thread stack"); assert_eq!(rewritten["native_plugin"], true); pop_scope( @@ -656,6 +658,123 @@ async fn sdk_cdylib_registers_tool_request_intercept() { activation.clear(); } +#[tokio::test] +async fn native_v3_async_registration_supports_all_middleware_kinds() { + let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; + let fixture = build_fixture_plugin(); + let manifest_ref = write_manifest_with_plugin_id_and_symbol( + &fixture, + "fixture_async", + "nemo_relay_fixture_async_entry", + ); + + let activation = load_native_plugins([NativePluginLoadSpec { + plugin_id: "fixture_async".into(), + manifest_ref: manifest_ref.to_string_lossy().into_owned(), + }]) + .expect("v3 async native fixture should load"); + let mut config = PluginConfig::default(); + config.components.push(PluginComponentSpec { + kind: "fixture_async".into(), + enabled: true, + config: Map::new(), + }); + initialize_plugins_exact(config) + .await + .expect("v3 async native fixture should register"); + + let rewritten = tool_request_intercepts("async-tool", json!({"input": true})) + .await + .expect("v3 async request intercept should settle"); + assert_eq!(rewritten["input"], true); + assert_eq!(rewritten["native_async"], true); + + let duplicate = tool_request_intercepts("async-double", json!({"input": true})) + .await + .expect("duplicate v3 async settlement keeps the first result"); + assert_eq!(duplicate["native_async"], true); + + let executed = tool_call_execute( + ToolCallExecuteParams::builder() + .name("async-execution") + .args(json!({"input": true})) + .func(Arc::new(|args| Box::pin(async move { Ok(args) }))) + .build(), + ) + .await + .expect("v3 async execution intercept should continue with next"); + assert_eq!(executed["native_async_execution"], true); + + let llm_response = llm_call_execute( + LlmCallExecuteParams::builder() + .name("async-llm") + .request(LlmRequest { + headers: Map::new(), + content: json!({"prompt": "native async"}), + }) + .func(Arc::new(|_request| { + Box::pin(async move { Ok(json!({"content": "native async response"})) }) + })) + .build(), + ) + .await + .expect("v3 async LLM middleware should settle"); + assert_eq!(llm_response["content"], "native async response"); + flush_subscribers().expect("async native LLM events should flush"); + + let stream_chunks = Arc::new(Mutex::new(Vec::::new())); + let collected_chunks = stream_chunks.clone(); + let finalized_chunks = stream_chunks.clone(); + let mut stream = llm_stream_call_execute( + LlmStreamCallExecuteParams::builder() + .name("async-llm-stream") + .request(LlmRequest { + headers: Map::new(), + content: json!({"prompt": "native async stream"}), + }) + .func(Arc::new(|_request| { + Box::pin(async move { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok(json!({ + "content": "native async stream response" + }))]))) + }) + })) + .collector(Box::new(move |chunk| { + collected_chunks.lock().unwrap().push(chunk); + Ok(()) + })) + .finalizer(Box::new(move || { + Json::Array(finalized_chunks.lock().unwrap().clone()) + })) + .build(), + ) + .await + .expect("v3 async LLM stream middleware should settle"); + assert_eq!( + stream + .next() + .await + .expect("stream should contain a chunk") + .expect("stream chunk should succeed")["content"], + "native async stream response" + ); + assert!(stream.next().await.is_none()); + flush_subscribers().expect("async native LLM stream events should flush"); + + let pending = tokio::spawn(async { + tool_request_intercepts("async-pending", json!({"input": true})).await + }); + tokio::task::yield_now().await; + clear_plugin_configuration().expect("v3 async native fixture should clear while pending"); + let pending = pending + .await + .expect("pending v3 async task should not panic") + .expect("pending v3 async request intercept should settle after clear"); + assert_eq!(pending["native_async"], true); + + drop(activation); +} + #[tokio::test] async fn native_validation_diagnostics_prevent_initialization() { let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; @@ -752,7 +871,7 @@ async fn native_tool_execution_rejects_null_malformed_and_error_outcomes() { } #[tokio::test] -async fn native_event_sanitizer_callback_errors_clear_observability_fields() { +async fn native_event_sanitizer_callback_errors_preserve_observability_fields() { let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; let fixture = build_fixture_plugin(); let manifest_ref = @@ -793,8 +912,8 @@ async fn native_event_sanitizer_callback_errors_clear_observability_fields() { let captured_events = events.lock().unwrap().clone(); let event = find_event(&captured_events, "native-event-sanitize-error", None); - assert_eq!(event.data(), None); - assert_eq!(event.metadata(), None); + assert_eq!(event.data(), Some(&json!({ "secret": true }))); + assert_eq!(event.metadata(), Some(&json!({ "secret": true }))); deregister_subscriber("native_event_sanitizer_error_capture") .expect("test subscriber should deregister"); @@ -1218,6 +1337,7 @@ async fn plugin_host_activation_owns_configuration_until_clear() { .any(|kind| kind == "fixture_native") ); let rewritten = tool_request_intercepts("host-owned-tool", json!({ "input": true })) + .await .expect("host-owned intercept should run"); assert_eq!(rewritten["native_plugin"], true); @@ -1238,6 +1358,7 @@ async fn plugin_host_activation_owns_configuration_until_clear() { .any(|kind| kind == "fixture_native") ); let unchanged = tool_request_intercepts("host-owned-tool", json!({ "input": true })) + .await .expect("cleared intercept chain should be empty"); assert_eq!(unchanged, json!({ "input": true })); } @@ -1393,6 +1514,7 @@ async fn plugin_host_clear_allows_an_in_flight_native_callback_to_finish() { .clear() .expect("host should clear while a callback snapshot remains in flight"); let unchanged = tool_request_intercepts("after-clear", json!({ "input": true })) + .await .expect("new calls should observe the cleared registries"); assert_eq!(unchanged, json!({ "input": true })); diff --git a/crates/core/tests/integration/pipeline_tests.rs b/crates/core/tests/integration/pipeline_tests.rs index 1bc192b4b..9711974de 100644 --- a/crates/core/tests/integration/pipeline_tests.rs +++ b/crates/core/tests/integration/pipeline_tests.rs @@ -9,6 +9,9 @@ use std::sync::{Arc, Mutex}; +mod test_support; +use test_support::ready; + use futures::StreamExt; use serde_json::json; @@ -359,7 +362,7 @@ async fn test_decode_runs_before_intercepts() { false, Arc::new(move |_name, req, annotated| { *cap.lock().unwrap() = Some(annotated.clone()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -414,7 +417,7 @@ async fn test_encode_runs_after_intercepts() { let mut ann = annotated.unwrap(); ann.model = Some("modified".into()); req.headers.insert("x-codec-route".into(), json!("blue")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, Some(ann), )) @@ -517,7 +520,7 @@ async fn anthropic_issue_501_round_trips_and_applies_annotated_edits() { request .headers .insert("x-annotation-seen".into(), json!("yes")); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + ready(LlmRequestInterceptOutcome::new(request, Some(annotated))) }), ) .unwrap(); @@ -571,8 +574,10 @@ async fn test_codec_rejects_raw_content_mutation_before_lifecycle() { false, Arc::new(|_name, mut request, annotated| { request.content["model"] = json!("raw-model-edit"); - Ok(LlmRequestInterceptOutcome::new(request, annotated) - .with_pending_mark(PendingMarkSpec::builder().name("must.not.emit").build())) + ready( + LlmRequestInterceptOutcome::new(request, annotated) + .with_pending_mark(PendingMarkSpec::builder().name("must.not.emit").build()), + ) }), ) .unwrap(); @@ -583,7 +588,7 @@ async fn test_codec_rejects_raw_content_mutation_before_lifecycle() { false, Arc::new(move |_name, request, annotated| { *later_called.lock().unwrap() = true; - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + ready(LlmRequestInterceptOutcome::new(request, annotated)) }), ) .unwrap(); @@ -630,7 +635,9 @@ async fn test_codec_rejects_missing_annotation_before_lifecycle() { "codec_missing_annotation", 1, false, - Arc::new(|_name, request, _annotated| Ok(LlmRequestInterceptOutcome::new(request, None))), + Arc::new(|_name, request, _annotated| { + ready(LlmRequestInterceptOutcome::new(request, None)) + }), ) .unwrap(); @@ -680,7 +687,7 @@ async fn test_stream_codec_rejects_raw_content_mutation_before_lifecycle() { false, Arc::new(|_name, mut request, annotated| { request.content["model"] = json!("raw-stream-edit"); - Ok(LlmRequestInterceptOutcome::new(request, annotated)) + ready(LlmRequestInterceptOutcome::new(request, annotated)) }), ) .unwrap(); @@ -741,7 +748,7 @@ async fn test_annotated_intercept_receives_both() { false, Arc::new(move |_name, req, annotated| { *cp.lock().unwrap() = Some((req.clone(), annotated.clone())); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -799,7 +806,7 @@ async fn test_canonical_intercept_with_and_without_codec() { Arc::new(move |_name, mut req, annotated| { *lc1.lock().unwrap() = true; req.headers.insert("x-legacy".into(), json!("was-here")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -845,7 +852,7 @@ async fn test_canonical_intercept_with_and_without_codec() { Arc::new(move |_name, mut req, annotated| { *lc2.lock().unwrap() = true; req.headers.insert("x-legacy-2".into(), json!("also-here")); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -903,7 +910,7 @@ async fn test_stream_path_also_decodes() { false, Arc::new(move |_name, req, annotated| { *ca.lock().unwrap() = Some(annotated.clone()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -930,6 +937,7 @@ async fn test_stream_path_also_decodes() { // Consume the stream to trigger full pipeline while let Some(_chunk) = stream.next().await {} + stream.close().await.unwrap(); // Assert decode was called let dl = decode_log.lock().unwrap(); @@ -969,7 +977,7 @@ async fn test_shared_helper_both_paths() { false, Arc::new(move |_name, req, annotated| { *acc.lock().unwrap() += 1; - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1048,7 +1056,7 @@ async fn test_explicit_codec_param_overrides() { if let Some(ref ann) = annotated { *cm.lock().unwrap() = ann.model.clone(); } - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1097,7 +1105,7 @@ async fn test_encode_merge_not_replace() { Arc::new(|_name, req, annotated| { let mut ann = annotated.unwrap(); ann.model = Some("new_model".into()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, Some(ann), )) @@ -1168,7 +1176,7 @@ async fn test_unified_chain_priority_order() { false, Arc::new(move |_name, req, annotated| { cl1.lock().unwrap().push("legacy_p10".into()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1183,7 +1191,7 @@ async fn test_unified_chain_priority_order() { false, Arc::new(move |_name, req, annotated| { cl2.lock().unwrap().push("annotated_p5".into()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1233,7 +1241,7 @@ async fn test_no_codec_annotated_intercept_receives_none() { false, Arc::new(move |_name, req, annotated| { *ca.lock().unwrap() = Some(annotated.clone()); - Ok(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( + ready(nemo_relay::api::llm::LlmRequestInterceptOutcome::new( req, annotated, )) }), @@ -1448,7 +1456,7 @@ async fn test_response_codec_annotation_uses_sanitized_managed_response() { register_llm_sanitize_response_guardrail( "sanitize_resp_codec_annotation", 1, - Arc::new(|_response, _context| Some(make_openai_chat_response("Sanitized"))), + Arc::new(|_response, _context| ready(Some(make_openai_chat_response("Sanitized")))), ) .unwrap(); @@ -1649,10 +1657,10 @@ async fn test_request_codec_annotation_uses_sanitized_start_payload() { "sanitize_req_codec_annotation", 1, Arc::new(|request, _context| { - Some(LlmRequest { + ready(Some(LlmRequest { headers: request.headers, content: make_openai_chat_request("Sanitized").content, - }) + })) }), ) .unwrap(); @@ -1733,6 +1741,7 @@ async fn test_stream_response_codec_populates_annotated_response() { // Drain the stream to trigger finalization and END event while let Some(_chunk) = stream.next().await {} + stream.close().await.unwrap(); let captured = captured_events_snapshot(&events); let end_event = captured @@ -1818,6 +1827,7 @@ async fn managed_buffered_and_streaming_close_price_the_committed_route_not_resp while let Some(item) = stream.next().await { item.unwrap(); } + stream.close().await.unwrap(); llm_call_execute( LlmCallExecuteParams::builder() @@ -1916,7 +1926,7 @@ async fn test_stream_response_codec_annotation_uses_sanitized_aggregated_respons register_llm_sanitize_response_guardrail( "stream_sanitize_resp_codec_annotation", 1, - Arc::new(|_response, _context| Some(make_openai_chat_response("Sanitized"))), + Arc::new(|_response, _context| ready(Some(make_openai_chat_response("Sanitized")))), ) .unwrap(); @@ -1939,6 +1949,7 @@ async fn test_stream_response_codec_annotation_uses_sanitized_aggregated_respons .unwrap(); while let Some(_chunk) = stream.next().await {} + stream.close().await.unwrap(); let captured = captured_events_snapshot(&events); let end_event = captured diff --git a/crates/core/tests/integration/scope_local_tests.rs b/crates/core/tests/integration/scope_local_tests.rs index 551d39dd7..ceb17e821 100644 --- a/crates/core/tests/integration/scope_local_tests.rs +++ b/crates/core/tests/integration/scope_local_tests.rs @@ -84,7 +84,7 @@ fn test_scope_local_guardrail_registration_and_execution() { args.as_object_mut() .unwrap() .insert("scope_sanitized".into(), json!(true)); - args + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -166,7 +166,7 @@ async fn test_auto_cleanup_on_scope_pop() { args.as_object_mut() .unwrap() .insert("ephemeral".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -234,7 +234,7 @@ async fn test_priority_merge_global_and_scope_local() { args.as_object_mut() .unwrap() .insert("p10".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -250,7 +250,7 @@ async fn test_priority_merge_global_and_scope_local() { args.as_object_mut() .unwrap() .insert("p30".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -267,7 +267,7 @@ async fn test_priority_merge_global_and_scope_local() { args.as_object_mut() .unwrap() .insert("p20".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -326,7 +326,7 @@ fn test_name_coexistence_global_and_scope_local() { 1, Arc::new(move |_name, args| { c1.fetch_add(1, Ordering::SeqCst); - args + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -339,7 +339,7 @@ fn test_name_coexistence_global_and_scope_local() { 2, Arc::new(move |_name, args| { c2.fetch_add(1, Ordering::SeqCst); - args + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -354,6 +354,7 @@ fn test_name_coexistence_global_and_scope_local() { .unwrap(); // Both guardrails with the same name ran. + flush_subscribers().unwrap(); assert_eq!(count.load(Ordering::SeqCst), 2); // Cleanup @@ -400,7 +401,7 @@ async fn test_scope_isolation_between_stacks() { args.as_object_mut() .unwrap() .insert("agent".into(), json!("a")); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -426,7 +427,7 @@ async fn test_scope_isolation_between_stacks() { args.as_object_mut() .unwrap() .insert("agent".into(), json!("b")); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -506,7 +507,7 @@ async fn test_nested_scope_inheritance() { args.as_object_mut() .unwrap() .insert("global".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -530,7 +531,7 @@ async fn test_nested_scope_inheritance() { args.as_object_mut() .unwrap() .insert("scope_a".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -555,7 +556,7 @@ async fn test_nested_scope_inheritance() { args.as_object_mut() .unwrap() .insert("scope_b".into(), json!(true)); - Ok(args) + Box::pin(async move { Ok(args) }) }), ) .unwrap(); @@ -702,11 +703,13 @@ async fn test_scope_local_conditional_execution_guardrail() { "tool_blocker", 1, Arc::new(|name, _args| { - if name == "banned_tool" { - Ok(Some("banned_tool is not allowed in this scope".to_string())) - } else { - Ok(None) - } + Box::pin(async move { + if name == "banned_tool" { + Ok(Some("banned_tool is not allowed in this scope".to_string())) + } else { + Ok(None) + } + }) }), ) .unwrap(); diff --git a/crates/core/tests/integration/subscriber_dispatcher_tests.rs b/crates/core/tests/integration/subscriber_dispatcher_tests.rs index e83fc7d6d..3ec5a3fca 100644 --- a/crates/core/tests/integration/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/integration/subscriber_dispatcher_tests.rs @@ -6,11 +6,15 @@ use std::sync::{Arc, Mutex, mpsc}; use std::time::Duration; +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; static TEST_MUTEX: Mutex<()> = Mutex::new(()); @@ -126,3 +130,80 @@ 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.name().to_string()) + }), + ) + .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(); + + emit_mark("unsanitized-fallback"); + flush_subscribers().unwrap(); + + assert_eq!( + observed.lock().unwrap().as_slice(), + ["unsanitized-fallback"] + ); + deregister_mark_sanitize_guardrail("fail-open-mark-sanitizer").unwrap(); + deregister_subscriber("fail-open-sanitizer-subscriber").unwrap(); +} + +#[test] +fn 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.name().to_string()) + }), + ) + .unwrap(); + register_mark_sanitize_guardrail( + "panic-mark-sanitizer", + 10, + Arc::new(|_, _| Box::pin(async { panic!("intentional event-sanitizer panic") })), + ) + .unwrap(); + + emit_mark("panic-fallback"); + flush_subscribers().unwrap(); + + assert_eq!(observed.lock().unwrap().as_slice(), ["panic-fallback"]); + deregister_mark_sanitize_guardrail("panic-mark-sanitizer").unwrap(); + deregister_subscriber("panic-sanitizer-subscriber").unwrap(); +} diff --git a/crates/core/tests/integration/test_support.rs b/crates/core/tests/integration/test_support.rs new file mode 100644 index 000000000..bf1507f0f --- /dev/null +++ b/crates/core/tests/integration/test_support.rs @@ -0,0 +1,18 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::future::Future; +use std::pin::Pin; + +pub fn ready( + value: T, +) -> Pin> + Send>> { + Box::pin(async move { Ok(value) }) +} + +#[allow(dead_code)] +pub fn ready_result( + value: nemo_relay::error::Result, +) -> Pin> + Send>> { + Box::pin(async move { value }) +} diff --git a/crates/core/tests/integration/worker_plugin_tests.rs b/crates/core/tests/integration/worker_plugin_tests.rs index 4387b048a..b7eb75611 100644 --- a/crates/core/tests/integration/worker_plugin_tests.rs +++ b/crates/core/tests/integration/worker_plugin_tests.rs @@ -76,6 +76,7 @@ async fn plugin_host_activation_owns_worker_lifecycle() { .any(|kind| kind == "fixture_worker") ); let rewritten = tool_request_intercepts("worker-host-tool", json!({ "input": true })) + .await .expect("worker host intercept should run"); assert_eq!(rewritten["worker_plugin"], true); @@ -86,6 +87,7 @@ async fn plugin_host_activation_owns_worker_lifecycle() { .any(|kind| kind == "fixture_worker") ); let unchanged = tool_request_intercepts("worker-host-tool", json!({ "input": true })) + .await .expect("cleared worker intercept chain should be empty"); assert_eq!(unchanged, json!({ "input": true })); } @@ -109,6 +111,7 @@ async fn plugin_host_clear_surfaces_worker_shutdown_failure_and_releases_safe_ow .expect("worker plugin host should activate"); tool_request_intercepts("terminate-worker", json!({ "input": true })) + .await .expect_err("fixture worker should terminate during callback"); let error = activation .clear() @@ -159,6 +162,7 @@ async fn rust_worker_registers_and_invokes_all_current_surfaces() { .expect("outer scope should push"); let outer_uuid = outer.uuid; let rewritten = tool_request_intercepts("demo_tool", json!({ "input": "value" })) + .await .expect("worker request intercept should run"); let tool_result = tool_call_execute( ToolCallExecuteParams::builder() @@ -473,6 +477,7 @@ async fn worker_request_intercept_callback_error_surfaces_to_host() { .await; let error = tool_request_intercepts("demo_tool", json!({ "input": "value" })) + .await .expect_err("worker callback error should surface"); assert!( error @@ -1120,6 +1125,7 @@ async fn python_worker_host_runtime_mark_and_mutated_request_round_trip() { cleanup.subscriber_name = Some(subscriber_name); let rewritten = tool_request_intercepts("lookup", json!({ "query": "relay" })) + .await .expect("Python callback should emit a mark and return its mutation"); assert_eq!( rewritten["_nemo_relay_plugin"]["tag"], diff --git a/crates/core/tests/unit/context_tests.rs b/crates/core/tests/unit/context_tests.rs index 49e57b2cb..48853939d 100644 --- a/crates/core/tests/unit/context_tests.rs +++ b/crates/core/tests/unit/context_tests.rs @@ -46,7 +46,7 @@ fn scope_stack_tracks_scope_local_registries_and_subscribers() { priority: 10, payload: RequestIntercept { break_chain: false, - callable: Arc::new(|_, value| Ok(value)), + callable: Arc::new(|_, value| Box::pin(async move { Ok(value) })), }, }) .unwrap(); @@ -221,15 +221,17 @@ fn merge_helpers_preserve_global_and_scope_local_priority_order() { assert_eq!(merged_exec, vec![("local", 1), ("global", 15)]); } -#[test] -fn conditional_guardrail_snapshots_keep_names_and_callbacks_after_deregister() { +#[tokio::test] +async fn conditional_guardrail_snapshots_keep_names_and_callbacks_after_deregister() { let mut state = NemoRelayContextState::new(); state .tool_conditional_execution_guardrails .register(Guardrail { name: "snapshot_guardrail".to_string(), priority: 1, - payload: Arc::new(|name, _args| Ok(Some(format!("{name} blocked")))), + payload: Arc::new(|name, _args| { + Box::pin(async move { Ok(Some(format!("{name} blocked"))) }) + }), }) .unwrap(); @@ -257,6 +259,7 @@ fn conditional_guardrail_snapshots_keep_names_and_callbacks_after_deregister() { None, None, ) + .await .unwrap(); assert_eq!(rejection.as_deref(), Some("snapshot_target blocked")); @@ -320,12 +323,14 @@ fn context_state_supports_extensions_events_and_builders() { content: json!({"messages": []}), }; let entries = state.llm_sanitize_request_entries(&[]); - let sanitized = NemoRelayContextState::llm_sanitize_request_snapshot_chain( - request.clone(), - crate::api::runtime::LlmSanitizeRequestContext::default(), - &entries, - ) - .expect("an empty sanitizer chain must retain the request"); + let sanitized = tokio::runtime::Runtime::new() + .unwrap() + .block_on(NemoRelayContextState::llm_sanitize_request_snapshot_chain( + request.clone(), + crate::api::runtime::LlmSanitizeRequestContext::default(), + &entries, + )) + .expect("an empty sanitizer chain must retain the request"); assert!(sanitized.headers.is_empty()); let events = Arc::new(Mutex::new(Vec::::new())); diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index 7754238dc..da3d9da7e 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -461,6 +461,7 @@ async fn callback_helpers_cover_worker_response_edges() { RegistrationSurface::MarkSanitizeGuardrail, &event, ) + .await .expect_err("invalid event sanitizer fields should fail"); assert!( error @@ -474,6 +475,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), LlmSanitizeRequestContext::default(), ) + .await .expect_err("invalid LLM JSON result should fail"); assert!(error.to_string().contains("invalid type")); @@ -484,6 +486,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), None, ) + .await .expect_err("invalid LLM intercept request should fail"); assert!( error @@ -498,6 +501,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), None, ) + .await .expect_err("legacy outcome schema should fail"); assert!( error @@ -512,6 +516,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), None, ) + .await .expect_err("invalid annotated request should fail"); assert!( error @@ -521,6 +526,7 @@ async fn callback_helpers_cover_worker_response_edges() { let error = callback .invoke_llm_request_intercept("llm_intercept_error", "model", valid_llm_request(), None) + .await .expect_err("LLM intercept worker error should surface"); assert!(error.to_string().contains("worker.failed: boom")); @@ -531,6 +537,7 @@ async fn callback_helpers_cover_worker_response_edges() { valid_llm_request(), None, ) + .await .expect_err("unexpected LLM intercept result should fail"); assert!( error @@ -586,6 +593,7 @@ async fn llm_worker_sanitizers_forward_codec_context_and_omission() { valid_llm_request(), LlmSanitizeRequestContext::with_identity(identity.clone()), ) + .await .expect("empty worker result must represent request omission") .is_none() ); @@ -596,6 +604,7 @@ async fn llm_worker_sanitizers_forward_codec_context_and_omission() { json!({"secret": "value"}), LlmSanitizeResponseContext::with_identity(identity), ) + .await .expect("empty worker result must represent response omission") .is_none() ); @@ -774,6 +783,7 @@ async fn llm_worker_codec_capabilities_are_active_only_during_sanitizer_invocati request, LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), ) + .await .expect("request sanitizer must succeed") .is_none() ); @@ -793,6 +803,7 @@ async fn llm_worker_codec_capabilities_are_active_only_during_sanitizer_invocati response, LlmSanitizeResponseContext::for_response_codec(Some(codec)), ) + .await .expect_err("worker sanitizer error must surface"); assert!(error.to_string().contains("worker.failed: boom")); @@ -1324,6 +1335,7 @@ async fn install_registrations_covers_registry_error_edges() { } #[tokio::test(flavor = "multi_thread")] +#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. async fn installed_callbacks_apply_surface_specific_fallbacks() { struct RuntimeCleanup { registrations: Option, @@ -1423,9 +1435,9 @@ async fn installed_callbacks_apply_surface_specific_fallbacks() { ] { let entries = NemoRelayContextState::event_sanitize_entries(registry, &[]); let sanitized = - NemoRelayContextState::event_sanitize_snapshot_chain(event.clone(), &entries); - assert_eq!(sanitized.data(), None); - assert_eq!(sanitized.metadata(), None); + NemoRelayContextState::event_sanitize_snapshot_chain(event.clone(), &entries).await; + assert_eq!(sanitized.data(), event.data()); + assert_eq!(sanitized.metadata(), event.metadata()); } let entries = state.tool_sanitize_request_entries(&[]); @@ -1434,7 +1446,8 @@ async fn installed_callbacks_apply_surface_specific_fallbacks() { "tool", tool_request.clone(), &entries, - ), + ) + .await, tool_request ); let entries = state.tool_sanitize_response_entries(&[]); @@ -1443,28 +1456,29 @@ async fn installed_callbacks_apply_surface_specific_fallbacks() { "tool", tool_response.clone(), &entries, - ), + ) + .await, tool_response ); let entries = state.llm_sanitize_request_entries(&[]); - assert!( + assert_eq!( NemoRelayContextState::llm_sanitize_request_snapshot_chain( llm_request.clone(), crate::api::runtime::LlmSanitizeRequestContext::default(), &entries, ) - .is_none(), - "a worker request sanitizer failure must omit the observability payload" + .await, + Some(llm_request), ); let entries = state.llm_sanitize_response_entries(&[]); - assert!( + assert_eq!( NemoRelayContextState::llm_sanitize_response_snapshot_chain( llm_response.clone(), crate::api::runtime::LlmSanitizeResponseContext::default(), &entries, ) - .is_none(), - "a worker response sanitizer failure must omit the observability payload" + .await, + Some(llm_response), ); } crate::api::subscriber::flush_subscribers().expect("subscriber callback should flush"); diff --git a/crates/core/tests/unit/llm_api_tests.rs b/crates/core/tests/unit/llm_api_tests.rs index 65c6e5b17..843fe4bf5 100644 --- a/crates/core/tests/unit/llm_api_tests.rs +++ b/crates/core/tests/unit/llm_api_tests.rs @@ -295,7 +295,7 @@ fn credential_headers_are_removed_before_request_sanitizers_and_event_emission() 1, Arc::new(move |request, _context| { sanitizer_capture.lock().unwrap().push(request.clone()); - Some(request) + Box::pin(async move { Ok(Some(request)) }) }), ) .unwrap(); @@ -431,13 +431,13 @@ fn sanitization_invalidates_manual_annotations_without_a_codec() { register_llm_sanitize_request_guardrail( "manual-annotation-invalidation-request", 1, - Arc::new(|_request, _context| Some(redacted_request())), + Arc::new(|_request, _context| Box::pin(async { Ok(Some(redacted_request())) })), ) .unwrap(); register_llm_sanitize_response_guardrail( "manual-annotation-invalidation-response", 1, - Arc::new(|_response, _context| Some(redacted_response())), + Arc::new(|_response, _context| Box::pin(async { Ok(Some(redacted_response())) })), ) .unwrap(); @@ -500,13 +500,13 @@ fn no_op_sanitizers_keep_manual_annotations() { register_llm_sanitize_request_guardrail( "manual-annotation-noop-request", 1, - Arc::new(|request, _context| Some(request)), + Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); register_llm_sanitize_response_guardrail( "manual-annotation-noop-response", 1, - Arc::new(|response, _context| Some(response)), + Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); @@ -558,13 +558,13 @@ fn sanitization_regenerates_annotations_with_active_codecs() { register_llm_sanitize_request_guardrail( "active-codec-annotation-regeneration-request", 1, - Arc::new(|_request, _context| Some(redacted_request())), + Arc::new(|_request, _context| Box::pin(async { Ok(Some(redacted_request())) })), ) .unwrap(); register_llm_sanitize_response_guardrail( "active-codec-annotation-regeneration-response", 1, - Arc::new(|_response, _context| Some(redacted_response())), + Arc::new(|_response, _context| Box::pin(async { Ok(Some(redacted_response())) })), ) .unwrap(); @@ -647,7 +647,7 @@ fn buffered_null_fallback_is_sanitized_before_emission() { 1, Arc::new(move |response, _context| { sanitizer_inputs.lock().unwrap().push(response); - Some(Json::Null) + Box::pin(async { Ok(Some(Json::Null)) }) }), ) .unwrap(); @@ -676,7 +676,7 @@ fn buffered_null_fallback_is_sanitized_before_emission() { register_llm_sanitize_response_guardrail( "buffered-null-fallback-redacted", 1, - Arc::new(|_response, _context| Some(redacted_response())), + Arc::new(|_response, _context| Box::pin(async { Ok(Some(redacted_response())) })), ) .unwrap(); let handle = create_llm_handle( @@ -760,7 +760,7 @@ fn streaming_null_fallback_is_sanitized_before_emission() { 1, Arc::new(move |response, _context| { sanitizer_inputs.lock().unwrap().push(response); - Some(Json::Null) + Box::pin(async { Ok(Some(Json::Null)) }) }), ) .unwrap(); @@ -791,7 +791,7 @@ fn streaming_null_fallback_is_sanitized_before_emission() { register_llm_sanitize_response_guardrail( "streaming-null-fallback-redacted", 1, - Arc::new(|_response, _context| Some(redacted_response())), + Arc::new(|_response, _context| Box::pin(async { Ok(Some(redacted_response())) })), ) .unwrap(); runtime.block_on(async { @@ -1525,11 +1525,12 @@ fn failed_managed_calls_sanitize_fallback_end_data() { "failed-managed-call-sanitization", 1, Arc::new(move |response, context| { - sanitizer_inputs - .lock() - .unwrap() - .push((response, context.codec().clone())); - Some(redacted_response()) + let codec = context.codec().clone(); + let sanitizer_inputs = Arc::clone(&sanitizer_inputs); + Box::pin(async move { + sanitizer_inputs.lock().unwrap().push((response, codec)); + Ok(Some(redacted_response())) + }) }), ) .unwrap(); diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index beec626af..a29955c17 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -257,10 +257,135 @@ fn native_string_and_json_helpers_cover_abi_boundaries() { assert_eq!(host_api.abi_version, NEMO_RELAY_NATIVE_ABI_VERSION); assert_eq!( host_api.struct_size, - std::mem::size_of::() + std::mem::size_of::() ); } +#[test] +fn native_async_next_abi_runs_tool_llm_and_stream_continuations() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let cases: Vec<(NativeAsyncNextInner, Json, Json)> = vec![ + ( + NativeAsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + json!({"tool": true}), + json!({"result": {"tool": true}, "pending_marks": []}), + ), + ( + NativeAsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + serde_json::to_value(LlmRequest { + headers: Map::new(), + content: json!({"llm": true}), + }) + .unwrap(), + json!({"llm": true}), + ), + ( + NativeAsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![ + Ok(json!({"chunk": 1})), + Ok(json!({"chunk": 2})), + ]))) + }) + })), + serde_json::to_value(LlmRequest { + headers: Map::new(), + content: json!({"stream": true}), + }) + .unwrap(), + json!([{"chunk": 1}, {"chunk": 2}]), + ), + ]; + + for (inner, invocation, expected) in cases { + let next = Arc::new(NativeAsyncNext { + inner, + runtime: runtime.handle().clone(), + _callback_user_data: None, + }); + let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invocation = native_string_from_json(&invocation).unwrap(); + assert_eq!( + unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, + NemoRelayStatus::Ok + ); + assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); + unsafe { + native_string_free(invocation); + native_async_next_release(next_ref); + native_async_completion_release(completion_ref); + } + } +} + +#[test] +fn native_async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlement() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invalid = native_string("not-json"); + assert_eq!( + unsafe { native_async_completion_resolve_json(completion_ref, invalid) }, + NemoRelayStatus::InvalidJson + ); + let value = native_string(r#"{"ok":true}"#); + assert_eq!( + unsafe { native_async_completion_resolve_json(completion_ref, value) }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { native_async_completion_resolve_json(completion_ref, value) }, + NemoRelayStatus::InvalidArg + ); + assert_eq!( + runtime.block_on(receiver).unwrap().unwrap(), + json!({"ok": true}) + ); + unsafe { + native_string_free(invalid); + native_string_free(value); + native_async_completion_release(completion_ref); + } + + let (sender, _receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(true), + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + assert!(unsafe { native_async_completion_is_cancelled(completion_ref) }); + assert!(unsafe { native_async_completion_is_cancelled(ptr::null()) }); + assert_eq!( + unsafe { native_async_completion_reject(completion_ref, ptr::null()) }, + NemoRelayStatus::InvalidArg + ); + unsafe { native_async_completion_release(completion_ref) }; +} + #[test] fn native_timestamp_scope_type_and_error_mappings_cover_variants() { assert_eq!(optional_timestamp_from_native(ptr::null()).unwrap(), None); diff --git a/crates/core/tests/unit/plugin_tests.rs b/crates/core/tests/unit/plugin_tests.rs index b39abb9d2..f7c85186a 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)) + }) }), ) }) @@ -644,27 +646,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(); } @@ -688,25 +692,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(); } @@ -960,8 +966,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, @@ -973,7 +984,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(); @@ -1476,69 +1487,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:?}"), } @@ -1546,13 +1566,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(); @@ -1564,22 +1589,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( @@ -1587,7 +1636,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(); @@ -1597,7 +1646,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:", @@ -1606,14 +1655,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:", ); @@ -1621,14 +1670,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:", ); @@ -1636,14 +1685,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:", ); @@ -1651,14 +1700,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:", ); @@ -1666,25 +1715,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:", ); @@ -1731,14 +1784,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:", ); @@ -1769,18 +1827,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| {})) @@ -1790,42 +1852,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, @@ -1844,8 +1910,13 @@ fn test_plugin_registration_context_maps_deregistration_errors() { }), ) .unwrap(); - ctx.register_tool_request_intercept("tool-request", 1, false, Arc::new(|_name, args| Ok(args))) - .unwrap(); + ctx.register_tool_request_intercept( + "tool-request", + 1, + false, + Arc::new(|_name, args| Box::pin(async move { Ok(args) })), + ) + .unwrap(); ctx.register_tool_execution_intercept( "tool-exec", 1, diff --git a/crates/core/tests/unit/shared_tests.rs b/crates/core/tests/unit/shared_tests.rs index 9d17ce188..1de4a9ed7 100644 --- a/crates/core/tests/unit/shared_tests.rs +++ b/crates/core/tests/unit/shared_tests.rs @@ -168,8 +168,9 @@ fn stale_process_runtime_owner_is_reclaimed() { reset_global(); } -#[test] -fn test_run_request_intercepts_with_codec_none_and_codec_paths() { +#[tokio::test] +#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. +async fn test_run_request_intercepts_with_codec_none_and_codec_paths() { let _guard = lock_runtime_owner(); reset_global(); @@ -178,11 +179,13 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { 1, false, Arc::new(|_name, mut request, annotated| { - assert!(annotated.is_none()); - request.headers.insert("x-no-codec".into(), json!(true)); - let mut annotated = SharedTestCodec.decode(&request)?; - annotated.model = Some("interceptor-model".into()); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + Box::pin(async move { + assert!(annotated.is_none()); + request.headers.insert("x-no-codec".into(), json!(true)); + let mut annotated = SharedTestCodec.decode(&request)?; + annotated.model = Some("interceptor-model".into()); + Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + }) }), ) .unwrap(); @@ -196,6 +199,7 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { }, None, ) + .await .unwrap(); assert_eq!( request_without_codec.headers.get("x-no-codec"), @@ -215,10 +219,12 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { 1, false, Arc::new(|_name, mut request, annotated| { - let mut annotated = annotated.expect("codec should provide annotated request"); - annotated.model = Some("intercepted-model".into()); - request.headers.insert("x-codec".into(), json!(true)); - Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + Box::pin(async move { + let mut annotated = annotated.expect("codec should provide annotated request"); + annotated.model = Some("intercepted-model".into()); + request.headers.insert("x-codec".into(), json!(true)); + Ok(LlmRequestInterceptOutcome::new(request, Some(annotated))) + }) }), ) .unwrap(); @@ -233,6 +239,7 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { }, Some(codec), ) + .await .unwrap(); assert_eq!( @@ -259,8 +266,9 @@ fn test_run_request_intercepts_with_codec_none_and_codec_paths() { reset_global(); } -#[test] -fn managed_request_chain_records_contributions_incrementally_while_standalone_retains_them() { +#[tokio::test] +#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. +async fn managed_request_chain_records_contributions_incrementally_while_standalone_retains_them() { let _guard = lock_runtime_owner(); reset_global(); @@ -269,11 +277,12 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re 1, false, Arc::new(|_name, request, annotated| { - Ok( - LlmRequestInterceptOutcome::new(request, annotated).with_optimization_contribution( - LlmOptimizationContribution::new("accepted", "custom"), - ), - ) + Box::pin(async move { + Ok(LlmRequestInterceptOutcome::new(request, annotated) + .with_optimization_contribution(LlmOptimizationContribution::new( + "accepted", "custom", + ))) + }) }), ) .unwrap(); @@ -282,14 +291,13 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re 2, false, Arc::new(|_name, request, annotated| { - Ok( - LlmRequestInterceptOutcome::new(request, annotated).with_optimization_contribution( - LlmOptimizationContribution::new( + Box::pin(async move { + Ok(LlmRequestInterceptOutcome::new(request, annotated) + .with_optimization_contribution(LlmOptimizationContribution::new( "x".repeat(MAX_LLM_OPTIMIZATION_CONTRIBUTION_BYTES), "custom", - ), - ), - ) + ))) + }) }), ) .unwrap(); @@ -302,6 +310,7 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re }, None, ) + .await .unwrap(); assert_eq!(standalone.3.len(), 2); assert!(standalone.3.iter().all(|item| item.sequence.is_none())); @@ -316,6 +325,7 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re None, &recorder, ) + .await .unwrap(); assert!(managed.3.is_empty()); let recorded = recorder.unemitted(); @@ -326,8 +336,9 @@ fn managed_request_chain_records_contributions_incrementally_while_standalone_re reset_global(); } -#[test] -fn test_run_request_intercepts_injects_dynamo_agent_lineage() { +#[tokio::test] +#[allow(clippy::await_holding_lock)] // Serializes access to global runtime state. +async fn test_run_request_intercepts_injects_dynamo_agent_lineage() { let _guard = lock_runtime_owner(); reset_global(); @@ -369,6 +380,7 @@ fn test_run_request_intercepts_injects_dynamo_agent_lineage() { }, None, ) + .await .unwrap(); assert_eq!( request.headers.get(DYNAMO_SESSION_ID_HEADER_KEY), @@ -396,6 +408,7 @@ fn test_run_request_intercepts_injects_dynamo_agent_lineage() { }, Some(Arc::new(SharedTestCodec)), ) + .await .unwrap(); assert_eq!( request_with_codec.headers.get(DYNAMO_SESSION_ID_HEADER_KEY), @@ -433,6 +446,7 @@ fn test_run_request_intercepts_injects_dynamo_agent_lineage() { }, None, ) + .await .unwrap(); assert_eq!( request.headers.get(DYNAMO_SESSION_ID_HEADER_KEY), @@ -475,6 +489,7 @@ fn test_run_request_intercepts_injects_dynamo_agent_lineage() { }, None, ) + .await .unwrap(); assert!(!request.headers.contains_key(DYNAMO_SESSION_ID_HEADER_KEY)); assert!( diff --git a/crates/ffi/build.rs b/crates/ffi/build.rs index 013b53194..d7f0d3d32 100644 --- a/crates/ffi/build.rs +++ b/crates/ffi/build.rs @@ -5,6 +5,7 @@ fn main() { let crate_dir = std::env::var("CARGO_MANIFEST_DIR").unwrap(); + validate_async_registration_parity(&crate_dir); let config = cbindgen::Config::from_file(format!("{crate_dir}/cbindgen.toml")) .expect("Unable to read cbindgen.toml"); @@ -13,6 +14,153 @@ fn main() { .with_config(config) .generate() { - bindings.write_to_file(format!("{crate_dir}/nemo_relay.h")); + let header_path = format!("{crate_dir}/nemo_relay.h"); + bindings.write_to_file(&header_path); + // cbindgen intentionally does not expand declarative macros. Keep the + // macro-generated async registration functions in the generated C ABI. + let header = std::fs::read_to_string(&header_path).expect("read generated FFI header"); + let marker = "\n#endif /* NEMO_RELAY_H */\n"; + assert!( + header.contains(marker), + "generated FFI header is missing its NEMO_RELAY_H closing guard" + ); + let header = header.replacen( + marker, + &format!("\n{}\n#endif /* NEMO_RELAY_H */\n", ASYNC_REGISTRATIONS), + 1, + ); + std::fs::write(header_path, header).expect("write generated FFI header"); } } + +#[derive(Debug, PartialEq, Eq)] +struct AsyncPrototype<'a> { + name: &'a str, + parameters: Vec<&'a str>, +} + +/// cbindgen does not expand the declarative registration macros. Keep the +/// handwritten C declarations checked against macro-generated exports, their +/// complete parameter lists, and ordering. +fn validate_async_registration_parity(crate_dir: &str) { + const REGISTRATION_SOURCES: &[&str] = &[ + "src/api/event_registry.rs", + "src/api/llm_registry.rs", + "src/api/scope_registry.rs", + "src/api/tool_registry.rs", + ]; + + let mut exported = Vec::new(); + for source in REGISTRATION_SOURCES { + println!("cargo:rerun-if-changed={source}"); + let source_path = format!("{crate_dir}/{source}"); + let contents = std::fs::read_to_string(&source_path) + .unwrap_or_else(|error| panic!("read {source_path}: {error}")); + exported.extend( + contents + .split(|character: char| !character.is_ascii_alphanumeric() && character != '_') + .filter(|token| token.starts_with("nemo_relay_") && token.ends_with("_async")) + .map(str::to_owned), + ); + } + exported.sort(); + exported.dedup(); + + let mut declared = ASYNC_REGISTRATIONS + .lines() + .filter_map(parse_async_prototype) + .collect::>(); + declared.sort_by(|left, right| left.name.cmp(right.name)); + for duplicates in declared.windows(2) { + assert_ne!( + duplicates[0].name, duplicates[1].name, + "ASYNC_REGISTRATIONS contains duplicate declaration for {}", + duplicates[0].name + ); + } + let declared_names = declared + .iter() + .map(|prototype| prototype.name.to_owned()) + .collect::>(); + assert_eq!( + declared_names, exported, + "ASYNC_REGISTRATIONS must declare exactly the async Rust FFI exports" + ); + + for prototype in declared { + assert!( + exported + .binary_search_by(|name| name.as_str().cmp(prototype.name)) + .is_ok(), + "async declaration for {} is not a Rust FFI export", + prototype.name + ); + assert_eq!( + prototype, + expected_async_prototype(prototype.name), + "async declaration for {} has a mismatched C prototype", + prototype.name + ); + } +} + +fn parse_async_prototype(line: &str) -> Option> { + let line = line.strip_prefix("NemoRelayStatus ")?; + let (name, parameters) = line.split_once('(')?; + let parameters = parameters.strip_suffix(");")?; + Some(AsyncPrototype { + name, + parameters: parameters.split(", ").collect(), + }) +} + +fn expected_async_prototype(name: &str) -> AsyncPrototype<'_> { + let mut parameters = Vec::new(); + if name.starts_with("nemo_relay_scope_") { + parameters.push("const char *scope_uuid"); + } + parameters.extend(["const char *name", "int32_t priority"]); + if name.contains("request_intercept_async") { + parameters.push("bool break_chain"); + } + parameters.push(if name.contains("execution_intercept_async") { + "NemoRelayAsyncInterceptCb cb" + } else { + "NemoRelayAsyncJsonCb cb" + }); + parameters.extend(["void *user_data", "NemoRelayFreeFn free_fn"]); + AsyncPrototype { name, parameters } +} + +const ASYNC_REGISTRATIONS: &str = r#" +/* Completion-based async middleware registrations generated from Rust macros. */ +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion); +NemoRelayStatus nemo_relay_register_mark_sanitize_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_scope_sanitize_start_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_scope_sanitize_end_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_mark_sanitize_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_scope_sanitize_start_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_scope_sanitize_end_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_sanitize_request_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_sanitize_response_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_sanitize_request_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_sanitize_response_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_stream_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_stream_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +"#; diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index be8822bc4..8a3f00170 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -132,6 +132,21 @@ enum NemoRelayScopeType { }; typedef int32_t NemoRelayScopeType; +/** + * Indicates whether an async callback settled its completion before returning. + */ +enum NemoRelayAsyncCallbackState { + /** + * The callback called a resolve/reject function before returning. + */ + NEMO_RELAY_ASYNC_CALLBACK_STATE_COMPLETE = 0, + /** + * The callback retained the completion and will settle it later. + */ + NEMO_RELAY_ASYNC_CALLBACK_STATE_PENDING = 1, +}; +typedef uint32_t NemoRelayAsyncCallbackState; + /** * Opaque owned adaptive runtime handle. */ @@ -240,6 +255,16 @@ typedef struct FfiThreadScopeStackBinding FfiThreadScopeStackBinding; */ typedef struct FfiToolHandle FfiToolHandle; +/** + * One-shot completion passed to asynchronous C callbacks. + */ +typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; + +/** + * Runtime-owned asynchronous `next` continuation for execution intercepts. + */ +typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; + typedef struct Option_NemoRelayCollectorCb Option_NemoRelayCollectorCb; typedef struct Option_NemoRelayFinalizerCb Option_NemoRelayFinalizerCb; @@ -439,6 +464,26 @@ typedef char *(*NemoRelayToolExecInterceptCb)(void *user_data, */ typedef char *(*NemoRelayToolExecCb)(void *user_data, const char *args_json); +/** + * Completion-based execution-intercept callback. + */ +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncInterceptCb)(void *user_data, + const char *invocation_json, + const struct NemoRelayAsyncNext *next, + const struct NemoRelayAsyncCompletion *completion); + +/** + * Result callback used by channel/future-style async `next` wrappers. + * + * Invoked on a Tokio runtime worker thread, not necessarily the thread that + * called `nemo_relay_async_next_invoke_callback`; `user_data` must therefore + * be safe for cross-thread use. `value_json` and `error_message` are borrowed + * for the duration of the callback only. + */ +typedef void (*NemoRelayAsyncNextResultCb)(void *user_data, + const char *value_json, + const char *error_message); + /** * Run the registered tool request intercept chain on the given arguments. * @@ -2548,6 +2593,15 @@ NemoRelayStatus nemo_relay_tool_call_execute(const char *name, const char *metadata_json, char **out); +/** + * Register a completion-based asynchronous tool execution intercept. + */ +NemoRelayStatus nemo_relay_register_tool_execution_intercept_async(const char *name, + int32_t priority, + NemoRelayAsyncInterceptCb cb, + void *user_data, + NemoRelayFreeFn free_fn); + /** * Register a tool conditional execution guardrail. The callback decides whether * a tool call should proceed. Returns an error message to reject, or null to allow. @@ -2610,6 +2664,48 @@ NemoRelayStatus nemo_relay_register_tool_execution_intercept(const char *name, */ NemoRelayStatus nemo_relay_deregister_tool_execution_intercept(const char *name); +/** + * Release the callback-owned async `next` reference after a pending intercept. + */ +void nemo_relay_async_next_release(const struct NemoRelayAsyncNext *next); + +/** + * Invoke the next execution layer and settle `completion` with its result. + */ +NemoRelayStatus nemo_relay_async_next_invoke(const struct NemoRelayAsyncNext *next, + const char *invocation_json, + const struct NemoRelayAsyncCompletion *completion); + +/** + * Invoke the next execution layer and report its result through a callback. + */ +NemoRelayStatus nemo_relay_async_next_invoke_callback(const struct NemoRelayAsyncNext *next, + const char *invocation_json, + NemoRelayAsyncNextResultCb callback, + void *user_data); + +/** + * Resolve an async C callback with owned JSON. + */ +NemoRelayStatus nemo_relay_async_completion_resolve_json(const struct NemoRelayAsyncCompletion *completion, + const char *value_json); + +/** + * Reject an async C callback with an error message. + */ +NemoRelayStatus nemo_relay_async_completion_reject(const struct NemoRelayAsyncCompletion *completion, + const char *message); + +/** + * Returns whether an async completion's invocation has been cancelled. + */ +bool nemo_relay_async_completion_is_cancelled(const struct NemoRelayAsyncCompletion *completion); + +/** + * Release the callback-owned completion reference after a pending invocation. + */ +void nemo_relay_async_completion_release(const struct NemoRelayAsyncCompletion *completion); + /** * Free a C string previously returned by any `nemo_relay_*` accessor function. * Passing null is a safe no-op. @@ -3113,4 +3209,36 @@ char *nemo_relay_event_annotated_request(const struct FfiEvent *ptr); */ char *nemo_relay_event_annotated_response(const struct FfiEvent *ptr); + +/* Completion-based async middleware registrations generated from Rust macros. */ +typedef NemoRelayAsyncCallbackState (*NemoRelayAsyncJsonCb)(void *user_data, const char *invocation_json, const struct NemoRelayAsyncCompletion *completion); +NemoRelayStatus nemo_relay_register_mark_sanitize_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_scope_sanitize_start_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_scope_sanitize_end_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_mark_sanitize_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_scope_sanitize_start_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_scope_sanitize_end_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_sanitize_request_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_sanitize_response_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_tool_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_sanitize_request_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_sanitize_response_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_conditional_execution_guardrail_async(const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_request_intercept_async(const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_register_llm_stream_execution_intercept_async(const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_sanitize_request_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_sanitize_response_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_conditional_execution_guardrail_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_request_intercept_async(const char *scope_uuid, const char *name, int32_t priority, bool break_chain, NemoRelayAsyncJsonCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_tool_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); +NemoRelayStatus nemo_relay_scope_register_llm_stream_execution_intercept_async(const char *scope_uuid, const char *name, int32_t priority, NemoRelayAsyncInterceptCb cb, void *user_data, NemoRelayFreeFn free_fn); + #endif /* NEMO_RELAY_H */ diff --git a/crates/ffi/src/api/event_registry.rs b/crates/ffi/src/api/event_registry.rs index 13af2b8d8..4badaa381 100644 --- a/crates/ffi/src/api/event_registry.rs +++ b/crates/ffi/src/api/event_registry.rs @@ -2,8 +2,9 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayEventSanitizeCb, NemoRelayFreeFn, NemoRelayStatus, c_char, c_str_to_string, - clear_last_error, core_registry_api, set_last_error, status_from_error, wrap_event_sanitize_fn, + NemoRelayAsyncJsonCb, NemoRelayEventSanitizeCb, NemoRelayFreeFn, NemoRelayStatus, c_char, + c_str_to_string, clear_last_error, core_registry_api, set_last_error, status_from_error, + wrap_async_event_sanitize_fn, wrap_event_sanitize_fn, }; #[derive(Clone, Copy)] @@ -43,6 +44,66 @@ unsafe fn register_global( .unwrap_or_else(|error| status_from_error(&error)) } +unsafe fn register_global_async( + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + surface: Surface, +) -> NemoRelayStatus { + clear_last_error(); + let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + let result = match surface { + Surface::Mark => { + core_registry_api::register_mark_sanitize_guardrail(&name, priority, callback) + } + Surface::Start => { + core_registry_api::register_scope_sanitize_start_guardrail(&name, priority, callback) + } + Surface::End => { + core_registry_api::register_scope_sanitize_end_guardrail(&name, priority, callback) + } + }; + result + .map(|()| NemoRelayStatus::Ok) + .unwrap_or_else(|error| status_from_error(&error)) +} + +macro_rules! async_event_registration { + ($name:ident, $surface:expr) => { + /// Register a completion-based asynchronous event sanitizer. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name( + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + unsafe { register_global_async(name, priority, cb, user_data, free_fn, $surface) } + } + }; +} + +async_event_registration!( + nemo_relay_register_mark_sanitize_guardrail_async, + Surface::Mark +); +async_event_registration!( + nemo_relay_register_scope_sanitize_start_guardrail_async, + Surface::Start +); +async_event_registration!( + nemo_relay_register_scope_sanitize_end_guardrail_async, + Surface::End +); + unsafe fn deregister_global(name: *const c_char, surface: Surface) -> NemoRelayStatus { clear_last_error(); let name = match c_str_to_string(name) { @@ -102,6 +163,74 @@ unsafe fn register_scope( .unwrap_or_else(|error| status_from_error(&error)) } +unsafe fn register_scope_async( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + surface: Surface, +) -> NemoRelayStatus { + clear_last_error(); + let uuid = match parse_scope_uuid(scope_uuid) { + Ok(uuid) => uuid, + Err(status) => return status, + }; + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + let callback = wrap_async_event_sanitize_fn(cb, user_data, free_fn); + let result = match surface { + Surface::Mark => core_registry_api::scope_register_mark_sanitize_guardrail( + &uuid, &name, priority, callback, + ), + Surface::Start => core_registry_api::scope_register_scope_sanitize_start_guardrail( + &uuid, &name, priority, callback, + ), + Surface::End => core_registry_api::scope_register_scope_sanitize_end_guardrail( + &uuid, &name, priority, callback, + ), + }; + result + .map(|()| NemoRelayStatus::Ok) + .unwrap_or_else(|error| status_from_error(&error)) +} + +macro_rules! scope_async_event_registration { + ($name:ident, $surface:expr) => { + /// Register a scope-local completion-based asynchronous event sanitizer. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + unsafe { + register_scope_async(scope_uuid, name, priority, cb, user_data, free_fn, $surface) + } + } + }; +} + +scope_async_event_registration!( + nemo_relay_scope_register_mark_sanitize_guardrail_async, + Surface::Mark +); +scope_async_event_registration!( + nemo_relay_scope_register_scope_sanitize_start_guardrail_async, + Surface::Start +); +scope_async_event_registration!( + nemo_relay_scope_register_scope_sanitize_end_guardrail_async, + Surface::End +); + unsafe fn deregister_scope( scope_uuid: *const c_char, name: *const c_char, diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index 38883ecee..60d1d0053 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -2,14 +2,107 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayEventSubscriberCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, - NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, - NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, c_char, c_str_to_string, clear_last_error, - core_registry_api, core_subscriber_api, status_from_error, wrap_event_subscriber, - wrap_llm_conditional_fn, wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, - wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_intercept_fn, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayEventSubscriberCb, NemoRelayFreeFn, + NemoRelayLlmConditionalCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, + NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, c_char, + c_str_to_string, clear_last_error, core_registry_api, core_subscriber_api, status_from_error, + wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, + wrap_async_llm_request_intercept_fn, wrap_async_llm_sanitize_request_fn, + wrap_async_llm_sanitize_response_fn, wrap_async_llm_stream_execution_intercept_fn, + wrap_event_subscriber, wrap_llm_conditional_fn, wrap_llm_exec_intercept_fn, + wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, + wrap_llm_stream_exec_intercept_fn, }; +macro_rules! async_llm_registration { + ($fn_name:ident, $register:path, $wrapper:path $(, $break_chain:ident)?) => { + /// Register a completion-based asynchronous LLM middleware callback. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $fn_name( + name: *const c_char, + priority: i32, + $( $break_chain: bool, )? + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register( + &name, + priority, + $( $break_chain, )? + $wrapper(cb, user_data, free_fn), + ) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +async_llm_registration!( + nemo_relay_register_llm_sanitize_request_guardrail_async, + core_registry_api::register_llm_sanitize_request_guardrail, + wrap_async_llm_sanitize_request_fn +); + +macro_rules! async_llm_execution_registration { + ($name:ident, $register:path, $wrapper:path) => { + /// Register a completion-based asynchronous LLM execution intercept. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name( + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register(&name, priority, $wrapper(cb, user_data, free_fn)) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +async_llm_execution_registration!( + nemo_relay_register_llm_execution_intercept_async, + core_registry_api::register_llm_execution_intercept, + wrap_async_llm_execution_intercept_fn +); +async_llm_execution_registration!( + nemo_relay_register_llm_stream_execution_intercept_async, + core_registry_api::register_llm_stream_execution_intercept, + wrap_async_llm_stream_execution_intercept_fn +); +async_llm_registration!( + nemo_relay_register_llm_sanitize_response_guardrail_async, + core_registry_api::register_llm_sanitize_response_guardrail, + wrap_async_llm_sanitize_response_fn +); +async_llm_registration!( + nemo_relay_register_llm_conditional_execution_guardrail_async, + core_registry_api::register_llm_conditional_execution_guardrail, + wrap_async_llm_conditional_fn +); +async_llm_registration!( + nemo_relay_register_llm_request_intercept_async, + core_registry_api::register_llm_request_intercept, + wrap_async_llm_request_intercept_fn, + break_chain +); + // --------------------------------------------------------------------------- // LLM guardrail registrations // --------------------------------------------------------------------------- diff --git a/crates/ffi/src/api/mod.rs b/crates/ffi/src/api/mod.rs index 42f40bdca..6e53afd8b 100644 --- a/crates/ffi/src/api/mod.rs +++ b/crates/ffi/src/api/mod.rs @@ -14,17 +14,23 @@ use std::sync::{Arc, OnceLock}; use std::time::Duration; use crate::callable::{ - NemoRelayCodecDecodeFn, NemoRelayCodecEncodeFn, NemoRelayCollectorCb, NemoRelayEventSanitizeCb, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayCodecDecodeFn, + NemoRelayCodecEncodeFn, NemoRelayCollectorCb, NemoRelayEventSanitizeCb, NemoRelayEventSubscriberCb, NemoRelayFinalizerCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, NemoRelayLlmExecCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayPluginRegisterCb, NemoRelayPluginValidateCb, NemoRelayToolConditionalCb, NemoRelayToolExecCb, - NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, wrap_codec_fn, wrap_collector_fn, - wrap_event_sanitize_fn, wrap_event_subscriber, wrap_finalizer_fn, wrap_llm_conditional_fn, - wrap_llm_exec_fn, wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, - wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_fn, - wrap_llm_stream_exec_intercept_fn, wrap_tool_conditional_fn, wrap_tool_exec_fn, - wrap_tool_exec_intercept_fn, wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, + NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, wrap_async_event_sanitize_fn, + wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, + wrap_async_llm_request_intercept_fn, wrap_async_llm_sanitize_request_fn, + wrap_async_llm_sanitize_response_fn, wrap_async_llm_stream_execution_intercept_fn, + wrap_async_tool_conditional_fn, wrap_async_tool_execution_intercept_fn, + wrap_async_tool_json_fn, wrap_codec_fn, wrap_collector_fn, wrap_event_sanitize_fn, + wrap_event_subscriber, wrap_finalizer_fn, wrap_llm_conditional_fn, wrap_llm_exec_fn, + wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, + wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_fn, wrap_llm_stream_exec_intercept_fn, + wrap_tool_conditional_fn, wrap_tool_exec_fn, wrap_tool_exec_intercept_fn, + wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, }; use crate::convert::{ c_str_to_json, c_str_to_opt_json, c_str_to_string, json_to_c_string, nemo_relay_string_free, @@ -99,6 +105,15 @@ fn tokio_runtime() -> &'static Runtime { }) } +fn block_on_sync_ffi(future: impl Future>) -> FlowResult { + if tokio::runtime::Handle::try_current().is_ok() { + return Err(nemo_relay::error::FlowError::Internal( + "synchronous FFI middleware helpers cannot run on a Tokio runtime thread; use the completion-based async registration API".into(), + )); + } + tokio_runtime().block_on(future) +} + // --------------------------------------------------------------------------- // Standalone middleware chains // --------------------------------------------------------------------------- @@ -141,7 +156,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 +194,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 +248,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 +411,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..aebf09cd1 100644 --- a/crates/ffi/src/api/scope_registry.rs +++ b/crates/ffi/src/api/scope_registry.rs @@ -2,15 +2,19 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayEventSubscriberCb, NemoRelayFreeFn, NemoRelayLlmConditionalCb, - NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, NemoRelayLlmSanitizeRequestCb, - NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, NemoRelayToolConditionalCb, - NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, c_char, c_str_to_string, - clear_last_error, core_registry_api, core_subscriber_api, set_last_error, status_from_error, - wrap_event_subscriber, wrap_llm_conditional_fn, wrap_llm_exec_intercept_fn, - wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, wrap_llm_sanitize_response_fn, - wrap_llm_stream_exec_intercept_fn, wrap_tool_conditional_fn, wrap_tool_exec_intercept_fn, - wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayEventSubscriberCb, NemoRelayFreeFn, + NemoRelayLlmConditionalCb, NemoRelayLlmExecInterceptCb, NemoRelayLlmRequestInterceptCb, + NemoRelayLlmSanitizeRequestCb, NemoRelayLlmSanitizeResponseCb, NemoRelayStatus, + NemoRelayToolConditionalCb, NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, c_char, + c_str_to_string, clear_last_error, core_registry_api, core_subscriber_api, set_last_error, + status_from_error, wrap_async_llm_conditional_fn, wrap_async_llm_execution_intercept_fn, + wrap_async_llm_request_intercept_fn, wrap_async_llm_sanitize_request_fn, + wrap_async_llm_sanitize_response_fn, wrap_async_llm_stream_execution_intercept_fn, + wrap_async_tool_conditional_fn, wrap_async_tool_execution_intercept_fn, + wrap_async_tool_json_fn, wrap_event_subscriber, wrap_llm_conditional_fn, + wrap_llm_exec_intercept_fn, wrap_llm_request_intercept_fn, wrap_llm_sanitize_request_fn, + wrap_llm_sanitize_response_fn, wrap_llm_stream_exec_intercept_fn, wrap_tool_conditional_fn, + wrap_tool_exec_intercept_fn, wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, }; // --------------------------------------------------------------------------- @@ -26,6 +30,132 @@ fn parse_scope_uuid(scope_uuid: *const c_char) -> Result { + /// Register a scope-local completion-based asynchronous middleware callback. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $fn_name( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + $( $break_chain: bool, )? + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let uuid = match parse_scope_uuid(scope_uuid) { + Ok(uuid) => uuid, + Err(status) => return status, + }; + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register( + &uuid, + &name, + priority, + $( $break_chain, )? + $wrapper(cb, user_data, free_fn), + ) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +scope_async_registration!( + nemo_relay_scope_register_tool_sanitize_request_guardrail_async, + core_registry_api::scope_register_tool_sanitize_request_guardrail, + wrap_async_tool_json_fn +); +scope_async_registration!( + nemo_relay_scope_register_tool_sanitize_response_guardrail_async, + core_registry_api::scope_register_tool_sanitize_response_guardrail, + wrap_async_tool_json_fn +); +scope_async_registration!( + nemo_relay_scope_register_tool_conditional_execution_guardrail_async, + core_registry_api::scope_register_tool_conditional_execution_guardrail, + wrap_async_tool_conditional_fn +); +scope_async_registration!( + nemo_relay_scope_register_tool_request_intercept_async, + core_registry_api::scope_register_tool_request_intercept, + wrap_async_tool_json_fn, + break_chain +); +scope_async_registration!( + nemo_relay_scope_register_llm_sanitize_request_guardrail_async, + core_registry_api::scope_register_llm_sanitize_request_guardrail, + wrap_async_llm_sanitize_request_fn +); +scope_async_registration!( + nemo_relay_scope_register_llm_sanitize_response_guardrail_async, + core_registry_api::scope_register_llm_sanitize_response_guardrail, + wrap_async_llm_sanitize_response_fn +); +scope_async_registration!( + nemo_relay_scope_register_llm_conditional_execution_guardrail_async, + core_registry_api::scope_register_llm_conditional_execution_guardrail, + wrap_async_llm_conditional_fn +); +scope_async_registration!( + nemo_relay_scope_register_llm_request_intercept_async, + core_registry_api::scope_register_llm_request_intercept, + wrap_async_llm_request_intercept_fn, + break_chain +); + +macro_rules! scope_async_execution_registration { + ($fn_name:ident, $register:path, $wrapper:path) => { + /// Register a scope-local completion-based asynchronous execution intercept. + #[allow(clippy::missing_safety_doc)] + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $fn_name( + scope_uuid: *const c_char, + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let uuid = match parse_scope_uuid(scope_uuid) { + Ok(uuid) => uuid, + Err(status) => return status, + }; + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match $register(&uuid, &name, priority, $wrapper(cb, user_data, free_fn)) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +scope_async_execution_registration!( + nemo_relay_scope_register_tool_execution_intercept_async, + core_registry_api::scope_register_tool_execution_intercept, + wrap_async_tool_execution_intercept_fn +); +scope_async_execution_registration!( + nemo_relay_scope_register_llm_execution_intercept_async, + core_registry_api::scope_register_llm_execution_intercept, + wrap_async_llm_execution_intercept_fn +); +scope_async_execution_registration!( + nemo_relay_scope_register_llm_stream_execution_intercept_async, + core_registry_api::scope_register_llm_stream_execution_intercept, + wrap_async_llm_stream_execution_intercept_fn +); + macro_rules! ffi_scope_guardrail_tool_api { ($(#[$reg_doc:meta])* $register_name:ident, $(#[$dereg_doc:meta])* $deregister_name:ident, diff --git a/crates/ffi/src/api/tool_registry.rs b/crates/ffi/src/api/tool_registry.rs index 5d5cefb40..fd1f2e580 100644 --- a/crates/ffi/src/api/tool_registry.rs +++ b/crates/ffi/src/api/tool_registry.rs @@ -2,12 +2,89 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - NemoRelayFreeFn, NemoRelayStatus, NemoRelayToolConditionalCb, NemoRelayToolExecInterceptCb, - NemoRelayToolSanitizeCb, c_char, c_str_to_string, clear_last_error, core_registry_api, - status_from_error, wrap_tool_conditional_fn, wrap_tool_exec_intercept_fn, + NemoRelayAsyncInterceptCb, NemoRelayAsyncJsonCb, NemoRelayFreeFn, NemoRelayStatus, + NemoRelayToolConditionalCb, NemoRelayToolExecInterceptCb, NemoRelayToolSanitizeCb, c_char, + c_str_to_string, clear_last_error, core_registry_api, status_from_error, + wrap_async_tool_conditional_fn, wrap_async_tool_execution_intercept_fn, + wrap_async_tool_json_fn, wrap_tool_conditional_fn, wrap_tool_exec_intercept_fn, wrap_tool_request_intercept_fn, wrap_tool_sanitize_fn, }; +macro_rules! async_tool_json_registration { + ($name:ident, $register:path, $wrapper:path $(, $break_chain:ident)?) => { + /// Register a completion-based asynchronous tool middleware callback. + #[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. + #[unsafe(no_mangle)] + pub unsafe extern "C" fn $name( + name: *const c_char, + priority: i32, + $( $break_chain: bool, )? + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, + ) -> NemoRelayStatus { + clear_last_error(); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + let callback = $wrapper(cb, user_data, free_fn); + match $register(&name, priority, $( $break_chain, )? callback) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } + } + }; +} + +async_tool_json_registration!( + nemo_relay_register_tool_sanitize_request_guardrail_async, + core_registry_api::register_tool_sanitize_request_guardrail, + wrap_async_tool_json_fn +); +async_tool_json_registration!( + nemo_relay_register_tool_sanitize_response_guardrail_async, + core_registry_api::register_tool_sanitize_response_guardrail, + wrap_async_tool_json_fn +); +async_tool_json_registration!( + nemo_relay_register_tool_conditional_execution_guardrail_async, + core_registry_api::register_tool_conditional_execution_guardrail, + wrap_async_tool_conditional_fn +); + +async_tool_json_registration!( + nemo_relay_register_tool_request_intercept_async, + core_registry_api::register_tool_request_intercept, + wrap_async_tool_json_fn, + break_chain +); + +/// Register a completion-based asynchronous tool execution intercept. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_register_tool_execution_intercept_async( + name: *const c_char, + priority: i32, + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> NemoRelayStatus { + clear_last_error(); + let name = match c_str_to_string(name) { + Ok(name) => name, + Err(status) => return status, + }; + match core_registry_api::register_tool_execution_intercept( + &name, + priority, + wrap_async_tool_execution_intercept_fn(cb, user_data, free_fn), + ) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_error(&error), + } +} + // --------------------------------------------------------------------------- // Tool guardrail registrations // --------------------------------------------------------------------------- diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index f327e01d2..bb9bcbe90 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -19,14 +19,17 @@ use std::ffi::{CStr, CString}; use std::future::Future; use std::pin::Pin; +use std::ptr; use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; use libc::c_char; use nemo_relay::api::runtime::{ - EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmConditionalFn, LlmExecutionNextFn, - LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, - LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionNextFn, ToolConditionalFn, - ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmConditionalFn, LlmExecutionFn, + LlmExecutionNextFn, LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, + LlmSanitizeRequestFn, LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, + LlmStreamExecutionNextFn, ToolConditionalFn, ToolExecutionFn, ToolExecutionNextFn, + ToolInterceptFn, ToolSanitizeFn, }; use serde_json::Value as Json; use tokio_stream::StreamExt; @@ -50,6 +53,401 @@ use crate::types::{FfiEvent, FfiLLMRequest, FfiPluginContext}; /// Called when the runtime no longer needs the associated callback. pub type NemoRelayFreeFn = Option; +/// Indicates whether an async callback settled its completion before returning. +#[repr(u32)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum NemoRelayAsyncCallbackState { + /// The callback called a resolve/reject function before returning. + Complete = 0, + /// The callback retained the completion and will settle it later. + Pending = 1, +} + +/// One-shot completion passed to asynchronous C callbacks. +pub struct NemoRelayAsyncCompletion { + sender: std::sync::Mutex>>>, + cancelled: AtomicBool, +} + +/// Generic completion-based middleware callback. +/// +/// `invocation_json` is borrowed for the duration of the call. The completion +/// has one callback-owned reference. A callback returning `Complete` need not +/// release it; a callback returning `Pending` must eventually settle and call +/// `nemo_relay_async_completion_release`. +pub type NemoRelayAsyncJsonCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState; + +/// Runtime-owned asynchronous `next` continuation for execution intercepts. +pub struct NemoRelayAsyncNext { + inner: AsyncNextInner, + runtime: tokio::runtime::Handle, +} + +enum AsyncNextInner { + Tool(ToolExecutionNextFn), + Llm(LlmExecutionNextFn), + LlmStream(LlmStreamExecutionNextFn), +} + +/// Completion-based execution-intercept callback. +pub type NemoRelayAsyncInterceptCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState; + +/// Result callback used by channel/future-style async `next` wrappers. +/// +/// Invoked on a Tokio runtime worker thread, not necessarily the thread that +/// called `nemo_relay_async_next_invoke_callback`; `user_data` must therefore +/// be safe for cross-thread use. `value_json` and `error_message` are borrowed +/// for the duration of the callback only. +pub type NemoRelayAsyncNextResultCb = unsafe extern "C" fn( + user_data: *mut libc::c_void, + value_json: *const c_char, + error_message: *const c_char, +); + +struct CompletionWait { + completion: Arc, + receiver: tokio::sync::oneshot::Receiver>, +} + +impl Drop for CompletionWait { + fn drop(&mut self) { + self.completion.cancelled.store(true, Ordering::Release); + } +} + +async fn invoke_async_json( + cb: NemoRelayAsyncJsonCb, + user_data: Arc, + invocation: Json, +) -> Result { + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let callback_ref = Arc::into_raw(completion.clone()); + let invocation = json_to_c_string(&invocation); + let state = unsafe { cb(user_data.ptr, invocation, callback_ref) }; + unsafe { nemo_relay_string_free_internal(invocation) }; + if state == NemoRelayAsyncCallbackState::Complete { + unsafe { drop(Arc::from_raw(callback_ref)) }; + if completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .is_some() + { + return Err(FlowError::Internal( + "async C callback returned Complete without settling".into(), + )); + } + } + let mut wait = CompletionWait { + completion, + receiver, + }; + (&mut wait.receiver) + .await + .map_err(|_| FlowError::Internal("async C callback dropped without settling".into()))? +} + +async fn invoke_async_intercept( + cb: NemoRelayAsyncInterceptCb, + user_data: Arc, + invocation: Json, + next: AsyncNextInner, +) -> Result { + let runtime = tokio::runtime::Handle::try_current().map_err(|error| { + FlowError::Internal(format!( + "async C intercept requires a Tokio runtime: {error}" + )) + })?; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let callback_ref = Arc::into_raw(completion.clone()); + let next = Arc::new(NemoRelayAsyncNext { + inner: next, + runtime, + }); + let next_ref = Arc::into_raw(next); + let invocation = json_to_c_string(&invocation); + let state = unsafe { cb(user_data.ptr, invocation, next_ref, callback_ref) }; + unsafe { nemo_relay_string_free_internal(invocation) }; + if state == NemoRelayAsyncCallbackState::Complete { + unsafe { drop(Arc::from_raw(callback_ref)) }; + unsafe { drop(Arc::from_raw(next_ref)) }; + if completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .is_some() + { + return Err(FlowError::Internal( + "async C intercept returned Complete without settling".into(), + )); + } + } + let mut wait = CompletionWait { + completion, + receiver, + }; + (&mut wait.receiver) + .await + .map_err(|_| FlowError::Internal("async C intercept dropped without settling".into()))? +} + +/// Release the callback-owned async `next` reference after a pending intercept. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_next_release(next: *const NemoRelayAsyncNext) { + if !next.is_null() { + unsafe { drop(Arc::from_raw(next)) }; + } +} + +/// Invoke the next execution layer and settle `completion` with its result. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_next_invoke( + next: *const NemoRelayAsyncNext, + invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayStatus { + let Some(next) = (unsafe { next.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if completion.is_null() { + return NemoRelayStatus::NullPointer; + } + let Some(invocation) = c_str_to_json(invocation_json) else { + return NemoRelayStatus::InvalidJson; + }; + unsafe { Arc::increment_strong_count(completion) }; + let completion = unsafe { Arc::from_raw(completion) }; + let future: Pin> + Send>> = match &next.inner { + AsyncNextInner::Tool(next) => { + let next = next.clone(); + Box::pin(async move { + let outcome = next(invocation).await?; + serde_json::to_value(outcome) + .map_err(|error| FlowError::Internal(error.to_string())) + }) + } + AsyncNextInner::Llm(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(error) => { + return { + let _ = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); + NemoRelayStatus::InvalidJson + }; + } + }; + let next = next.clone(); + Box::pin(async move { next(request).await }) + } + AsyncNextInner::LlmStream(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(error) => { + return { + let _ = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|sender| sender.send(Err(FlowError::Internal(error.to_string())))); + NemoRelayStatus::InvalidJson + }; + } + }; + let next = next.clone(); + Box::pin(async move { + let mut stream = next(request).await?; + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk?); + } + Ok(Json::Array(chunks)) + }) + } + }; + next.runtime.spawn(async move { + let result = future.await; + if let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + { + let _ = sender.send(result); + } + }); + NemoRelayStatus::Ok +} + +/// Invoke the next execution layer and report its result through a callback. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_next_invoke_callback( + next: *const NemoRelayAsyncNext, + invocation_json: *const c_char, + callback: NemoRelayAsyncNextResultCb, + user_data: *mut libc::c_void, +) -> NemoRelayStatus { + let Some(next) = (unsafe { next.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + let Some(invocation) = c_str_to_json(invocation_json) else { + return NemoRelayStatus::InvalidJson; + }; + let future: Pin> + Send>> = match &next.inner { + AsyncNextInner::Tool(next) => { + let next = next.clone(); + Box::pin(async move { + serde_json::to_value(next(invocation).await?) + .map_err(|error| FlowError::Internal(error.to_string())) + }) + } + AsyncNextInner::Llm(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(_) => return NemoRelayStatus::InvalidJson, + }; + let next = next.clone(); + Box::pin(async move { next(request).await }) + } + AsyncNextInner::LlmStream(next) => { + let request = match serde_json::from_value(invocation) { + Ok(request) => request, + Err(_) => return NemoRelayStatus::InvalidJson, + }; + let next = next.clone(); + Box::pin(async move { + let mut stream = next(request).await?; + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk?); + } + Ok(Json::Array(chunks)) + }) + } + }; + let user_data = user_data as usize; + next.runtime.spawn(async move { + match future.await { + Ok(value) => { + let value = json_to_c_string(&value); + unsafe { callback(user_data as *mut libc::c_void, value, ptr::null()) }; + unsafe { nemo_relay_string_free_internal(value) }; + } + Err(error) => { + let error = CString::new(error.to_string()).unwrap_or_default(); + unsafe { callback(user_data as *mut libc::c_void, ptr::null(), error.as_ptr()) }; + } + } + }); + NemoRelayStatus::Ok +} + +/// Resolve an async C callback with owned JSON. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_completion_resolve_json( + completion: *const NemoRelayAsyncCompletion, + value_json: *const c_char, +) -> NemoRelayStatus { + let Some(completion) = (unsafe { completion.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if completion.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let Some(value) = c_str_to_json(value_json) else { + return NemoRelayStatus::InvalidJson; + }; + let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + else { + return NemoRelayStatus::InvalidArg; + }; + let _ = sender.send(Ok(value)); + NemoRelayStatus::Ok +} + +/// Reject an async C callback with an error message. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_completion_reject( + completion: *const NemoRelayAsyncCompletion, + message: *const c_char, +) -> NemoRelayStatus { + let Some(completion) = (unsafe { completion.as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + if completion.cancelled.load(Ordering::Acquire) { + return NemoRelayStatus::InvalidArg; + } + let message = if message.is_null() { + "async C callback rejected".to_string() + } else { + unsafe { CStr::from_ptr(message) } + .to_string_lossy() + .into_owned() + }; + let Some(sender) = completion + .sender + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + else { + return NemoRelayStatus::InvalidArg; + }; + let _ = sender.send(Err(FlowError::Internal(message))); + NemoRelayStatus::Ok +} + +/// Returns whether an async completion's invocation has been cancelled. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_completion_is_cancelled( + completion: *const NemoRelayAsyncCompletion, +) -> bool { + unsafe { completion.as_ref() } + .is_none_or(|completion| completion.cancelled.load(Ordering::Acquire)) +} + +/// Release the callback-owned completion reference after a pending invocation. +#[allow(clippy::missing_safety_doc)] // The shared C ABI safety contract applies. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_async_completion_release( + completion: *const NemoRelayAsyncCompletion, +) { + if !completion.is_null() { + unsafe { drop(Arc::from_raw(completion)) }; + } +} + /// Callback for tool request/response sanitization guardrails and intercepts. /// Receives tool name and arguments as JSON, returns sanitized arguments as JSON. /// The returned string must be allocated with `malloc` or equivalent. @@ -310,6 +708,249 @@ fn make_user_data( }) } +/// Wrap a completion-based C tool sanitizer or request intercept. +pub fn wrap_async_tool_json_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> ToolSanitizeFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |name: String, value: Json| { + let user_data = user_data.clone(); + Box::pin(invoke_async_json( + cb, + user_data, + serde_json::json!({"name": name, "value": value}), + )) + }) +} + +/// Wrap a completion-based C tool conditional guardrail. +pub fn wrap_async_tool_conditional_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> ToolConditionalFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |name: String, value: Json| { + let user_data = user_data.clone(); + Box::pin(async move { + match invoke_async_json( + cb, + user_data, + serde_json::json!({"name": name, "value": value}), + ) + .await? + { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => Err(FlowError::Internal(format!( + "async conditional callback returned {other}; expected string or null" + ))), + } + }) + }) +} + +/// Wrap a completion-based C event sanitizer. +pub fn wrap_async_event_sanitize_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> EventSanitizeFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |event: Event, fields: EventSanitizeFields| { + let user_data = user_data.clone(); + Box::pin(async move { + let value = invoke_async_json( + cb, + user_data, + serde_json::json!({"event": event, "fields": fields}), + ) + .await?; + serde_json::from_value(value) + .map_err(|error| FlowError::Internal(format!("invalid event fields: {error}"))) + }) + }) +} + +/// Wrap a completion-based C LLM conditional guardrail. +pub fn wrap_async_llm_conditional_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmConditionalFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |request: LlmRequest| { + let user_data = user_data.clone(); + Box::pin(async move { + match invoke_async_json(cb, user_data, serde_json::json!({"request": request})).await? { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => Err(FlowError::Internal(format!( + "async conditional callback returned {other}; expected string or null" + ))), + } + }) + }) +} + +/// Wrap a completion-based C LLM request sanitizer. +pub fn wrap_async_llm_sanitize_request_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmSanitizeRequestFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new( + move |request: LlmRequest, context: LlmSanitizeRequestContext| { + let user_data = user_data.clone(); + let codec = format!("{:?}", context.codec()); + Box::pin(async move { + let value = invoke_async_json( + cb, + user_data, + serde_json::json!({"request": request, "context": {"codec": codec}}), + ) + .await?; + if value.is_null() { + Ok(None) + } else { + serde_json::from_value(value) + .map(Some) + .map_err(|error| FlowError::Internal(error.to_string())) + } + }) + }, + ) +} + +/// Wrap a completion-based C LLM response sanitizer. +pub fn wrap_async_llm_sanitize_response_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmSanitizeResponseFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { + let user_data = user_data.clone(); + let codec = format!("{:?}", context.codec()); + Box::pin(async move { + let value = invoke_async_json( + cb, + user_data, + serde_json::json!({"response": response, "context": {"codec": codec}}), + ) + .await?; + Ok((!value.is_null()).then_some(value)) + }) + }) +} + +/// Wrap a completion-based C LLM request intercept. +pub fn wrap_async_llm_request_intercept_fn( + cb: NemoRelayAsyncJsonCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmRequestInterceptFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new( + move |name: String, request: LlmRequest, annotated: Option| { + let user_data = user_data.clone(); + Box::pin(async move { + let value = invoke_async_json( + cb, + user_data, + serde_json::json!({ + "name": name, + "request": request, + "annotated": annotated, + }), + ) + .await?; + serde_json::from_value(value).map_err(|error| { + FlowError::Internal(format!("invalid LLM request intercept outcome: {error}")) + }) + }) + }, + ) +} + +/// Wrap a completion-based C tool execution intercept. +pub fn wrap_async_tool_execution_intercept_fn( + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> ToolExecutionFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new(move |name: &str, args: Json, next: ToolExecutionNextFn| { + let user_data = user_data.clone(); + let invocation = serde_json::json!({"name": name, "value": args}); + Box::pin(async move { + let value = + invoke_async_intercept(cb, user_data, invocation, AsyncNextInner::Tool(next)) + .await?; + serde_json::from_value(value) + .map_err(|error| FlowError::Internal(format!("invalid tool outcome: {error}"))) + }) + }) +} + +/// Wrap a completion-based C LLM execution intercept. +pub fn wrap_async_llm_execution_intercept_fn( + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmExecutionFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new( + move |name: &str, request: LlmRequest, next: LlmExecutionNextFn| { + let user_data = user_data.clone(); + let invocation = serde_json::json!({"name": name, "request": request}); + Box::pin(invoke_async_intercept( + cb, + user_data, + invocation, + AsyncNextInner::Llm(next), + )) + }, + ) +} + +/// Wrap a completion-based C LLM stream execution intercept. +/// +/// The completion ABI resolves one JSON value, so a stream intercept must +/// resolve to an array of chunks. Relay replays that array as a stream after +/// completion; incremental chunk delivery is not available through this ABI. +pub fn wrap_async_llm_stream_execution_intercept_fn( + cb: NemoRelayAsyncInterceptCb, + user_data: *mut libc::c_void, + free_fn: NemoRelayFreeFn, +) -> LlmStreamExecutionFn { + let user_data = make_user_data(user_data, free_fn); + Arc::new( + move |name: &str, request: LlmRequest, next: LlmStreamExecutionNextFn| { + let user_data = user_data.clone(); + let invocation = serde_json::json!({"name": name, "request": request}); + Box::pin(async move { + let value = invoke_async_intercept( + cb, + user_data, + invocation, + AsyncNextInner::LlmStream(next), + ) + .await?; + let chunks = value.as_array().cloned().ok_or_else(|| { + FlowError::Internal("async stream intercept must resolve to an array".into()) + })?; + Ok(LlmJsonStream::new(tokio_stream::iter( + chunks.into_iter().map(Ok), + ))) + }) + }, + ) +} + // --------------------------------------------------------------------------- // Wrapper functions: C callback -> core trait objects // --------------------------------------------------------------------------- @@ -321,14 +962,17 @@ pub fn wrap_tool_sanitize_fn( free_fn: NemoRelayFreeFn, ) -> ToolSanitizeFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |name: &str, args: Json| { - let c_name = CString::new(name).unwrap_or_default(); - let c_args = json_to_c_string(&args); - let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; - unsafe { nemo_relay_string_free_internal(c_args) }; - let result = ptr_to_json(result_ptr); - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |name: String, args: Json| { + let ud = ud.clone(); + Box::pin(async move { + let c_name = CString::new(name).unwrap_or_default(); + let c_args = json_to_c_string(&args); + let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; + unsafe { nemo_relay_string_free_internal(c_args) }; + let result = ptr_to_json(result_ptr); + unsafe { nemo_relay_string_free_internal(result_ptr) }; + Ok(result) + }) }) } @@ -339,22 +983,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 +1012,19 @@ pub fn wrap_tool_request_intercept_fn( free_fn: NemoRelayFreeFn, ) -> ToolInterceptFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |name: &str, args: Json| { - clear_last_error(); - let c_name = CString::new(name).unwrap_or_default(); - let c_args = json_to_c_string(&args); - let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; - unsafe { nemo_relay_string_free_internal(c_args) }; - let result = - json_result_from_ptr(result_ptr, "tool request intercept callback returned null"); - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |name: String, args: Json| { + let ud = ud.clone(); + Box::pin(async move { + clear_last_error(); + let c_name = CString::new(name).unwrap_or_default(); + let c_args = json_to_c_string(&args); + let result_ptr = unsafe { cb(ud.ptr, c_name.as_ptr(), c_args) }; + unsafe { nemo_relay_string_free_internal(c_args) }; + let result = + json_result_from_ptr(result_ptr, "tool request intercept callback returned null"); + unsafe { nemo_relay_string_free_internal(result_ptr) }; + result + }) }) } @@ -617,64 +1267,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 +1341,85 @@ pub fn wrap_llm_sanitize_request_fn( let ud = make_user_data(user_data, free_fn); Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { + let ud = ud.clone(); + Box::pin(async move { + clear_last_error(); + let (codec_kind, codec_id) = match ffi_codec_identity(context.codec()) { + Ok(identity) => identity, + Err(error) => { + set_last_error(&error.to_string()); + return 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 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 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 = c_str_to_json(result_ptr); + unsafe { + nemo_relay_string_free_internal(response_json); + if result_ptr != response_json { + nemo_relay_string_free_internal(result_ptr); + } } - } - result + Ok(result) + }) }) } @@ -790,20 +1449,23 @@ pub fn wrap_llm_conditional_fn( free_fn: NemoRelayFreeFn, ) -> LlmConditionalFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |request: &LlmRequest| { - clear_last_error(); - let ffi_req = FfiLLMRequest(request.clone()); - let result_ptr = unsafe { cb(ud.ptr, &ffi_req) }; - let result = if result_ptr.is_null() { - match last_error_message() { - Some(message) => Err(FlowError::Internal(message)), - None => Ok(None), - } - } else { - Ok(ptr_to_opt_string(result_ptr)) - }; - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |request: LlmRequest| { + let ud = ud.clone(); + Box::pin(async move { + clear_last_error(); + let ffi_req = FfiLLMRequest(request); + let result_ptr = unsafe { cb(ud.ptr, &ffi_req) }; + let result = if result_ptr.is_null() { + match last_error_message() { + Some(message) => Err(FlowError::Internal(message)), + None => Ok(None), + } + } else { + Ok(ptr_to_opt_string(result_ptr)) + }; + unsafe { nemo_relay_string_free_internal(result_ptr) }; + result + }) }) } @@ -918,14 +1580,18 @@ pub fn wrap_event_sanitize_fn( free_fn: NemoRelayFreeFn, ) -> EventSanitizeFn { let ud = make_user_data(user_data, free_fn); - Arc::new(move |event: &Event, fields: EventSanitizeFields| { - let ffi_event = FfiEvent(event.clone()); - let fields_json = json_to_c_string(&serde_json::to_value(&fields).unwrap_or(Json::Null)); - let result_ptr = unsafe { cb(ud.ptr, &ffi_event, fields_json) }; - unsafe { nemo_relay_string_free_internal(fields_json) }; - let result = serde_json::from_value(ptr_to_json(result_ptr)).unwrap_or_default(); - unsafe { nemo_relay_string_free_internal(result_ptr) }; - result + Arc::new(move |event: Event, fields: EventSanitizeFields| { + let ud = ud.clone(); + Box::pin(async move { + let ffi_event = FfiEvent(event); + let fields_json = + json_to_c_string(&serde_json::to_value(&fields).unwrap_or(Json::Null)); + let result_ptr = unsafe { cb(ud.ptr, &ffi_event, fields_json) }; + unsafe { nemo_relay_string_free_internal(fields_json) }; + let result = serde_json::from_value(ptr_to_json(result_ptr)).unwrap_or_default(); + unsafe { nemo_relay_string_free_internal(result_ptr) }; + Ok(result) + }) }) } diff --git a/crates/ffi/tests/integration/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index 140f089ce..4f6e90e10 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -661,6 +661,168 @@ fn scope_stack_api_round_trip() { unsafe { nemo_relay_scope_stack_free(stack) }; } +#[test] +fn scope_stack_propagation_and_thread_binding_validate_all_ffi_inputs() { + let _guard = TEST_MUTEX.lock().unwrap(); + + assert_eq!( + unsafe { nemo_relay_capture_propagation_context_json(ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { + nemo_relay_capture_propagation_context_with_root_json(ptr::null(), ptr::null_mut()) + }, + NemoRelayStatus::NullPointer + ); + + let mut inherited_context = ptr::null_mut(); + assert_eq!( + unsafe { nemo_relay_capture_propagation_context_json(&mut inherited_context) }, + NemoRelayStatus::Ok + ); + let inherited_context = unsafe { take_string(inherited_context) }.unwrap(); + assert_eq!( + serde_json::from_str::(&inherited_context).unwrap()["version"], + json!(1) + ); + + let root_uuid = cstring("018f13f0-7c1a-7a80-8000-000000000001"); + let mut rooted_context = ptr::null_mut(); + assert_eq!( + unsafe { + nemo_relay_capture_propagation_context_with_root_json( + root_uuid.as_ptr(), + &mut rooted_context, + ) + }, + NemoRelayStatus::Ok + ); + let rooted_context = unsafe { take_string(rooted_context) }.unwrap(); + assert_eq!( + serde_json::from_str::(&rooted_context).unwrap()["root_uuid"], + json!("018f13f0-7c1a-7a80-8000-000000000001") + ); + + let invalid_root = cstring("not-a-uuid"); + let mut context_json = ptr::null_mut(); + assert_eq!( + unsafe { + nemo_relay_capture_propagation_context_with_root_json( + invalid_root.as_ptr(), + &mut context_json, + ) + }, + NemoRelayStatus::InvalidArg + ); + assert!(context_json.is_null()); + + assert_eq!( + unsafe { + nemo_relay_scope_stack_create_from_propagation_json(ptr::null(), ptr::null_mut()) + }, + NemoRelayStatus::NullPointer + ); + let invalid_context = cstring("not-json"); + let mut stack = ptr::null_mut(); + assert_eq!( + unsafe { + nemo_relay_scope_stack_create_from_propagation_json( + invalid_context.as_ptr(), + &mut stack, + ) + }, + NemoRelayStatus::InvalidJson + ); + assert!(stack.is_null()); + + let rooted_context = cstring(&rooted_context); + assert_eq!( + unsafe { + nemo_relay_scope_stack_create_from_propagation_json(rooted_context.as_ptr(), &mut stack) + }, + NemoRelayStatus::Ok + ); + assert!(!stack.is_null()); + unsafe { nemo_relay_scope_stack_free(stack) }; + + assert_eq!( + unsafe { nemo_relay_scope_stack_set_thread(ptr::null()) }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { nemo_relay_scope_stack_capture_thread(ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { nemo_relay_scope_stack_restore_thread(ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + + let stack = unsafe { fresh_scope_stack() }; + let mut binding = ptr::null_mut(); + assert_eq!( + unsafe { nemo_relay_scope_stack_capture_thread(&mut binding) }, + NemoRelayStatus::Ok + ); + assert!(!binding.is_null()); + assert_eq!( + unsafe { nemo_relay_scope_stack_restore_thread(binding) }, + NemoRelayStatus::Ok + ); + unsafe { nemo_relay_scope_stack_free(stack) }; +} + +#[test] +fn observability_component_helpers_serialize_defaults_and_validate_inputs() { + let _guard = TEST_MUTEX.lock().unwrap(); + + let kind = unsafe { take_string(api::nemo_relay_observability_plugin_kind()) }.unwrap(); + assert_eq!(kind, "observability"); + + assert_eq!( + unsafe { api::nemo_relay_observability_default_config_json(ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + let mut default_config = ptr::null_mut(); + assert_eq!( + unsafe { api::nemo_relay_observability_default_config_json(&mut default_config) }, + NemoRelayStatus::Ok + ); + assert!(unsafe { returned_json(default_config) }.is_object()); + + assert_eq!( + unsafe { + api::nemo_relay_observability_component_spec_json(ptr::null(), true, ptr::null_mut()) + }, + NemoRelayStatus::NullPointer + ); + let mut component = ptr::null_mut(); + assert_eq!( + unsafe { + api::nemo_relay_observability_component_spec_json(ptr::null(), true, &mut component) + }, + NemoRelayStatus::Ok + ); + let component = unsafe { returned_json(component) }; + assert_eq!(component["kind"], json!("observability")); + assert_eq!(component["enabled"], json!(true)); + + let invalid_config = cstring("not-json"); + let mut rejected = ptr::null_mut(); + assert_eq!( + unsafe { + api::nemo_relay_observability_component_spec_json( + invalid_config.as_ptr(), + false, + &mut rejected, + ) + }, + NemoRelayStatus::InvalidJson + ); + assert!(rejected.is_null()); +} + #[test] fn llm_request_accessors_round_trip() { let headers = cstring(r#"{"x-trace":"1"}"#); diff --git a/crates/ffi/tests/integration/callable_extra_tests.rs b/crates/ffi/tests/integration/callable_extra_tests.rs index b176f4988..69a341b96 100644 --- a/crates/ffi/tests/integration/callable_extra_tests.rs +++ b/crates/ffi/tests/integration/callable_extra_tests.rs @@ -4,10 +4,19 @@ //! Integration tests for callable extra in the NeMo Relay FFI crate. use super::*; +use std::future::Future; use std::ptr; use tokio_stream::StreamExt; +fn resolve(future: impl Future) -> T { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap() + .block_on(future) +} + unsafe extern "C" fn tool_conditional_error_cb( _user_data: *mut libc::c_void, _name: *const c_char, @@ -172,7 +181,7 @@ fn test_callable_extra_trampoline_and_helper_paths() { .unwrap(); let conditional = wrap_tool_conditional_fn(tool_conditional_error_cb, ptr::null_mut(), None); - let conditional_err = conditional("tool", &json!({})).unwrap_err(); + let conditional_err = resolve(conditional("tool".into(), json!({}))).unwrap_err(); assert!( conditional_err .to_string() @@ -243,7 +252,7 @@ fn test_callable_extra_request_intercept_and_codec_paths() { let intercept_error = wrap_llm_request_intercept_fn(llm_request_intercept_status_error_cb, ptr::null_mut(), None); - let err = intercept_error("llm", request.clone(), None).unwrap_err(); + let err = resolve(intercept_error("llm".into(), request.clone(), None)).unwrap_err(); assert!( err.to_string() .contains("request intercept callback failed") @@ -254,7 +263,7 @@ fn test_callable_extra_request_intercept_and_codec_paths() { ptr::null_mut(), None, ); - let err = intercept_null("llm", request.clone(), None).unwrap_err(); + let err = resolve(intercept_null("llm".into(), request.clone(), None)).unwrap_err(); assert!(err.to_string().contains("null out_outcome_json")); let intercept_invalid_annotated = wrap_llm_request_intercept_fn( @@ -262,22 +271,28 @@ fn test_callable_extra_request_intercept_and_codec_paths() { ptr::null_mut(), None, ); - let err = intercept_invalid_annotated("llm", request.clone(), None).unwrap_err(); + let err = resolve(intercept_invalid_annotated( + "llm".into(), + request.clone(), + None, + )) + .unwrap_err(); assert!( err.to_string() .contains("invalid LLM request intercept outcome JSON") ); let sanitize = wrap_llm_sanitize_request_fn(llm_request_passthrough_cb, ptr::null_mut(), None); - let sanitized = sanitize( + let sanitized = resolve(sanitize( request.clone(), nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ) + )) + .unwrap() .expect("non-null sanitizer result"); assert_eq!(sanitized.content, request.content); let conditional = wrap_llm_conditional_fn(llm_conditional_error_cb, ptr::null_mut(), None); - let conditional_err = conditional(&request).unwrap_err(); + let conditional_err = resolve(conditional(request.clone())).unwrap_err(); assert!( conditional_err .to_string() @@ -359,12 +374,16 @@ fn test_sanitizer_context_resolves_directional_ffi_codecs() { "preserve": true }), }; - let sanitized = - wrap_llm_sanitize_request_fn(llm_request_codec_round_trip_cb, ptr::null_mut(), None)( - request.clone(), - LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), - ) - .expect("codec round trip returns a request"); + let sanitized = resolve(wrap_llm_sanitize_request_fn( + llm_request_codec_round_trip_cb, + ptr::null_mut(), + None, + )( + request.clone(), + LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), + )) + .unwrap() + .expect("codec round trip returns a request"); assert_eq!(sanitized.content, request.content); let response = json!({ @@ -376,11 +395,15 @@ fn test_sanitizer_context_resolves_directional_ffi_codecs() { "finish_reason": "stop" }] }); - let sanitized = - wrap_llm_sanitize_response_fn(llm_response_codec_decode_cb, ptr::null_mut(), None)( - response.clone(), - LlmSanitizeResponseContext::for_response_codec(Some(codec)), - ) - .expect("codec decode returns a response"); + let sanitized = resolve(wrap_llm_sanitize_response_fn( + llm_response_codec_decode_cb, + ptr::null_mut(), + None, + )( + response.clone(), + LlmSanitizeResponseContext::for_response_codec(Some(codec)), + )) + .unwrap() + .expect("codec decode returns a response"); assert_eq!(sanitized, response); } diff --git a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs index 3c0381eb6..29676608e 100644 --- a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs +++ b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs @@ -13,6 +13,279 @@ struct EnvGuard { original: Option, } +unsafe extern "C" fn async_json_registration_callback( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _completion: *const callable::NemoRelayAsyncCompletion, +) -> callable::NemoRelayAsyncCallbackState { + callable::NemoRelayAsyncCallbackState::Pending +} + +unsafe extern "C" fn async_intercept_registration_callback( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _next: *const callable::NemoRelayAsyncNext, + _completion: *const callable::NemoRelayAsyncCompletion, +) -> callable::NemoRelayAsyncCallbackState { + callable::NemoRelayAsyncCallbackState::Pending +} + +#[test] +fn test_ffi_async_registration_entrypoints_cover_global_and_scope_surfaces() { + let _lock = TEST_MUTEX.lock().unwrap_or_else(|e| e.into_inner()); + reset_globals(); + + macro_rules! global_json { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + name.as_ptr(), + 0, + async_json_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(unsafe { $deregister(name.as_ptr()) }, NemoRelayStatus::Ok); + }}; + } + macro_rules! global_json_with_break_chain { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + name.as_ptr(), + 0, + false, + async_json_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(unsafe { $deregister(name.as_ptr()) }, NemoRelayStatus::Ok); + }}; + } + macro_rules! global_intercept { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + name.as_ptr(), + 0, + async_intercept_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(unsafe { $deregister(name.as_ptr()) }, NemoRelayStatus::Ok); + }}; + } + + global_json!( + nemo_relay_register_mark_sanitize_guardrail_async, + nemo_relay_deregister_mark_sanitize_guardrail + ); + global_json!( + nemo_relay_register_scope_sanitize_start_guardrail_async, + nemo_relay_deregister_scope_sanitize_start_guardrail + ); + global_json!( + nemo_relay_register_scope_sanitize_end_guardrail_async, + nemo_relay_deregister_scope_sanitize_end_guardrail + ); + global_json!( + nemo_relay_register_tool_sanitize_request_guardrail_async, + nemo_relay_deregister_tool_sanitize_request_guardrail + ); + global_json!( + nemo_relay_register_tool_sanitize_response_guardrail_async, + nemo_relay_deregister_tool_sanitize_response_guardrail + ); + global_json!( + nemo_relay_register_tool_conditional_execution_guardrail_async, + nemo_relay_deregister_tool_conditional_execution_guardrail + ); + global_json_with_break_chain!( + nemo_relay_register_tool_request_intercept_async, + nemo_relay_deregister_tool_request_intercept + ); + global_intercept!( + nemo_relay_register_tool_execution_intercept_async, + nemo_relay_deregister_tool_execution_intercept + ); + global_json!( + nemo_relay_register_llm_sanitize_request_guardrail_async, + nemo_relay_deregister_llm_sanitize_request_guardrail + ); + global_json!( + nemo_relay_register_llm_sanitize_response_guardrail_async, + nemo_relay_deregister_llm_sanitize_response_guardrail + ); + global_json!( + nemo_relay_register_llm_conditional_execution_guardrail_async, + nemo_relay_deregister_llm_conditional_execution_guardrail + ); + global_json_with_break_chain!( + nemo_relay_register_llm_request_intercept_async, + nemo_relay_deregister_llm_request_intercept + ); + global_intercept!( + nemo_relay_register_llm_execution_intercept_async, + nemo_relay_deregister_llm_execution_intercept + ); + global_intercept!( + nemo_relay_register_llm_stream_execution_intercept_async, + nemo_relay_deregister_llm_stream_execution_intercept + ); + + let _stack = unsafe { fresh_scope_stack() }; + let mut scope = ptr::null_mut(); + assert_eq!( + unsafe { nemo_relay_get_handle(&mut scope) }, + NemoRelayStatus::Ok + ); + let scope_uuid = cstring( + &unsafe { take_string(nemo_relay_scope_handle_uuid(scope)) } + .expect("root scope UUID should exist"), + ); + + macro_rules! scope_json { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + scope_uuid.as_ptr(), + name.as_ptr(), + 0, + async_json_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { $deregister(scope_uuid.as_ptr(), name.as_ptr()) }, + NemoRelayStatus::Ok + ); + }}; + } + macro_rules! scope_json_with_break_chain { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + scope_uuid.as_ptr(), + name.as_ptr(), + 0, + false, + async_json_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { $deregister(scope_uuid.as_ptr(), name.as_ptr()) }, + NemoRelayStatus::Ok + ); + }}; + } + macro_rules! scope_intercept { + ($register:ident, $deregister:ident) => {{ + let name = cstring(&unique_name(stringify!($register))); + assert_eq!( + unsafe { + $register( + scope_uuid.as_ptr(), + name.as_ptr(), + 0, + async_intercept_registration_callback, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { $deregister(scope_uuid.as_ptr(), name.as_ptr()) }, + NemoRelayStatus::Ok + ); + }}; + } + + scope_json!( + nemo_relay_scope_register_mark_sanitize_guardrail_async, + nemo_relay_scope_deregister_mark_sanitize_guardrail + ); + scope_json!( + nemo_relay_scope_register_scope_sanitize_start_guardrail_async, + nemo_relay_scope_deregister_scope_sanitize_start_guardrail + ); + scope_json!( + nemo_relay_scope_register_scope_sanitize_end_guardrail_async, + nemo_relay_scope_deregister_scope_sanitize_end_guardrail + ); + scope_json!( + nemo_relay_scope_register_tool_sanitize_request_guardrail_async, + nemo_relay_scope_deregister_tool_sanitize_request_guardrail + ); + scope_json!( + nemo_relay_scope_register_tool_sanitize_response_guardrail_async, + nemo_relay_scope_deregister_tool_sanitize_response_guardrail + ); + scope_json!( + nemo_relay_scope_register_tool_conditional_execution_guardrail_async, + nemo_relay_scope_deregister_tool_conditional_execution_guardrail + ); + scope_json_with_break_chain!( + nemo_relay_scope_register_tool_request_intercept_async, + nemo_relay_scope_deregister_tool_request_intercept + ); + scope_intercept!( + nemo_relay_scope_register_tool_execution_intercept_async, + nemo_relay_scope_deregister_tool_execution_intercept + ); + scope_json!( + nemo_relay_scope_register_llm_sanitize_request_guardrail_async, + nemo_relay_scope_deregister_llm_sanitize_request_guardrail + ); + scope_json!( + nemo_relay_scope_register_llm_sanitize_response_guardrail_async, + nemo_relay_scope_deregister_llm_sanitize_response_guardrail + ); + scope_json!( + nemo_relay_scope_register_llm_conditional_execution_guardrail_async, + nemo_relay_scope_deregister_llm_conditional_execution_guardrail + ); + scope_json_with_break_chain!( + nemo_relay_scope_register_llm_request_intercept_async, + nemo_relay_scope_deregister_llm_request_intercept + ); + scope_intercept!( + nemo_relay_scope_register_llm_execution_intercept_async, + nemo_relay_scope_deregister_llm_execution_intercept + ); + scope_intercept!( + nemo_relay_scope_register_llm_stream_execution_intercept_async, + nemo_relay_scope_deregister_llm_stream_execution_intercept + ); + unsafe { nemo_relay_scope_handle_free(scope) }; +} + impl EnvGuard { fn set(key: &'static str, value: &str) -> Self { let original = std::env::var_os(key); diff --git a/crates/ffi/tests/unit/api/registry_tests.rs b/crates/ffi/tests/unit/api/registry_tests.rs index efc28e2a0..f66c86403 100644 --- a/crates/ffi/tests/unit/api/registry_tests.rs +++ b/crates/ffi/tests/unit/api/registry_tests.rs @@ -169,6 +169,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(); @@ -269,8 +272,6 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!(*lock_unpoisoned(plugin_frees()), 7); - let invalid_uuid = cstring("not-a-uuid"); let invalid_name = cstring("invalid-scope-event-sanitizer"); assert_eq!( @@ -330,6 +331,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ); nemo_relay_scope_handle_free(owner); assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_eq!(*lock_unpoisoned(plugin_frees()), 7); let events = lock_unpoisoned(event_log()); let invalid_callback_event = events .iter() diff --git a/crates/ffi/tests/unit/callable_private_tests.rs b/crates/ffi/tests/unit/callable_private_tests.rs index a2684dd2f..f2aced5ee 100644 --- a/crates/ffi/tests/unit/callable_private_tests.rs +++ b/crates/ffi/tests/unit/callable_private_tests.rs @@ -5,6 +5,35 @@ use super::*; +unsafe extern "C" fn complete_without_settling( + _user_data: *mut libc::c_void, + _invocation_json: *const c_char, + _completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + NemoRelayAsyncCallbackState::Complete +} + +unsafe extern "C" fn send_next_result( + user_data: *mut libc::c_void, + value_json: *const c_char, + error_message: *const c_char, +) { + let sender = unsafe { + Box::from_raw( + user_data.cast::>>(), + ) + }; + let result = if error_message.is_null() { + serde_json::from_str(unsafe { CStr::from_ptr(value_json) }.to_str().unwrap()) + .map_err(|error| error.to_string()) + } else { + Err(unsafe { CStr::from_ptr(error_message) } + .to_string_lossy() + .into_owned()) + }; + let _ = sender.send(result); +} + #[test] fn test_callable_private_helper_paths() { clear_last_error(); @@ -17,3 +46,221 @@ fn test_callable_private_helper_paths() { assert_eq!(ptr_to_opt_string(raw), Some("ffi-string".into())); unsafe { nemo_relay_string_free_internal(raw) }; } + +#[test] +fn async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlements() { + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + let invalid_json = CString::new("not-json").unwrap(); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion_ref, invalid_json.as_ptr()) }, + NemoRelayStatus::InvalidJson + ); + let value = CString::new(r#"{"ok":true}"#).unwrap(); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion_ref, value.as_ptr()) }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion_ref, value.as_ptr()) }, + NemoRelayStatus::InvalidArg + ); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + assert_eq!( + runtime.block_on(receiver).unwrap().unwrap(), + serde_json::json!({"ok": true}) + ); + unsafe { nemo_relay_async_completion_release(completion_ref) }; + + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + assert_eq!( + unsafe { nemo_relay_async_completion_reject(completion_ref, std::ptr::null()) }, + NemoRelayStatus::Ok + ); + assert!( + runtime + .block_on(receiver) + .unwrap() + .unwrap_err() + .to_string() + .contains("async C callback rejected") + ); + unsafe { nemo_relay_async_completion_release(completion_ref) }; + + let (sender, _receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(true), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + assert!(unsafe { nemo_relay_async_completion_is_cancelled(completion_ref) }); + assert!(unsafe { nemo_relay_async_completion_is_cancelled(std::ptr::null()) }); + assert_eq!( + unsafe { nemo_relay_async_completion_reject(completion_ref, std::ptr::null()) }, + NemoRelayStatus::InvalidArg + ); + unsafe { nemo_relay_async_completion_release(completion_ref) }; +} + +#[test] +fn async_callback_wrappers_reject_complete_callbacks_without_settlement() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + + let error = runtime + .block_on(invoke_async_json( + complete_without_settling, + Arc::new(UserData { + ptr: std::ptr::null_mut(), + free_fn: None, + }), + serde_json::json!({}), + )) + .unwrap_err(); + assert!( + error + .to_string() + .contains("returned Complete without settling") + ); +} + +#[test] +fn async_next_invocation_supports_tool_llm_and_stream_continuations() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + + let cases: Vec<(AsyncNextInner, CString, serde_json::Value)> = vec![ + ( + AsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + CString::new(r#"{"tool":true}"#).unwrap(), + serde_json::json!({"tool": true}), + ), + ( + AsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + CString::new( + serde_json::to_string(&LlmRequest { + headers: serde_json::Map::new(), + content: serde_json::json!({"llm": true}), + }) + .unwrap(), + ) + .unwrap(), + serde_json::json!({"llm": true}), + ), + ( + AsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![ + Ok(serde_json::json!({"chunk": 1})), + Ok(serde_json::json!({"chunk": 2})), + ]))) + }) + })), + CString::new( + serde_json::to_string(&LlmRequest { + headers: serde_json::Map::new(), + content: serde_json::json!({"stream": true}), + }) + .unwrap(), + ) + .unwrap(), + serde_json::json!([{"chunk": 1}, {"chunk": 2}]), + ), + ]; + + for (inner, invocation, expected) in cases { + let next = Arc::new(NemoRelayAsyncNext { + inner, + runtime: runtime.handle().clone(), + }); + let next_ref = Arc::into_raw(next); + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NemoRelayAsyncCompletion { + sender: std::sync::Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + }); + let completion_ref = Arc::into_raw(Arc::clone(&completion)); + assert_eq!( + unsafe { nemo_relay_async_next_invoke(next_ref, invocation.as_ptr(), completion_ref) }, + NemoRelayStatus::Ok + ); + assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); + unsafe { + nemo_relay_async_next_release(next_ref); + nemo_relay_async_completion_release(completion_ref); + } + } +} + +#[test] +fn async_next_callback_reports_tool_llm_and_stream_results() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let cases: Vec<(AsyncNextInner, CString, Json)> = vec![ + ( + AsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + CString::new(r#"{"tool":true}"#).unwrap(), + serde_json::json!({"tool": true}), + ), + ( + AsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + CString::new(r#"{"headers":{},"content":{"llm":true}}"#).unwrap(), + serde_json::json!({"llm": true}), + ), + ( + AsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( + serde_json::json!({"stream": true}), + )]))) + }) + })), + CString::new(r#"{"headers":{},"content":{}}"#).unwrap(), + serde_json::json!([{ "stream": true }]), + ), + ]; + for (inner, invocation, expected) in cases { + let next = Arc::new(NemoRelayAsyncNext { + inner, + runtime: runtime.handle().clone(), + }); + let next_ref = Arc::into_raw(next); + let (sender, receiver) = + tokio::sync::oneshot::channel::>(); + assert_eq!( + unsafe { + nemo_relay_async_next_invoke_callback( + next_ref, + invocation.as_ptr(), + send_next_result, + Box::into_raw(Box::new(sender)).cast(), + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); + unsafe { nemo_relay_async_next_release(next_ref) }; + } +} diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index c58098b4a..2baadaae1 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -4,6 +4,7 @@ //! Unit tests for callable in the NeMo Relay FFI crate. use super::*; +use std::future::Future; use std::sync::atomic::{AtomicUsize, Ordering}; use nemo_relay::api::event::{Event, EventSanitizeFields}; @@ -16,12 +17,260 @@ extern "C" fn free_arc_counter(user_data: *mut libc::c_void) { counter.fetch_add(1, Ordering::SeqCst); } +unsafe extern "C" fn async_json_passthrough_callback( + user_data: *mut libc::c_void, + invocation_json: *const c_char, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + let kind = unsafe { *(user_data.cast::()) }; + let invocation: Json = serde_json::from_str( + unsafe { CStr::from_ptr(invocation_json) } + .to_str() + .expect("invocation must be UTF-8"), + ) + .expect("invocation must be JSON"); + let value = match kind { + 0 => invocation["value"].clone(), + 1 | 3 => Json::Null, + 2 | 4 => invocation["request"].clone(), + 5 => invocation["response"].clone(), + 6 => json!({ + "request": invocation["request"], + "annotated_request": invocation["annotated"], + "pending_marks": [], + "optimization_contributions": [], + }), + 7 => invocation["fields"].clone(), + _ => unreachable!("test callback kind must be known"), + }; + let value = CString::new(value.to_string()).expect("JSON has no NUL"); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, + NemoRelayStatus::Ok + ); + NemoRelayAsyncCallbackState::Complete +} + +fn async_callback_user_data(kind: usize) -> *mut libc::c_void { + Box::into_raw(Box::new(kind)).cast() +} + +unsafe extern "C" fn free_async_callback_user_data(user_data: *mut libc::c_void) { + unsafe { drop(Box::from_raw(user_data.cast::())) }; +} + +unsafe extern "C" fn async_next_callback( + _user_data: *mut libc::c_void, + invocation_json: *const c_char, + next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + let invocation: Json = serde_json::from_str( + unsafe { CStr::from_ptr(invocation_json) } + .to_str() + .expect("invocation must be UTF-8"), + ) + .expect("invocation must be JSON"); + let value = invocation + .get("value") + .or_else(|| invocation.get("request")) + .expect("intercept invocation must carry a value") + .to_string(); + let value = CString::new(value).expect("JSON has no NUL"); + assert_eq!( + unsafe { nemo_relay_async_next_invoke(next, value.as_ptr(), completion) }, + NemoRelayStatus::Ok + ); + unsafe { + nemo_relay_async_next_release(next); + nemo_relay_async_completion_release(completion); + } + NemoRelayAsyncCallbackState::Pending +} + +unsafe extern "C" fn async_tool_outcome_callback( + _user_data: *mut libc::c_void, + invocation_json: *const c_char, + _next: *const NemoRelayAsyncNext, + completion: *const NemoRelayAsyncCompletion, +) -> NemoRelayAsyncCallbackState { + let invocation: Json = serde_json::from_str( + unsafe { CStr::from_ptr(invocation_json) } + .to_str() + .expect("invocation must be UTF-8"), + ) + .expect("invocation must be JSON"); + let value = + CString::new(json!({"result": invocation["value"], "pending_marks": []}).to_string()) + .expect("JSON has no NUL"); + assert_eq!( + unsafe { nemo_relay_async_completion_resolve_json(completion, value.as_ptr()) }, + NemoRelayStatus::Ok + ); + NemoRelayAsyncCallbackState::Complete +} + +#[test] +fn async_callback_wrappers_cover_all_middleware_shapes() { + let tool_json = wrap_async_tool_json_fn( + async_json_passthrough_callback, + async_callback_user_data(0), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(tool_json("tool".into(), json!({"value": true}))).unwrap(), + json!({"value": true}) + ); + + let tool_conditional = wrap_async_tool_conditional_fn( + async_json_passthrough_callback, + async_callback_user_data(1), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(tool_conditional("tool".into(), json!({}))).unwrap(), + None + ); + + let llm_conditional = wrap_async_llm_conditional_fn( + async_json_passthrough_callback, + async_callback_user_data(3), + Some(free_async_callback_user_data), + ); + assert_eq!(resolve(llm_conditional(make_request())).unwrap(), None); + + let request_sanitizer = wrap_async_llm_sanitize_request_fn( + async_json_passthrough_callback, + async_callback_user_data(4), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(request_sanitizer( + make_request(), + nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), + )) + .unwrap(), + Some(make_request()) + ); + + let response_sanitizer = wrap_async_llm_sanitize_response_fn( + async_json_passthrough_callback, + async_callback_user_data(5), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(response_sanitizer( + json!({"response": true}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), + )) + .unwrap(), + Some(json!({"response": true})) + ); + + let request_intercept = wrap_async_llm_request_intercept_fn( + async_json_passthrough_callback, + async_callback_user_data(6), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(request_intercept("llm".into(), make_request(), None)) + .unwrap() + .request, + make_request() + ); + + let event = Event::Scope(nemo_relay::api::event::ScopeEvent::new( + nemo_relay::api::event::BaseEvent::builder() + .name("async-event") + .build(), + nemo_relay::api::event::ScopeCategory::Start, + Vec::new(), + nemo_relay::api::event::EventCategory::llm(), + None, + )); + let fields = EventSanitizeFields::builder() + .data(json!({"safe": true})) + .build(); + let event_sanitizer = wrap_async_event_sanitize_fn( + async_json_passthrough_callback, + async_callback_user_data(7), + Some(free_async_callback_user_data), + ); + assert_eq!( + resolve(event_sanitizer(event, fields.clone())).unwrap(), + fields + ); +} + +#[test] +fn async_execution_wrappers_continue_tool_and_llm_calls() { + let tool_intercept = wrap_async_tool_execution_intercept_fn( + async_tool_outcome_callback, + std::ptr::null_mut(), + None, + ); + let tool_next: ToolExecutionNextFn = Arc::new(|args| Box::pin(async move { Ok(args) })); + assert_eq!( + resolve(tool_intercept("tool", json!({"ok": true}), tool_next)) + .unwrap() + .result, + json!({"ok": true}) + ); + + let llm_intercept = + wrap_async_llm_execution_intercept_fn(async_next_callback, std::ptr::null_mut(), None); + let llm_next: LlmExecutionNextFn = + Arc::new(|request| Box::pin(async move { Ok(json!({"model": request.content["model"]})) })); + assert_eq!( + resolve(llm_intercept("llm", make_request(), llm_next)).unwrap(), + json!({"model": "test-model"}) + ); +} + +#[test] +fn async_stream_execution_wrapper_collects_the_continued_stream() { + let intercept = wrap_async_llm_stream_execution_intercept_fn( + async_next_callback, + std::ptr::null_mut(), + None, + ); + let next: LlmStreamExecutionNextFn = Arc::new(|request| { + Box::pin(async move { + Ok(nemo_relay::api::runtime::LlmJsonStream::new( + tokio_stream::iter(vec![ + Ok(json!({"model": request.content["model"], "chunk": 1})), + Ok(json!({"chunk": 2})), + ]), + )) + }) + }); + + let mut stream = resolve(intercept("llm", make_request(), next)).unwrap(); + assert_eq!( + resolve(async { stream.next().await.unwrap().unwrap() }), + json!({"model": "test-model", "chunk": 1}) + ); + assert_eq!( + resolve(async { stream.next().await.unwrap().unwrap() }), + json!({"chunk": 2}) + ); + assert!(resolve(async { stream.next().await }).is_none()); +} + fn user_data_counter() -> (*mut libc::c_void, Arc) { let counter = Arc::new(AtomicUsize::new(0)); let ptr = Box::into_raw(Box::new(counter.clone())) as *mut libc::c_void; (ptr, counter) } +fn resolve(future: impl Future) -> T { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap() + .block_on(future) +} + unsafe extern "C" fn tool_sanitize_cb( user_data: *mut libc::c_void, name: *const c_char, @@ -328,7 +577,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 +587,11 @@ fn test_wrap_tool_request_and_conditional_callbacks() { let wrapped_conditional = wrap_tool_conditional_fn(tool_conditional_cb, std::ptr::null_mut(), None); assert_eq!( - wrapped_conditional("tool", &json!({"block": true})).unwrap(), + resolve(wrapped_conditional("tool".into(), json!({"block": true}))).unwrap(), Some("blocked".into()) ); assert_eq!( - wrapped_conditional("tool", &json!({"block": false})).unwrap(), + resolve(wrapped_conditional("tool".into(), json!({"block": false}))).unwrap(), None ); } @@ -411,56 +660,60 @@ fn test_wrap_tool_exec_and_intercept_callbacks() { fn test_wrap_llm_request_response_and_conditional_callbacks() { let request_intercept = wrap_llm_request_intercept_fn(llm_request_intercept_cb, std::ptr::null_mut(), None); - let outcome = request_intercept("llm", make_request(), None).unwrap(); + let outcome = resolve(request_intercept("llm".into(), make_request(), None)).unwrap(); assert_eq!(outcome.request.content["intercepted"], json!(true)); let sanitize_request = wrap_llm_sanitize_request_fn(llm_request_null_cb, std::ptr::null_mut(), None); assert_eq!( - sanitize_request( + resolve(sanitize_request( make_request(), nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ), + )) + .unwrap(), None ); let alias_request = wrap_llm_sanitize_request_fn(llm_request_alias_cb, std::ptr::null_mut(), None); assert_eq!( - alias_request( + resolve(alias_request( make_request(), nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ), + )) + .unwrap(), Some(make_request()) ); let conditional = wrap_llm_conditional_fn(llm_conditional_cb, std::ptr::null_mut(), None); assert_eq!( - conditional(&LlmRequest { + resolve(conditional(LlmRequest { headers: serde_json::Map::new(), content: json!({"block": true}), - }) + })) .unwrap(), Some("blocked llm".into()) ); - assert_eq!(conditional(&make_request()).unwrap(), None); + assert_eq!(resolve(conditional(make_request())).unwrap(), None); let wrapped_response = wrap_llm_sanitize_response_fn(json_cb, std::ptr::null_mut(), None); assert_eq!( - wrapped_response( + resolve(wrapped_response( json!({"value": 2}), nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ) + )) + .unwrap() .unwrap()["wrapped"], json!(true) ); let alias_response = wrap_llm_sanitize_response_fn(json_alias_cb, std::ptr::null_mut(), None); assert_eq!( - alias_response( + resolve(alias_response( json!({"value": 2}), nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ), + )) + .unwrap(), Some(json!({"value": 2})) ); @@ -468,10 +721,11 @@ fn test_wrap_llm_request_response_and_conditional_callbacks() { let malformed_response = wrap_llm_sanitize_response_fn(callback, std::ptr::null_mut(), None); assert_eq!( - malformed_response( + resolve(malformed_response( json!({"secret": "must be omitted"}), nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ), + )) + .unwrap(), None ); } @@ -484,14 +738,17 @@ fn test_llm_sanitizers_fail_closed_for_runtime_codec_ids_with_embedded_nul() { let request_sanitizer = wrap_llm_sanitize_request_fn(llm_request_alias_cb, std::ptr::null_mut(), None); - assert_eq!( - request_sanitizer( - make_request(), - nemo_relay::api::runtime::LlmSanitizeRequestContext::with_identity( - runtime_identity.clone(), - ), + let request_error = resolve(request_sanitizer( + make_request(), + nemo_relay::api::runtime::LlmSanitizeRequestContext::with_identity( + runtime_identity.clone(), ), - None + )) + .expect_err("an embedded runtime codec ID must fail the callback wrapper"); + assert!( + request_error + .to_string() + .contains("runtime codec ID contains an embedded NUL") ); assert!( last_error_message() @@ -501,12 +758,15 @@ fn test_llm_sanitizers_fail_closed_for_runtime_codec_ids_with_embedded_nul() { let response_sanitizer = wrap_llm_sanitize_response_fn(json_alias_cb, std::ptr::null_mut(), None); - assert_eq!( - response_sanitizer( - json!({"secret": "must be omitted"}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), - ), - None + let response_error = resolve(response_sanitizer( + json!({"secret": "must be omitted"}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::with_identity(runtime_identity), + )) + .expect_err("an embedded runtime codec ID must fail the callback wrapper"); + assert!( + response_error + .to_string() + .contains("runtime codec ID contains an embedded NUL") ); assert!( last_error_message() @@ -542,7 +802,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 +916,7 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { .build(); let (user_data, sanitize_calls) = user_data_counter(); let sanitizer = wrap_event_sanitize_fn(event_sanitize_cb, user_data, Some(free_arc_counter)); - let sanitized = sanitizer(&event, original_fields.clone()); + let sanitized = resolve(sanitizer(event.clone(), original_fields.clone())).unwrap(); assert_eq!(sanitized.data, Some(json!({"safe": true}))); assert_eq!( sanitized @@ -667,12 +932,12 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { let invalid = wrap_event_sanitize_fn(invalid_event_sanitize_cb, std::ptr::null_mut(), None); assert_eq!( - invalid(&event, original_fields.clone()), + resolve(invalid(event.clone(), original_fields.clone())).unwrap(), EventSanitizeFields::default() ); let null = wrap_event_sanitize_fn(null_event_sanitize_cb, std::ptr::null_mut(), None); assert_eq!( - null(&event, original_fields.clone()), + resolve(null(event, original_fields.clone())).unwrap(), EventSanitizeFields::default() ); diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index c45b3326d..07b497b43 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() { @@ -781,7 +779,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)?; @@ -825,7 +825,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)?; @@ -865,9 +867,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)?; @@ -913,9 +915,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)?; @@ -959,9 +961,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)?; @@ -1001,9 +1003,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)?; @@ -1043,12 +1045,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)?; @@ -1172,12 +1175,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)?; @@ -1409,40 +1413,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 @@ -1467,89 +1437,9 @@ impl PersistentJsFunction { } } -fn core_event_fields( - fields: EventSanitizeFields, -) -> Option { - Some(nemo_relay::api::event::EventSanitizeFields { - data: fields.data, - category_profile: fields - .category_profile - .map(serde_json::from_value) - .transpose() - .ok()?, - metadata: fields.metadata, - }) -} - -fn js_event_fields(fields: &nemo_relay::api::event::EventSanitizeFields) -> EventSanitizeFields { - EventSanitizeFields { - data: fields.data.clone(), - category_profile: fields - .category_profile - .as_ref() - .and_then(|value| serde_json::to_value(value).ok()), - metadata: fields.metadata.clone(), - } -} - fn node_event_sanitize_fn(env: &Env, func: &JsFunction) -> napi::Result { - let callback = callable::safe_middleware_callback(env, func)?; - let direct = Arc::new(PersistentJsFunction::new(env, &callback)?); - let register_thread = std::thread::current().id(); - let mut tsfn = callback.create_threadsafe_function( - 0, - |ctx: napi::threadsafe_function::ThreadSafeCallContext<(Json, Json)>| { - Ok(vec![ctx.value.0, ctx.value.1]) - }, - )?; - tsfn.unref(env)?; - let background = callable::wrap_js_event_sanitize_fn(tsfn); - Ok(Arc::new(move |event, fields| { - if std::thread::current().id() == register_thread { - let event_json = match event.try_to_json_value() { - Ok(event_json) => event_json, - Err(error) => { - record_callback_error(format!( - "nemo_relay: failed to serialize JS event sanitizer context: {error}" - )); - return nemo_relay::api::event::EventSanitizeFields::default(); - } - }; - let sanitized = (|| -> FlowResult<_> { - let value = direct - .call_event_sanitize(event_json, js_event_fields(&fields)) - .map_err(|error| { - FlowError::Internal(format!( - "nemo_relay: JS event sanitizer callback failed: {error}" - )) - })?; - let value = callable::unwrap_middleware_result( - value, - "nemo_relay: JS event sanitizer callback failed", - )?; - let fields = event_sanitize_fields_from_json(value).map_err(|error| { - FlowError::Internal(format!( - "nemo_relay: JS event sanitizer callback failed: invalid JS event sanitizer result: {error}" - )) - })?; - core_event_fields(fields).ok_or_else(|| { - FlowError::Internal( - "nemo_relay: JS event sanitizer callback failed: invalid JS event sanitizer result" - .to_string(), - ) - }) - })(); - match sanitized { - Ok(sanitized) => sanitized, - Err(error) => { - record_callback_error(error.to_string()); - nemo_relay::api::event::EventSanitizeFields::default() - } - } - } else { - background(event, fields) - } - })) + let callback = Arc::new(crate::promise_call::PromiseAwareFn::new(env, func)?); + Ok(callable::wrap_js_event_sanitize_promise_fn(callback)) } type NodeLlmCodec = ( @@ -1882,15 +1772,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. @@ -2792,8 +2689,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( @@ -2801,7 +2700,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<()> { @@ -2847,8 +2746,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])* @@ -2903,13 +2807,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) } @@ -2939,8 +2850,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])* @@ -3028,15 +2949,16 @@ pub fn register_llm_sanitize_request_guardrail( name: String, priority: i32, #[napi( - ts_arg_type = "(request: Json, context: import('./plugin').LlmSanitizeRequestContext) => Json | null" + ts_arg_type = "(request: Json, context: import('./plugin').LlmSanitizeRequestContext) => Json | null | Promise" )] guardrail: JsFunction, ) -> Result<()> { - let callback = middleware_llm_sanitize_request_callback_tsfn(&env, &guardrail)?; core_registry_api::register_llm_sanitize_request_guardrail( &name, priority, - callable::wrap_js_llm_sanitize_request_fn(callback), + callable::wrap_js_llm_sanitize_request_promise_fn(Arc::new(PromiseAwareFn::new( + &env, &guardrail, + )?)), ) .map_err(to_napi_err) } @@ -3061,15 +2983,16 @@ pub fn register_llm_sanitize_response_guardrail( name: String, priority: i32, #[napi( - ts_arg_type = "(response: Json, context: import('./plugin').LlmSanitizeResponseContext) => Json | null" + ts_arg_type = "(response: Json, context: import('./plugin').LlmSanitizeResponseContext) => Json | null | Promise" )] guardrail: JsFunction, ) -> Result<()> { - let callback = middleware_llm_sanitize_response_callback_tsfn(&env, &guardrail)?; core_registry_api::register_llm_sanitize_response_guardrail( &name, priority, - callable::wrap_js_llm_sanitize_response_fn(callback), + callable::wrap_js_llm_sanitize_response_promise_fn(Arc::new(PromiseAwareFn::new( + &env, &guardrail, + )?)), ) .map_err(to_napi_err) } @@ -3092,13 +3015,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) } @@ -3128,16 +3056,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) } @@ -3262,16 +3194,17 @@ pub fn deregister_subscriber(name: String) -> Result { core_subscriber_api::deregister_subscriber(&name).map_err(to_napi_err) } -/// Wait for native subscriber callbacks queued before this call to finish. +/// Return a Promise that resolves when native subscriber callbacks queued +/// before this call finish. /// /// Call this function outside native subscriber callbacks. A re-entrant call returns without /// waiting to avoid blocking the dispatcher, so callbacks later in the same dispatch snapshot can /// still run. /// -/// JavaScript subscribers are queued through Node's `ThreadsafeFunction`; callers that -/// need JS callback side effects should await an event-loop tick after this returns. +/// JavaScript subscribers are queued through Node's `ThreadsafeFunction`. Awaiting this +/// Promise does not block the Node event loop while Promise-returning event sanitizers settle. #[napi] -pub fn flush_subscribers() -> Result<()> { +pub async fn flush_subscribers() -> Result<()> { core_subscriber_api::flush_subscribers().map_err(to_napi_err) } @@ -3283,8 +3216,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( @@ -3293,7 +3228,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<()> { @@ -3351,8 +3286,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])* @@ -3420,7 +3361,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) } @@ -3459,9 +3404,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])* @@ -3564,7 +3519,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<()> { @@ -3574,9 +3529,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) } @@ -3608,7 +3563,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<()> { @@ -3618,9 +3573,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) } @@ -3658,7 +3613,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) } @@ -3695,19 +3654,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) } @@ -3883,7 +3846,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 }, @@ -3923,6 +3890,7 @@ pub fn llm_request_intercepts(env: Env, name: String, request: Json) -> Result Result { async move { TASK_SCOPE_STACK .scope(scope_stack, async move { - core_llm_api::llm_conditional_execution(&llm_request).map_err(to_napi_err) + core_llm_api::llm_conditional_execution(&llm_request) + .await + .map_err(to_napi_err) }) .await }, diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index ba8abafc2..b40aba073 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -207,6 +207,318 @@ fn recv_middleware_option_string_result( } } +async fn await_middleware_json_result( + rx: tokio::sync::oneshot::Receiver, + error_prefix: &str, +) -> Result { + let value = rx + .await + .map_err(|error| FlowError::Internal(format!("{error_prefix}: {error}")))?; + unwrap_middleware_result(value, error_prefix) +} + +async fn await_middleware_json_or_value( + rx: tokio::sync::oneshot::Receiver, + error_prefix: &str, + fallback: Json, +) -> Json { + match await_middleware_json_result(rx, error_prefix).await { + Ok(value) => value, + Err(error) => { + record_callback_error(error.to_string()); + fallback + } + } +} + +async fn await_middleware_option_string_result( + rx: tokio::sync::oneshot::Receiver, + error_prefix: &str, +) -> Result> { + match await_middleware_json_result(rx, error_prefix).await? { + Json::Null => Ok(None), + Json::String(value) => Ok(Some(value)), + other => Err(FlowError::Internal(format!( + "{error_prefix}: expected string or null, got {other:?}", + ))), + } +} + +/// Wrap a Promise-aware JS `(name, args) => string | null` tool guardrail. +pub fn wrap_js_tool_conditional_promise_fn(func: Arc) -> ToolConditionalFn { + Arc::new(move |name: String, args: Json| { + let func = func.clone(); + Box::pin(async move { + let value = func + .call_spread(vec![Json::String(name), args]) + .await + .inspect_err(|error| record_callback_error(error.to_string()))?; + match value { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => { + let error = FlowError::Internal(format!( + "JS tool conditional callback failed: expected string or null, got {other:?}" + )); + record_callback_error(error.to_string()); + Err(error) + } + } + }) + }) +} + +/// Wrap a Promise-aware JS `(name, args) => Json` tool request intercept. +pub fn wrap_js_tool_request_intercept_promise_fn(func: Arc) -> ToolInterceptFn { + Arc::new(move |name: String, args: Json| { + let func = func.clone(); + Box::pin(async move { + func.call_spread(vec![Json::String(name), args]) + .await + .inspect_err(|error| record_callback_error(error.to_string())) + }) + }) +} + +/// Wrap a Promise-aware JS tool sanitizer. +pub fn wrap_js_tool_sanitize_promise_fn(func: Arc) -> ToolSanitizeFn { + Arc::new(move |name: String, value: Json| { + let func = func.clone(); + Box::pin(async move { + func.call_spread(vec![Json::String(name), value]) + .await + .inspect_err(|error| { + record_callback_error(error.to_string()); + }) + }) + }) +} + +/// Wrap a Promise-aware JS LLM request sanitizer. +pub fn wrap_js_llm_sanitize_request_promise_fn(func: Arc) -> LlmSanitizeRequestFn { + Arc::new( + move |request: LlmRequest, context: LlmSanitizeRequestContext| { + let func = func.clone(); + Box::pin(async move { + let request = serde_json::to_value(request).unwrap_or(Json::Null); + let context = js_llm_sanitize_request_context(&context); + let value = func + .call_spread_with_arg0(Box::new(move |env| { + let mut args = env.create_array_with_length(2)?; + let request = unsafe { + JsUnknown::from_raw_unchecked( + env.raw(), + Json::to_napi_value(env.raw(), request)?, + ) + }; + args.set_element(0, request)?; + args.set_element( + 1, + js_llm_sanitize_request_context_to_napi(env, context)?, + )?; + Ok(js_object_to_unknown(env, args)) + })) + .await + .inspect_err(|error| { + record_callback_error(error.to_string()); + })?; + if value.is_null() { + Ok(None) + } else { + serde_json::from_value(value) + .map(Some) + .map_err(|error| { + let error = FlowError::Internal(format!( + "JS LLM sanitize request callback failed: failed to deserialize LlmRequest: {error}" + )); + record_callback_error(error.to_string()); + error + }) + } + }) + }, + ) +} + +/// Wrap a Promise-aware JS LLM response sanitizer. +pub fn wrap_js_llm_sanitize_response_promise_fn( + func: Arc, +) -> LlmSanitizeResponseFn { + Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { + let func = func.clone(); + Box::pin(async move { + let context = js_llm_sanitize_response_context(&context); + let value = func + .call_spread_with_arg0(Box::new(move |env| { + let mut args = env.create_array_with_length(2)?; + let response = unsafe { + JsUnknown::from_raw_unchecked( + env.raw(), + Json::to_napi_value(env.raw(), response)?, + ) + }; + args.set_element(0, response)?; + args.set_element(1, js_llm_sanitize_response_context_to_napi(env, context)?)?; + Ok(js_object_to_unknown(env, args)) + })) + .await + .inspect_err(|error| { + record_callback_error(error.to_string()); + })?; + Ok((!value.is_null()).then_some(value)) + }) + }) +} + +/// Wrap a Promise-aware JS `(request) => string | null` LLM guardrail. +pub fn wrap_js_llm_conditional_promise_fn(func: Arc) -> LlmConditionalFn { + Arc::new(move |request: LlmRequest| { + let func = func.clone(); + Box::pin(async move { + let value = func + .call(serde_json::to_value(request).unwrap_or(Json::Null)) + .await + .inspect_err(|error| record_callback_error(error.to_string()))?; + match value { + Json::Null => Ok(None), + Json::String(reason) => Ok(Some(reason)), + other => { + let error = FlowError::Internal(format!( + "JS LLM conditional callback failed: expected string or null, got {other:?}" + )); + record_callback_error(error.to_string()); + Err(error) + } + } + }) + }) +} + +/// Wrap a Promise-aware JS LLM request intercept. +pub fn wrap_js_llm_request_intercept_promise_fn( + func: Arc, +) -> LlmRequestInterceptFn { + Arc::new( + move |name: String, request: LlmRequest, annotated: Option| { + let func = func.clone(); + Box::pin(async move { + let value = func + .call(serde_json::json!({ + "name": name, + "request": request, + "annotated": annotated, + })) + .await + .inspect_err(|error| { + record_callback_error(error.to_string()); + })?; + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct JsOutcome { + request: LlmRequest, + #[serde(default)] + annotated: Option, + #[serde(default)] + pending_marks: Vec, + #[serde(default)] + optimization_contributions: Vec, + } + let outcome: JsOutcome = serde_json::from_value(value).map_err(|error| { + let error = FlowError::Internal(format!( + "invalid JS LLM request intercept outcome: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + Ok(LlmRequestInterceptOutcome { + request: outcome.request, + annotated_request: outcome.annotated, + pending_marks: outcome.pending_marks.into_iter().map(Into::into).collect(), + optimization_contributions: outcome.optimization_contributions, + }) + }) + }, + ) +} + +/// Wrap a Promise-aware JS event sanitizer. +/// +/// Event sanitizers run on Relay's serial publication dispatcher, not on the +/// JavaScript registration thread. Waiting here therefore preserves synchronous +/// scope/mark APIs while allowing the JavaScript callback to settle a Promise +/// on the Node event loop. +pub fn wrap_js_event_sanitize_promise_fn(func: Arc) -> EventSanitizeFn { + Arc::new(move |event: Event, fields: CoreEventSanitizeFields| { + let func = func.clone(); + Box::pin(async move { + let event_json = JsEvent::try_from_event(&event) + .map(JsEvent::into_json) + .map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer context: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let js_fields = EventSanitizeFields { + data: fields.data, + category_profile: fields + .category_profile + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer category profile: {error}" + )); + record_callback_error(error.to_string()); + error + })?, + metadata: fields.metadata, + }; + let value = func + .call_spread(vec![ + event_json, + serde_json::to_value(js_fields).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer fields: {error}" + )); + record_callback_error(error.to_string()); + error + })?, + ]) + .await + .inspect_err(|error| { + // Scope and mark publication happens on the dispatcher + // thread. Preserve the event (the core fails open) while + // making the binding-visible failure available to Node. + record_callback_error(error.to_string()); + })?; + let fields = event_sanitize_fields_from_json(value).map_err(|error| { + let error = + FlowError::Internal(format!("invalid JS event sanitizer result: {error}")); + record_callback_error(error.to_string()); + error + })?; + let category_profile = fields + .category_profile + .map(serde_json::from_value) + .transpose() + .map_err(|error| { + let error = + FlowError::Internal(format!("invalid JS event sanitizer result: {error}")); + record_callback_error(error.to_string()); + error + })?; + Ok(CoreEventSanitizeFields { + data: fields.data, + category_profile, + metadata: fields.metadata, + }) + }) + }) +} + fn recv_json_or_null(rx: std::sync::mpsc::Receiver, error_prefix: &str) -> Json { rx.recv().unwrap_or_else(|e| { record_callback_error(format!("{error_prefix}: {e}")); @@ -249,28 +561,25 @@ pub fn wrap_js_tool_fn( func: ThreadsafeFunction<(String, Json), ErrorStrategy::Fatal>, ) -> ToolSanitizeFn { let func = Arc::new(func); - Arc::new(move |name: &str, args: Json| { + Arc::new(move |name: String, args: Json| { let func = func.clone(); - let name = name.to_string(); - let fallback = args.clone(); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - (name, args), - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - record_callback_error(format!( - "nemo_relay: failed to queue JS tool callback: {status:?}" - )); - return fallback; - } - // TODO: This closure returns Json (not Result), so we cannot propagate - // errors through the type system. Log the error so failures are not silent. - recv_middleware_json_or_value(rx, "nemo_relay: JS tool callback failed", fallback) + Box::pin(async move { + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + (name, args), + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS tool callback: {status:?}" + ))); + } + await_middleware_json_result(rx, "nemo_relay: JS tool callback failed").await + }) }) } @@ -279,25 +588,25 @@ pub fn wrap_js_tool_conditional_fn( func: ThreadsafeFunction<(String, Json), ErrorStrategy::Fatal>, ) -> ToolConditionalFn { let func = Arc::new(func); - Arc::new(move |name: &str, args: &Json| { + Arc::new(move |name: String, args: Json| { let func = func.clone(); - let name = name.to_string(); - let args = args.clone(); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - (name, args), - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - return Err(FlowError::Internal(format!( - "failed to queue JS tool conditional callback: {status:?}", - ))); - } - recv_middleware_option_string_result(rx, "JS tool conditional callback failed") + Box::pin(async move { + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + (name, args), + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS tool conditional callback: {status:?}", + ))); + } + await_middleware_option_string_result(rx, "JS tool conditional callback failed").await + }) }) } @@ -306,24 +615,25 @@ pub fn wrap_js_tool_request_intercept_fn( func: ThreadsafeFunction<(String, Json), ErrorStrategy::Fatal>, ) -> ToolInterceptFn { let func = Arc::new(func); - Arc::new(move |name: &str, args: Json| { + Arc::new(move |name: String, args: Json| { let func = func.clone(); - let name = name.to_string(); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - (name, args), - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - return Err(FlowError::Internal(format!( - "failed to queue JS tool callback: {status:?}", - ))); - } - recv_middleware_json_result(rx, "JS tool callback failed") + Box::pin(async move { + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + (name, args), + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS tool callback: {status:?}", + ))); + } + await_middleware_json_result(rx, "JS tool callback failed").await + }) }) } @@ -367,57 +677,57 @@ pub fn wrap_js_llm_request_intercept_fn( ) -> LlmRequestInterceptFn { let func = Arc::new(func); Arc::new( - move |name: &str, - request: LlmRequest, - annotated: Option| - -> Result { + move |name: String, request: LlmRequest, annotated: Option| { let func = func.clone(); - let req_json = serde_json::to_value(&request).unwrap_or(Json::Null); - let annotated_json = annotated - .as_ref() - .map(|a| serde_json::to_value(a).unwrap_or(Json::Null)) - .unwrap_or(Json::Null); - let arg = serde_json::json!({ - "name": name, - "request": req_json, - "annotated": annotated_json, - }); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - arg, - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - return Err(FlowError::Internal(format!( - "failed to queue JS LLM request intercept callback: {status:?}", - ))); - } - let result = - recv_middleware_json_result(rx, "JS LLM request intercept callback failed")?; + Box::pin(async move { + let req_json = serde_json::to_value(&request).unwrap_or(Json::Null); + let annotated_json = annotated + .as_ref() + .map(|a| serde_json::to_value(a).unwrap_or(Json::Null)) + .unwrap_or(Json::Null); + let arg = serde_json::json!({ + "name": name, + "request": req_json, + "annotated": annotated_json, + }); + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + arg, + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS LLM request intercept callback: {status:?}", + ))); + } + let result = + await_middleware_json_result(rx, "JS LLM request intercept callback failed") + .await?; - #[derive(Deserialize)] - #[serde(rename_all = "camelCase")] - struct JsOutcome { - request: LlmRequest, - #[serde(default)] - annotated: Option, - #[serde(default)] - pending_marks: Vec, - #[serde(default)] - optimization_contributions: Vec, - } - let outcome: JsOutcome = serde_json::from_value(result).map_err(|e| { - FlowError::Internal(format!("invalid JS LLM request intercept outcome: {e}")) - })?; - Ok(LlmRequestInterceptOutcome { - request: outcome.request, - annotated_request: outcome.annotated, - pending_marks: outcome.pending_marks.into_iter().map(Into::into).collect(), - optimization_contributions: outcome.optimization_contributions, + #[derive(Deserialize)] + #[serde(rename_all = "camelCase")] + struct JsOutcome { + request: LlmRequest, + #[serde(default)] + annotated: Option, + #[serde(default)] + pending_marks: Vec, + #[serde(default)] + optimization_contributions: Vec, + } + let outcome: JsOutcome = serde_json::from_value(result).map_err(|e| { + FlowError::Internal(format!("invalid JS LLM request intercept outcome: {e}")) + })?; + Ok(LlmRequestInterceptOutcome { + request: outcome.request, + annotated_request: outcome.annotated, + pending_marks: outcome.pending_marks.into_iter().map(Into::into).collect(), + optimization_contributions: outcome.optimization_contributions, + }) }) }, ) @@ -431,11 +741,66 @@ pub fn wrap_js_llm_sanitize_request_fn( let func = Arc::new(func); Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { - let context = js_llm_sanitize_request_context(&context); - let request = serde_json::to_value(request).unwrap_or(Json::Null); - let (tx, rx) = std::sync::mpsc::channel(); + let func = func.clone(); + Box::pin(async move { + let context = js_llm_sanitize_request_context(&context); + let request = serde_json::to_value(request).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS LLM sanitize request: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let (tx, rx) = tokio::sync::oneshot::channel(); + if func.call_with_return_value( + (request.clone(), context), + ThreadsafeFunctionCallMode::Blocking, + move |value: Option| { + let _ = tx.send(callback_json(value)); + Ok(()) + }, + ) != napi::Status::Ok + { + record_callback_error( + "nemo_relay: failed to queue JS LLM sanitize request callback", + ); + return Err(FlowError::Internal( + "failed to queue JS LLM sanitize request callback".into(), + )); + } + let value = await_middleware_json_result( + rx, + "nemo_relay: JS LLM request sanitizer callback failed", + ) + .await + .inspect_err(|error| record_callback_error(error.to_string()))?; + if value.is_null() { + return Ok(None); + } + serde_json::from_value(value) + .map(Some) + .map_err(|error| FlowError::Internal(format!( + "JS LLM sanitize request callback failed: failed to deserialize LlmRequest: {error}" + ))) + .inspect_err(|error| record_callback_error(error.to_string())) + }) + }, + ) +} + +/// Wrap a JS function for LLM response sanitization. The callback receives +/// `(response, context)`; returning `null` omits the event payload. +pub fn wrap_js_llm_sanitize_response_fn( + func: ThreadsafeFunction<(Json, JsLlmSanitizeResponseContext), ErrorStrategy::Fatal>, +) -> LlmSanitizeResponseFn { + let func = Arc::new(func); + Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { + let func = func.clone(); + Box::pin(async move { + let context = js_llm_sanitize_response_context(&context); + let (tx, rx) = tokio::sync::oneshot::channel(); if func.call_with_return_value( - (request.clone(), context), + (response, context), ThreadsafeFunctionCallMode::Blocking, move |value: Option| { let _ = tx.send(callback_json(value)); @@ -444,58 +809,20 @@ pub fn wrap_js_llm_sanitize_request_fn( ) != napi::Status::Ok { record_callback_error( - "nemo_relay: failed to queue JS LLM sanitize request callback", + "nemo_relay: failed to queue JS LLM sanitize response callback", ); - return None; + return Err(FlowError::Internal( + "failed to queue JS LLM sanitize response callback".into(), + )); } - let value = recv_middleware_json_or_value( + let value = await_middleware_json_result( rx, - "nemo_relay: JS LLM request sanitizer callback failed", - Json::Null, - ); - if value.is_null() { - return None; - } - serde_json::from_value(value).map_or_else( - |error| { - record_callback_error(format!( - "nemo_relay: JS LLM sanitize request callback failed: failed to deserialize LlmRequest: {error}" - )); - None - }, - Some, - ) - }, - ) -} - -/// Wrap a JS function for LLM response sanitization. The callback receives -/// `(response, context)`; returning `null` omits the event payload. -pub fn wrap_js_llm_sanitize_response_fn( - func: ThreadsafeFunction<(Json, JsLlmSanitizeResponseContext), ErrorStrategy::Fatal>, -) -> LlmSanitizeResponseFn { - let func = Arc::new(func); - Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { - let context = js_llm_sanitize_response_context(&context); - let (tx, rx) = std::sync::mpsc::channel(); - if func.call_with_return_value( - (response, context), - ThreadsafeFunctionCallMode::Blocking, - move |value: Option| { - let _ = tx.send(callback_json(value)); - Ok(()) - }, - ) != napi::Status::Ok - { - record_callback_error("nemo_relay: failed to queue JS LLM sanitize response callback"); - return None; - } - let value = recv_middleware_json_or_value( - rx, - "nemo_relay: JS LLM response sanitizer callback failed", - Json::Null, - ); - Some(value).and_then(|value| (!value.is_null()).then_some(value)) + "nemo_relay: JS LLM response sanitizer callback failed", + ) + .await + .inspect_err(|error| record_callback_error(error.to_string()))?; + Ok((!value.is_null()).then_some(value)) + }) }) } @@ -649,24 +976,26 @@ pub fn wrap_js_llm_conditional_fn( func: ThreadsafeFunction, ) -> LlmConditionalFn { let func = Arc::new(func); - Arc::new(move |request: &LlmRequest| { + Arc::new(move |request: LlmRequest| { let func = func.clone(); - let req_json = serde_json::to_value(request).unwrap_or(Json::Null); - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - req_json, - ThreadsafeFunctionCallMode::Blocking, - move |val: Option| { - let _ = tx.send(callback_json(val)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - return Err(FlowError::Internal(format!( - "failed to queue JS LLM conditional callback: {status:?}", - ))); - } - recv_middleware_option_string_result(rx, "JS LLM conditional callback failed") + Box::pin(async move { + let req_json = serde_json::to_value(request).unwrap_or(Json::Null); + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + req_json, + ThreadsafeFunctionCallMode::Blocking, + move |val: Option| { + let _ = tx.send(callback_json(val)); + Ok(()) + }, + ); + if status != napi::Status::Ok { + return Err(FlowError::Internal(format!( + "failed to queue JS LLM conditional callback: {status:?}", + ))); + } + await_middleware_option_string_result(rx, "JS LLM conditional callback failed").await + }) }) } @@ -783,72 +1112,93 @@ pub fn wrap_js_event_sanitize_fn( func: ThreadsafeFunction<(Json, Json), ErrorStrategy::Fatal>, ) -> EventSanitizeFn { let func = Arc::new(func); - Arc::new(move |event: &Event, fields: CoreEventSanitizeFields| { - let event_json = match JsEvent::try_from_event(event) { - Ok(event) => event.into_json(), - Err(error) => { + Arc::new(move |event: Event, fields: CoreEventSanitizeFields| { + let func = func.clone(); + Box::pin(async move { + let event_json = match JsEvent::try_from_event(&event) { + Ok(event) => event.into_json(), + Err(error) => { + record_callback_error(format!( + "nemo_relay: failed to serialize JS event sanitizer context: {error}" + )); + return Err(FlowError::Internal(error.to_string())); + } + }; + let js_fields = EventSanitizeFields { + data: fields.data.clone(), + category_profile: fields + .category_profile + .as_ref() + .map(serde_json::to_value) + .transpose() + .map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer category profile: {error}" + )); + record_callback_error(error.to_string()); + error + })?, + metadata: fields.metadata.clone(), + }; + let js_fields = serde_json::to_value(js_fields).map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JS event sanitizer fields: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let (tx, rx) = tokio::sync::oneshot::channel(); + let status = func.call_with_return_value( + (event_json, js_fields), + ThreadsafeFunctionCallMode::Blocking, + move |value: Option| { + let _ = tx.send(callback_json(value)); + Ok(()) + }, + ); + if status != napi::Status::Ok { record_callback_error(format!( - "nemo_relay: failed to serialize JS event sanitizer context: {error}" + "nemo_relay: failed to queue JS event sanitizer callback: {status:?}" )); - return CoreEventSanitizeFields::default(); + return Err(FlowError::Internal(format!( + "failed to queue JS event sanitizer callback: {status:?}" + ))); } - }; - let js_fields = EventSanitizeFields { - data: fields.data.clone(), - category_profile: fields - .category_profile - .as_ref() - .and_then(|value| serde_json::to_value(value).ok()), - metadata: fields.metadata.clone(), - }; - let (tx, rx) = std::sync::mpsc::channel(); - let status = func.call_with_return_value( - ( - event_json, - serde_json::to_value(js_fields).unwrap_or(Json::Null), - ), - ThreadsafeFunctionCallMode::Blocking, - move |value: Option| { - let _ = tx.send(callback_json(value)); - Ok(()) - }, - ); - if status != napi::Status::Ok { - record_callback_error(format!( - "nemo_relay: failed to queue JS event sanitizer callback: {status:?}" - )); - return CoreEventSanitizeFields::default(); - } - let sanitized = (|| -> Result<_> { - let result = - recv_middleware_json_result(rx, "nemo_relay: JS event sanitizer callback failed")?; - let result = event_sanitize_fields_from_json(result).map_err(|error| { - FlowError::Internal(format!( - "nemo_relay: invalid JS event sanitizer result: {error}" - )) - })?; - let category_profile = result - .category_profile - .map(serde_json::from_value) - .transpose() - .map_err(|error| { + let sanitized: Result = async { + let result = await_middleware_json_result( + rx, + "nemo_relay: JS event sanitizer callback failed", + ) + .await?; + let result = event_sanitize_fields_from_json(result).map_err(|error| { FlowError::Internal(format!( "nemo_relay: invalid JS event sanitizer result: {error}" )) })?; - Ok(CoreEventSanitizeFields { - data: result.data, - category_profile, - metadata: result.metadata, - }) - })(); - match sanitized { - Ok(sanitized) => sanitized, - Err(error) => { - record_callback_error(error.to_string()); - CoreEventSanitizeFields::default() + let category_profile = result + .category_profile + .map(serde_json::from_value) + .transpose() + .map_err(|error| { + FlowError::Internal(format!( + "nemo_relay: invalid JS event sanitizer result: {error}" + )) + })?; + Ok(CoreEventSanitizeFields { + data: result.data, + category_profile, + metadata: result.metadata, + }) } - } + .await; + match sanitized { + Ok(sanitized) => Ok(sanitized), + Err(error) => { + record_callback_error(error.to_string()); + Err(error) + } + } + }) }) } diff --git a/crates/node/src/callback_factory.rs b/crates/node/src/callback_factory.rs index 891373dbf..48b301883 100644 --- a/crates/node/src/callback_factory.rs +++ b/crates/node/src/callback_factory.rs @@ -66,14 +66,26 @@ const CALLBACK_FACTORIES_SOURCE: &str = r#"(() => { }, promise(fn) { - return function __nemo_relay_promise_wrapper(error, arg0, next, resolve, reject) { + return function __nemo_relay_promise_wrapper(error, arg0, spread, next, resolve, reject) { if (error != null) { reject(error); return; } Promise.resolve().then(() => ( - next === undefined ? fn(arg0) : fn(arg0, next) - )).then((value) => jsonValue(value === undefined ? null : value)).then(resolve, reject); + next === undefined + ? (spread ? fn(...arg0) : fn(arg0)) + : (spread ? fn(...arg0, next) : fn(arg0, next)) + )).then((value) => jsonValue(value === undefined ? null : value)).then(resolve, (error) => { + let message = 'unknown error'; + try { + if (typeof error === 'string') { + message = error; + } else if (error != null && typeof error.message === 'string') { + message = error.message; + } + } catch {} + reject(message); + }); }; }, }; diff --git a/crates/node/src/promise_call.rs b/crates/node/src/promise_call.rs index bb5435207..cdc15bada 100644 --- a/crates/node/src/promise_call.rs +++ b/crates/node/src/promise_call.rs @@ -53,6 +53,7 @@ enum PrimaryArg { struct CallArgs { arg0: PrimaryArg, + spread: bool, next: Option, completion: CallCompletion, } @@ -76,19 +77,6 @@ impl CallCompletion { } } -fn rejection_message( - string_result: napi::Result, - object_message_result: Option>, -) -> String { - if let Ok(value) = string_result { - value - } else if let Some(message_result) = object_message_result { - message_result.unwrap_or_else(|_| "unknown error".to_string()) - } else { - "unknown error".to_string() - } -} - fn closed_tsfn_error() -> FlowError { FlowError::Internal("PromiseAwareFn threadsafe function closed".into()) } @@ -168,12 +156,12 @@ fn build_completion_unknowns( })?; let reject = env.create_function_from_closure("__nemo_relay_reject", move |ctx| { - let message = rejection_message( - ctx.get::(0), - ctx.get::(0) - .ok() - .map(|value| value.get_named_property::("message")), - ); + // Do not invoke arbitrary `error.message` getters here. A throwing + // getter used to escape this callback and abort the N-API call rather + // than settling the middleware future as a rejection. + let message = ctx + .get::(0) + .unwrap_or_else(|_| "unknown error".to_string()); completion.send(Err(FlowError::Internal(message))); ctx.env.get_undefined() })?; @@ -208,7 +196,13 @@ impl PromiseAwareFn { PrimaryArg::Build(build) => build(&ctx.env)?, }; - let args = vec![arg0, next, resolve, reject]; + let spread = unsafe { + JsUnknown::from_raw_unchecked( + ctx.env.raw(), + ctx.env.get_boolean(ctx.value.spread)?.raw(), + ) + }; + let args = vec![arg0, spread, next, resolve, reject]; Ok(args) })?; @@ -222,7 +216,17 @@ impl PromiseAwareFn { /// Call the JS function with the given args and await the result. pub async fn call(&self, args: Json) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), None).await + self.call_inner(PrimaryArg::Json(args), false, None).await + } + + /// Call a JavaScript callback with several JSON arguments. + /// + /// This retains the normal callback shape for middleware such as tool + /// guardrails, whose public contract is `(name, payload)` rather than a + /// single envelope object. + pub async fn call_spread(&self, args: Vec) -> FlowResult { + self.call_inner(PrimaryArg::Json(Json::Array(args)), true, None) + .await } /// Call the JS function with a builder-constructed first argument and await @@ -232,13 +236,20 @@ impl PromiseAwareFn { /// cannot cross the threadsafe-function boundary as plain JSON, such as a /// `#[napi]` class instance. pub async fn call_with_arg0(&self, build_arg0: Arg0Builder) -> FlowResult { - self.call_inner(PrimaryArg::Build(build_arg0), None).await + self.call_inner(PrimaryArg::Build(build_arg0), false, None) + .await + } + + /// Call a JavaScript callback with builder-constructed spread arguments. + pub async fn call_spread_with_arg0(&self, build_arg0: Arg0Builder) -> FlowResult { + self.call_inner(PrimaryArg::Build(build_arg0), true, None) + .await } /// Call the JS function with a middleware-style `next(arg)` callback that /// resolves to a JSON result. pub async fn call_with_json_next(&self, args: Json, next: JsonNextFn) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), Some(NextFn::Json(next))) + self.call_inner(PrimaryArg::Json(args), false, Some(NextFn::Json(next))) .await } @@ -249,7 +260,7 @@ impl PromiseAwareFn { args: Json, next: JsonStreamNextFn, ) -> FlowResult { - self.call_inner(PrimaryArg::Json(args), Some(NextFn::Stream(next))) + self.call_inner(PrimaryArg::Json(args), false, Some(NextFn::Stream(next))) .await } @@ -260,7 +271,12 @@ impl PromiseAwareFn { } } - async fn call_inner(&self, arg0: PrimaryArg, next: Option) -> FlowResult { + async fn call_inner( + &self, + arg0: PrimaryArg, + spread: bool, + next: Option, + ) -> FlowResult { let (sender, receiver) = tokio::sync::oneshot::channel(); let tsfn = self .tsfn @@ -272,6 +288,7 @@ impl PromiseAwareFn { let status = tsfn.call( Ok(CallArgs { arg0, + spread, next, completion: CallCompletion::new(sender), }), diff --git a/crates/node/tests/callback_error_tests.mjs b/crates/node/tests/callback_error_tests.mjs index 83e684a5f..c05eccef9 100644 --- a/crates/node/tests/callback_error_tests.mjs +++ b/crates/node/tests/callback_error_tests.mjs @@ -67,11 +67,11 @@ describe('callback error helpers', () => { } }); - it('closed tool sanitize callbacks preserve the original payload and record the queue failure', () => { + it('closed tool sanitize callbacks preserve the original payload and record the queue failure', async () => { const args = { value: 1, }; - const result = __testClosedToolCallback( + const result = await __testClosedToolCallback( () => ({ ok: true, }), diff --git a/crates/node/tests/event_sanitizers_tests.mjs b/crates/node/tests/event_sanitizers_tests.mjs index 812a2a0ef..e06eee228 100644 --- a/crates/node/tests/event_sanitizers_tests.mjs +++ b/crates/node/tests/event_sanitizers_tests.mjs @@ -25,10 +25,10 @@ async function waitFor(events, count) { assert.ok(events.length >= count, `expected ${count} events, received ${events.length}`); } -function assertSanitizerFieldsCleared(event) { - assert.equal(event.data, null); - assert.equal(event.category_profile, null); - assert.equal(event.metadata, null); +function assertSanitizerFieldsPreserved(event, expectedData, expectedMetadata = expectedData) { + assert.deepEqual(event.data, expectedData); + assert.equal(event.category_profile?.subtype, 'seeded'); + assert.deepEqual(event.metadata, expectedMetadata); } async function initializeWithoutDiscoveredPluginConfig(config) { @@ -107,13 +107,35 @@ describe('event sanitizer registries', () => { assert.ok(lifecycle.every((event) => event.category_profile.subtype === 'sanitized')); }); - it('fails closed and records invalid direct sanitizer results', async () => { + it('awaits Promise-returning mark sanitizers without making event() asynchronous', async () => { + const events = capture('node-event-sanitize-promise-sub'); + let settled = false; + lib.registerMarkSanitizeGuardrail('node-event-promise', 0, async (_event, fields) => { + await new Promise((resolve) => setImmediate(resolve)); + settled = true; + return { ...fields, data: { sanitized: true } }; + }); + try { + const result = lib.event('promise-checkpoint', null, { raw: true }); + assert.equal(result, undefined); + assert.equal(settled, false); + await lib.flushSubscribers(); + await waitFor(events, 1); + } finally { + lib.deregisterMarkSanitizeGuardrail('node-event-promise'); + lib.deregisterSubscriber('node-event-sanitize-promise-sub'); + } + assert.equal(settled, true); + assert.deepEqual(events.at(-1).data, { sanitized: true }); + }); + + it('fails open and records invalid sanitizer results', async () => { const events = capture('node-event-sanitize-invalid-sub'); const invalidResults = { scalar: () => 'invalid', emptyObject: () => ({}), array: () => [], - promise: () => Promise.resolve({ data: { changed: true } }), + promise: () => Promise.resolve([]), }; try { for (const [kind, sanitizer] of Object.entries(invalidResults)) { @@ -135,8 +157,8 @@ describe('event sanitizer registries', () => { 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'); @@ -163,12 +185,12 @@ describe('event sanitizer registries', () => { assert.equal(start.metadata.background, true); }); - it('fails closed and records invalid thread-safe sanitizer results', async () => { + it('fails open and records invalid queued sanitizer results', async () => { const events = capture('node-event-sanitize-background-invalid-sub'); const invalidResults = { emptyObject: () => ({}), array: () => [], - promise: () => Promise.resolve({ data: { changed: true } }), + promise: () => Promise.resolve([]), }; try { for (const [kind, sanitizer] of Object.entries(invalidResults)) { @@ -193,7 +215,7 @@ describe('event sanitizer registries', () => { const start = events.find( (event) => event.kind === 'scope' && event.name === name && event.scope_category === 'start', ); - assertSanitizerFieldsCleared(start); + assertSanitizerFieldsPreserved(start, { kept: kind }); assert.match(lib.getLastCallbackError(), /invalid JS event sanitizer result/); } } finally { @@ -201,7 +223,7 @@ describe('event sanitizer registries', () => { } }); - it('fails closed when a thread-safe sanitizer throws', async () => { + it('fails open when a queued sanitizer throws', async () => { const events = capture('node-event-sanitize-background-throw-sub'); lib.clearLastCallbackError(); lib.registerScopeSanitizeStartGuardrail('node-background-throw-seed', -1, (_event, fields) => ({ @@ -220,7 +242,7 @@ describe('event sanitizer registries', () => { const start = events.find( (event) => event.kind === 'scope' && event.name === 'background-throw-tool' && event.scope_category === 'start', ); - assertSanitizerFieldsCleared(start); + assertSanitizerFieldsPreserved(start, { kept: true }); assert.match(lib.getLastCallbackError() ?? '', /background sanitizer boom/i); } finally { lib.deregisterScopeSanitizeStartGuardrail('node-background-throw-seed'); @@ -289,7 +311,7 @@ describe('event sanitizer registries', () => { assert.deepEqual(marks.cleared.data, { raw: true }); }); - it('fails closed when a plugin-owned sanitizer throws', async () => { + it('fails open when a plugin-owned sanitizer throws', async () => { const kind = `node.test.event-sanitize-throw.${Date.now()}`; const events = capture('node-event-sanitize-plugin-throw-sub'); plugin.register(kind, { @@ -314,7 +336,7 @@ describe('event sanitizer registries', () => { lib.event('plugin-throw', null, { raw: true }, { raw: true }); 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..bcab47e47 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -57,6 +57,17 @@ async function flushSubscriberCallbacks() { } } +async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { + flushSubscribers(); + const deadline = Date.now() + timeoutMs; + while (!predicate()) { + if (Date.now() >= deadline) { + throw new Error('timed out waiting for subscriber callbacks'); + } + await new Promise((resolve) => setImmediate(resolve)); + } +} + function makeNative() { return { headers: {}, @@ -291,7 +302,10 @@ describe('LLM execute', () => { /llm status failure/, ); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks( + () => events.some((e) => e.name === 'exec_status_ok_llm' && e.scope_category === 'end') + && events.some((e) => e.name === 'exec_status_error_llm' && e.scope_category === 'end'), + ); const okEnd = events.find( (e) => e.name === 'exec_status_ok_llm' && e.kind === 'scope' && e.category === 'llm' && e.scope_category === 'end', @@ -386,7 +400,11 @@ describe('LLM guardrails', () => { assert.deepEqual(result, { ok: true }); assert.equal(requestContextChecked, true); assert.equal(responseContextChecked, true); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks( + () => + events.some((event) => event.name === 'contextual_sanitize_llm' && event.scope_category === 'start') && + events.some((event) => event.name === 'contextual_sanitize_llm' && event.scope_category === 'end'), + ); const start = events.find( (event) => event.name === 'contextual_sanitize_llm' && event.scope_category === 'start', ); @@ -718,7 +736,7 @@ describe('LLM guardrails', () => { } }); - it('sanitize request guardrail failures omit the payload and remain usable', async () => { + it('sanitize request guardrail failures preserve the payload and remain usable', async () => { const events = []; clearLastCallbackError(); registerSubscriber('node_llm_san_req_throw_sub', (event) => events.push(event)); @@ -728,7 +746,16 @@ describe('LLM guardrails', () => { try { const request = makeNative(); await llmCallExecute('llm_san_req_throw', request, () => ({ ok: true }), null, null, null, null, null); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks( + () => + events.some( + (event) => + event.name === 'llm_san_req_throw' && + event.kind === 'scope' && + event.category === 'llm' && + event.scope_category === 'start', + ), + ); const start = events.find( (event) => event.name === 'llm_san_req_throw' && @@ -736,8 +763,8 @@ describe('LLM guardrails', () => { event.category === 'llm' && event.scope_category === 'start', ); - assert.equal(start.data, null); - assert.match(getLastCallbackError() ?? '', /JavaScript callback threw/i); + assert.deepEqual(start.data, { headers: request.headers, content: request.content }); + assert.match(getLastCallbackError() ?? '', /(unknown error|callback)/i); deregisterLlmSanitizeRequestGuardrail('node_llm_san_req_throw'); const result = await llmCallExecute( @@ -840,7 +867,7 @@ describe('LLM guardrails', () => { } }); - it('sanitize response guardrail failures omit the payload and remain usable', async () => { + it('sanitize response guardrail failures preserve the payload and remain usable', async () => { const events = []; clearLastCallbackError(); registerSubscriber('node_llm_san_resp_throw_sub', (event) => events.push(event)); @@ -850,7 +877,15 @@ describe('LLM guardrails', () => { try { const response = { ok: true }; await llmCallExecute('llm_san_resp_throw', makeNative(), () => response, null, null, null, null, null); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => + events.some( + (event) => + event.name === 'llm_san_resp_throw' && + event.kind === 'scope' && + event.category === 'llm' && + event.scope_category === 'end', + ), + ); const end = events.find( (event) => event.name === 'llm_san_resp_throw' && @@ -858,7 +893,7 @@ describe('LLM guardrails', () => { event.category === 'llm' && event.scope_category === 'end', ); - assert.equal(end.data, null); + assert.deepEqual(end.data, response); assert.match(getLastCallbackError() ?? '', /response sanitizer boom/i); deregisterLlmSanitizeResponseGuardrail('node_llm_san_resp_throw'); @@ -885,6 +920,19 @@ describe('LLM guardrails', () => { deregisterLlmConditionalExecutionGuardrail('node_llm_cond'); }); + it('conditional guardrail awaits a Promise result', async () => { + registerLlmConditionalExecutionGuardrail('node_llm_cond_promise', 10, async () => { + await new Promise((resolve) => setImmediate(resolve)); + return null; + }); + try { + const result = await llmCallExecute('llm_cond_promise', makeNative(), () => ({ ok: true }), null, null, null, null, null); + assert.deepEqual(result, { ok: true }); + } finally { + deregisterLlmConditionalExecutionGuardrail('node_llm_cond_promise'); + } + }); + it('conditional guardrail treats implicit undefined as allow', async () => { registerLlmConditionalExecutionGuardrail('node_llm_cond_undefined', 10, () => undefined); try { @@ -1021,6 +1069,28 @@ describe('LLM intercepts', () => { deregisterLlmRequestIntercept('node_llm_req_mod'); }); + it('request intercept awaits a Promise result', async () => { + registerLlmRequestIntercept('node_llm_req_promise', 10, false, async ({ request, annotated }) => { + await new Promise((resolve) => setImmediate(resolve)); + return { request: { ...request, content: { ...request.content, promised: true } }, annotated }; + }); + try { + const result = await llmCallExecute( + 'llm_req_promise', + makeNative(), + (request) => ({ promised: request.content.promised }), + null, + null, + null, + null, + null, + ); + assert.deepEqual(result, { promised: true }); + } finally { + deregisterLlmRequestIntercept('node_llm_req_promise'); + } + }); + it('request intercept throws a catchable error without terminating Node', async () => { registerLlmRequestIntercept('node_llm_req_throw', 10, false, () => { throw new Error('llm request intercept boom'); @@ -1340,11 +1410,21 @@ describe('LLM intercepts', () => { deregisterLlmRequestIntercept('node_llm_req_helper'); }); - it('generated request-intercept declarations preserve the open optimization kind', () => { + it('generated request-intercept declarations reference the canonical open optimization type', () => { const declarations = readFileSync(new URL('../index.d.ts', import.meta.url), 'utf8'); + const pluginDeclarations = readFileSync(new URL('../plugin.d.ts', import.meta.url), 'utf8'); const openKind = "kind: 'input_compression' | 'model_routing' | (string & {})"; - assert.equal(declarations.split(openKind).length - 1, 3); + assert.equal(declarations.split(openKind).length - 1, 1); + assert.equal(pluginDeclarations.split(openKind).length - 1, 1); + assert.match( + declarations, + /registerLlmRequestIntercept\([^\n]*import\('\.\/plugin'\)\.LlmRequestInterceptOutcome/, + ); + assert.match( + declarations, + /scopeRegisterLlmRequestIntercept\([^\n]*import\('\.\/plugin'\)\.LlmRequestInterceptOutcome/, + ); }); it('generated LLM sanitizer declarations expose directional codec contexts', () => { diff --git a/crates/node/tests/scope_tests.mjs b/crates/node/tests/scope_tests.mjs index a4cb41aed..f04593ccd 100644 --- a/crates/node/tests/scope_tests.mjs +++ b/crates/node/tests/scope_tests.mjs @@ -29,9 +29,13 @@ function rejectWithPrimitive(value) { return Promise.reject(value); } -async function flushSubscriberCallbacks() { +async function waitForSubscriberCallbacks(predicate, timeoutMs = 15000) { flushSubscribers(); - for (let i = 0; i < 10; i += 1) { + const deadline = Date.now() + timeoutMs; + while (!predicate()) { + if (Date.now() >= deadline) { + throw new Error('timed out waiting for subscriber callbacks'); + } await new Promise((resolve) => setImmediate(resolve)); } } @@ -104,7 +108,7 @@ describe('Scope operations', () => { try { const scope = pushScope('pop_metadata_scope', ScopeType.Agent, null, null, null, { a: 1, b: 2, c: 3 }); popScope(scope, null, null, { c: 3.5, d: 4 }); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => events.some((e) => e.name === 'pop_metadata_scope' && e.scope_category === 'end')); const end = events.find( (e) => e.name === 'pop_metadata_scope' && e.kind === 'scope' && e.scope_category === 'end', @@ -208,7 +212,7 @@ describe('withScope', () => { await withScope('with_scope_ok_status', ScopeType.Function, () => ({ ok: true }), null, null, null, { caller: 'node', }); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => events.some((e) => e.name === 'with_scope_ok_status' && e.scope_category === 'end')); const end = events.find( (e) => e.name === 'with_scope_ok_status' && e.kind === 'scope' && e.scope_category === 'end', @@ -260,7 +264,7 @@ describe('withScope', () => { }), /node status failure/, ); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => events.some((e) => e.name === 'with_scope_error_status' && e.scope_category === 'end')); const end = events.find( (e) => e.name === 'with_scope_error_status' && e.kind === 'scope' && e.scope_category === 'end', @@ -355,8 +359,7 @@ describe('Subscribers', () => { try { const scope = pushScope('sub_test', ScopeType.Agent, null, null); popScope(scope); - await flushSubscriberCallbacks(); - assert.ok(events.length > 0, 'Expected at least one event'); + await waitForSubscriberCallbacks(() => events.length > 0); } finally { deregisterSubscriber('node_event_collector'); } @@ -369,7 +372,7 @@ describe('Subscribers', () => { 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 +387,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 +410,7 @@ describe('Subscribers', () => { }, null, ); - await flushSubscriberCallbacks(); + await waitForSubscriberCallbacks(() => events.some((e) => e.kind === 'mark')); const found = events.some((e) => e.kind === 'mark'); assert.ok(found, 'Expected a Mark event'); } finally { diff --git a/crates/node/tests/tools_tests.mjs b/crates/node/tests/tools_tests.mjs index 125475561..b81581a69 100644 --- a/crates/node/tests/tools_tests.mjs +++ b/crates/node/tests/tools_tests.mjs @@ -608,6 +608,33 @@ describe('Tool guardrails', () => { deregisterToolConditionalExecutionGuardrail('node_tool_cond'); }); + it('conditional guardrail awaits a Promise result', async () => { + registerToolConditionalExecutionGuardrail('node_tool_cond_promise', 10, async () => { + await new Promise((resolve) => setImmediate(resolve)); + return null; + }); + try { + const result = await toolCallExecute('tool_cond_promise', { ok: true }, (args) => args, null, null, null, null); + assert.deepEqual(result, { ok: true }); + } finally { + deregisterToolConditionalExecutionGuardrail('node_tool_cond_promise'); + } + }); + + it('conditional guardrail propagates a rejected Promise', async () => { + registerToolConditionalExecutionGuardrail('node_tool_cond_reject', 10, async () => { + throw new Error('guardrail rejected promise'); + }); + try { + await assert.rejects( + () => toolCallExecute('tool_cond_reject', {}, () => ({ should_not: 'run' }), null, null, null, null), + /guardrail rejected promise/i, + ); + } finally { + deregisterToolConditionalExecutionGuardrail('node_tool_cond_reject'); + } + }); + it('conditional guardrail treats implicit undefined as allow', async () => { registerToolConditionalExecutionGuardrail('node_tool_cond_undefined', 10, () => undefined); try { @@ -782,6 +809,41 @@ describe('Tool intercepts', () => { deregisterToolRequestIntercept('node_tool_req_mod'); }); + it('request intercept awaits a Promise result', async () => { + registerToolRequestIntercept('node_tool_req_promise', 10, false, async (_name, args) => { + await new Promise((resolve) => setImmediate(resolve)); + return { ...args, promised: true }; + }); + try { + const result = await toolCallExecute( + 'tool_req_promise', + { original: true }, + (args) => args, + null, + null, + null, + null, + ); + assert.deepEqual(result, { original: true, promised: true }); + } finally { + deregisterToolRequestIntercept('node_tool_req_promise'); + } + }); + + it('request intercept propagates a rejected Promise', async () => { + registerToolRequestIntercept('node_tool_req_reject', 10, false, async () => { + throw new Error('request intercept rejected promise'); + }); + try { + await assert.rejects( + () => toolCallExecute('tool_req_reject', {}, () => ({ should_not: 'run' }), null, null, null, null), + /request intercept rejected promise/i, + ); + } finally { + deregisterToolRequestIntercept('node_tool_req_reject'); + } + }); + it('request intercept throws a catchable error without terminating Node', async () => { registerToolRequestIntercept('node_tool_req_throw', 10, false, () => { throw new Error('tool request intercept boom'); diff --git a/crates/pii-redaction/src/builtin.rs b/crates/pii-redaction/src/builtin.rs index 6e6d1dc5c..1f005cb73 100644 --- a/crates/pii-redaction/src/builtin.rs +++ b/crates/pii-redaction/src/builtin.rs @@ -454,12 +454,15 @@ impl CompiledBuiltinBackend { } pub(super) fn tool_sanitize_callback(backend: CompiledBuiltinBackend) -> ToolSanitizeFn { - Arc::new( - move |_name: &str, payload: Json| match backend.trajectory.as_ref() { - Some(trajectory) => trajectory.sanitize_tool_payload(payload), - None => backend.sanitize_json_preorder_dfs(payload), - }, - ) + Arc::new(move |_name: String, payload: Json| { + let backend = backend.clone(); + Box::pin(async move { + Ok(match backend.trajectory.as_ref() { + Some(trajectory) => trajectory.sanitize_tool_payload(payload), + None => backend.sanitize_json_preorder_dfs(payload), + }) + }) + }) } pub(super) fn event_sanitize_callback(backend: CompiledBuiltinBackend) -> EventSanitizeFn { @@ -479,40 +482,43 @@ fn event_sanitize_callback_with_scope_categories( scope_categories: Option<(bool, bool)>, ) -> EventSanitizeFn { Arc::new(move |event, mut fields| { - if scope_categories.is_some_and(|(sanitize_llm, sanitize_tool)| { - matches!(event, Event::Scope(_)) + let backend = backend.clone(); + Box::pin(async move { + if scope_categories.is_some_and(|(sanitize_llm, sanitize_tool)| { + matches!(event, Event::Scope(_)) + && event + .category() + .is_some_and(|category| match category.as_str() { + "llm" => !sanitize_llm, + "tool" => !sanitize_tool, + _ => false, + }) + }) { + return Ok(fields); + } + + if let Some(trajectory) = backend.trajectory.as_ref() { + return Ok(trajectory.sanitize_event_fields(&event, fields)); + } + let specialized_scope = matches!(event, Event::Scope(_)) && event .category() - .is_some_and(|category| match category.as_str() { - "llm" => !sanitize_llm, - "tool" => !sanitize_tool, - _ => false, - }) - }) { - return fields; - } - - if let Some(trajectory) = backend.trajectory.as_ref() { - return trajectory.sanitize_event_fields(event, fields); - } - let specialized_scope = matches!(event, Event::Scope(_)) - && event - .category() - .is_some_and(|category| matches!(category.as_str(), "tool" | "llm")); - - if !specialized_scope { - fields.data = fields - .data - .map(|data| backend.sanitize_json_preorder_dfs(data)); - fields.category_profile = fields.category_profile.and_then(|profile| { - sanitize_serializable_with_backend::(&backend, profile).ok() - }); - } + .is_some_and(|category| matches!(category.as_str(), "tool" | "llm")); + + if !specialized_scope { + fields.data = fields + .data + .map(|data| backend.sanitize_json_preorder_dfs(data)); + fields.category_profile = fields.category_profile.and_then(|profile| { + sanitize_serializable_with_backend::(&backend, profile).ok() + }); + } - fields.metadata = fields - .metadata - .map(|metadata| backend.sanitize_json_preorder_dfs(metadata)); - fields + fields.metadata = fields + .metadata + .map(|metadata| backend.sanitize_json_preorder_dfs(metadata)); + Ok(fields) + }) }) } @@ -520,41 +526,44 @@ pub(super) fn llm_sanitize_request_callback( backend: CompiledBuiltinBackend, ) -> LlmSanitizeRequestFn { Arc::new(move |mut request: LlmRequest, context| { - if let Some(trajectory) = backend.trajectory.as_ref() { - request.headers = trajectory - .sanitize_tool_payload(Json::Object(request.headers)) - .as_object() - .cloned() - .unwrap_or_default(); - request.content = trajectory.sanitize_provider_payload(request.content); - return Some(request); - } - request.headers = backend.sanitize_request_headers(request.headers); - if backend.target_paths.is_empty() { - request.content = backend.sanitize_json_preorder_dfs(request.content); - return Some(request); - } - let resolved = context.resolve_codec(); - let fallback = if resolved.is_none() { - backend - .selected_surface(context.codec()) - .map(build_request_codec) - } else { - None - }; - let Some(codec) = resolved.as_deref().or(fallback.as_deref()) else { - log_llm_payload_omitted("request", context.codec(), "no usable request codec"); - return None; - }; - let sanitized = backend.sanitize_request_with_codec(codec, &request); - if sanitized.is_none() { - log_llm_payload_omitted( - "request", - context.codec(), - "codec decode, sanitize, or encode failure", - ); - } - sanitized + let backend = backend.clone(); + Box::pin(async move { + if let Some(trajectory) = backend.trajectory.as_ref() { + request.headers = trajectory + .sanitize_tool_payload(Json::Object(request.headers)) + .as_object() + .cloned() + .unwrap_or_default(); + request.content = trajectory.sanitize_provider_payload(request.content); + return Ok(Some(request)); + } + request.headers = backend.sanitize_request_headers(request.headers); + if backend.target_paths.is_empty() { + request.content = backend.sanitize_json_preorder_dfs(request.content); + return Ok(Some(request)); + } + let resolved = context.resolve_codec(); + let fallback = if resolved.is_none() { + backend + .selected_surface(context.codec()) + .map(build_request_codec) + } else { + None + }; + let Some(codec) = resolved.as_deref().or(fallback.as_deref()) else { + log_llm_payload_omitted("request", context.codec(), "no usable request codec"); + return Ok(None); + }; + let sanitized = backend.sanitize_request_with_codec(codec, &request); + if sanitized.is_none() { + log_llm_payload_omitted( + "request", + context.codec(), + "codec decode, sanitize, or encode failure", + ); + } + Ok(sanitized) + }) }) } @@ -562,49 +571,52 @@ pub(super) fn llm_sanitize_response_callback( backend: CompiledBuiltinBackend, ) -> LlmSanitizeResponseFn { Arc::new(move |payload: Json, context| { - if let Some(trajectory) = backend.trajectory.as_ref() { - return Some(trajectory.sanitize_provider_payload(payload)); - } - if backend.target_paths.is_empty() { - return Some(backend.sanitize_json_preorder_dfs(payload)); - } - if matches!(context.codec(), LlmCodecIdentity::None) - && !backend.uses_compatible_legacy_response_codec(&payload) - { - log_llm_payload_omitted( - "response", - context.codec(), - "no active response codec or compatible legacy codec", - ); - return None; - } - let Some(surface) = backend.selected_surface(context.codec()) else { - log_llm_payload_omitted( - "response", - context.codec(), - "no recognized response codec surface", - ); - return None; - }; - let resolved = context.resolve_codec(); - let fallback = if resolved.is_none() { - Some(build_response_codec(surface)) - } else { - None - }; - let Some(codec) = resolved.as_deref().or(fallback.as_deref()) else { - log_llm_payload_omitted("response", context.codec(), "no usable response codec"); - return None; - }; - let sanitized = backend.sanitize_response_with_codec(codec, surface, payload); - if sanitized.is_none() { - log_llm_payload_omitted( - "response", - context.codec(), - "codec decode, sanitize, or encode failure", - ); - } - sanitized + let backend = backend.clone(); + Box::pin(async move { + if let Some(trajectory) = backend.trajectory.as_ref() { + return Ok(Some(trajectory.sanitize_provider_payload(payload))); + } + if backend.target_paths.is_empty() { + return Ok(Some(backend.sanitize_json_preorder_dfs(payload))); + } + if matches!(context.codec(), LlmCodecIdentity::None) + && !backend.uses_compatible_legacy_response_codec(&payload) + { + log_llm_payload_omitted( + "response", + context.codec(), + "no active response codec or compatible legacy codec", + ); + return Ok(None); + } + let Some(surface) = backend.selected_surface(context.codec()) else { + log_llm_payload_omitted( + "response", + context.codec(), + "no recognized response codec surface", + ); + return Ok(None); + }; + let resolved = context.resolve_codec(); + let fallback = if resolved.is_none() { + Some(build_response_codec(surface)) + } else { + None + }; + let Some(codec) = resolved.as_deref().or(fallback.as_deref()) else { + log_llm_payload_omitted("response", context.codec(), "no usable response codec"); + return Ok(None); + }; + let sanitized = backend.sanitize_response_with_codec(codec, surface, payload); + if sanitized.is_none() { + log_llm_payload_omitted( + "response", + context.codec(), + "codec decode, sanitize, or encode failure", + ); + } + Ok(sanitized) + }) }) } diff --git a/crates/pii-redaction/tests/unit/component_tests.rs b/crates/pii-redaction/tests/unit/component_tests.rs index cd7ad4026..7e375c821 100644 --- a/crates/pii-redaction/tests/unit/component_tests.rs +++ b/crates/pii-redaction/tests/unit/component_tests.rs @@ -296,8 +296,8 @@ impl LlmCodec for IdentifiedRequestCodec { } } -#[test] -fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs() { +#[tokio::test] +async fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".to_string(), @@ -324,7 +324,9 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( }, LlmSanitizeRequestContext::for_request_codec(Some(Arc::new(OpenAIResponsesCodec))), ) - .expect("the active OpenAI Responses codec must override the legacy fallback"); + .await + .expect("the active OpenAI Responses codec must override the legacy fallback") + .expect("the active OpenAI Responses codec must retain the payload"); assert_eq!( active_request.content["input"][0]["content"][0]["text"], json!("[REDACTED]") @@ -350,7 +352,9 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( inner: OpenAIResponsesCodec, }))), ) - .expect("an active runtime or opaque request codec must remain usable"); + .await + .expect("an active runtime or opaque request codec must remain usable") + .expect("an active runtime or opaque request codec must retain the payload"); assert_eq!( active_request.content["input"][0]["content"][0]["text"], json!("[REDACTED]") @@ -373,14 +377,19 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( BuiltinLlmCodec::OpenAiResponses, )), ) - .expect("the active OpenAI Responses codec must override the legacy fallback"); + .await + .expect("the active OpenAI Responses codec must override the legacy fallback") + .expect("the active OpenAI Responses codec must retain the payload"); assert_eq!( active_responses["output"][0]["content"][0]["text"], json!("[REDACTED]") ); assert!( - sanitize_response(responses_payload.clone(), no_codec_context()).is_none(), + sanitize_response(responses_payload.clone(), no_codec_context()) + .await + .expect("sanitizer callback must succeed") + .is_none(), "an incompatible configured fallback codec must omit a normalized payload" ); @@ -389,6 +398,8 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( responses_payload, LlmSanitizeResponseContext::with_identity(LlmCodecIdentity::Opaque), ) + .await + .expect("sanitizer callback must succeed") .is_none(), "a normalized-path policy must omit an unknown active provider payload" ); @@ -403,13 +414,15 @@ fn normalized_llm_paths_use_the_active_codec_and_fail_closed_for_unknown_codecs( "com.example.chat.v1".to_owned(), )), ) + .await + .expect("sanitizer callback must succeed") .is_none(), "a normalized-path policy must omit a runtime codec until it has a compatible projection" ); } -#[test] -fn normalized_llm_paths_omit_payloads_when_legacy_codec_decode_fails() { +#[tokio::test] +async fn normalized_llm_paths_omit_payloads_when_legacy_codec_decode_fails() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".to_string(), @@ -432,17 +445,22 @@ fn normalized_llm_paths_omit_payloads_when_legacy_codec_decode_fails() { }, no_codec_request_context(), ) + .await + .expect("sanitizer callback must succeed") .is_none(), "a shallow legacy surface match must not enable a raw-payload fallback" ); assert!( - sanitize_response(json!({"choices": "sk-response-secret"}), no_codec_context()).is_none(), + sanitize_response(json!({"choices": "sk-response-secret"}), no_codec_context()) + .await + .expect("sanitizer callback must succeed") + .is_none(), "a legacy response codec failure must omit the payload instead of emitting raw content" ); } -#[test] -fn normalized_openai_chat_api_specific_policy_omits_multiple_choices() { +#[tokio::test] +async fn normalized_openai_chat_api_specific_policy_omits_multiple_choices() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "remove".to_string(), @@ -475,12 +493,14 @@ fn normalized_openai_chat_api_specific_policy_omits_multiple_choices() { BuiltinLlmCodec::OpenAiChat, )), ) + .await + .expect("sanitizer callback must succeed") .is_none() ); } -#[test] -fn normalized_llm_paths_use_configured_anthropic_codec_without_a_system_message() { +#[tokio::test] +async fn normalized_llm_paths_use_configured_anthropic_codec_without_a_system_message() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".to_string(), @@ -504,7 +524,9 @@ fn normalized_llm_paths_use_configured_anthropic_codec_without_a_system_message( }, no_codec_request_context(), ) - .expect("the configured Anthropic codec must sanitize a valid message-only request"); + .await + .expect("the configured Anthropic codec must sanitize a valid message-only request") + .expect("the configured Anthropic codec must retain the payload"); assert_eq!( sanitized.content["messages"][0]["content"], @@ -512,8 +534,8 @@ fn normalized_llm_paths_use_configured_anthropic_codec_without_a_system_message( ); } -#[test] -fn trajectory_preset_redacts_chat_content_without_erasing_request_structure() { +#[tokio::test] +async fn trajectory_preset_redacts_chat_content_without_erasing_request_structure() { let callback = crate::builtin::llm_sanitize_request_callback(trajectory_backend( Some("openai_chat"), "preserve", @@ -548,6 +570,8 @@ fn trajectory_preset_redacts_chat_content_without_erasing_request_structure() { "person_name": "Alice Example" }), }, no_codec_request_context()) + .await + .unwrap() .unwrap(); assert_eq!(request.content["model"], "claude-sonnet-4-6"); @@ -595,8 +619,8 @@ fn trajectory_preset_redacts_chat_content_without_erasing_request_structure() { ); } -#[test] -fn trajectory_preset_preserves_response_analytics_and_redacts_response_content() { +#[tokio::test] +async fn trajectory_preset_preserves_response_analytics_and_redacts_response_content() { let callback = crate::builtin::llm_sanitize_response_callback(trajectory_backend( Some("openai_chat"), "preserve", @@ -617,6 +641,8 @@ fn trajectory_preset_preserves_response_analytics_and_redacts_response_content() }), no_codec_context(), ) + .await + .unwrap() .unwrap(); assert_eq!(sanitized["id"], "chatcmpl_1"); @@ -643,8 +669,8 @@ fn trajectory_preset_preserves_response_analytics_and_redacts_response_content() assert_eq!(sanitized["cost"]["total"], 1.25); } -#[test] -fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { +#[tokio::test] +async fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { let responses_request = crate::builtin::llm_sanitize_request_callback(trajectory_backend( Some("openai_responses"), "preserve", @@ -657,6 +683,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { "max_output_tokens": 100 }), }, no_codec_request_context()) + .await + .unwrap() .unwrap(); assert_eq!(responses_request.content["model"], "gpt-5"); assert_eq!(responses_request.content["input"][0]["role"], "user"); @@ -681,6 +709,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { }), no_codec_context(), ) + .await + .unwrap() .unwrap(); assert_eq!(responses_response["id"], "resp_1"); assert_eq!(responses_response["status"], "completed"); @@ -706,6 +736,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { "max_tokens": 128 }), }, no_codec_request_context()) + .await + .unwrap() .unwrap(); assert_eq!(anthropic_request.content["model"], "claude-sonnet-4-6"); assert_eq!(anthropic_request.content["system"], "[REDACTED]"); @@ -735,6 +767,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { "stop_reason": "end_turn", "usage": {"input_tokens": 12, "output_tokens": 6, "cache_read_input_tokens": 8} }), no_codec_context()) + .await + .unwrap() .unwrap(); assert_eq!(anthropic_response["id"], "msg_1"); assert_eq!(anthropic_response["role"], "assistant"); @@ -745,8 +779,8 @@ fn trajectory_preset_covers_responses_and_anthropic_provider_shapes() { assert_eq!(anthropic_response["usage"]["cache_read_input_tokens"], 8); } -#[test] -fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { +#[tokio::test] +async fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { let callback = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); let chunk = Event::Mark(MarkEvent::new( BaseEvent::builder().name("llm.chunk").build(), @@ -754,7 +788,7 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { Some(CategoryProfile::builder().subtype("llm.chunk").build()), )); let sanitized = callback( - &chunk, + chunk.clone(), EventSanitizeFields { data: Some(json!({ "chunk_index": 2, @@ -764,7 +798,9 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { category_profile: chunk.category_profile().cloned(), metadata: None, }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.data.as_ref().unwrap()["chunk_index"], 2); assert_eq!( sanitized.data.as_ref().unwrap()["event_type"], @@ -787,7 +823,7 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { ), )); let sanitized = callback( - &optimization, + optimization.clone(), EventSanitizeFields { data: Some(json!({ "producer": "neutral.router", @@ -803,7 +839,9 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { category_profile: optimization.category_profile().cloned(), metadata: None, }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.data.as_ref().unwrap()["producer"], "neutral.router" @@ -829,7 +867,7 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { None, )); let sanitized = callback( - &nested_agent, + nested_agent, EventSanitizeFields { data: Some(json!({ "request_id": "request-1", @@ -839,7 +877,9 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { category_profile: None, metadata: Some(json!({"parent_scope_id": "scope-1", "note": "private note"})), }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.data.as_ref().unwrap()["request_id"], "request-1"); assert_eq!( sanitized.data.as_ref().unwrap()["instruction"], @@ -860,8 +900,8 @@ fn trajectory_preset_redacts_known_marks_and_nested_scope_content() { assert_eq!(sanitized.metadata.as_ref().unwrap()["note"], "[REDACTED]"); } -#[test] -fn trajectory_preset_preserves_trusted_scope_metadata_only() { +#[tokio::test] +async fn trajectory_preset_preserves_trusted_scope_metadata_only() { let callback = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); let metadata = json!({ "nemo_relay_scope_role": "turn", @@ -922,13 +962,15 @@ fn trajectory_preset_preserves_trusted_scope_metadata_only() { None, )); let sanitized = callback( - &event, + event, EventSanitizeFields { data: None, category_profile: None, metadata: Some(metadata.clone()), }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.metadata, Some(expected_metadata.clone())); } @@ -940,7 +982,7 @@ fn trajectory_preset_preserves_trusted_scope_metadata_only() { None, )); let sanitized = callback( - &malformed, + malformed, EventSanitizeFields { data: None, category_profile: None, @@ -951,7 +993,9 @@ fn trajectory_preset_preserves_trusted_scope_metadata_only() { "provider_payload_exact": "private context" })), }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.metadata, Some(json!({ @@ -968,21 +1012,23 @@ fn trajectory_preset_preserves_trusted_scope_metadata_only() { Some(CategoryProfile::builder().subtype("llm.chunk").build()), )); let sanitized = callback( - &mark, + mark.clone(), EventSanitizeFields { data: None, category_profile: mark.category_profile().cloned(), metadata: Some(json!({"harness": "codex", "source": "hook"})), }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.metadata, Some(json!({"harness": "[REDACTED]", "source": "[REDACTED]"})) ); } -#[test] -fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { +#[tokio::test] +async fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { let event = Event::Mark(MarkEvent::new( BaseEvent::builder().name("neutral.plugin.evidence").build(), Some(EventCategory::custom()), @@ -999,11 +1045,14 @@ fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { }; let preserve = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); - assert_eq!(preserve(&event, fields.clone()), fields); + assert_eq!( + preserve(event.clone(), fields.clone()).await.unwrap(), + fields + ); let redact = crate::builtin::event_sanitize_callback(trajectory_backend(None, "redact_all_leaves")); - let sanitized = redact(&event, fields); + let sanitized = redact(event, fields).await.unwrap(); assert_eq!( sanitized.data.unwrap(), json!({ @@ -1016,8 +1065,8 @@ fn trajectory_custom_mark_policy_is_explicit_and_shape_preserving() { assert_eq!(profile.extra["opaque"]["label"], "[REDACTED]"); } -#[test] -fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations() { +#[tokio::test] +async fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations() { let callback = crate::builtin::event_sanitize_callback(trajectory_backend(None, "preserve")); let annotated_response: nemo_relay::codec::response::AnnotatedLlmResponse = serde_json::from_value(json!({ @@ -1068,7 +1117,7 @@ fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations None, )); let sanitized = callback( - &event, + event, EventSanitizeFields { data: Some(json!({"already": "sanitized by the response callback"})), category_profile: Some( @@ -1079,7 +1128,9 @@ fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations ), metadata: None, }, - ); + ) + .await + .unwrap(); let profile = sanitized.category_profile.unwrap(); assert_eq!(profile.model_name.as_deref(), Some("claude-sonnet-4-6")); @@ -1116,8 +1167,8 @@ fn trajectory_profile_preserves_typed_llm_accounting_while_redacting_annotations ); } -#[test] -fn preserved_custom_marks_remain_eligible_for_a_later_email_profile() { +#[tokio::test] +async fn preserved_custom_marks_remain_eligible_for_a_later_email_profile() { let event = Event::Mark(MarkEvent::new( BaseEvent::builder().name("neutral.plugin.evidence").build(), Some(EventCategory::custom()), @@ -1141,7 +1192,8 @@ fn preserved_custom_marks_remain_eligible_for_a_later_email_profile() { .unwrap(), ); - let sanitized = email(&event, trajectory(&event, fields)); + let fields = trajectory(event.clone(), fields).await.unwrap(); + let sanitized = email(event, fields).await.unwrap(); assert_eq!(sanitized.data.as_ref().unwrap()["owner"], "[REDACTED]"); assert_eq!(sanitized.data.as_ref().unwrap()["score"], 0.9); assert_eq!( @@ -1365,7 +1417,11 @@ fn local_profile_registrations_receive_generated_namespaces() { reset_runtime(); register_local_backend_provider(Arc::new(|_, ctx| { - ctx.register_mark_sanitize_guardrail("shared", 100, Arc::new(|_, fields| fields)) + ctx.register_mark_sanitize_guardrail( + "shared", + 100, + Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), + ) })) .unwrap(); @@ -1430,8 +1486,8 @@ fn failed_later_profile_rolls_back_earlier_profile_registrations() { deregister_subscriber("pii-profile-rollback").unwrap(); } -#[test] -fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() { +#[tokio::test] +async fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".into(), @@ -1449,7 +1505,7 @@ fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() None, )); let sanitized = callback( - &event, + event, EventSanitizeFields { data: Some(json!({"email": "person@example.com"})), category_profile: Some( @@ -1459,7 +1515,9 @@ fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() ), metadata: Some(json!({"owner": "person@example.com"})), }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.data.unwrap()["email"], "[REDACTED]"); assert_eq!( sanitized.category_profile.unwrap().subtype.as_deref(), @@ -1468,8 +1526,8 @@ fn event_sanitizer_transforms_data_category_profile_and_metadata_independently() assert_eq!(sanitized.metadata.unwrap()["owner"], "[REDACTED]"); } -#[test] -fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() { +#[tokio::test] +async fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "redact".into(), @@ -1493,13 +1551,15 @@ fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() .subtype("person@example.com") .build(); let sanitized = callback( - &event, + event, EventSanitizeFields { data: Some(json!({"content": "person@example.com"})), category_profile: Some(original_profile.clone()), metadata: Some(json!({"owner": "person@example.com"})), }, - ); + ) + .await + .unwrap(); assert_eq!( sanitized.data.unwrap()["content"], @@ -1511,8 +1571,8 @@ fn llm_and_tool_scope_metadata_is_sanitized_without_reprocessing_typed_fields() } } -#[test] -fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { +#[tokio::test] +async fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "redact".into(), @@ -1547,13 +1607,15 @@ fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { .subtype("person@example.com") .build(); let sanitized = callback( - &event, + event, EventSanitizeFields { data: Some(json!({"content": "person@example.com"})), category_profile: Some(original_profile.clone()), metadata: Some(json!({"owner": "person@example.com"})), }, - ); + ) + .await + .unwrap(); assert_eq!(sanitized.data.unwrap()["content"], "person@example.com"); assert_eq!(sanitized.category_profile.unwrap(), original_profile); @@ -1562,8 +1624,8 @@ fn scope_event_sanitizer_respects_enabled_llm_and_tool_surfaces() { } } -#[test] -fn event_sanitizer_discards_category_profile_when_sanitization_fails() { +#[tokio::test] +async fn event_sanitizer_discards_category_profile_when_sanitization_fails() { let backend = crate::builtin::CompiledBuiltinBackend::new( BuiltinBackendConfig { action: "regex_replace".into(), @@ -1581,7 +1643,7 @@ fn event_sanitizer_discards_category_profile_when_sanitization_fails() { None, )); let sanitized = callback( - &event, + event, EventSanitizeFields { data: None, category_profile: Some(CategoryProfile { @@ -1593,7 +1655,9 @@ fn event_sanitizer_discards_category_profile_when_sanitization_fails() { }), metadata: None, }, - ); + ) + .await + .unwrap(); assert!(sanitized.category_profile.is_none()); } diff --git a/crates/plugin/README.md b/crates/plugin/README.md index 7f3dd6c42..0b68a888e 100644 --- a/crates/plugin/README.md +++ b/crates/plugin/README.md @@ -42,8 +42,12 @@ the dynamic-library boundary on the stable C-compatible ABI. - **`PluginContext`**: Component-scoped registration APIs for middleware and subscribers. - **`PluginRuntime`**: Typed helpers for Relay-owned scopes and marks. -- **Stable native ABI v1**: C-compatible host and plugin tables behind the - safe Rust authoring interface. +- **Stable native ABI v3**: C-compatible host and plugin tables behind the + safe Rust authoring interface, with a v2-compatible prefix for existing + plugins. +- **Raw async middleware**: Completion-based raw registrations for plugins + that need asynchronous guardrails, intercepts, or event sanitizers. Typed + Rust callbacks remain synchronous convenience APIs. ## Installation diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index 727e09aa2..5f5500072 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -36,7 +36,15 @@ use serde::{Serialize, de::DeserializeOwned}; use serde_json::Map; /// Native plugin ABI version supported by this crate. -pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 2; +/// +/// Version 3 reserves the native async middleware extension. Hosts retain a +/// version-2 table for already-built plugins during entry-point negotiation. +pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 3; +/// ABI version that introduced completion-based asynchronous middleware. +pub const NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE: u32 = 3; + +/// Legacy native plugin ABI accepted by Relay hosts for compatibility. +pub const NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY: u32 = 2; /// Built-in LLM codec identities available to native plugins. #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, serde::Deserialize)] @@ -754,6 +762,133 @@ pub struct NemoRelayNativeHostApiV1 { ) -> NemoRelayStatus, } +/// Middleware surface selected by the native async registration hook. +/// +/// The host only exposes this through the ABI-v3 extension table. It keeps +/// every asynchronous callback shape uniform while allowing the host to +/// deserialize the surface-specific invocation and result payloads. +#[repr(u32)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum NemoRelayNativeAsyncMiddlewareKind { + /// Tool start-event request sanitizer. + ToolSanitizeRequest = 0, + /// Tool end-event response sanitizer. + ToolSanitizeResponse = 1, + /// Tool execution admission guardrail. + ToolConditionalExecution = 2, + /// Tool request rewrite intercept. + ToolRequestIntercept = 3, + /// Tool execution intercept with a continuation. + ToolExecutionIntercept = 4, + /// LLM start-event request sanitizer. + LlmSanitizeRequest = 5, + /// LLM end-event response sanitizer. + LlmSanitizeResponse = 6, + /// LLM execution admission guardrail. + LlmConditionalExecution = 7, + /// LLM request rewrite intercept. + LlmRequestIntercept = 8, + /// LLM execution intercept with a continuation. + LlmExecutionIntercept = 9, + /// Streaming LLM execution intercept with a continuation. + LlmStreamExecutionIntercept = 10, + /// Mark event sanitizer. + MarkSanitize = 11, + /// Scope-start event sanitizer. + ScopeSanitizeStart = 12, + /// Scope-end event sanitizer. + ScopeSanitizeEnd = 13, +} + +/// Indicates whether an asynchronous native callback settled before returning. +#[repr(u32)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum NemoRelayNativeAsyncCallbackState { + /// The callback settled its completion before returning. + Complete = 0, + /// The callback retained its completion for later settlement. + Pending = 1, +} + +/// Opaque one-shot completion retained by a pending native callback. +#[repr(C)] +pub struct NemoRelayNativeAsyncCompletion { + _private: [u8; 0], + _marker: PhantomData<(*mut u8, PhantomPinned)>, +} + +/// Opaque native execution continuation supplied only to execution intercepts. +#[repr(C)] +pub struct NemoRelayNativeAsyncNext { + _private: [u8; 0], + _marker: PhantomData<(*mut u8, PhantomPinned)>, +} + +/// Completion-based native middleware callback. +/// +/// `invocation_json` is borrowed for the call. A callback that returns +/// [`NemoRelayNativeAsyncCallbackState::Pending`] owns one completion +/// reference and must settle it then call the v3 `async_completion_release` +/// hook. When `next` is non-null, the callback owns that handle for the +/// invocation and must call `async_next_release` after its final use. `next` +/// is null for non-execution middleware. +pub type NemoRelayNativeAsyncMiddlewareCb = + unsafe extern "C" fn( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, + ) -> NemoRelayNativeAsyncCallbackState; + +/// ABI-v3 host extension appended to [`NemoRelayNativeHostApiV1`]. +/// +/// Its first field is the complete v1/v2 table, so legacy plugins can keep +/// treating the pointer as a [`NemoRelayNativeHostApiV1`]. +#[repr(C)] +#[derive(Clone, Copy)] +pub struct NemoRelayNativeHostApiV3 { + /// Compatibility prefix for ABI-v1/v2 plugins. + pub v1: NemoRelayNativeHostApiV1, + /// Resolves an async callback completion with a JSON value. + pub async_completion_resolve_json: unsafe extern "C" fn( + completion: *const NemoRelayNativeAsyncCompletion, + value_json: *const NemoRelayNativeString, + ) -> NemoRelayStatus, + /// Rejects an async callback completion with a UTF-8 message. + pub async_completion_reject: unsafe extern "C" fn( + completion: *const NemoRelayNativeAsyncCompletion, + message: *const NemoRelayNativeString, + ) -> NemoRelayStatus, + /// Returns true after the awaiting runtime has cancelled the invocation. + pub async_completion_is_cancelled: + unsafe extern "C" fn(completion: *const NemoRelayNativeAsyncCompletion) -> bool, + /// Releases the callback-owned reference after a pending completion settles. + pub async_completion_release: + unsafe extern "C" fn(completion: *const NemoRelayNativeAsyncCompletion), + /// Invokes an execution continuation and settles a supplied completion. + pub async_next_invoke: unsafe extern "C" fn( + next: *const NemoRelayNativeAsyncNext, + invocation_json: *const NemoRelayNativeString, + completion: *const NemoRelayNativeAsyncCompletion, + ) -> NemoRelayStatus, + /// Releases the callback-owned continuation reference for a pending callback. + pub async_next_release: unsafe extern "C" fn(next: *const NemoRelayNativeAsyncNext), + /// Registers any completion-based asynchronous middleware surface. + pub plugin_context_register_async_middleware: unsafe extern "C" fn( + ctx: *mut NemoRelayNativePluginContext, + kind: NemoRelayNativeAsyncMiddlewareKind, + name: *const NemoRelayNativeString, + priority: i32, + break_chain: bool, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, + ) -> NemoRelayStatus, +} + +unsafe impl Send for NemoRelayNativeHostApiV3 {} +unsafe impl Sync for NemoRelayNativeHostApiV3 {} + // The host API table is immutable after construction. Function pointers and // the null-terminated version string pointer are safe to share across threads. unsafe impl Send for NemoRelayNativeHostApiV1 {} @@ -2231,6 +2366,47 @@ impl<'a> PluginContext<'a> { }) } + /// Registers completion-based asynchronous middleware through the ABI-v3 + /// extension table. + /// + /// Plugins built against older hosts receive [`NemoRelayStatus::InvalidArg`] + /// instead of attempting to read beyond the legacy host table. + /// + /// # Safety + /// `cb`, `user_data`, and `free_fn` must remain valid until the host + /// deregisters the callback or invokes `free_fn`. A callback returning + /// `Pending` must settle and release its completion/next references. + #[allow(clippy::too_many_arguments)] // Mirrors the native C ABI registration callback. + pub unsafe fn register_async_middleware_raw( + &mut self, + kind: NemoRelayNativeAsyncMiddlewareKind, + name: &str, + priority: i32, + break_chain: bool, + cb: NemoRelayNativeAsyncMiddlewareCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, + ) -> NemoRelayStatus { + if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE + || self.host.struct_size < std::mem::size_of::() + { + return NemoRelayStatus::InvalidArg; + } + let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV3) }; + self.with_name(name, |_, name| unsafe { + (host.plugin_context_register_async_middleware)( + self.raw, + kind, + name, + priority, + break_chain, + cb, + user_data, + free_fn, + ) + }) + } + fn with_name( &self, name: &str, diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index fb3562bd3..865f05e16 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -300,8 +300,8 @@ static LLM_REQUEST_INTERCEPT_REGISTRATION: Mutex(), test_host().struct_size diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index f5c2ccd51..80b9d672e 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1313,12 +1313,42 @@ fn deregister_llm_stream_execution_intercept(name: &str) -> PyResult { #[pyfunction] fn tool_request_intercepts<'py>( py: Python<'py>, - name: &str, + name: String, args: &Bound<'py, PyAny>, -) -> PyResult> { +) -> PyResult> { let args_json = py_to_json(args)?; - let result = core_tool_api::tool_request_intercepts(name, args_json).map_err(to_py_err)?; - json_to_py(py, &result) + // Preserve the established synchronous helper behavior when no Python + // event loop is active. Awaitable middleware is supported from async + // callers below; a synchronous caller can continue using direct + // callbacks without manufacturing an asyncio loop. + if py + .import("asyncio")? + .call_method0("get_running_loop") + .is_err() + { + let scope_stack = current_scope_stack_handle(); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; + let result = runtime + .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { + core_tool_api::tool_request_intercepts(&name, args_json).await + })) + .map_err(to_py_err)?; + return json_to_py(py, &result).map(|value| value.into_bound(py)); + } + let scope_stack = current_scope_stack_handle(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + TASK_SCOPE_STACK + .scope(scope_stack, async move { + let result = core_tool_api::tool_request_intercepts(&name, args_json) + .await + .map_err(to_py_err)?; + Python::attach(|py| json_to_py(py, &result)) + }) + .await + }) } /// Run the registered tool conditional execution guardrail chain. @@ -1329,9 +1359,39 @@ fn tool_request_intercepts<'py>( /// name: Tool name. /// args: Tool arguments (any JSON-serializable object). #[pyfunction] -fn tool_conditional_execution(name: &str, args: &Bound<'_, PyAny>) -> PyResult<()> { +fn tool_conditional_execution<'py>( + py: Python<'py>, + name: String, + args: &Bound<'py, PyAny>, +) -> PyResult> { let args_json = py_to_json(args)?; - core_tool_api::tool_conditional_execution(name, &args_json).map_err(to_py_err) + if py + .import("asyncio")? + .call_method0("get_running_loop") + .is_err() + { + let scope_stack = current_scope_stack_handle(); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; + runtime + .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { + core_tool_api::tool_conditional_execution(&name, &args_json).await + })) + .map_err(to_py_err)?; + return Ok(py.None().into_bound(py)); + } + let scope_stack = current_scope_stack_handle(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + TASK_SCOPE_STACK + .scope(scope_stack, async move { + core_tool_api::tool_conditional_execution(&name, &args_json) + .await + .map_err(to_py_err) + }) + .await + }) } /// Run the registered LLM request intercept chain on the given request. @@ -1344,12 +1404,43 @@ fn tool_conditional_execution(name: &str, args: &Bound<'_, PyAny>) -> PyResult<( /// Returns: /// The (possibly transformed) ``LlmRequest``. #[pyfunction] -fn llm_request_intercepts( - name: &str, +fn llm_request_intercepts<'py>( + py: Python<'py>, + name: String, request: PyLLMRequest, -) -> PyResult { - let result = core_llm_api::llm_request_intercepts(name, request.inner).map_err(to_py_err)?; - Ok(crate::py_types::PyLLMRequestInterceptOutcome { inner: result }) +) -> PyResult> { + if py + .import("asyncio")? + .call_method0("get_running_loop") + .is_err() + { + let scope_stack = current_scope_stack_handle(); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; + let result = runtime + .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { + core_llm_api::llm_request_intercepts(&name, request.inner).await + })) + .map_err(to_py_err)?; + return Py::new( + py, + crate::py_types::PyLLMRequestInterceptOutcome { inner: result }, + ) + .map(|value| value.into_bound(py).into_any()); + } + let scope_stack = current_scope_stack_handle(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + TASK_SCOPE_STACK + .scope(scope_stack, async move { + let result = core_llm_api::llm_request_intercepts(&name, request.inner) + .await + .map_err(to_py_err)?; + Ok(crate::py_types::PyLLMRequestInterceptOutcome { inner: result }) + }) + .await + }) } /// Run the registered LLM conditional execution guardrail chain. @@ -1359,8 +1450,37 @@ fn llm_request_intercepts( /// Args: /// request: An ``LlmRequest`` object. #[pyfunction] -fn llm_conditional_execution(request: PyLLMRequest) -> PyResult<()> { - core_llm_api::llm_conditional_execution(&request.inner).map_err(to_py_err) +fn llm_conditional_execution<'py>( + py: Python<'py>, + request: PyLLMRequest, +) -> PyResult> { + if py + .import("asyncio")? + .call_method0("get_running_loop") + .is_err() + { + let scope_stack = current_scope_stack_handle(); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .map_err(|error| to_py_err(FlowError::Internal(error.to_string())))?; + runtime + .block_on(TASK_SCOPE_STACK.scope(scope_stack, async move { + core_llm_api::llm_conditional_execution(&request.inner).await + })) + .map_err(to_py_err)?; + return Ok(py.None().into_bound(py)); + } + let scope_stack = current_scope_stack_handle(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + TASK_SCOPE_STACK + .scope(scope_stack, async move { + core_llm_api::llm_conditional_execution(&request.inner) + .await + .map_err(to_py_err) + }) + .await + }) } // --------------------------------------------------------------------------- diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 57fffd5d5..f5cb02e55 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -39,7 +39,7 @@ use tokio_stream::Stream; use tokio_stream::wrappers::ReceiverStream; use nemo_relay::api::event::{Event, EventSanitizeFields}; -use nemo_relay::api::llm::{LlmRequest, LlmRequestInterceptOutcome}; +use nemo_relay::api::llm::LlmRequest; use nemo_relay::api::tool::ToolExecutionInterceptOutcome; use nemo_relay::codec::request::AnnotatedLlmRequest as AnnotatedLLMRequest; use nemo_relay::codec::response::AnnotatedLlmResponse as AnnotatedLLMResponse; @@ -408,71 +408,69 @@ fn stream_from_async_iter(async_iter: Py) -> FlowResult { /// Wrap a Python callable `(str, Json) -> Json` for tool sanitize/intercept fns. pub fn wrap_py_tool_fn(py_fn: Py) -> ToolSanitizeFn { - Arc::new(move |name: &str, args: Json| { - Python::attach(|py| { - let py_args = match json_to_py(py, &args) { - Ok(v) => v, - Err(e) => { - eprintln!("nemo_relay: json_to_py failed in tool fn for '{name}': {e}"); - return args.clone(); - } - }; - let result = match py_fn.call1(py, (name, py_args)) { - Ok(v) => v, - Err(e) => { - eprintln!("nemo_relay: Python tool callable failed for '{name}': {e}"); - return args.clone(); - } - }; - py_to_json(result.bind(py)).unwrap_or_else(|e| { - eprintln!("nemo_relay: py_to_json failed in tool fn for '{name}': {e}"); - args.clone() - }) + let py_fn = Arc::new(py_fn); + Arc::new(move |name: String, args: Json| { + let py_fn = py_fn.clone(); + Box::pin(async move { + resolve_json_or_future(Python::attach(|py| { + let py_args = json_to_py(py, &args) + .map_err(|e| FlowError::Internal(format!("tool json_to_py failed: {e}")))?; + let result = py_fn.call1(py, (name, py_args)).map_err(|e| { + FlowError::Internal(format!("Python tool callback failed: {e}")) + })?; + split_json_or_future(py, result) + })) + .await }) }) } /// Wrap a Python callable `(str, Json) -> Optional[str]` for tool conditional guardrails. pub fn wrap_py_tool_conditional_fn(py_fn: Py) -> ToolConditionalFn { - Arc::new(move |name: &str, args: &Json| { - Python::attach(|py| { - let py_args = json_to_py(py, args).map_err(|e| { - FlowError::Internal(format!( - "tool conditional json_to_py failed for '{name}': {e}" - )) - })?; - let result = py_fn.call1(py, (name, py_args)).map_err(|e| { - FlowError::Internal(format!( - "Python tool conditional callable failed for '{name}': {e}" - )) - })?; - let bound = result.bind(py); - if bound.is_none() { - Ok(None) - } else { - bound.extract::().map(Some).map_err(|e| { - FlowError::Internal(format!( - "tool conditional guardrail for '{name}' returned unexpected type (expected str or None): {e}" - )) - }) - } + let py_fn = Arc::new(py_fn); + Arc::new(move |name: String, args: Json| { + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let py_args = + json_to_py(py, &args).map_err(|e| FlowError::Internal(e.to_string()))?; + let result = py_fn + .call1(py, (name, py_args)) + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + let bound = result.bind(py); + if bound.is_none() { + Ok(None) + } else { + bound.extract::().map(Some).map_err(|e| { + FlowError::Internal(format!( + "tool conditional guardrail returned unexpected type: {e}" + )) + }) + } + }) }) }) } /// Wrap a Python callable `(str, Json) -> Json` for tool request intercepts. pub fn wrap_py_tool_request_intercept_fn(py_fn: Py) -> ToolInterceptFn { - Arc::new(move |name: &str, args: Json| { - Python::attach(|py| { - let py_args = json_to_py(py, &args).map_err(|e| { - FlowError::Internal(format!("tool callback json_to_py failed for '{name}': {e}")) - })?; - let result = py_fn.call1(py, (name, py_args)).map_err(|e| { - FlowError::Internal(format!("Python tool callable failed for '{name}': {e}")) - })?; - py_to_json(result.bind(py)).map_err(|e| { - FlowError::Internal(format!("tool callback py_to_json failed for '{name}': {e}")) - }) + let py_fn = Arc::new(py_fn); + Arc::new(move |name: String, args: Json| { + let py_fn = py_fn.clone(); + Box::pin(async move { + resolve_json_or_future(Python::attach(|py| { + let py_args = + json_to_py(py, &args).map_err(|e| FlowError::Internal(e.to_string()))?; + let result = py_fn + .call1(py, (name, py_args)) + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_json_or_future(py, result) + })) + .await }) }) } @@ -815,30 +813,38 @@ pub fn wrap_py_llm_stream_exec_intercept_fn( /// Wrap a Python callable `(LlmRequest, LlmSanitizeRequestContext) -> Optional`. fn wrap_py_llm_sanitize_request_callback(py_fn: Py) -> LlmSanitizeRequestFn { + let py_fn = Arc::new(py_fn); Arc::new( move |request: LlmRequest, context: LlmSanitizeRequestContext| { - Python::attach(|py| { - let py_context = PyLlmSanitizeRequestContext { inner: context }; - let py_request = PyLLMRequest { inner: request }; - let result = match py_fn.call1(py, (py_request, py_context)) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: LLM sanitize request callable failed: {error}"); - return None; - } - }; - if result.is_none(py) { - return None; - } - match result.extract::(py) { - Ok(request) => Some(request.inner), - Err(error) => { - eprintln!( - "nemo_relay: LLM sanitize request callable returned unexpected type: {error}" - ); - None + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let result = py_fn + .call1( + py, + ( + PyLLMRequest { inner: request }, + PyLlmSanitizeRequestContext { inner: context }, + ), + ) + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + if result.is_none(py) { + Ok(None) + } else { + result + .extract::(py) + .map(|request| Some(request.inner)) + .map_err(|error| { + FlowError::Internal(format!( + "LLM sanitize request returned unexpected type: {error}" + )) + }) } - } + }) }) }, ) @@ -846,24 +852,29 @@ fn wrap_py_llm_sanitize_request_callback(py_fn: Py) -> LlmSanitizeRequest /// Wrap a Python callable `(LlmRequest) -> Optional[str]` for LLM conditional guardrails. pub fn wrap_py_llm_conditional_fn(py_fn: Py) -> LlmConditionalFn { - Arc::new(move |request: &LlmRequest| { - Python::attach(|py| { - let py_req = PyLLMRequest { - inner: request.clone(), - }; - let result = py_fn.call1(py, (py_req,)).map_err(|e| { - FlowError::Internal(format!("LLM conditional guardrail callable failed: {e}")) - })?; - let bound = result.bind(py); - if bound.is_none() { - Ok(None) - } else { - bound.extract::().map(Some).map_err(|e| { - FlowError::Internal(format!( - "LLM conditional guardrail returned unexpected type (expected str or None): {e}" - )) - }) - } + let py_fn = Arc::new(py_fn); + Arc::new(move |request: LlmRequest| { + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let result = py_fn + .call1(py, (PyLLMRequest { inner: request },)) + .map_err(|e| FlowError::Internal(e.to_string()))?; + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + let bound = result.bind(py); + if bound.is_none() { + Ok(None) + } else { + bound.extract::().map(Some).map_err(|e| { + FlowError::Internal(format!( + "LLM conditional guardrail returned unexpected type: {e}" + )) + }) + } + }) }) }) } @@ -875,42 +886,45 @@ pub fn wrap_py_llm_conditional_fn(py_fn: Py) -> LlmConditionalFn { /// When ``annotated`` is present, request content is read-only and provider-body /// edits must be made through the returned annotation; headers remain writable. pub fn wrap_py_llm_request_intercept_fn(py_fn: Py) -> LlmRequestInterceptFn { + let py_fn = Arc::new(py_fn); Arc::new( - move |name: &str, - request: LlmRequest, - annotated: Option| - -> FlowResult { - Python::attach(|py| { - let py_req = PyLLMRequest { - inner: request.clone(), - }; - let py_ann: Py = match annotated { - Some(ann) => { - let wrapper = PyAnnotatedLLMRequest { inner: ann }; - wrapper - .into_pyobject(py) - .map_err(|e| { - FlowError::Internal(format!( - "Failed to convert AnnotatedLLMRequest to Python: {e}" - )) - })? - .into_any() - .unbind() - } - None => py.None(), - }; - let result = py_fn.call1(py, (name, py_req, py_ann)).map_err(|e| { - FlowError::Internal(format!("LLM request intercept callable failed: {e}")) - })?; + move |name: String, request: LlmRequest, annotated: Option| { + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let py_req = PyLLMRequest { inner: request }; + let py_ann: Py = match annotated { + Some(ann) => { + let wrapper = PyAnnotatedLLMRequest { inner: ann }; + wrapper + .into_pyobject(py) + .map_err(|e| { + FlowError::Internal(format!( + "Failed to convert AnnotatedLLMRequest to Python: {e}" + )) + })? + .into_any() + .unbind() + } + None => py.None(), + }; + let result = py_fn.call1(py, (name, py_req, py_ann)).map_err(|e| { + FlowError::Internal(format!("LLM request intercept callable failed: {e}")) + })?; - result - .extract::(py) - .map(|value| value.inner) - .map_err(|e| { - FlowError::Internal(format!( - "LLM request intercept must return LLMRequestInterceptOutcome: {e}" - )) - }) + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + result + .extract::(py) + .map(|value| value.inner) + .map_err(|e| { + FlowError::Internal(format!( + "LLM request intercept must return LLMRequestInterceptOutcome: {e}" + )) + }) + }) }) }, ) @@ -1014,33 +1028,29 @@ pub fn wrap_py_finalizer_fn(py_fn: Py) -> Box Json + Send /// Wrap a Python callable `(Json, LlmSanitizeResponseContext) -> Optional[Json]`. fn wrap_py_llm_sanitize_response_callback(py_fn: Py) -> LlmSanitizeResponseFn { + let py_fn = Arc::new(py_fn); Arc::new(move |response: Json, context: LlmSanitizeResponseContext| { - Python::attach(|py| { - let py_context = PyLlmSanitizeResponseContext { inner: context }; - let py_response = match json_to_py(py, &response) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: json_to_py failed in LLM sanitize response: {error}"); - return None; - } - }; - let result = match py_fn.call1(py, (py_response, py_context)) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: LLM sanitize response callable failed: {error}"); - return None; - } - }; - if result.is_none(py) { - return None; - } - match py_to_json(result.bind(py)) { - Ok(response) => Some(response), - Err(error) => { - eprintln!("nemo_relay: py_to_json failed in LLM sanitize response: {error}"); - None + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = resolve_py_object_or_future(Python::attach(|py| { + let py_context = PyLlmSanitizeResponseContext { inner: context }; + let py_response = json_to_py(py, &response) + .map_err(|error| FlowError::Internal(error.to_string()))?; + let result = py_fn + .call1(py, (py_response, py_context)) + .map_err(|error| FlowError::Internal(error.to_string()))?; + split_py_object_or_future(py, result) + })) + .await?; + Python::attach(|py| { + if result.is_none(py) { + Ok(None) + } else { + py_to_json(result.bind(py)) + .map(Some) + .map_err(|error| FlowError::Internal(error.to_string())) } - } + }) }) }) } @@ -1084,61 +1094,71 @@ pub fn wrap_py_event_subscriber(py_fn: Py) -> EventSubscriberFn { /// Wrap a Python callable ``(Event, EventSanitizeFields) -> EventSanitizeFields``. pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { - Arc::new(move |event: &Event, fields: EventSanitizeFields| { - Python::attach(|py| { - let py_event = match event { - Event::Scope(inner) => Py::new( - py, - crate::py_types::PyScopeEvent { - inner: inner.clone(), - }, - ) - .map(|value| value.into_any()), - Event::Mark(inner) => Py::new( - py, - crate::py_types::PyMarkEvent { - inner: inner.clone(), - }, - ) - .map(|value| value.into_any()), - }; - let py_event = match py_event { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: failed to convert event sanitizer context: {error}"); - return EventSanitizeFields::default(); - } - }; - let fields_json = match serde_json::to_value(&fields) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: failed to serialize event sanitizer fields: {error}"); - return EventSanitizeFields::default(); - } - }; - let py_fields = match json_to_py(py, &fields_json) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: failed to convert event sanitizer fields: {error}"); - return EventSanitizeFields::default(); - } - }; - let result = match py_fn.call1(py, (py_event, py_fields)) { - Ok(value) => value, - Err(error) => { - eprintln!("nemo_relay: Python event sanitizer callable failed: {error}"); - return EventSanitizeFields::default(); - } - }; - py_to_json(result.bind(py)) - .ok() - .and_then(|value| serde_json::from_value(value).ok()) - .unwrap_or_else(|| { - eprintln!( - "nemo_relay: event sanitizer must return data, category_profile, and metadata fields" - ); - EventSanitizeFields::default() - }) + let py_fn = Arc::new(py_fn); + Arc::new(move |event: Event, fields: EventSanitizeFields| { + let py_fn = py_fn.clone(); + Box::pin(async move { + let result = Python::attach( + |py| -> FlowResult, PyValueFuture>> { + let py_event = match &event { + Event::Scope(inner) => Py::new( + py, + crate::py_types::PyScopeEvent { + inner: inner.clone(), + }, + ) + .map(|value| value.into_any()), + Event::Mark(inner) => Py::new( + py, + crate::py_types::PyMarkEvent { + inner: inner.clone(), + }, + ) + .map(|value| value.into_any()), + }; + let py_event = match py_event { + Ok(value) => value, + Err(error) => { + eprintln!( + "nemo_relay: failed to convert event sanitizer context: {error}" + ); + return Err(FlowError::Internal(error.to_string())); + } + }; + let fields_json = match serde_json::to_value(&fields) { + Ok(value) => value, + Err(error) => { + eprintln!( + "nemo_relay: failed to serialize event sanitizer fields: {error}" + ); + return Err(FlowError::Internal(error.to_string())); + } + }; + let py_fields = match json_to_py(py, &fields_json) { + Ok(value) => value, + Err(error) => { + eprintln!( + "nemo_relay: failed to convert event sanitizer fields: {error}" + ); + return Err(FlowError::Internal(error.to_string())); + } + }; + let result = py_fn + .call1(py, (py_event, py_fields)) + .map_err(|error| FlowError::Internal(error.to_string()))?; + split_py_object_or_future(py, result) + }, + ); + let result = resolve_py_object_or_future(result).await?; + Python::attach(|py| { + py_to_json(result.bind(py)) + .map_err(|error| FlowError::Internal(error.to_string())) + .and_then(|value| { + serde_json::from_value(value).map_err(|error| { + FlowError::Internal(format!("invalid event sanitizer result: {error}")) + }) + }) + }) }) }) } diff --git a/crates/python/tests/coverage/coverage_tests.rs b/crates/python/tests/coverage/coverage_tests.rs index 67cebd3ba..2fa805abf 100644 --- a/crates/python/tests/coverage/coverage_tests.rs +++ b/crates/python/tests/coverage/coverage_tests.rs @@ -661,51 +661,68 @@ def event_fail(event): "#, ); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); let tool_ok = wrap_py_tool_fn(module.getattr("tool_ok").unwrap().unbind()); assert_eq!( - tool_ok("demo", json!({"x": 1})), + runtime + .block_on(tool_ok("demo".to_string(), json!({"x": 1}))) + .unwrap(), json!({"seen": 1, "name": "demo"}) ); let tool_fail = wrap_py_tool_fn(module.getattr("tool_fail").unwrap().unbind()); - assert_eq!(tool_fail("demo", json!({"x": 1})), json!({"x": 1})); + assert!( + runtime + .block_on(tool_fail("demo".to_string(), json!({"x": 1}))) + .is_err() + ); let tool_cond = wrap_py_tool_conditional_fn(module.getattr("tool_cond_bad").unwrap().unbind()); + let error = runtime + .block_on(tool_cond("demo".to_string(), json!({"x": 1}))) + .unwrap_err(); assert!( - tool_cond("demo", &json!({"x": 1})) - .unwrap_err() - .to_string() - .contains("expected str or None") + error.to_string().contains("unexpected type"), + "unexpected tool conditional error: {error}" ); let request = make_request(); let llm_sanitize = wrap_py_llm_sanitize_request_fn(module.getattr("llm_sanitize_bad").unwrap().unbind()) .unwrap(); - assert_eq!( - llm_sanitize( - request.clone(), - nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), - ), - None + assert!( + runtime + .block_on(llm_sanitize( + request.clone(), + nemo_relay::api::runtime::LlmSanitizeRequestContext::default(), + )) + .is_err() ); let llm_cond = wrap_py_llm_conditional_fn(module.getattr("llm_cond_bad").unwrap().unbind()); assert!( - llm_cond(&request) + runtime + .block_on(llm_cond(request.clone())) .unwrap_err() .to_string() - .contains("expected str or None") + .contains("unexpected type") ); let llm_cond_none = wrap_py_llm_conditional_fn(module.getattr("llm_cond_none").unwrap().unbind()); - assert_eq!(llm_cond_none(&request).unwrap(), None); + assert_eq!( + runtime.block_on(llm_cond_none(request.clone())).unwrap(), + None + ); let llm_req = wrap_py_llm_request_intercept_fn(module.getattr("llm_req_bad").unwrap().unbind()); assert!( - llm_req("demo", request.clone(), None) + runtime + .block_on(llm_req("demo".to_string(), request.clone(), None)) .unwrap_err() .to_string() .contains("intercept callable failed") @@ -714,21 +731,21 @@ def event_fail(event): let tool_req = wrap_py_tool_request_intercept_fn(module.getattr("tool_fail").unwrap().unbind()); assert!( - tool_req("demo", json!({"x": 1})) - .unwrap_err() - .to_string() - .contains("Python tool callable failed") + runtime + .block_on(tool_req("demo".to_string(), json!({"x": 1}))) + .is_err() ); let llm_resp = wrap_py_llm_sanitize_response_fn(module.getattr("llm_resp_fail").unwrap().unbind()) .unwrap(); - assert_eq!( - llm_resp( - json!({"ok": true}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), - ), - None + assert!( + runtime + .block_on(llm_resp( + json!({"ok": true}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::default(), + )) + .is_err() ); let mut collector = diff --git a/crates/python/tests/coverage/py_api_coverage_tests.rs b/crates/python/tests/coverage/py_api_coverage_tests.rs index 0bfe38573..c1cc50293 100644 --- a/crates/python/tests/coverage/py_api_coverage_tests.rs +++ b/crates/python/tests/coverage/py_api_coverage_tests.rs @@ -472,18 +472,31 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute ) .unwrap(); - let tool_intercepted = - tool_request_intercepts(py, "demo-tool", &py_dict(py, json!({"value": 1}))).unwrap(); + let tool_intercepted = tool_request_intercepts( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": 1})), + ) + .unwrap(); assert_eq!( - crate::convert::py_to_json(tool_intercepted.bind(py)).unwrap(), + crate::convert::py_to_json(&tool_intercepted).unwrap(), json!({"value": 3}) ); - tool_conditional_execution("demo-tool", &py_dict(py, json!({"value": 1}))).unwrap(); + tool_conditional_execution( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": 1})), + ) + .unwrap(); assert!( - tool_conditional_execution("demo-tool", &py_dict(py, json!({"value": -1}))) - .unwrap_err() - .to_string() - .contains("blocked") + tool_conditional_execution( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": -1})) + ) + .unwrap_err() + .to_string() + .contains("blocked") ); let llm_request = PyLLMRequest { @@ -492,7 +505,10 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute content: json!({"messages": [{"role": "user", "content": "hello"}], "model": "demo-model"}), }, }; - let intercepted_request = llm_request_intercepts("demo-llm", llm_request.clone()).unwrap(); + let intercepted_request = + llm_request_intercepts(py, "demo-llm".to_string(), llm_request.clone()).unwrap(); + let intercepted_request: PyRef<'_, crate::py_types::PyLLMRequestInterceptOutcome> = + intercepted_request.extract().unwrap(); assert_eq!( intercepted_request .inner @@ -501,14 +517,17 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute .get("x-intercepted"), Some(&json!("1")) ); - llm_conditional_execution(llm_request.clone()).unwrap(); + llm_conditional_execution(py, llm_request.clone()).unwrap(); assert!( - llm_conditional_execution(PyLLMRequest { - inner: nemo_relay::api::llm::LlmRequest { - headers: serde_json::Map::new(), - content: json!({"messages": [], "model": "blocked"}), - }, - }) + llm_conditional_execution( + py, + PyLLMRequest { + inner: nemo_relay::api::llm::LlmRequest { + headers: serde_json::Map::new(), + content: json!({"messages": [], "model": "blocked"}), + }, + } + ) .unwrap_err() .to_string() .contains("blocked") diff --git a/crates/python/tests/coverage/py_callable_coverage_tests.rs b/crates/python/tests/coverage/py_callable_coverage_tests.rs index fe815cc1b..89a382f88 100644 --- a/crates/python/tests/coverage/py_callable_coverage_tests.rs +++ b/crates/python/tests/coverage/py_callable_coverage_tests.rs @@ -153,7 +153,13 @@ class RaisingResponseCodec: "model": "codec-model" })) .unwrap(); - let outcome = request_intercept("llm", make_request(), Some(annotated.clone())).unwrap(); + let outcome = runtime + .block_on(request_intercept( + "llm".to_string(), + make_request(), + Some(annotated.clone()), + )) + .unwrap(); assert_eq!( outcome.annotated_request.unwrap().last_user_message(), Some("annotated") @@ -163,7 +169,12 @@ class RaisingResponseCodec: module.getattr("request_bad_annotated").unwrap().unbind(), ); assert!( - bad_request_intercept("llm", make_request(), Some(annotated)) + runtime + .block_on(bad_request_intercept( + "llm".to_string(), + make_request(), + Some(annotated), + )) .unwrap_err() .to_string() .contains("must return LLMRequestInterceptOutcome") @@ -173,7 +184,12 @@ class RaisingResponseCodec: module.getattr("request_short_tuple").unwrap().unbind(), ); assert!( - short_request_intercept("llm", make_request(), None) + runtime + .block_on(short_request_intercept( + "llm".to_string(), + make_request(), + None + )) .unwrap_err() .to_string() .contains("must return LLMRequestInterceptOutcome") @@ -189,12 +205,13 @@ class RaisingResponseCodec: let llm_response = wrap_py_llm_sanitize_response_fn(module.getattr("llm_resp_bad_json").unwrap().unbind()) .unwrap(); - assert_eq!( - llm_response( - json!({"ok": true}), - nemo_relay::api::runtime::LlmSanitizeResponseContext::default() - ), - None + assert!( + runtime + .block_on(llm_response( + json!({"ok": true}), + nemo_relay::api::runtime::LlmSanitizeResponseContext::default() + )) + .is_err() ); let bad_codec = PyLlmCodecWrapper { @@ -683,23 +700,31 @@ def invalid(event, fields): metadata: Some(json!({"secret": true})), }; - let sanitized = wrap_py_event_sanitize_fn(module.getattr("sanitize").unwrap().unbind())( - &event, - fields.clone(), - ); + let runtime = tokio::runtime::Runtime::new().unwrap(); + let sanitized = runtime + .block_on(wrap_py_event_sanitize_fn( + module.getattr("sanitize").unwrap().unbind(), + )(event.clone(), fields.clone())) + .unwrap(); assert_eq!(sanitized.data, Some(json!({"safe": "checkpoint"}))); assert_eq!(sanitized.metadata, None); - let raised = wrap_py_event_sanitize_fn(module.getattr("raises").unwrap().unbind())( - &event, - fields.clone(), - ); - assert_eq!(raised, EventSanitizeFields::default()); - - let invalid = wrap_py_event_sanitize_fn(module.getattr("invalid").unwrap().unbind())( - &event, - fields.clone(), + let raised = runtime + .block_on(wrap_py_event_sanitize_fn( + module.getattr("raises").unwrap().unbind(), + )(event.clone(), fields.clone())) + .unwrap_err(); + assert!(raised.to_string().contains("sanitize boom")); + + let invalid = runtime + .block_on(wrap_py_event_sanitize_fn( + module.getattr("invalid").unwrap().unbind(), + )(event, fields.clone())) + .unwrap_err(); + assert!( + invalid + .to_string() + .contains("invalid event sanitizer result") ); - assert_eq!(invalid, EventSanitizeFields::default()); }); } diff --git a/docs/about-nemo-relay/concepts/middleware.mdx b/docs/about-nemo-relay/concepts/middleware.mdx index 0afcd2d56..afd1ae0cf 100644 --- a/docs/about-nemo-relay/concepts/middleware.mdx +++ b/docs/about-nemo-relay/concepts/middleware.mdx @@ -19,6 +19,19 @@ events. NeMo Relay applies each surface at a specific lifecycle point. Middleware is organized by lifecycle meaning rather than as one undifferentiated hook system. +## Asynchronous Callbacks + +All middleware families accept asynchronous callbacks. Rust callbacks return a +future; Python callbacks may return a value or an awaitable; and Node callbacks +may return a value or a Promise. Relay awaits entries sequentially in priority +order, so later callbacks observe earlier middleware output. + +Managed execution and standalone conditional/request-intercept helpers are +asynchronous because their result depends on middleware completion. Manual +lifecycle APIs (`tool_call`, `tool_call_end`, `llm_call`, and `llm_call_end`) +remain synchronous: they create or close their handle immediately and queue +observability work rather than awaiting it. + ## Registration Levels Middleware and subscribers can be registered at different levels depending on their @@ -125,6 +138,19 @@ context. For the callback contract and binding APIs, refer to Sanitize guardrails are observability-oriented. They do not rewrite the real arguments passed to the callback or the real value returned to the caller. +## Queued Event Publication + +Scope operations, marks, and manual tool/LLM lifecycle calls never become +awaitable because an event sanitizer is asynchronous. At emission time Relay +snapshots the event, visible sanitizer chain, and subscribers, then places the +work on a serial dispatcher. The dispatcher awaits sanitizers and publishes the +event later in FIFO order. + +Subscriber and exporter delivery is therefore delayed, while start/end/mark +order is preserved. Closing a scope or deregistering middleware after emission +does not affect queued snapshots. Sanitizer failures fail open: Relay records +the callback failure and publishes the last valid event snapshot. + ## Managed Execution Order For managed execution, NeMo Relay applies middleware and emits lifecycle events diff --git a/docs/build-plugins/dynamic-plugins/native-dynamic/about.mdx b/docs/build-plugins/dynamic-plugins/native-dynamic/about.mdx index a4c162368..8bb88a328 100644 --- a/docs/build-plugins/dynamic-plugins/native-dynamic/about.mdx +++ b/docs/build-plugins/dynamic-plugins/native-dynamic/about.mdx @@ -1,6 +1,6 @@ --- title: "Native Dynamic Plugins (Rust)" -description: "Build in-process Rust shared-library plugins against the NeMo Relay Native ABI v2." +description: "Build in-process Rust shared-library plugins against the NeMo Relay Native ABI v3." position: 10 --- {/* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. @@ -116,10 +116,12 @@ path, then replace `` with that library's SHA-256 digest. Use Native Plugin](/build-plugins/dynamic-plugins/native-dynamic/rust-native-plugin-example) for a complete example with validation, middleware, scopes, and configuration schema support. -## Native ABI v2 +## Native ABI v3 -The host passes a `NemoRelayNativeHostApiV1` table to the entry symbol. The -plugin returns a `NemoRelayNativePluginV1` descriptor: +The entry symbol receives a `*const NemoRelayNativeHostApiV1` pointer. It +points at the v1 prefix of a v3 `NemoRelayNativeHostApiV3` table; check +`abi_version` and `struct_size` before casting. The plugin returns a +`NemoRelayNativePluginV1` descriptor: ```rust extern "C" fn nemo_relay_register_plugin( @@ -128,6 +130,35 @@ extern "C" fn nemo_relay_register_plugin( ) -> NemoRelayStatus ``` +The v3 host table retains the frozen legacy prefix and appends a +completion-based asynchronous middleware extension. An entry that rejects the +v3 table with `InvalidArg` is retried with the legacy table. Rust plugins using the +typed `NativePlugin` APIs continue to work unchanged. Raw ABI plugins can use +`PluginContext::register_async_middleware_raw` when a callback must complete +later. The callback receives a JSON invocation, an optional continuation for +execution intercepts, and a one-shot completion handle. + +Return `Complete` after resolving or rejecting the completion before the +callback returns. Return `Pending` only when retaining the completion; settle +it exactly once, call `async_completion_release`, and release an async `next` +handle after use. The host marks a completion cancelled when the awaiting +runtime work is dropped; late and duplicate settlement is rejected safely. + +Event sanitizers registered through this extension still run on Relay's serial +publication dispatcher. Scope and mark emission remain synchronous and their +sanitized events are delivered later in emission order. + +The v3 completion and continuation ABI settles one JSON value. Consequently, +an async native LLM stream execution intercept currently receives and returns +the complete JSON array of chunks: Relay buffers the provider stream before +replaying it to the caller. It is not an incremental streaming transport and +does not provide per-chunk backpressure. Use a synchronous native stream +intercept or a worker plugin when first-token latency is required. + +Legacy v1/v2 middleware callbacks are synchronous and run on the runtime's +execution path. They must not block on I/O; use the v3 completion-based API for +long-running work. + Text and JSON data cross this boundary as host-owned `NemoRelayNativeString` handles. ABI structs also carry scalars, opaque handles, callback pointers, and plugin-owned `user_data`. Do not pass Rust diff --git a/docs/reference/event-sanitizers.mdx b/docs/reference/event-sanitizers.mdx index 8ded00de5..6ff89e7e7 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. @@ -174,17 +190,23 @@ activation fails. ## Experimental C and Go Bindings -The source-first C API uses `NemoRelayEventSanitizeCb`. It provides global, -scope-local, and plugin-context registration functions for all three surfaces. -Global names start with `nemo_relay_register_`, and scope-local names start -with `nemo_relay_scope_register_`. - -The Go binding provides `EventSanitizeFields`, `EventSanitizeFunc`, global -`Register*SanitizeGuardrail` helpers, scope-local -`ScopeRegister*SanitizeGuardrail` helpers, and the same methods on -`PluginContext`. The `guardrails` package provides shorter aliases. Because a -returned `EventSanitizeFields` replaces all three fields, copy the supplied -value and modify only the fields that should change. +The source-first C API retains `NemoRelayEventSanitizeCb` and adds parallel +completion-based async registration APIs. An async callback returns `Complete` +or `Pending` and settles its one-shot completion handle with resolve or reject; +there is no implicit timeout. A callback that returns `Pending` must settle the +handle exactly once, or serial event publication remains blocked. Relay cancels +the handle when the invocation is abandoned; late or duplicate settlement after +cancellation is rejected safely. After resolving or rejecting a retained +completion, call `nemo_relay_async_completion_release` to release the +callback-owned reference. Global names start with `nemo_relay_register_`, and +scope-local names start with `nemo_relay_scope_register_`. + +The Go binding provides `EventSanitizeFields`, `EventSanitizeFunc`, and +`AsyncMiddlewareFunc` variants for global and scope-local event sanitizers. +Async Go callbacks receive a `context.Context`; Relay cancels it when the +invocation is abandoned. Because a returned `EventSanitizeFields` replaces all +three fields, copy the supplied value and modify only the fields that should +change. ## Related Topics diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index edfeada6f..83fdd207c 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -1,6 +1,6 @@ --- title: "Migration Guides" -description: "Upgrade NeMo Relay integrations and migrate LLM sanitizer callbacks, plugins, workers, and PII policy." +description: "Upgrade NeMo Relay integrations and migrate async middleware, plugins, workers, LLM sanitizers, and PII policy." position: 6 --- {/* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. @@ -13,17 +13,64 @@ intervening release in sequence. ## Upgrade to NeMo Relay 0.7 -NeMo Relay 0.7 changes the LLM observability sanitizer contract across -in-process bindings, native plugins, raw C FFI consumers, and worker plugins. -Complete the following migrations before you run an existing sanitizer with a -0.7 host. +NeMo Relay 0.7 makes the Rust middleware callback contract asynchronous and +adds awaitable middleware support across the in-process bindings, native +plugins, raw C FFI consumers, and worker plugins. It also changes the LLM +observability sanitizer contract. Complete the following migrations before you +run existing middleware or a sanitizer with a 0.7 host. -Do not deploy a 0.6 sanitizer plugin or worker against a 0.7 host. The LLM -callback signature, native ABI layout, and worker invocation schema changed. -NeMo Relay does not adapt one-argument LLM sanitizer callbacks. +Do not deploy a 0.6 plugin or worker against a 0.7 host. The middleware +callback contract, LLM callback signature, native ABI layout, and worker +invocation schema changed. NeMo Relay does not adapt synchronous Rust +middleware callbacks or one-argument LLM sanitizer callbacks. +### Migrate Middleware Callbacks + +The following callback families are now asynchronous: conditional execution +guardrails, request intercepts, execution intercepts, tool and LLM sanitizers, +and event sanitizers. Relay awaits each registered callback sequentially in +priority order. A callback that rejects or returns an error preserves the +existing error behavior for its middleware family. + +| Surface | 0.6 Callback | 0.7 Callback | +| --- | --- | --- | +| Rust | `Fn(...) -> Result` | `Fn(...) -> Pin> + Send>>` | +| Python | Direct return value | Direct return value or awaitable | +| Node.js | Direct return value | Direct return value or `Promise` | +| Go / raw C FFI | Synchronous callback | Existing synchronous callback, or the new `Async` / completion-based registration API | + +For Rust, wrap the existing result in a ready async future, or use an async +block when the callback needs to await work: + +```rust +use std::sync::Arc; + +use nemo_relay::api::registry::register_tool_conditional_execution_guardrail; + +register_tool_conditional_execution_guardrail( + "policy", + 10, + Arc::new(|_name, _args| { + Box::pin(async move { + // Await policy I/O here when needed. + Ok(None) // Return Some(reason) to block execution. + }) + }), +)?; +``` + +Python and Node.js registration names are unchanged. Mark a Python callback +`async def`, or return a Promise from Node.js, only when it needs asynchronous +work; existing direct-value callbacks remain supported. + +Scope lifecycle and mark emission remain synchronous. `push_scope`, +`pop_scope`, and mark APIs snapshot the event and visible sanitizer/subscriber +chain, then enqueue sanitization and publication on a serial dispatcher. Event +subscribers and exporters therefore receive sanitized events later, in emission +order. Do not add `await` to scope or mark emission calls. + ### Update LLM Sanitizer Callbacks The registration names remain unchanged for global, plugin-context, and @@ -49,9 +96,11 @@ register_llm_sanitize_request_guardrail( "redact-request", 10, Arc::new(|request, context| { - let _active_codec = context.resolve_codec(); - // Apply policy, using _active_codec when normalized access is required. - Some(request) + Box::pin(async move { + let _active_codec = context.resolve_codec(); + // Apply policy, using _active_codec when normalized access is required. + Ok(Some(request)) + }) }), )?; ``` @@ -150,15 +199,15 @@ For complete in-process examples, refer to ### Migrate Worker Sanitizers -All Rust worker sanitizer registrations now require callbacks that return -futures. This change applies to mark, scope-start, scope-end, tool-request, -tool-response, LLM-request, and LLM-response sanitizers. Conditional guardrails -and request intercepts keep their existing synchronous contracts. +All Rust worker middleware registrations now require callbacks that return +futures. This includes conditional guardrails, request and execution intercepts, +mark and scope event sanitizers, tool request/response sanitizers, and LLM +request/response sanitizers. Python worker middleware can return either an +immediate value or an awaitable. Python LLM sanitizers must still accept both +the payload and directional context. -Update Rust worker callbacks to use `async move` and return `Result` from the -future. Python worker sanitizers can return either an immediate value or an -awaitable, but Python LLM sanitizers must still accept both the payload and -directional context. +Update Rust worker callbacks to return `Box::pin(async move { ... })` and +resolve to `Result` from the future. @@ -167,13 +216,15 @@ directional context. ctx.register_llm_sanitize_request_guardrail( "redact-request", 10, - |request, context| async move { - if let Some(codec) = context.resolve_codec() { - let annotated = codec.decode(&request).await?; - let request = codec.encode(&annotated, &request).await?; - return Ok(Some(request)); - } - Ok(Some(request)) + |request, context| { + Box::pin(async move { + if let Some(codec) = context.resolve_codec() { + let annotated = codec.decode(&request).await?; + let request = codec.encode(&annotated, &request).await?; + return Ok(Some(request)); + } + Ok(Some(request)) + }) }, ); ``` @@ -220,13 +271,14 @@ wrong-direction capability IDs. ### Rebuild Native and Raw FFI Plugins -NeMo Relay 0.7 uses native ABI v2. Recompile native plugins against the 0.7 +NeMo Relay 0.7 uses native ABI v3. Recompile native plugins against the 0.7 `nemo-relay-plugin` crate and rebuild raw FFI consumers against the generated 0.7 header. -If you already built a plugin against an earlier 0.7 ABI v2 prerelease, rebuild -it again. The ABI version remains 2, but the prerelease LLM sanitizer callback -slots and context layouts changed before release. +The v3 table preserves the v2 prefix, and Relay retries a legacy v2 table when +loading a plugin that rejects v3. Rebuild anyway if a plugin uses raw ABI +callbacks: v3 adds completion-based async middleware registration, async +execution continuations, and explicit cancellation/late-settlement behavior. The plugin manifest value remains `compat.native_api = "1"`. This manifest contract version is separate from the host ABI version; do not change it to @@ -243,7 +295,7 @@ after the callback returns. Release host-owned output strings with the standard host string release operation. For the complete ABI contract, refer to -[Native ABI v2](/build-plugins/dynamic-plugins/native-dynamic/about#native-abi-v2). +[Native ABI v3](/build-plugins/dynamic-plugins/native-dynamic/about#native-abi-v3). ### Update PII Redaction Configuration diff --git a/go/nemo_relay/adaptive_runtime_test.go b/go/nemo_relay/adaptive_runtime_test.go index 51db86a5d..dca103b69 100644 --- a/go/nemo_relay/adaptive_runtime_test.go +++ b/go/nemo_relay/adaptive_runtime_test.go @@ -169,3 +169,69 @@ func TestSetLatencySensitivityRejectsInvalidValue(t *testing.T) { t.Fatal("expected SetLatencySensitivity(0) to fail") } } + +func TestAdaptiveRuntimeRejectsNilAndShutDownHandles(t *testing.T) { + var nilRuntime *AdaptiveRuntime + if err := nilRuntime.Register(); err == nil { + t.Fatal("expected nil Register to fail") + } + if err := nilRuntime.Deregister(); err == nil { + t.Fatal("expected nil Deregister to fail") + } + if err := nilRuntime.Shutdown(); err == nil { + t.Fatal("expected nil Shutdown to fail") + } + if err := nilRuntime.WaitForIdle(); err == nil { + t.Fatal("expected nil WaitForIdle to fail") + } + if _, err := nilRuntime.Report(); err == nil { + t.Fatal("expected nil Report to fail") + } + if err := nilRuntime.BindScope(nil); err == nil { + t.Fatal("expected nil BindScope to fail") + } + if _, err := nilRuntime.BuildCacheRequestFacts(CacheRequestFactsInput{}); err == nil { + t.Fatal("expected nil BuildCacheRequestFacts to fail") + } + + runtime, err := NewAdaptiveRuntime(NewAdaptiveConfig()) + if err != nil { + t.Fatalf(newAdaptiveRuntimeFailedMsg, err) + } + if err := runtime.Shutdown(); err != nil { + t.Fatalf("Shutdown failed: %v", err) + } + if err := runtime.Register(); err == nil { + t.Fatal("expected Register after Shutdown to fail") + } + if err := runtime.WaitForIdle(); err == nil { + t.Fatal("expected WaitForIdle after Shutdown to fail") + } +} + +func TestAdaptiveRuntimeHelpersRejectInvalidInputs(t *testing.T) { + if _, err := BuildCacheTelemetryEvent(CacheTelemetryEventInput{ + Provider: "unsupported", + RequestID: "not-a-uuid", + }); err == nil { + t.Fatal("expected invalid telemetry input to fail") + } + + runtime, err := NewAdaptiveRuntime(testAdaptiveRuntimeConfig("openai")) + if err != nil { + t.Fatalf(newAdaptiveRuntimeFailedMsg, err) + } + defer runtime.Shutdown() + if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{ + Provider: "unsupported", + RequestID: "not-a-uuid", + AnnotatedRequest: json.RawMessage(`{}`), + }); err == nil { + t.Fatal("expected invalid cache request facts input to fail") + } + if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{ + AnnotatedRequest: json.RawMessage(`not-json`), + }); err == nil { + t.Fatal("expected malformed annotated request JSON to fail before the FFI call") + } +} diff --git a/go/nemo_relay/async_middleware_test.go b/go/nemo_relay/async_middleware_test.go new file mode 100644 index 000000000..ee501908b --- /dev/null +++ b/go/nemo_relay/async_middleware_test.go @@ -0,0 +1,224 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +package nemo_relay + +import ( + "context" + "encoding/json" + "errors" + "strings" + "testing" +) + +func asyncMiddlewareNoop(context.Context, json.RawMessage) (any, error) { + return nil, nil +} + +func asyncExecutionNoop(context.Context, json.RawMessage, AsyncNext) (any, error) { + return nil, nil +} + +func TestAsyncMiddlewareGlobalRegistrationParity(t *testing.T) { + registrations := []struct { + name string + register func(string) error + deregister func(string) error + }{ + {"mark", func(name string) error { return RegisterMarkSanitizeGuardrailAsync(name, 0, asyncMiddlewareNoop) }, DeregisterMarkSanitizeGuardrail}, + {"scope-start", func(name string) error { return RegisterScopeSanitizeStartGuardrailAsync(name, 0, asyncMiddlewareNoop) }, DeregisterScopeSanitizeStartGuardrail}, + {"scope-end", func(name string) error { return RegisterScopeSanitizeEndGuardrailAsync(name, 0, asyncMiddlewareNoop) }, DeregisterScopeSanitizeEndGuardrail}, + {"tool-sanitize-request", func(name string) error { + return RegisterToolSanitizeRequestGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterToolSanitizeRequestGuardrail}, + {"tool-sanitize-response", func(name string) error { + return RegisterToolSanitizeResponseGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterToolSanitizeResponseGuardrail}, + {"tool-conditional", func(name string) error { + return RegisterToolConditionalExecutionGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterToolConditionalExecutionGuardrail}, + {"tool-request", func(name string) error { return RegisterToolRequestInterceptAsync(name, 0, false, asyncMiddlewareNoop) }, DeregisterToolRequestIntercept}, + {"tool-execution", func(name string) error { return RegisterToolExecutionInterceptAsync(name, 0, asyncExecutionNoop) }, DeregisterToolExecutionIntercept}, + {"llm-sanitize-request", func(name string) error { return RegisterLlmSanitizeRequestGuardrailAsync(name, 0, asyncMiddlewareNoop) }, DeregisterLlmSanitizeRequestGuardrail}, + {"llm-sanitize-response", func(name string) error { + return RegisterLlmSanitizeResponseGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterLlmSanitizeResponseGuardrail}, + {"llm-conditional", func(name string) error { + return RegisterLlmConditionalExecutionGuardrailAsync(name, 0, asyncMiddlewareNoop) + }, DeregisterLlmConditionalExecutionGuardrail}, + {"llm-request", func(name string) error { return RegisterLlmRequestInterceptAsync(name, 0, false, asyncMiddlewareNoop) }, DeregisterLlmRequestIntercept}, + {"llm-execution", func(name string) error { return RegisterLlmExecutionInterceptAsync(name, 0, asyncExecutionNoop) }, DeregisterLlmExecutionIntercept}, + {"llm-stream-execution", func(name string) error { return RegisterLlmStreamExecutionInterceptAsync(name, 0, asyncExecutionNoop) }, DeregisterLlmStreamExecutionIntercept}, + } + + for _, registration := range registrations { + t.Run(registration.name, func(t *testing.T) { + name := "go-async-global-" + registration.name + if err := registration.register(name); err != nil { + t.Fatalf("register: %v", err) + } + if err := registration.deregister(name); err != nil { + t.Fatalf("deregister: %v", err) + } + }) + } +} + +func TestAsyncMiddlewareScopeLocalRegistrationParity(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + handle, err := PushScope("async-registration-owner", ScopeTypeAgent) + if err != nil { + t.Fatalf("push scope: %v", err) + } + defer func() { + if err := PopScope(handle); err != nil { + t.Fatalf("pop scope: %v", err) + } + }() + + scopeUUID := handle.UUID() + registrations := []struct { + name string + register func(string) error + deregister func(string) error + }{ + {"mark", func(name string) error { + return ScopeRegisterMarkSanitizeGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterMarkSanitizeGuardrail(scopeUUID, name) }}, + {"scope-start", func(name string) error { + return ScopeRegisterScopeSanitizeStartGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterScopeSanitizeStartGuardrail(scopeUUID, name) }}, + {"scope-end", func(name string) error { + return ScopeRegisterScopeSanitizeEndGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterScopeSanitizeEndGuardrail(scopeUUID, name) }}, + {"tool-sanitize-request", func(name string) error { + return ScopeRegisterToolSanitizeRequestGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterToolSanitizeRequestGuardrail(scopeUUID, name) }}, + {"tool-sanitize-response", func(name string) error { + return ScopeRegisterToolSanitizeResponseGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterToolSanitizeResponseGuardrail(scopeUUID, name) }}, + {"tool-conditional", func(name string) error { + return ScopeRegisterToolConditionalExecutionGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterToolConditionalExecutionGuardrail(scopeUUID, name) }}, + {"tool-request", func(name string) error { + return ScopeRegisterToolRequestInterceptAsync(scopeUUID, name, 0, false, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterToolRequestIntercept(scopeUUID, name) }}, + {"tool-execution", func(name string) error { + return ScopeRegisterToolExecutionInterceptAsync(scopeUUID, name, 0, asyncExecutionNoop) + }, func(name string) error { return ScopeDeregisterToolExecutionIntercept(scopeUUID, name) }}, + {"llm-sanitize-request", func(name string) error { + return ScopeRegisterLlmSanitizeRequestGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterLlmSanitizeRequestGuardrail(scopeUUID, name) }}, + {"llm-sanitize-response", func(name string) error { + return ScopeRegisterLlmSanitizeResponseGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterLlmSanitizeResponseGuardrail(scopeUUID, name) }}, + {"llm-conditional", func(name string) error { + return ScopeRegisterLlmConditionalExecutionGuardrailAsync(scopeUUID, name, 0, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterLlmConditionalExecutionGuardrail(scopeUUID, name) }}, + {"llm-request", func(name string) error { + return ScopeRegisterLlmRequestInterceptAsync(scopeUUID, name, 0, false, asyncMiddlewareNoop) + }, func(name string) error { return ScopeDeregisterLlmRequestIntercept(scopeUUID, name) }}, + {"llm-execution", func(name string) error { + return ScopeRegisterLlmExecutionInterceptAsync(scopeUUID, name, 0, asyncExecutionNoop) + }, func(name string) error { return ScopeDeregisterLlmExecutionIntercept(scopeUUID, name) }}, + {"llm-stream-execution", func(name string) error { + return ScopeRegisterLlmStreamExecutionInterceptAsync(scopeUUID, name, 0, asyncExecutionNoop) + }, func(name string) error { return ScopeDeregisterLlmStreamExecutionIntercept(scopeUUID, name) }}, + } + + for _, registration := range registrations { + name := "go-async-local-" + registration.name + if err := registration.register(name); err != nil { + t.Fatalf("register %s: %v", registration.name, err) + } + if err := registration.deregister(name); err != nil { + t.Fatalf("deregister %s: %v", registration.name, err) + } + } + }) +} + +func TestAsyncToolMiddlewareCompletionAndNext(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + if err := RegisterToolConditionalExecutionGuardrailAsync("go-async-tool-conditional", 0, asyncMiddlewareNoop); err != nil { + t.Fatalf("register conditional: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolConditionalExecutionGuardrail("go-async-tool-conditional") }) + + if err := RegisterToolExecutionInterceptAsync("go-async-tool-execution", 0, + func(ctx context.Context, invocation json.RawMessage, next AsyncNext) (any, error) { + var payload struct { + Value json.RawMessage `json:"value"` + } + if err := json.Unmarshal(invocation, &payload); err != nil { + return nil, err + } + result, err := next(ctx, payload.Value) + if err != nil { + return nil, err + } + return map[string]json.RawMessage{"result": result}, nil + }, + ); err != nil { + t.Fatalf("register execution intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolExecutionIntercept("go-async-tool-execution") }) + + result, err := ToolCallExecute("go-async-tool", json.RawMessage(`{"value":1}`), func(args json.RawMessage) (json.RawMessage, error) { + return args, nil + }) + if err != nil { + t.Fatalf("tool call execute: %v", err) + } + if string(result) != `{"value":1}` { + t.Fatalf("tool result = %s, want original result", result) + } + }) +} + +func TestAsyncToolMiddlewarePropagatesCallbackAndNextErrors(t *testing.T) { + runTestWithScopeStack(t, func(t *testing.T) { + const conditionalName = "go-async-tool-conditional-error" + if err := RegisterToolConditionalExecutionGuardrailAsync(conditionalName, 0, + func(context.Context, json.RawMessage) (any, error) { + return nil, errors.New("conditional callback failed") + }, + ); err != nil { + t.Fatalf("register conditional: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolConditionalExecutionGuardrail(conditionalName) }) + + _, err := ToolCallExecute("go-async-tool-conditional-error", json.RawMessage(`{}`), func(json.RawMessage) (json.RawMessage, error) { + return json.RawMessage(`{}`), nil + }) + if err == nil || !strings.Contains(err.Error(), "conditional callback failed") { + t.Fatalf("conditional error = %v, want callback failure", err) + } + + if err := DeregisterToolConditionalExecutionGuardrail(conditionalName); err != nil { + t.Fatalf("deregister conditional: %v", err) + } + const executionName = "go-async-tool-next-error" + if err := RegisterToolExecutionInterceptAsync(executionName, 0, + func(ctx context.Context, invocation json.RawMessage, next AsyncNext) (any, error) { + var payload struct { + Value json.RawMessage `json:"value"` + } + if err := json.Unmarshal(invocation, &payload); err != nil { + return nil, err + } + return next(ctx, payload.Value) + }, + ); err != nil { + t.Fatalf("register execution intercept: %v", err) + } + t.Cleanup(func() { _ = DeregisterToolExecutionIntercept(executionName) }) + + _, err = ToolCallExecute("go-async-tool-next-error", json.RawMessage(`{}`), func(json.RawMessage) (json.RawMessage, error) { + return nil, errors.New("tool implementation failed") + }) + if err == nil || !strings.Contains(err.Error(), "tool implementation failed") { + t.Fatalf("next error = %v, want implementation failure", err) + } + }) +} diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index 8096c7447..f7837ab2b 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -50,6 +50,16 @@ typedef char* (*NemoRelayLlmSanitizeResponseCb)(void* user_data, const char* res typedef void (*NemoRelayEventSubscriberFn)(void* user_data, const FfiEvent* event); typedef char* (*NemoRelayEventSanitizeFn)(void* user_data, const FfiEvent* event, const char* fields_json); typedef struct FfiPluginContext FfiPluginContext; +typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; +typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; +typedef void (*NemoRelayAsyncNextResultCb)(void*, const char*, const char*); +extern int32_t nemo_relay_async_completion_resolve_json(const NemoRelayAsyncCompletion*, const char*); +extern int32_t nemo_relay_async_completion_reject(const NemoRelayAsyncCompletion*, const char*); +extern bool nemo_relay_async_completion_is_cancelled(const NemoRelayAsyncCompletion*); +extern void nemo_relay_async_completion_release(const NemoRelayAsyncCompletion*); +extern int32_t nemo_relay_async_next_invoke_callback(const NemoRelayAsyncNext*, const char*, NemoRelayAsyncNextResultCb, void*); +extern void nemo_relay_async_next_release(const NemoRelayAsyncNext*); +extern void goAsyncNextResultTrampoline(void*, char*, char*); // Middleware chain next function types typedef char* (*NemoRelayToolExecNextFn)(const char* args_json, void* next_ctx); @@ -87,10 +97,12 @@ typedef NemoRelayCodecEncodeCb NemoRelayCodecEncodeFn; import "C" import ( + "context" "encoding/json" "errors" "sync" "sync/atomic" + "time" "unsafe" ) @@ -167,6 +179,43 @@ type ToolSanitizeFunc func(name string, args json.RawMessage) json.RawMessage // message string to reject the call. type ToolConditionalFunc func(name string, args json.RawMessage) *string +// AsyncMiddlewareFunc is the common completion-based middleware callback. +// The JSON envelope identifies the middleware family and invocation fields. +type AsyncMiddlewareFunc func(ctx context.Context, invocation json.RawMessage) (any, error) + +// AsyncNext invokes the remaining execution chain and returns its eventual result. +type AsyncNext func(ctx context.Context, invocation json.RawMessage) (json.RawMessage, error) + +// AsyncExecutionInterceptFunc is an asynchronous execution intercept with an awaitable next helper. +type AsyncExecutionInterceptFunc func(ctx context.Context, invocation json.RawMessage, next AsyncNext) (any, error) + +const asyncCallbackPending = C.uint32_t(1) + +func contextForCompletion(completion *C.NemoRelayAsyncCompletion) (context.Context, func()) { + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + var doneOnce sync.Once + go func() { + ticker := time.NewTicker(10 * time.Millisecond) + defer ticker.Stop() + for { + select { + case <-done: + return + case <-ticker.C: + if bool(C.nemo_relay_async_completion_is_cancelled(completion)) { + cancel() + return + } + } + } + }() + return ctx, func() { + doneOnce.Do(func() { close(done) }) + cancel() + } +} + // ToolExecutionFunc is a callback that executes a tool call, receiving the // arguments as JSON and returning the result JSON or an error. type ToolExecutionFunc func(args json.RawMessage) (json.RawMessage, error) @@ -590,6 +639,137 @@ func goToolSanitizeTrampoline(userData unsafe.Pointer, name *C.char, argsJSON *C return C.CString(string(result)) } +//export goAsyncMiddlewareTrampoline +func goAsyncMiddlewareTrampoline(userData unsafe.Pointer, invocationJSON *C.char, completion *C.NemoRelayAsyncCompletion) C.uint32_t { + fn, ok := lookupClosure(userData).(AsyncMiddlewareFunc) + if !ok { + message := C.CString("nemo_relay: async middleware callback is not registered") + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) + C.nemo_relay_async_completion_release(completion) + return asyncCallbackPending + } + invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) + go func() { + defer C.nemo_relay_async_completion_release(completion) + ctx, cancel := contextForCompletion(completion) + defer cancel() + value, err := fn(ctx, invocation) + if err != nil { + message := C.CString(err.Error()) + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) + return + } + encoded, err := json.Marshal(value) + if err != nil { + message := C.CString(err.Error()) + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) + return + } + result := C.CString(string(encoded)) + defer C.free(unsafe.Pointer(result)) + C.nemo_relay_async_completion_resolve_json(completion, result) + }() + return asyncCallbackPending +} + +type asyncNextResult struct { + value json.RawMessage + err error +} + +//export goAsyncNextResultTrampoline +func goAsyncNextResultTrampoline(userData unsafe.Pointer, valueJSON *C.char, errorMessage *C.char) { + ch, ok := lookupClosure(userData).(chan asyncNextResult) + if !ok { + return + } + defer unregisterClosure(userData) + if errorMessage != nil { + select { + case ch <- asyncNextResult{err: errors.New(C.GoString(errorMessage))}: + default: + } + return + } + select { + case ch <- asyncNextResult{value: append(json.RawMessage(nil), []byte(C.GoString(valueJSON))...)}: + default: + } +} + +//export goAsyncExecutionInterceptTrampoline +func goAsyncExecutionInterceptTrampoline(userData unsafe.Pointer, invocationJSON *C.char, next *C.NemoRelayAsyncNext, completion *C.NemoRelayAsyncCompletion) C.uint32_t { + fn, ok := lookupClosure(userData).(AsyncExecutionInterceptFunc) + if !ok { + message := C.CString("nemo_relay: async execution intercept callback is not registered") + defer C.free(unsafe.Pointer(message)) + C.nemo_relay_async_completion_reject(completion, message) + C.nemo_relay_async_completion_release(completion) + C.nemo_relay_async_next_release(next) + return asyncCallbackPending + } + invocation := append(json.RawMessage(nil), []byte(C.GoString(invocationJSON))...) + go func() { + defer C.nemo_relay_async_completion_release(completion) + ctx, cancel := contextForCompletion(completion) + defer cancel() + var nextMu sync.RWMutex + nextOpen := true + defer func() { + nextMu.Lock() + nextOpen = false + nextMu.Unlock() + C.nemo_relay_async_next_release(next) + }() + nextFn := func(ctx context.Context, payload json.RawMessage) (json.RawMessage, error) { + nextMu.RLock() + defer nextMu.RUnlock() + if !nextOpen { + return nil, context.Canceled + } + ch := make(chan asyncNextResult, 1) + token := registerClosure(ch) + cPayload := C.CString(string(payload)) + status := C.nemo_relay_async_next_invoke_callback( + next, cPayload, + (C.NemoRelayAsyncNextResultCb)(C.goAsyncNextResultTrampoline), token, + ) + C.free(unsafe.Pointer(cPayload)) + if err := checkStatus(status); err != nil { + unregisterClosure(token) + return nil, err + } + select { + case result := <-ch: + return result.value, result.err + case <-ctx.Done(): + return nil, ctx.Err() + } + } + value, err := fn(ctx, invocation, nextFn) + if err != nil { + message := C.CString(err.Error()) + C.nemo_relay_async_completion_reject(completion, message) + C.free(unsafe.Pointer(message)) + return + } + encoded, err := json.Marshal(value) + if err != nil { + message := C.CString(err.Error()) + C.nemo_relay_async_completion_reject(completion, message) + C.free(unsafe.Pointer(message)) + return + } + result := C.CString(string(encoded)) + C.nemo_relay_async_completion_resolve_json(completion, result) + C.free(unsafe.Pointer(result)) + }() + return asyncCallbackPending +} + //export goToolConditionalTrampoline func goToolConditionalTrampoline(userData unsafe.Pointer, name *C.char, argsJSON *C.char) *C.char { fn := lookupClosure(userData).(ToolConditionalFunc) diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index 8f4c510d9..d255f907f 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -44,6 +44,10 @@ typedef struct NemoRelayLlmSanitizeRequestContext { uint32_t codec_kind; const c typedef struct NemoRelayLlmSanitizeResponseContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeResponseCodec* codec; } NemoRelayLlmSanitizeResponseContext; typedef void (*NemoRelayFreeFn)(void* user_data); +typedef struct NemoRelayAsyncCompletion NemoRelayAsyncCompletion; +typedef struct NemoRelayAsyncNext NemoRelayAsyncNext; +typedef uint32_t (*NemoRelayAsyncJsonCb)(void*, const char*, const NemoRelayAsyncCompletion*); +typedef uint32_t (*NemoRelayAsyncInterceptCb)(void*, const char*, const NemoRelayAsyncNext*, const NemoRelayAsyncCompletion*); // Core API extern int32_t nemo_relay_get_handle(FfiScopeHandle** out); @@ -121,46 +125,57 @@ extern void nemo_relay_set_last_error_message(const char* msg); // Tool guardrails typedef char* (*NemoRelayToolSanitizeFn)(void* user_data, const char* name, const char* args_json); extern int32_t nemo_relay_register_tool_sanitize_request_guardrail(const char* name, int32_t priority, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_tool_sanitize_request_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_sanitize_request_guardrail(const char* name); extern int32_t nemo_relay_register_tool_sanitize_response_guardrail(const char* name, int32_t priority, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_tool_sanitize_response_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_sanitize_response_guardrail(const char* name); typedef char* (*NemoRelayToolConditionalFn)(void* user_data, const char* name, const char* args_json); extern int32_t nemo_relay_register_tool_conditional_execution_guardrail(const char* name, int32_t priority, NemoRelayToolConditionalFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_tool_conditional_execution_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_conditional_execution_guardrail(const char* name); // Tool intercepts extern int32_t nemo_relay_register_tool_request_intercept(const char* name, int32_t priority, _Bool break_chain, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_tool_request_intercept_async(const char*, int32_t, _Bool, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_request_intercept(const char* name); // Middleware chain intercept callback types (must be declared before use in externs) typedef char* (*NemoRelayToolExecNextFn)(const char* args_json, void* next_ctx); typedef char* (*NemoRelayToolExecInterceptCb)(void* user_data, const char* args_json, NemoRelayToolExecNextFn next_fn, void* next_ctx); extern int32_t nemo_relay_register_tool_execution_intercept(const char* name, int32_t priority, NemoRelayToolExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_register_tool_execution_intercept_async(const char*, int32_t, NemoRelayAsyncInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_tool_execution_intercept(const char* name); // LLM guardrails typedef FfiLLMRequest* (*NemoRelayLlmSanitizeRequestCb)(void* user_data, const FfiLLMRequest* request, NemoRelayLlmSanitizeRequestContext context); extern int32_t nemo_relay_register_llm_sanitize_request_guardrail(const char* name, int32_t priority, NemoRelayLlmSanitizeRequestCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_llm_sanitize_request_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_sanitize_request_guardrail(const char* name); typedef char* (*NemoRelayLlmSanitizeResponseCb)(void* user_data, const char* response_json, NemoRelayLlmSanitizeResponseContext context); extern int32_t nemo_relay_register_llm_sanitize_response_guardrail(const char* name, int32_t priority, NemoRelayLlmSanitizeResponseCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_llm_sanitize_response_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_sanitize_response_guardrail(const char* name); typedef char* (*NemoRelayLlmConditionalCb)(void* user_data, const FfiLLMRequest* request); extern int32_t nemo_relay_register_llm_conditional_execution_guardrail(const char* name, int32_t priority, NemoRelayLlmConditionalCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_llm_conditional_execution_guardrail_async(const char*, int32_t, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_conditional_execution_guardrail(const char* name); // LLM intercepts typedef int32_t (*NemoRelayLlmRequestInterceptCb)(void* user_data, const char* name, const FfiLLMRequest* request, const char* annotated_json, char** out_outcome_json); extern int32_t nemo_relay_register_llm_request_intercept(const char* name, int32_t priority, _Bool break_chain, NemoRelayLlmRequestInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_llm_request_intercept_async(const char*, int32_t, _Bool, NemoRelayAsyncJsonCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_request_intercept(const char* name); typedef char* (*NemoRelayLlmExecNextFn)(const char* native_json, void* next_ctx); typedef char* (*NemoRelayLlmExecInterceptCb)(void* user_data, const char* native_json, NemoRelayLlmExecNextFn next_fn, void* next_ctx); extern int32_t nemo_relay_register_llm_execution_intercept(const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_register_llm_execution_intercept_async(const char*, int32_t, NemoRelayAsyncInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_execution_intercept(const char* name); extern int32_t nemo_relay_register_llm_stream_execution_intercept(const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_register_llm_stream_execution_intercept_async(const char*, int32_t, NemoRelayAsyncInterceptCb, void*, NemoRelayFreeFn); extern int32_t nemo_relay_deregister_llm_stream_execution_intercept(const char* name); // Subscribers @@ -170,46 +185,63 @@ extern int32_t nemo_relay_register_subscriber(const char* name, NemoRelayEventSu extern int32_t nemo_relay_deregister_subscriber(const char* name); extern int32_t nemo_relay_flush_subscribers(void); extern int32_t nemo_relay_register_mark_sanitize_guardrail(const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_mark_sanitize_guardrail_async(const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_deregister_mark_sanitize_guardrail(const char* name); extern int32_t nemo_relay_register_scope_sanitize_start_guardrail(const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_scope_sanitize_start_guardrail_async(const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_deregister_scope_sanitize_start_guardrail(const char* name); extern int32_t nemo_relay_register_scope_sanitize_end_guardrail(const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_register_scope_sanitize_end_guardrail_async(const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_deregister_scope_sanitize_end_guardrail(const char* name); // Scope-local tool guardrails extern int32_t nemo_relay_scope_register_mark_sanitize_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_mark_sanitize_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_mark_sanitize_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_scope_sanitize_start_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_scope_sanitize_start_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_scope_sanitize_start_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_scope_sanitize_end_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayEventSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_scope_sanitize_end_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_scope_sanitize_end_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_tool_sanitize_request_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_tool_sanitize_request_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_sanitize_request_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_tool_sanitize_response_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_tool_sanitize_response_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_sanitize_response_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_tool_conditional_execution_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayToolConditionalFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_tool_conditional_execution_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_conditional_execution_guardrail(const char* scope_uuid, const char* name); // Scope-local tool intercepts extern int32_t nemo_relay_scope_register_tool_request_intercept(const char* scope_uuid, const char* name, int32_t priority, _Bool break_chain, NemoRelayToolSanitizeFn cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_tool_request_intercept_async(const char* scope_uuid, const char* name, int32_t priority, _Bool break_chain, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_request_intercept(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_tool_execution_intercept(const char* scope_uuid, const char* name, int32_t priority, NemoRelayToolExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_scope_register_tool_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_tool_execution_intercept(const char* scope_uuid, const char* name); // Scope-local LLM guardrails extern int32_t nemo_relay_scope_register_llm_sanitize_request_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmSanitizeRequestCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_sanitize_request_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_sanitize_request_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_sanitize_response_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmSanitizeResponseCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_sanitize_response_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_sanitize_response_guardrail(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_conditional_execution_guardrail(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmConditionalCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_conditional_execution_guardrail_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_conditional_execution_guardrail(const char* scope_uuid, const char* name); // Scope-local LLM intercepts extern int32_t nemo_relay_scope_register_llm_request_intercept(const char* scope_uuid, const char* name, int32_t priority, _Bool break_chain, NemoRelayLlmRequestInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); +extern int32_t nemo_relay_scope_register_llm_request_intercept_async(const char* scope_uuid, const char* name, int32_t priority, _Bool break_chain, NemoRelayAsyncJsonCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_request_intercept(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_execution_intercept(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_scope_register_llm_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_execution_intercept(const char* scope_uuid, const char* name); extern int32_t nemo_relay_scope_register_llm_stream_execution_intercept(const char* scope_uuid, const char* name, int32_t priority, NemoRelayLlmExecInterceptCb exec_cb, void* exec_user_data, NemoRelayFreeFn exec_free); +extern int32_t nemo_relay_scope_register_llm_stream_execution_intercept_async(const char* scope_uuid, const char* name, int32_t priority, NemoRelayAsyncInterceptCb cb, void* user_data, NemoRelayFreeFn free_fn); extern int32_t nemo_relay_scope_deregister_llm_stream_execution_intercept(const char* scope_uuid, const char* name); // Scope-local subscribers @@ -276,6 +308,8 @@ extern void nemo_relay_openinference_subscriber_free(void*); // Go trampoline forward declarations (defined via //export in callbacks.go) extern char* goToolSanitizeTrampoline(void*, const char*, const char*); +extern uint32_t goAsyncMiddlewareTrampoline(void*, const char*, const NemoRelayAsyncCompletion*); +extern uint32_t goAsyncExecutionInterceptTrampoline(void*, const char*, const NemoRelayAsyncNext*, const NemoRelayAsyncCompletion*); extern char* goEventSanitizeTrampoline(void*, const FfiEvent*, const char*); extern char* goToolConditionalTrampoline(void*, const char*, const char*); extern char* goToolExecTrampoline(void*, const char*); @@ -1206,11 +1240,34 @@ func registerEventSanitizer(name string, priority int32, fn EventSanitizeFunc, k return checkStatus(status) } +func registerAsyncEventSanitizer(name string, priority int32, fn AsyncMiddlewareFunc, kind int) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + callback := C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline) + free := C.NemoRelayFreeFn(C.goFreeTrampoline) + var status C.int32_t + switch kind { + case 0: + status = C.nemo_relay_register_mark_sanitize_guardrail_async(cName, C.int32_t(priority), callback, id, free) + case 1: + status = C.nemo_relay_register_scope_sanitize_start_guardrail_async(cName, C.int32_t(priority), callback, id, free) + default: + status = C.nemo_relay_register_scope_sanitize_end_guardrail_async(cName, C.int32_t(priority), callback, id, free) + } + return checkStatus(status) +} + // RegisterMarkSanitizeGuardrail registers a global mark event sanitizer. func RegisterMarkSanitizeGuardrail(name string, priority int32, fn EventSanitizeFunc) error { return registerEventSanitizer(name, priority, fn, 0) } +// RegisterMarkSanitizeGuardrailAsync registers an asynchronous global mark sanitizer. +func RegisterMarkSanitizeGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + return registerAsyncEventSanitizer(name, priority, fn, 0) +} + // DeregisterMarkSanitizeGuardrail removes a global mark event sanitizer. func DeregisterMarkSanitizeGuardrail(name string) error { cName := C.CString(name) @@ -1223,6 +1280,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) @@ -1235,6 +1297,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) @@ -1260,6 +1327,17 @@ func RegisterToolSanitizeRequestGuardrail(name string, priority int32, fn ToolSa )) } +// RegisterToolSanitizeRequestGuardrailAsync registers an asynchronous tool request sanitizer. +func RegisterToolSanitizeRequestGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_sanitize_request_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolSanitizeRequestGuardrail removes a previously registered tool // sanitize-request guardrail by name. Returns a NotFound error if no guardrail // with the given name is registered. @@ -1285,6 +1363,17 @@ func RegisterToolSanitizeResponseGuardrail(name string, priority int32, fn ToolS )) } +// RegisterToolSanitizeResponseGuardrailAsync registers an asynchronous tool response sanitizer. +func RegisterToolSanitizeResponseGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_sanitize_response_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolSanitizeResponseGuardrail removes a previously registered tool // sanitize-response guardrail by name. Returns a NotFound error if no guardrail // with the given name is registered. @@ -1312,6 +1401,17 @@ func RegisterToolConditionalExecutionGuardrail(name string, priority int32, fn T )) } +// RegisterToolConditionalExecutionGuardrailAsync registers an asynchronous tool guardrail. +func RegisterToolConditionalExecutionGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_conditional_execution_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolConditionalExecutionGuardrail removes a previously registered // tool conditional-execution guardrail by name. Returns a NotFound error if no // guardrail with the given name is registered. @@ -1338,6 +1438,18 @@ func RegisterToolRequestIntercept(name string, priority int32, breakChain bool, )) } +// RegisterToolRequestInterceptAsync registers an asynchronous tool request intercept. +func RegisterToolRequestInterceptAsync(name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_request_intercept_async( + cName, C.int32_t(priority), C._Bool(breakChain), + C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolRequestIntercept removes a previously registered tool request // intercept by name. func DeregisterToolRequestIntercept(name string) error { @@ -1362,6 +1474,18 @@ func RegisterToolExecutionIntercept(name string, priority int32, execFn ToolExec )) } +// RegisterToolExecutionInterceptAsync registers an asynchronous tool execution intercept. +func RegisterToolExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_tool_execution_intercept_async( + cName, C.int32_t(priority), + C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterToolExecutionIntercept removes a previously registered tool // execution intercept by name. func DeregisterToolExecutionIntercept(name string) error { @@ -1388,6 +1512,17 @@ func RegisterLlmSanitizeRequestGuardrail(name string, priority int32, fn LLMRequ )) } +// RegisterLlmSanitizeRequestGuardrailAsync registers an asynchronous LLM request sanitizer. +func RegisterLlmSanitizeRequestGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_sanitize_request_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmSanitizeRequestGuardrail removes a previously registered LLM // sanitize-request guardrail by name. func DeregisterLlmSanitizeRequestGuardrail(name string) error { @@ -1410,6 +1545,17 @@ func RegisterLlmSanitizeResponseGuardrail(name string, priority int32, fn LLMRes )) } +// RegisterLlmSanitizeResponseGuardrailAsync registers an asynchronous LLM response sanitizer. +func RegisterLlmSanitizeResponseGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_sanitize_response_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmSanitizeResponseGuardrail removes a previously registered LLM // sanitize-response guardrail by name. func DeregisterLlmSanitizeResponseGuardrail(name string) error { @@ -1436,6 +1582,17 @@ func RegisterLlmConditionalExecutionGuardrail(name string, priority int32, fn LL )) } +// RegisterLlmConditionalExecutionGuardrailAsync registers an asynchronous LLM guardrail. +func RegisterLlmConditionalExecutionGuardrailAsync(name string, priority int32, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_conditional_execution_guardrail_async( + cName, C.int32_t(priority), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmConditionalExecutionGuardrail removes a previously registered // LLM conditional-execution guardrail by name. func DeregisterLlmConditionalExecutionGuardrail(name string) error { @@ -1462,6 +1619,18 @@ func RegisterLlmRequestIntercept(name string, priority int32, breakChain bool, f )) } +// RegisterLlmRequestInterceptAsync registers an asynchronous LLM request intercept. +func RegisterLlmRequestInterceptAsync(name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_request_intercept_async( + cName, C.int32_t(priority), C._Bool(breakChain), + C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmRequestIntercept removes a previously registered LLM request // intercept by name. func DeregisterLlmRequestIntercept(name string) error { @@ -1486,6 +1655,18 @@ func RegisterLlmExecutionIntercept(name string, priority int32, execFn LLMExecut )) } +// RegisterLlmExecutionInterceptAsync registers an asynchronous LLM execution intercept. +func RegisterLlmExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_execution_intercept_async( + cName, C.int32_t(priority), + C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmExecutionIntercept removes a previously registered LLM // execution intercept by name. func DeregisterLlmExecutionIntercept(name string) error { @@ -1511,6 +1692,18 @@ func RegisterLlmStreamExecutionIntercept(name string, priority int32, execFn LLM )) } +// RegisterLlmStreamExecutionInterceptAsync registers an asynchronous streaming LLM intercept. +func RegisterLlmStreamExecutionInterceptAsync(name string, priority int32, fn AsyncExecutionInterceptFunc) error { + id := registerClosure(fn) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(C.nemo_relay_register_llm_stream_execution_intercept_async( + cName, C.int32_t(priority), + C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), + id, C.NemoRelayFreeFn(C.goFreeTrampoline), + )) +} + // DeregisterLlmStreamExecutionIntercept removes a previously registered LLM // stream execution intercept by name. func DeregisterLlmStreamExecutionIntercept(name string) error { @@ -2300,6 +2493,71 @@ func (s *OpenInferenceSubscriber) Close() { // Scope-local guardrail/intercept registration (Tool) // --------------------------------------------------------------------------- +func withScopeAsyncMiddleware(scopeUUID, name string, priority int32, fn any, call func(*C.char, *C.char, C.int32_t, unsafe.Pointer) C.int32_t) error { + id := registerClosure(fn) + cScopeUUID := C.CString(scopeUUID) + defer C.free(unsafe.Pointer(cScopeUUID)) + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + return checkStatus(call(cScopeUUID, cName, C.int32_t(priority), id)) +} + +// ScopeRegisterMarkSanitizeGuardrailAsync registers an asynchronous scope-local mark sanitizer. +func ScopeRegisterMarkSanitizeGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_mark_sanitize_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterScopeSanitizeStartGuardrailAsync registers an asynchronous scope-local start sanitizer. +func ScopeRegisterScopeSanitizeStartGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_scope_sanitize_start_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterScopeSanitizeEndGuardrailAsync registers an asynchronous scope-local end sanitizer. +func ScopeRegisterScopeSanitizeEndGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_scope_sanitize_end_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolSanitizeRequestGuardrailAsync registers an asynchronous scope-local tool request sanitizer. +func ScopeRegisterToolSanitizeRequestGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_sanitize_request_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolSanitizeResponseGuardrailAsync registers an asynchronous scope-local tool response sanitizer. +func ScopeRegisterToolSanitizeResponseGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_sanitize_response_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolConditionalExecutionGuardrailAsync registers an asynchronous scope-local tool guardrail. +func ScopeRegisterToolConditionalExecutionGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_conditional_execution_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolRequestInterceptAsync registers an asynchronous scope-local tool request intercept. +func ScopeRegisterToolRequestInterceptAsync(scopeUUID, name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_request_intercept_async(scope, name, priority, C._Bool(breakChain), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterToolExecutionInterceptAsync registers an asynchronous scope-local tool execution intercept. +func ScopeRegisterToolExecutionInterceptAsync(scopeUUID, name string, priority int32, fn AsyncExecutionInterceptFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_tool_execution_intercept_async(scope, name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + func registerScopeEventSanitizer(scopeUUID, name string, priority int32, fn EventSanitizeFunc, kind int) error { id := registerClosure(fn) cScopeUUID := C.CString(scopeUUID) @@ -2498,6 +2756,48 @@ func ScopeDeregisterToolExecutionIntercept(scopeUUID, name string) error { // Scope-local guardrail/intercept registration (LLM) // --------------------------------------------------------------------------- +// ScopeRegisterLlmSanitizeRequestGuardrailAsync registers an asynchronous scope-local LLM request sanitizer. +func ScopeRegisterLlmSanitizeRequestGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_sanitize_request_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmSanitizeResponseGuardrailAsync registers an asynchronous scope-local LLM response sanitizer. +func ScopeRegisterLlmSanitizeResponseGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_sanitize_response_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmConditionalExecutionGuardrailAsync registers an asynchronous scope-local LLM guardrail. +func ScopeRegisterLlmConditionalExecutionGuardrailAsync(scopeUUID, name string, priority int32, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_conditional_execution_guardrail_async(scope, name, priority, C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmRequestInterceptAsync registers an asynchronous scope-local LLM request intercept. +func ScopeRegisterLlmRequestInterceptAsync(scopeUUID, name string, priority int32, breakChain bool, fn AsyncMiddlewareFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_request_intercept_async(scope, name, priority, C._Bool(breakChain), C.NemoRelayAsyncJsonCb(C.goAsyncMiddlewareTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmExecutionInterceptAsync registers an asynchronous scope-local LLM execution intercept. +func ScopeRegisterLlmExecutionInterceptAsync(scopeUUID, name string, priority int32, fn AsyncExecutionInterceptFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_execution_intercept_async(scope, name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + +// ScopeRegisterLlmStreamExecutionInterceptAsync registers an asynchronous scope-local streaming LLM intercept. +func ScopeRegisterLlmStreamExecutionInterceptAsync(scopeUUID, name string, priority int32, fn AsyncExecutionInterceptFunc) error { + return withScopeAsyncMiddleware(scopeUUID, name, priority, fn, func(scope, name *C.char, priority C.int32_t, id unsafe.Pointer) C.int32_t { + return C.nemo_relay_scope_register_llm_stream_execution_intercept_async(scope, name, priority, C.NemoRelayAsyncInterceptCb(C.goAsyncExecutionInterceptTrampoline), id, C.NemoRelayFreeFn(C.goFreeTrampoline)) + }) +} + // ScopeRegisterLlmSanitizeRequestGuardrail registers a scope-local guardrail // that sanitizes LLM request data. func ScopeRegisterLlmSanitizeRequestGuardrail(scopeUUID, name string, priority int32, fn LLMRequestFunc) error { diff --git a/go/nemo_relay/optimization_test.go b/go/nemo_relay/optimization_test.go index 60712a8da..491899ce5 100644 --- a/go/nemo_relay/optimization_test.go +++ b/go/nemo_relay/optimization_test.go @@ -98,6 +98,26 @@ func TestLLMOptimizationContributionOmittedAppliedIsNonApplied(t *testing.T) { } } +func TestLLMOptimizationContributionRejectsMalformedAndUnknownWireShapes(t *testing.T) { + var contribution LLMOptimizationContribution + if err := json.Unmarshal([]byte(`not-json`), &contribution); err == nil { + t.Fatal("expected malformed optimization contribution JSON to fail") + } + if err := json.Unmarshal([]byte(`[]`), &contribution); err == nil { + t.Fatal("expected non-object optimization contribution JSON to fail") + } + + contribution = LLMOptimizationContribution{ + Producer: "test", + Kind: "custom", + PayloadSchema: &LLMOptimizationDataSchema{Name: "test", Version: "v1"}, + Payload: json.RawMessage(`not-json`), + } + if _, err := json.Marshal(contribution); err == nil { + t.Fatal("expected malformed payload JSON to fail") + } +} + func TestLLMRequestInterceptOptimizationContributionsRoundTrip(t *testing.T) { fixture, contribution := optimizationContributionFixture(t) const interceptName = "go_optimization_fixture" diff --git a/integrations/openclaw/test/live-smoke.test.ts b/integrations/openclaw/test/live-smoke.test.ts index 0b480cbdf..26318160d 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 465b3ec17..cd48e62a3 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -169,26 +169,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. @@ -200,7 +207,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 951c58592..bacd50d98 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -167,8 +167,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: @@ -181,7 +184,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: @@ -190,7 +193,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: @@ -203,7 +209,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: @@ -216,7 +225,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: @@ -225,7 +234,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: @@ -252,7 +261,7 @@ Exceptional flow: """ LlmRequestIntercept: TypeAlias = Callable[ [str, LLMRequest, AnnotatedLLMRequest | None], - LLMRequestInterceptOutcome, + LLMRequestInterceptOutcome | Awaitable[LLMRequestInterceptOutcome], ] """Request intercept callback that rewrites raw and annotated LLM requests. diff --git a/python/nemo_relay/_native.pyi b/python/nemo_relay/_native.pyi index f91184e42..67ea68364 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]]], @@ -1688,7 +1697,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: @@ -1696,14 +1705,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: @@ -1711,7 +1721,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 @@ -1719,7 +1730,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: @@ -1727,21 +1740,23 @@ def llm_request_intercepts(name: str, request: LLMRequest) -> LLMRequestIntercep request: Current LLM request. Returns: - Transformed request after all applicable request intercepts. + Transformed request directly outside an event loop, or an awaitable + resolving to it from an async caller. Exceptional flow: Callback exceptions and native middleware errors propagate unchanged. """ ... -def llm_conditional_execution(request: LLMRequest) -> None: +def llm_conditional_execution(request: LLMRequest) -> None | Awaitable[None]: """Run LLM conditional-execution guardrails. Args: request: LLM request to evaluate. Returns: - ``None`` when all guardrails allow execution. + ``None`` when all guardrails allow execution, directly outside an event + loop or through an awaitable from an async caller. Exceptional flow: Raises a native rejection error when a guardrail returns a rejection diff --git a/python/tests/test_adaptive.py b/python/tests/test_adaptive.py index c7e956e7b..899c3b2e5 100644 --- a/python/tests/test_adaptive.py +++ b/python/tests/test_adaptive.py @@ -222,7 +222,7 @@ async def test_adaptive_runtime_bind_scope_passes_through_without_state(self): ) with scope.scope("adaptive-runtime-translate", ScopeType.Agent) as handle: runtime.bind_scope(handle) - translated = llm.request_intercepts("anthropic", request) + translated = await llm.request_intercepts("anthropic", request) assert translated.request.content == { "messages": [{"role": "user", "content": "Hello"}], "system": "You are helpful.", diff --git a/python/tests/test_builtin_codecs.py b/python/tests/test_builtin_codecs.py index 156122583..c1fcef898 100644 --- a/python/tests/test_builtin_codecs.py +++ b/python/tests/test_builtin_codecs.py @@ -12,8 +12,6 @@ from typing import cast -import pytest - import nemo_relay from nemo_relay import ( AnnotatedLLMRequest, @@ -434,8 +432,8 @@ def sanitize_response(response, context): subscribers.deregister("test-manual-call-end-sanitized-response-codec") guardrails.deregister_llm_sanitize_response("test-call-end-codec-sanitizer") - def test_manual_call_end_response_codec_failure_raises_after_end_event(self): - """manual llm.call_end() surfaces response codec failures instead of dropping them.""" + def test_manual_call_end_response_codec_failure_defers_without_raising(self): + """manual llm.call_end() records deferred response codec failures without blocking.""" captured_events = [] def capture(event): @@ -448,8 +446,7 @@ def capture(event): "manual-codec-error-llm", LLMRequest({}, {"model": "gpt-4", "messages": []}), ) - with pytest.raises(RuntimeError, match="OpenAI Chat response decode"): - llm.call_end(handle, "malformed response", response_codec=OpenAIChatCodec()) + llm.call_end(handle, "malformed response", response_codec=OpenAIChatCodec()) subscribers.flush() end_events = [ diff --git a/python/tests/test_context_isolation.py b/python/tests/test_context_isolation.py index 49314e312..7ee5a7531 100644 --- a/python/tests/test_context_isolation.py +++ b/python/tests/test_context_isolation.py @@ -199,8 +199,8 @@ async def run_tool(owner): ) await asyncio.sleep(0) - args = nemo_relay.tools.request_intercepts("task-tool", {"owner": owner}) - nemo_relay.tools.conditional_execution("task-tool", args) + args = await nemo_relay.tools.request_intercepts("task-tool", {"owner": owner}) + await nemo_relay.tools.conditional_execution("task-tool", args) manual_handle = nemo_relay.tools.call(f"manual-tool-{owner}", args) await asyncio.sleep(0) @@ -258,9 +258,9 @@ def intercept(name, request, annotated): request = nemo_relay.LLMRequest({}, {"messages": [], "owner": owner}) await asyncio.sleep(0) - intercepted = nemo_relay.llm.request_intercepts("task-llm", request) + intercepted = await nemo_relay.llm.request_intercepts("task-llm", request) assert intercepted.request.content["intercepted_by"] == owner - nemo_relay.llm.conditional_execution(request) + await nemo_relay.llm.conditional_execution(request) manual_handle = nemo_relay.llm.call(f"manual-llm-{owner}", request) await asyncio.sleep(0) diff --git a/python/tests/test_event_sanitizers.py b/python/tests/test_event_sanitizers.py index 89d97d9d6..08234b70e 100644 --- a/python/tests/test_event_sanitizers.py +++ b/python/tests/test_event_sanitizers.py @@ -56,7 +56,7 @@ def second(event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitiz assert calls == [("checkpoint", {"secret": "raw"}), ("mark", {"stage": "first"})] -def test_mark_sanitizer_exception_clears_observability_fields(capture_events): +def test_mark_sanitizer_exception_preserves_observability_fields(capture_events): _capture_name, events = capture_events def raises(_event: nemo_relay.Event, _fields: EventSanitizeFields) -> EventSanitizeFields: @@ -69,7 +69,7 @@ def raises(_event: nemo_relay.Event, _fields: EventSanitizeFields) -> EventSanit finally: guardrails.deregister_mark_sanitize("python-mark-raises") - assert events[-1].data is None + assert events[-1].data == {"kept": True} assert events[-1].metadata is None diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index 9a3662be4..411629401 100644 --- a/python/tests/test_llm.py +++ b/python/tests/test_llm.py @@ -167,6 +167,7 @@ def sanitize_response(response, context): try: handle = llm.call("py_llm_structured_context", make_request()) llm.call_end(handle, {"response": "ok"}) + subscribers.flush() finally: guardrails.deregister_llm_sanitize_request("py_llm_structured_context_request") guardrails.deregister_llm_sanitize_response("py_llm_structured_context_response") @@ -275,7 +276,7 @@ def test_duplicate_raises(self): guardrails.register_llm_sanitize_request("py_llm_dup", 1, lambda r, context: r) guardrails.deregister_llm_sanitize_request("py_llm_dup") - def test_sanitize_request_callable_error_omits_observability_input(self): + def test_sanitize_request_callable_error_preserves_observability_input(self): events = [] subscribers.register("py_llm_sanitize_req_sub", lambda event: events.append(event)) guardrails.register_llm_sanitize_request( @@ -284,7 +285,10 @@ def test_sanitize_request_callable_error_omits_observability_input(self): lambda request, context: raise_runtime_error("boom"), ) try: - request = make_request() + request = LLMRequest( + {"authorization": "secret", "x-request-id": "safe"}, + make_request().content, + ) handle = llm.call("llm_sanitize_req_fail", request) llm.call_end(handle, {"ok": True}) finally: @@ -295,10 +299,10 @@ def test_sanitize_request_callable_error_omits_observability_input(self): subscribers.deregister("py_llm_sanitize_req_sub") start = _llm_event(events, "llm_sanitize_req_fail", "start") - assert start.data is None + assert start.data == {"headers": {"x-request-id": "safe"}, "content": request.content} assert start.annotated_request is None - def test_sanitize_request_invalid_return_omits_observability_input(self): + def test_sanitize_request_invalid_return_preserves_observability_input(self): events = [] subscribers.register("py_llm_sanitize_req_bad_sub", lambda event: events.append(event)) guardrails.register_llm_sanitize_request( @@ -307,7 +311,10 @@ def test_sanitize_request_invalid_return_omits_observability_input(self): cast(guardrails.LlmSanitizeRequestGuardrail, lambda request, context: object()), ) try: - request = make_request() + request = LLMRequest( + {"authorization": "secret", "x-request-id": "safe"}, + make_request().content, + ) handle = llm.call("llm_sanitize_req_bad", request) llm.call_end(handle, {"ok": True}) finally: @@ -318,10 +325,10 @@ def test_sanitize_request_invalid_return_omits_observability_input(self): subscribers.deregister("py_llm_sanitize_req_bad_sub") start = _llm_event(events, "llm_sanitize_req_bad", "start") - assert start.data is None + assert start.data == {"headers": {"x-request-id": "safe"}, "content": request.content} assert start.annotated_request is None - def test_sanitize_response_callable_error_omits_observability_output(self): + def test_sanitize_response_callable_error_preserves_observability_output(self): events = [] subscribers.register("py_llm_sanitize_resp_sub", lambda event: events.append(event)) guardrails.register_llm_sanitize_response( @@ -340,10 +347,10 @@ def test_sanitize_response_callable_error_omits_observability_output(self): subscribers.deregister("py_llm_sanitize_resp_sub") end = _llm_event(events, "llm_sanitize_resp_fail", "end") - assert end.data is None + assert end.data == {"ok": True} assert end.annotated_response is None - def test_sanitize_response_invalid_return_omits_observability_output(self): + def test_sanitize_response_invalid_return_preserves_observability_output(self): events = [] subscribers.register("py_llm_sanitize_resp_bad_sub", lambda event: events.append(event)) guardrails.register_llm_sanitize_response( @@ -362,7 +369,7 @@ def test_sanitize_response_invalid_return_omits_observability_output(self): subscribers.deregister("py_llm_sanitize_resp_bad_sub") end = _llm_event(events, "llm_sanitize_resp_bad", "end") - assert end.data is None + assert end.data == {"ok": True} assert end.annotated_response is None def test_sanitize_response_guardrail_accepts_scalar_json_payloads(self): @@ -398,7 +405,7 @@ def test_conditional_execution_invalid_return_type_raises(self): cast(guardrails.LlmConditionalExecutionGuardrail, lambda request: 123), ) try: - with pytest.raises(RuntimeError, match="expected str or None"): + with pytest.raises(RuntimeError, match="unexpected type"): llm.conditional_execution(make_request()) finally: guardrails.deregister_llm_conditional_execution("py_llm_cond_bad_type") @@ -410,7 +417,7 @@ def test_conditional_execution_callable_error_raises(self): lambda request: raise_runtime_error("boom"), ) try: - with pytest.raises(RuntimeError, match="callable failed"): + with pytest.raises(RuntimeError, match="RuntimeError: boom"): llm.conditional_execution(make_request()) finally: guardrails.deregister_llm_conditional_execution("py_llm_cond_error") diff --git a/python/tests/test_tools.py b/python/tests/test_tools.py index 3619902f3..a401eb73b 100644 --- a/python/tests/test_tools.py +++ b/python/tests/test_tools.py @@ -361,7 +361,7 @@ def test_duplicate_intercept_raises(self): def test_request_intercept_raises_on_exception(self): intercepts.register_tool_request("py_req_raise", 1, False, lambda n, a: raise_runtime_error("boom")) try: - with pytest.raises(RuntimeError, match="callable failed"): + with pytest.raises(RuntimeError, match="RuntimeError: boom"): tools.request_intercepts("raise_tool", {"value": 1}) finally: intercepts.deregister_tool_request("py_req_raise") @@ -374,7 +374,7 @@ def test_request_intercept_raises_on_unserializable_return(self): cast(intercepts.ToolRequestIntercept, lambda n, a: object()), ) try: - with pytest.raises(RuntimeError, match="py_to_json failed"): + with pytest.raises(RuntimeError, match="unsupported type object"): tools.request_intercepts("bad_return_tool", {"value": 1}) finally: intercepts.deregister_tool_request("py_req_bad_return") @@ -485,7 +485,7 @@ def test_conditional_execution_invalid_return_type_raises(self): cast(guardrails.ToolConditionalExecutionGuardrail, lambda name, args: 123), ) try: - with pytest.raises(RuntimeError, match="expected str or None"): + with pytest.raises(RuntimeError, match="unexpected type"): tools.conditional_execution("bad_type_tool", {}) finally: guardrails.deregister_tool_conditional_execution("py_cond_bad_type") @@ -497,7 +497,7 @@ def test_conditional_execution_callable_error_raises(self): lambda name, args: raise_runtime_error("boom"), ) try: - with pytest.raises(RuntimeError, match="callable failed"): + with pytest.raises(RuntimeError, match="RuntimeError: boom"): tools.conditional_execution("error_tool", {}) finally: guardrails.deregister_tool_conditional_execution("py_cond_error")