From b7651448a7aa5759831a02982fc0fecaa84d9e69 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Mon, 21 Sep 2026 15:28:05 -0400 Subject: [PATCH 01/22] enhancement: expose LLM codec context to workers Signed-off-by: Alex Fournier --- crates/core/src/api/llm.rs | 26 +- crates/core/src/api/registry.rs | 76 +- crates/core/src/api/runtime.rs | 5 + crates/core/src/api/runtime/callbacks.rs | 4 +- .../src/api/runtime/llm_execution_context.rs | 113 +++ crates/core/src/api/runtime/state.rs | 36 +- crates/core/src/context/registries.rs | 11 +- crates/core/src/plugin.rs | 104 +- crates/core/src/plugin/dynamic/worker.rs | 357 +++++-- .../tests/integration/worker_plugin_tests.rs | 249 ++++- .../core/tests/unit/dynamic_worker_tests.rs | 914 ++++++++++++++++-- crates/core/tests/unit/llm_api_tests.rs | 276 ++++++ crates/worker-proto/build.rs | 1 + .../nemo/relay/worker/v1/plugin_worker.proto | 11 + crates/worker-proto/tests/proto_tests.rs | 85 +- crates/worker/src/lib.rs | 212 +++- .../tests/unit/execution_context_tests.rs | 98 ++ crates/worker/tests/worker_sdk_tests.rs | 2 + justfile | 8 +- .../plugin/src/nemo_relay_plugin/__init__.py | 12 + python/plugin/src/nemo_relay_plugin/_api.py | 120 ++- .../plugin/test_public_api_docstrings.py | 2 + python/tests/plugin/test_worker_sdk.py | 89 ++ 23 files changed, 2568 insertions(+), 243 deletions(-) create mode 100644 crates/core/src/api/runtime/llm_execution_context.rs create mode 100644 crates/worker/tests/unit/execution_context_tests.rs diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index 21d2ddce5..d5d8a88b2 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -29,9 +29,9 @@ use crate::api::runtime::subscriber_dispatcher::{ dispatch_sanitized_event, dispatch_transformed_event, register_pending_publication, }; use crate::api::runtime::{ - EventSubscriberFn, LlmCollectorFn, LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, - LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn, - MiddlewareContinuationContext, + EventSubscriberFn, LlmCollectorFn, LlmExecutionCodecContext, LlmExecutionNextFn, + LlmFinalizerFn, LlmJsonStream, LlmSanitizeRequestContext, LlmSanitizeResponseContext, + LlmStreamExecutionNextFn, MiddlewareContinuationContext, }; use crate::api::runtime::{ScopeStackHandle, capture_trace_context, current_scope_stack}; use crate::api::scope::event; @@ -1739,6 +1739,7 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { ); let execution_name = name.clone(); let event_uuid = handle.uuid; + let execution_context = LlmExecutionCodecContext::for_codecs(request_codec, &response_codec); let execution = with_active_event_trace_context( event_uuid, Some(active_trace_context), @@ -1757,7 +1758,12 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { .read() .map_err(|error| FlowError::Internal(error.to_string()))? .registry_snapshot(&[RuntimeRegistrationKind::LlmExecutionIntercept]); - state.llm_build_execution_chain(&execution_name, func, &scope_local_refs) + state.llm_build_execution_chain( + &execution_name, + func, + &scope_local_refs, + execution_context, + ) }; execution(intercepted_request).await }), @@ -1966,6 +1972,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu let execution_name = name.clone(); let event_uuid = handle.uuid; let stream_started_at = Instant::now(); + let execution_context = LlmExecutionCodecContext::for_codecs(request_codec, &response_codec); let execution = with_active_event_trace_context( event_uuid, Some(active_trace_context), @@ -1984,12 +1991,17 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu .read() .map_err(|error| FlowError::Internal(error.to_string()))? .registry_snapshot(&[RuntimeRegistrationKind::LlmStreamExecutionIntercept]); - state.llm_stream_build_execution_chain(&execution_name, func, &scope_local_refs) + state.llm_stream_build_execution_chain( + &execution_name, + func, + &scope_local_refs, + execution_context, + ) }; - let execution_context = MiddlewareContinuationContext::capture(); + let continuation_context = MiddlewareContinuationContext::capture(); execution(intercepted_request) .await - .map(|stream| contextualize_stream(stream, execution_context)) + .map(|stream| contextualize_stream(stream, continuation_context)) }), ) .await; diff --git a/crates/core/src/api/registry.rs b/crates/core/src/api/registry.rs index 10f627e95..5d8e16f56 100644 --- a/crates/core/src/api/registry.rs +++ b/crates/core/src/api/registry.rs @@ -8,7 +8,10 @@ use crate::api::runtime::{ ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, LlmConditionalFn, LlmExecutionFn, LlmRequestInterceptFn, LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionFn, ToolConditionalFn, ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn, + adapt_llm_execution_fn, adapt_llm_stream_execution_fn, }; +#[cfg(feature = "worker-grpc")] +use crate::api::runtime::{ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn}; use crate::api::runtime::{current_scope_stack, global_context}; use crate::api::shared::ensure_runtime_owner; use crate::error::{FlowError, Result}; @@ -332,6 +335,15 @@ pub(crate) type Intercept = RegistryRecord>; /// A priority-ordered execution intercept registration record. pub(crate) type ExecutionIntercept = RegistryRecord; +macro_rules! adapt_execution_callable { + ($callable:expr) => { + $callable + }; + ($callable:expr, $adapter:path) => { + $adapter($callable) + }; +} + macro_rules! global_guardrail_registry_api { ( $(#[$register_meta:meta])* @@ -463,6 +475,7 @@ macro_rules! global_execution_registry_api { $deregister_name:ident, $field:ident, $fn_type:ty + $(, $adapter:path)? ) => { $(#[$register_meta])* /// @@ -485,7 +498,11 @@ macro_rules! global_execution_registry_api { .map_err(|error| FlowError::Internal(error.to_string()))?; state .$field - .register(ExecutionIntercept::new(name, priority, callable)) + .register(ExecutionIntercept::new( + name, + priority, + adapt_execution_callable!(callable $(, $adapter)?), + )) .map_err(FlowError::AlreadyExists) } @@ -660,6 +677,7 @@ macro_rules! scope_execution_registry_api { $deregister_name:ident, $field:ident, $fn_type:ty + $(, $adapter:path)? ) => { $(#[$register_meta])* /// @@ -690,7 +708,11 @@ macro_rules! scope_execution_registry_api { .ok_or_else(|| FlowError::NotFound(format!("scope {scope_uuid} not found")))?; registries .$field - .register(ExecutionIntercept::new(name, priority, callable)) + .register(ExecutionIntercept::new( + name, + priority, + adapt_execution_callable!(callable $(, $adapter)?), + )) .map_err(FlowError::AlreadyExists) } @@ -876,7 +898,8 @@ global_execution_registry_api!( /// Deregister a global LLM execution intercept. deregister_llm_execution_intercept, llm_execution_intercepts, - LlmExecutionFn + LlmExecutionFn, + adapt_llm_execution_fn ); global_execution_registry_api!( /// Register a global streaming LLM execution intercept. @@ -886,9 +909,48 @@ global_execution_registry_api!( /// Deregister a global streaming LLM execution intercept. deregister_llm_stream_execution_intercept, llm_stream_execution_intercepts, - LlmStreamExecutionFn + LlmStreamExecutionFn, + adapt_llm_stream_execution_fn ); +/// Register a global non-streaming LLM execution intercept that receives +/// Relay's private invocation context. +#[cfg(feature = "worker-grpc")] +pub(crate) fn register_contextual_llm_execution_intercept( + name: &str, + priority: i32, + callable: ContextualLlmExecutionFn, +) -> Result<()> { + ensure_runtime_owner()?; + let context = global_context(); + let mut state = context + .write() + .map_err(|error| FlowError::Internal(error.to_string()))?; + state + .llm_execution_intercepts + .register(ExecutionIntercept::new(name, priority, callable)) + .map_err(FlowError::AlreadyExists) +} + +/// Register a global streaming LLM execution intercept that receives Relay's +/// private invocation context. +#[cfg(feature = "worker-grpc")] +pub(crate) fn register_contextual_llm_stream_execution_intercept( + name: &str, + priority: i32, + callable: ContextualLlmStreamExecutionFn, +) -> Result<()> { + ensure_runtime_owner()?; + let context = global_context(); + let mut state = context + .write() + .map_err(|error| FlowError::Internal(error.to_string()))?; + state + .llm_stream_execution_intercepts + .register(ExecutionIntercept::new(name, priority, callable)) + .map_err(FlowError::AlreadyExists) +} + scope_guardrail_registry_api!( /// Register a scope-local mark event sanitizer. scope_register_mark_sanitize_guardrail, @@ -1044,7 +1106,8 @@ scope_execution_registry_api!( /// Deregister a scope-local LLM execution intercept. scope_deregister_llm_execution_intercept, llm_execution_intercepts, - LlmExecutionFn + LlmExecutionFn, + adapt_llm_execution_fn ); scope_execution_registry_api!( /// Register a scope-local streaming LLM execution intercept. @@ -1054,5 +1117,6 @@ scope_execution_registry_api!( /// Deregister a scope-local streaming LLM execution intercept. scope_deregister_llm_stream_execution_intercept, llm_stream_execution_intercepts, - LlmStreamExecutionFn + LlmStreamExecutionFn, + adapt_llm_stream_execution_fn ); diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index f0d085271..8bdeb8de6 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -6,6 +6,7 @@ pub mod callbacks; mod continuation_context; pub mod global; +mod llm_execution_context; pub mod scope_stack; pub mod state; pub mod subscriber_dispatcher; @@ -24,6 +25,10 @@ pub use continuation_context::MiddlewareContinuationContext; #[cfg(test)] pub(crate) use continuation_context::MiddlewareContinuationLease; pub use global::global_context; +pub(crate) use llm_execution_context::{ + ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn, LlmExecutionCodecContext, + adapt_llm_execution_fn, adapt_llm_stream_execution_fn, +}; pub(crate) use scope_stack::capture_trace_context; pub use scope_stack::{ PropagationContext, ScopeStack, ScopeStackHandle, TASK_SCOPE_STACK, ThreadScopeStackBinding, diff --git a/crates/core/src/api/runtime/callbacks.rs b/crates/core/src/api/runtime/callbacks.rs index efc6795aa..a583b325c 100644 --- a/crates/core/src/api/runtime/callbacks.rs +++ b/crates/core/src/api/runtime/callbacks.rs @@ -628,7 +628,9 @@ pub type LlmFinalizerFn = Box Json + Send>; /// # Returns /// A shared reference to a scope-local streaming execution registry. pub(crate) type LlmStreamExecutionRegistryRef<'a> = &'a crate::registry::SortedRegistry< - crate::api::registry::ExecutionIntercept, + crate::api::registry::ExecutionIntercept< + super::llm_execution_context::ContextualLlmStreamExecutionFn, + >, >; /// Slice of scope-local streaming execution registries. /// diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs new file mode 100644 index 000000000..5c5a730d6 --- /dev/null +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -0,0 +1,113 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Invocation-scoped codec context for internal LLM execution adapters. + +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + +use crate::api::llm::LlmRequest; +use crate::codec::traits::{LlmCodec, LlmResponseCodec}; +use crate::error::Result; +use crate::json::Json; + +use super::callbacks::{ + LlmExecutionFn, LlmExecutionNextFn, LlmJsonStream, LlmStreamExecutionFn, + LlmStreamExecutionNextFn, +}; +#[cfg(feature = "worker-grpc")] +use super::callbacks::{LlmSanitizeRequestContext, LlmSanitizeResponseContext}; + +/// Active request and response codecs for one managed LLM execution. +/// +/// Public execution-interceptor callbacks remain unchanged. Relay adapts them +/// into one private context-aware callback shape so language and process +/// bridges receive the active codecs explicitly. +/// +/// The context describes the codecs selected when the managed invocation was +/// created. Execution interceptors may rewrite payloads within that codec's +/// contract, but changing the provider wire format does not select a new codec. +/// A subsequent decode or encode will reject an incompatible payload rather +/// than silently infer another codec. +#[derive(Clone)] +pub(crate) struct LlmExecutionCodecContext { + #[cfg(feature = "worker-grpc")] + request: LlmSanitizeRequestContext, + #[cfg(feature = "worker-grpc")] + response: LlmSanitizeResponseContext, +} + +impl LlmExecutionCodecContext { + #[cfg(feature = "worker-grpc")] + pub(crate) fn new( + request: LlmSanitizeRequestContext, + response: LlmSanitizeResponseContext, + ) -> Self { + Self { request, response } + } + + pub(crate) fn for_codecs( + request_codec: Option>, + response_codec: &Option>, + ) -> Self { + #[cfg(feature = "worker-grpc")] + { + Self::new( + LlmSanitizeRequestContext::for_request_codec(request_codec), + LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), + ) + } + #[cfg(not(feature = "worker-grpc"))] + { + let _ = (request_codec, response_codec); + Self {} + } + } + + #[cfg(feature = "worker-grpc")] + pub(crate) fn request(&self) -> &LlmSanitizeRequestContext { + &self.request + } + + #[cfg(feature = "worker-grpc")] + pub(crate) fn response(&self) -> &LlmSanitizeResponseContext { + &self.response + } +} + +/// Private non-streaming execution callback used by Relay's registries. +pub(crate) type ContextualLlmExecutionFn = Arc< + dyn Fn( + &str, + LlmRequest, + LlmExecutionCodecContext, + LlmExecutionNextFn, + ) -> Pin> + Send>> + + Send + + Sync, +>; + +/// Private streaming execution callback used by Relay's registries. +pub(crate) type ContextualLlmStreamExecutionFn = Arc< + dyn Fn( + &str, + LlmRequest, + LlmExecutionCodecContext, + LlmStreamExecutionNextFn, + ) -> Pin> + Send>> + + Send + + Sync, +>; + +/// Adapt the stable public callback into Relay's private context-aware shape. +pub(crate) fn adapt_llm_execution_fn(callback: LlmExecutionFn) -> ContextualLlmExecutionFn { + Arc::new(move |name, request, _context, next| callback(name, request, next)) +} + +/// Adapt the stable public stream callback into Relay's private context-aware shape. +pub(crate) fn adapt_llm_stream_execution_fn( + callback: LlmStreamExecutionFn, +) -> ContextualLlmStreamExecutionFn { + Arc::new(move |name, request, _context, next| callback(name, request, next)) +} diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index 59d9c2a7b..4825a45af 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -32,17 +32,20 @@ use crate::api::registry::{ }; use crate::api::runtime::ScopeStackHandle; use crate::api::runtime::callbacks::{ - EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionFn, LlmExecutionNextFn, - LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, - LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, - LlmStreamExecutionNextFn, LlmStreamExecutionRegistryRefs, LlmStreamInner, ToolConditionalFn, - ToolExecutionContext, ToolExecutionFn, ToolExecutionNextFn, ToolExecutionOutcomeNextFn, - ToolInterceptFn, ToolSanitizeFn, + EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionNextFn, LlmJsonStream, + LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, + LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionNextFn, + LlmStreamExecutionRegistryRefs, LlmStreamInner, ToolConditionalFn, ToolExecutionContext, + ToolExecutionFn, ToolExecutionNextFn, ToolExecutionOutcomeNextFn, ToolInterceptFn, + ToolSanitizeFn, }; use crate::api::runtime::continuation_context::{ MiddlewareContinuationContext, MiddlewareContinuationGuard, MiddlewareContinuationLease, }; use crate::api::runtime::subscriber_dispatcher; +use crate::api::runtime::{ + ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn, LlmExecutionCodecContext, +}; use crate::api::scope::{CreateScopeHandleParams, EndScopeHandleParams, ScopeHandle, ScopeType}; use crate::api::shared::snapshot_event_sanitizers; use crate::api::tool::ToolHandle; @@ -243,10 +246,11 @@ pub struct NemoRelayContextState { /// Global LLM request intercepts that can rewrite or annotate requests. pub(crate) llm_request_intercepts: SortedRegistry>, /// Global non-streaming LLM execution intercepts that wrap callback execution. - pub(crate) llm_execution_intercepts: SortedRegistry>, + pub(crate) llm_execution_intercepts: + SortedRegistry>, /// Global streaming LLM execution intercepts that wrap stream-producing callbacks. pub(crate) llm_stream_execution_intercepts: - SortedRegistry>, + SortedRegistry>, /// Global lifecycle subscribers notified after runtime events are emitted. pub(crate) event_subscribers: HashMap, /// Whether LLM start events retain complete sanitized request payloads. @@ -1799,6 +1803,8 @@ impl NemoRelayContextState { /// intercepts. /// - `scope_locals`: Scope-local execution intercept registries collected /// from the active scope stack. + /// - `execution_context`: Invocation-scoped identities and capabilities for + /// the request and response codecs selected by the managed call. /// /// # Returns /// A composed [`LlmExecutionNextFn`] that wraps `default_fn` in every @@ -1807,7 +1813,8 @@ impl NemoRelayContextState { &self, name: &str, default_fn: LlmExecutionNextFn, - scope_locals: &[&SortedRegistry>], + scope_locals: &[&SortedRegistry>], + execution_context: LlmExecutionCodecContext, ) -> LlmExecutionNextFn { let matching = merge_execution_intercept_callables( &self.llm_execution_intercepts, @@ -1819,10 +1826,12 @@ impl NemoRelayContextState { for (callable, _) in matching.into_iter().rev() { let current_next = next.clone(); let current_name = name.clone(); + let current_context = execution_context.clone(); next = Arc::new(move |request| { let callable = callable.clone(); let current_next = current_next.clone(); let current_name = current_name.clone(); + let current_context = current_context.clone(); Box::pin(async move { let (continuation, continuation_guard) = MiddlewareContinuationLease::capture(); let raw_next: LlmExecutionNextFn = Arc::new(move |request| { @@ -1832,7 +1841,7 @@ impl NemoRelayContextState { async move { invocation?.invoke(move || current_next(request)).await }, ) }); - let result = callable(¤t_name, request, raw_next).await; + let result = callable(¤t_name, request, current_context, raw_next).await; drop(continuation_guard); result }) @@ -1850,6 +1859,8 @@ impl NemoRelayContextState { /// intercepts. /// - `scope_locals`: Scope-local execution intercept registries collected /// from the active scope stack. + /// - `execution_context`: Invocation-scoped identities and capabilities for + /// the request and response codecs selected by the managed call. /// /// # Returns /// A composed [`LlmStreamExecutionNextFn`] that wraps `default_fn` in every @@ -1859,6 +1870,7 @@ impl NemoRelayContextState { name: &str, default_fn: LlmStreamExecutionNextFn, scope_locals: LlmStreamExecutionRegistryRefs<'_>, + execution_context: LlmExecutionCodecContext, ) -> LlmStreamExecutionNextFn { let matching = merge_execution_intercept_callables( &self.llm_stream_execution_intercepts, @@ -1870,10 +1882,12 @@ impl NemoRelayContextState { for (callable, _) in matching.into_iter().rev() { let current_next = next.clone(); let current_name = name.clone(); + let current_context = execution_context.clone(); next = Arc::new(move |request| { let callable = callable.clone(); let current_next = current_next.clone(); let current_name = current_name.clone(); + let current_context = current_context.clone(); Box::pin(async move { let (continuation, continuation_guard) = MiddlewareContinuationLease::capture(); let raw_next: LlmStreamExecutionNextFn = Arc::new(move |request| { @@ -1886,7 +1900,7 @@ impl NemoRelayContextState { Ok(contextualize_stream(stream, context)) }) }); - let result = callable(¤t_name, request, raw_next).await; + let result = callable(¤t_name, request, current_context, raw_next).await; result.map(|stream| guard_stream_continuation(stream, continuation_guard)) }) }); diff --git a/crates/core/src/context/registries.rs b/crates/core/src/context/registries.rs index 32991a24d..5934fdf65 100644 --- a/crates/core/src/context/registries.rs +++ b/crates/core/src/context/registries.rs @@ -14,9 +14,9 @@ use crate::api::registry::{ runtime_registration_is_enabled, }; use crate::api::runtime::{ - EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionFn, LlmRequestInterceptFn, - LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionFn, ToolConditionalFn, - ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn, + ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn, EventSanitizeFn, EventSubscriberFn, + LlmConditionalFn, LlmRequestInterceptFn, LlmSanitizeRequestFn, LlmSanitizeResponseFn, + ToolConditionalFn, ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn, }; use crate::registry::SortedRegistry; @@ -55,10 +55,11 @@ pub(crate) struct ScopeLocalRegistries { /// LLM request intercepts that can rewrite or annotate requests. pub(crate) llm_request_intercepts: SortedRegistry>, /// Non-streaming LLM execution intercepts that wrap callback execution. - pub(crate) llm_execution_intercepts: SortedRegistry>, + pub(crate) llm_execution_intercepts: + SortedRegistry>, /// Streaming LLM execution intercepts that wrap stream-producing callbacks. pub(crate) llm_stream_execution_intercepts: - SortedRegistry>, + SortedRegistry>, /// Scope-local lifecycle subscribers visible while the owning scope is active. pub(crate) event_subscribers: HashMap, } diff --git a/crates/core/src/plugin.rs b/crates/core/src/plugin.rs index c79188d97..cc2f85e94 100644 --- a/crates/core/src/plugin.rs +++ b/crates/core/src/plugin.rs @@ -41,12 +41,18 @@ use crate::api::registry::{ register_tool_request_intercept, register_tool_sanitize_request_guardrail, register_tool_sanitize_response_guardrail, }; +#[cfg(feature = "worker-grpc")] +use crate::api::registry::{ + register_contextual_llm_execution_intercept, register_contextual_llm_stream_execution_intercept, +}; use crate::api::runtime::{ ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionFn, LlmRequestInterceptFn, LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionFn, ToolConditionalFn, ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn, }; +#[cfg(feature = "worker-grpc")] +use crate::api::runtime::{ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn}; use crate::api::subscriber::{deregister_subscriber, register_subscriber}; pub use nemo_relay_types::plugin::{ConfigDiagnostic, DiagnosticLevel}; @@ -876,26 +882,33 @@ impl PluginRegistrationContext { priority: i32, callback: LlmExecutionFn, ) -> Result<()> { - let qualified_name = self.qualify_name(name); - register_llm_execution_intercept(&qualified_name, priority, callback).map_err(|err| { - PluginError::RegistrationFailed(format!("llm execution intercept: {err}")) - })?; + self.register_execution_intercept( + name, + priority, + callback, + "llm execution intercept", + register_llm_execution_intercept, + deregister_llm_execution_intercept, + ) + } - let name_owned = qualified_name; - self.registrations.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_llm_execution_intercept(&name_owned) - .map(|_| ()) - .map_err(|err| { - PluginError::RegistrationFailed(format!( - "llm execution intercept deregistration failed: {err}" - )) - }) - }), - )); - Ok(()) + /// Registers an internal context-aware LLM execution intercept and records + /// its rollback closure. + #[cfg(feature = "worker-grpc")] + pub(crate) fn register_contextual_llm_execution_intercept( + &mut self, + name: &str, + priority: i32, + callback: ContextualLlmExecutionFn, + ) -> Result<()> { + self.register_execution_intercept( + name, + priority, + callback, + "llm execution intercept", + register_contextual_llm_execution_intercept, + deregister_llm_execution_intercept, + ) } /// Registers an LLM stream execution intercept and records its rollback closure. @@ -904,24 +917,57 @@ impl PluginRegistrationContext { name: &str, priority: i32, callback: LlmStreamExecutionFn, + ) -> Result<()> { + self.register_execution_intercept( + name, + priority, + callback, + "llm stream execution intercept", + register_llm_stream_execution_intercept, + deregister_llm_stream_execution_intercept, + ) + } + + /// Registers an internal context-aware streaming LLM execution intercept + /// and records its rollback closure. + #[cfg(feature = "worker-grpc")] + pub(crate) fn register_contextual_llm_stream_execution_intercept( + &mut self, + name: &str, + priority: i32, + callback: ContextualLlmStreamExecutionFn, + ) -> Result<()> { + self.register_execution_intercept( + name, + priority, + callback, + "llm stream execution intercept", + register_contextual_llm_stream_execution_intercept, + deregister_llm_stream_execution_intercept, + ) + } + + fn register_execution_intercept( + &mut self, + name: &str, + priority: i32, + callback: F, + kind: &'static str, + register: fn(&str, i32, F) -> crate::error::Result<()>, + deregister: fn(&str) -> crate::error::Result, ) -> Result<()> { let qualified_name = self.qualify_name(name); - register_llm_stream_execution_intercept(&qualified_name, priority, callback).map_err( - |err| PluginError::RegistrationFailed(format!("llm stream execution intercept: {err}")), - )?; + register(&qualified_name, priority, callback) + .map_err(|err| PluginError::RegistrationFailed(format!("{kind}: {err}")))?; let name_owned = qualified_name; self.registrations.push(PluginRegistration::new( "plugin", name_owned.clone(), Box::new(move || { - deregister_llm_stream_execution_intercept(&name_owned) - .map(|_| ()) - .map_err(|err| { - PluginError::RegistrationFailed(format!( - "llm stream execution intercept deregistration failed: {err}" - )) - }) + deregister(&name_owned).map(|_| ()).map_err(|err| { + PluginError::RegistrationFailed(format!("{kind} deregistration failed: {err}")) + }) }), )); Ok(()) diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index 622a379d9..ae65a49a3 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -11,6 +11,7 @@ use std::pin::Pin; use std::process::{Child, Command, Stdio}; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Condvar, Mutex}; +use std::task::{Context, Poll}; use std::time::Duration; use futures_util::FutureExt; @@ -26,7 +27,8 @@ use nemo_relay_worker_proto::v1::{ HandshakeResponse, HealthRequest, HostAck, InvokeRequest, InvokeResponse, JsonEnvelope, JsonResult, JsonValue, ListRuntimeRegistrationsRequest, ListRuntimeRegistrationsResponse, LlmCodecDecodeRequest, LlmCodecDecodeResponse, LlmCodecEncodeRequest, - LlmCodecIdentity as ProtoLlmCodecIdentity, LlmCodecKind, LlmInvocation, LlmNextRequest, + LlmCodecIdentity as ProtoLlmCodecIdentity, LlmCodecKind, + LlmExecutionCodecContext as ProtoLlmExecutionCodecContext, LlmInvocation, LlmNextRequest, LlmSanitizeRequestContext as ProtoLlmSanitizeRequestContext, LlmSanitizeResponseContext as ProtoLlmSanitizeResponseContext, LlmStreamNextRequest, LogLevel, LogRequest, PopScopeRequest, PushScopeRequest, PushScopeResponse, @@ -47,8 +49,8 @@ use nemo_relay_worker_proto::{ use serde_json::{Map, Value as Json}; use sha2::{Digest, Sha256}; use tokio::runtime::{Builder as RuntimeBuilder, Runtime}; -use tokio::sync::{mpsc, oneshot}; -use tokio_stream::StreamExt; +use tokio::sync::{mpsc, oneshot, watch}; +use tokio_stream::{Stream, StreamExt}; use tonic::transport::{Channel, Endpoint, Server}; use tonic::{Request, Response, Status}; use uuid::Uuid; @@ -81,10 +83,10 @@ use crate::api::runtime::subscriber_dispatcher::{ PublicationBuffer, capture_nested_publication_buffer, with_nested_publication_buffer, }; use crate::api::runtime::{ - EventMetadataInjectorFn, EventSanitizeFn, LlmCodecIdentity, LlmExecutionNextFn, LlmJsonStream, - LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn, - MiddlewareContinuationContext, ToolExecutionContext, ToolExecutionNextFn, current_scope_stack, - with_scope_stack, + EventMetadataInjectorFn, EventSanitizeFn, LlmCodecIdentity, LlmExecutionCodecContext, + LlmExecutionNextFn, LlmJsonStream, LlmSanitizeRequestContext, LlmSanitizeResponseContext, + LlmStreamExecutionNextFn, LlmStreamInner, MiddlewareContinuationContext, ToolExecutionContext, + ToolExecutionNextFn, current_scope_stack, with_scope_stack, }; use crate::api::scope::{ EmitMarkEventParams, PopScopeParams, PushScopeParams, ScopeAttributes, ScopeHandle, ScopeType, @@ -1411,6 +1413,7 @@ impl WorkerPluginInstance { ) -> crate::plugin::Result<()> { let name = registration.local_name.as_str(); let priority = registration.priority; + let include_codec_context = registration.llm_execution_codec_context; let instance = Arc::new(self.clone_for_callback()); let callback_name = name.to_owned(); match surface { @@ -1473,25 +1476,32 @@ impl WorkerPluginInstance { }) }), ), - RegistrationSurface::LlmExecutionIntercept => ctx.register_llm_execution_intercept( - name, - priority, - Arc::new(move |model_name, request, next| { - let instance = instance.clone(); - let callback_name = callback_name.clone(); - let model_name = model_name.to_owned(); - Box::pin(async move { - instance - .invoke_llm_execution(&callback_name, &model_name, request, next) - .await - }) - }), - ), + RegistrationSurface::LlmExecutionIntercept => ctx + .register_contextual_llm_execution_intercept( + name, + priority, + Arc::new(move |model_name, request, context, next| { + let instance = instance.clone(); + let callback_name = callback_name.clone(); + let model_name = model_name.to_owned(); + Box::pin(async move { + instance + .invoke_llm_execution( + &callback_name, + &model_name, + request, + include_codec_context.then_some(context), + next, + ) + .await + }) + }), + ), RegistrationSurface::LlmStreamExecutionIntercept => ctx - .register_llm_stream_execution_intercept( + .register_contextual_llm_stream_execution_intercept( name, priority, - Arc::new(move |model_name, request, next| { + Arc::new(move |model_name, request, context, next| { let instance = instance.clone(); let callback_name = callback_name.clone(); let model_name = model_name.to_owned(); @@ -1501,6 +1511,7 @@ impl WorkerPluginInstance { &callback_name, &model_name, request, + include_codec_context.then_some(context), next, ) .await @@ -1675,6 +1686,48 @@ impl Drop for WorkerInvocationGuard { } } +struct WorkerStreamCompletionSignal(watch::Sender); + +impl Drop for WorkerStreamCompletionSignal { + fn drop(&mut self) { + self.0.send_replace(true); + } +} + +struct WorkerForwardedLlmStream { + receiver: Option>>, + completion: watch::Receiver, +} + +impl Stream for WorkerForwardedLlmStream { + type Item = FlowResult; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + match this.receiver.as_mut() { + Some(receiver) => Pin::new(receiver).poll_next(cx), + None => Poll::Ready(None), + } + } +} + +impl LlmStreamInner for WorkerForwardedLlmStream { + fn close(self: Pin<&mut Self>) -> Pin> + Send + '_>> { + let this = self.get_mut(); + this.receiver.take(); + // Keep the stored observer reusable if this close future is cancelled. + let mut completion = this.completion.clone(); + Box::pin(async move { + while !*completion.borrow_and_update() { + if completion.changed().await.is_err() { + break; + } + } + Ok(()) + }) + } +} + impl WorkerPluginCallback { fn invoke_conditional_middleware( &self, @@ -1849,24 +1902,26 @@ impl WorkerPluginCallback { ), )), ); - let capability_id = context.resolve_codec().map(|codec| { - let capability_id = self - .host_state - .insert_request_codec(&invoke.invocation_id, codec); - let Some(invoke_request_payload::Payload::Llm(llm)) = invoke.payload.as_mut() else { - unreachable!("LLM sanitizer invocation must have an LLM payload"); - }; - let Some(llm_invocation::SanitizeContext::RequestSanitizeContext(context)) = - llm.sanitize_context.as_mut() - else { - unreachable!("request sanitizer invocation must have a request context"); - }; - context.codec_capability_id = Some(capability_id.clone()); - capability_id - }); - let _capability_guard = capability_id.as_ref().map(|capability_id| { - WorkerCodecCapabilityGuard::new(Arc::clone(&self.host_state), capability_id.clone()) - }); + let capability = context + .resolve_codec() + .map(|codec| -> FlowResult { + let capability = self + .host_state + .issue_request_codec(&invoke.invocation_id, codec)?; + let Some(invoke_request_payload::Payload::Llm(llm)) = invoke.payload.as_mut() + else { + unreachable!("LLM sanitizer invocation must have an LLM payload"); + }; + let Some(llm_invocation::SanitizeContext::RequestSanitizeContext(context)) = + llm.sanitize_context.as_mut() + else { + unreachable!("request sanitizer invocation must have a request context"); + }; + context.codec_capability_id = Some(capability.id().into()); + Ok(capability) + }) + .transpose(); + let _capability = self.cleanup_after_setup_error(&invoke, capability)?; let response = self.invoke_async(invoke).await; optional_json_from_invoke_response(response?)? .map(serde_json::from_value) @@ -1899,24 +1954,26 @@ impl WorkerPluginCallback { ), )), ); - let capability_id = context.resolve_codec().map(|codec| { - let capability_id = self - .host_state - .insert_response_codec(&invoke.invocation_id, codec); - let Some(invoke_request_payload::Payload::Llm(llm)) = invoke.payload.as_mut() else { - unreachable!("LLM sanitizer invocation must have an LLM payload"); - }; - let Some(llm_invocation::SanitizeContext::ResponseSanitizeContext(context)) = - llm.sanitize_context.as_mut() - else { - unreachable!("response sanitizer invocation must have a response context"); - }; - context.codec_capability_id = Some(capability_id.clone()); - capability_id - }); - let _capability_guard = capability_id.as_ref().map(|capability_id| { - WorkerCodecCapabilityGuard::new(Arc::clone(&self.host_state), capability_id.clone()) - }); + let capability = context + .resolve_codec() + .map(|codec| -> FlowResult { + let capability = self + .host_state + .issue_response_codec(&invoke.invocation_id, codec)?; + let Some(invoke_request_payload::Payload::Llm(llm)) = invoke.payload.as_mut() + else { + unreachable!("LLM sanitizer invocation must have an LLM payload"); + }; + let Some(llm_invocation::SanitizeContext::ResponseSanitizeContext(context)) = + llm.sanitize_context.as_mut() + else { + unreachable!("response sanitizer invocation must have a response context"); + }; + context.codec_capability_id = Some(capability.id().into()); + Ok(capability) + }) + .transpose(); + let _capability = self.cleanup_after_setup_error(&invoke, capability)?; let response = self.invoke_async(invoke).await; optional_json_from_invoke_response(response?) } @@ -1981,12 +2038,13 @@ impl WorkerPluginCallback { registration_name: &str, model_name: &str, request: LlmRequest, + execution_context: Option, next: LlmExecutionNextFn, ) -> FlowResult { let continuation_id = self .host_state .insert_continuation(Continuation::llm(next))?; - let invoke = self.base_request( + let mut invoke = self.base_request( registration_name, RegistrationSurface::LlmExecutionIntercept, Some(continuation_id), @@ -1997,6 +2055,13 @@ impl WorkerPluginCallback { None, )), ); + let codec_capabilities = execution_context + .as_ref() + .map(|context| self.attach_llm_execution_codec_context(&mut invoke, context, true)) + .transpose(); + let _codec_capabilities = self + .cleanup_after_setup_error(&invoke, codec_capabilities)? + .unwrap_or_default(); json_from_invoke_response(self.invoke_async(invoke).await?) } @@ -2005,12 +2070,13 @@ impl WorkerPluginCallback { registration_name: &str, model_name: &str, request: LlmRequest, + execution_context: Option, next: LlmStreamExecutionNextFn, ) -> FlowResult { let continuation_id = self .host_state .insert_continuation(Continuation::llm_stream(next))?; - let invoke = self.base_request( + let mut invoke = self.base_request( registration_name, RegistrationSurface::LlmStreamExecutionIntercept, Some(continuation_id.clone()), @@ -2021,13 +2087,27 @@ impl WorkerPluginCallback { None, )), ); + let codec_capabilities = execution_context + .as_ref() + .map(|context| self.attach_llm_execution_codec_context(&mut invoke, context, false)) + .transpose(); + let codec_capabilities = self + .cleanup_after_setup_error(&invoke, codec_capabilities)? + .unwrap_or_default(); let mut client = self.client.clone(); let mut guard = WorkerInvocationGuard::new(self, &invoke); let (tx, rx) = mpsc::channel(16); let (next_ready_tx, next_ready_rx) = oneshot::channel(); + let (completion_tx, completion_rx) = watch::channel(false); self.runtime.spawn(async move { + let _completion = WorkerStreamCompletionSignal(completion_tx); + let _codec_capabilities = codec_capabilities; let result = tokio::select! { - result = worker_rpc(client.invoke_stream(worker_rpc_request(invoke))) => result, + // Stream setup can include awaiting the downstream provider through + // `next`, so the control-plane timeout must not cap it. Dropping the + // caller closes `rx`; the sibling branch then cancels the worker and + // releases the continuation and codec capabilities. + result = client.invoke_stream(worker_rpc_request(invoke)) => result, _ = tx.closed() => { guard.cancel("host stopped consuming the worker stream"); guard.finish(); @@ -2083,9 +2163,63 @@ impl WorkerPluginCallback { "worker stream invocation ended before the downstream stream opened".into(), ) })?; - Ok(LlmJsonStream::new( - tokio_stream::wrappers::ReceiverStream::new(rx), - )) + Ok(LlmJsonStream::from_closeable(WorkerForwardedLlmStream { + receiver: Some(tokio_stream::wrappers::ReceiverStream::new(rx)), + completion: completion_rx, + })) + } + + fn attach_llm_execution_codec_context( + &self, + invoke: &mut InvokeRequest, + context: &LlmExecutionCodecContext, + include_response_capability: bool, + ) -> FlowResult> { + let mut guards = Vec::with_capacity(2); + + let mut request = ProtoLlmSanitizeRequestContext { + codec: Some(codec_identity_to_proto(context.request().codec())), + codec_capability_id: None, + }; + if let Some(codec) = context.request().resolve_codec() { + let capability = self + .host_state + .issue_request_codec(&invoke.invocation_id, codec)?; + request.codec_capability_id = Some(capability.id().into()); + guards.push(capability); + } + + let mut response = ProtoLlmSanitizeResponseContext { + codec: Some(codec_identity_to_proto(context.response().codec())), + codec_capability_id: None, + }; + if include_response_capability && let Some(codec) = context.response().resolve_codec() { + let capability = self + .host_state + .issue_response_codec(&invoke.invocation_id, codec)?; + response.codec_capability_id = Some(capability.id().into()); + guards.push(capability); + } + + let Some(invoke_request_payload::Payload::Llm(llm)) = invoke.payload.as_mut() else { + unreachable!("LLM execution invocation must have an LLM payload"); + }; + llm.execution_codec_context = Some(Box::new(ProtoLlmExecutionCodecContext { + request: Some(request), + response: Some(response), + })); + Ok(guards) + } + + fn cleanup_after_setup_error( + &self, + invoke: &InvokeRequest, + result: FlowResult, + ) -> FlowResult { + result.inspect_err(|_| { + let mut guard = WorkerInvocationGuard::new(self, invoke); + guard.finish(); + }) } fn base_request( @@ -2121,9 +2255,15 @@ impl WorkerPluginCallback { async fn invoke_async(&self, request: InvokeRequest) -> FlowResult { let callback_name = request.registration_name.clone(); let surface = request.surface; - let result = self - .invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) - .await; + // A continuation-bearing callback can legitimately include downstream + // provider latency. The caller still owns cancellation through the + // invocation guard, but the control-plane timeout must not cap `next`. + let result = if request.continuation_id.is_empty() { + self.invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) + .await + } else { + self.invoke_async_without_timeout(request).await + }; if let Err(error) = &result { let surface_name = RegistrationSurface::try_from(surface) .map(|surface| surface.as_str_name()) @@ -2140,6 +2280,19 @@ impl WorkerPluginCallback { result } + async fn invoke_async_without_timeout( + &self, + request: InvokeRequest, + ) -> FlowResult { + let mut guard = WorkerInvocationGuard::new(self, &request); + let mut client = self.client.clone(); + let result = client.invoke(worker_rpc_request(request)).await; + guard.finish(); + result + .map(|response| response.into_inner()) + .map_err(|err| worker_status_to_flow("worker invoke failed", err)) + } + async fn invoke_async_with_timeout( &self, request: InvokeRequest, @@ -2339,11 +2492,8 @@ struct WorkerCodecCapabilityGuard { } impl WorkerCodecCapabilityGuard { - fn new(host_state: Arc, capability_id: String) -> Self { - Self { - host_state, - capability_id, - } + fn id(&self) -> &str { + &self.capability_id } } @@ -2523,30 +2673,43 @@ impl WorkerHostRuntimeState { Ok(handle) } - fn insert_request_codec(&self, invocation_id: &str, codec: Arc) -> String { - self.insert_codec(invocation_id, WorkerCodecDirection::Request(codec)) + fn issue_request_codec( + self: &Arc, + invocation_id: &str, + codec: Arc, + ) -> FlowResult { + self.issue_codec(invocation_id, WorkerCodecDirection::Request(codec)) } - fn insert_response_codec( - &self, + fn issue_response_codec( + self: &Arc, invocation_id: &str, codec: Arc, - ) -> String { - self.insert_codec(invocation_id, WorkerCodecDirection::Response(codec)) + ) -> FlowResult { + self.issue_codec(invocation_id, WorkerCodecDirection::Response(codec)) } - fn insert_codec(&self, invocation_id: &str, direction: WorkerCodecDirection) -> String { + fn issue_codec( + self: &Arc, + invocation_id: &str, + direction: WorkerCodecDirection, + ) -> FlowResult { let id = format!("codec-{}", Uuid::now_v7()); - if let Ok(mut codecs) = self.codecs.lock() { - codecs.insert( - id.clone(), - WorkerCodecCapability { - invocation_id: invocation_id.to_owned(), - direction, - }, - ); - } - id + let mut codecs = self + .codecs + .lock() + .map_err(|error| FlowError::Internal(format!("codec lock poisoned: {error}")))?; + codecs.insert( + id.clone(), + WorkerCodecCapability { + invocation_id: invocation_id.to_owned(), + direction, + }, + ); + Ok(WorkerCodecCapabilityGuard { + host_state: Arc::clone(self), + capability_id: id, + }) } fn remove_codec(&self, id: &str) { @@ -3611,6 +3774,7 @@ fn invoke_request_payload_llm( .as_ref() .map(|response| json_envelope_infallible(JSON_SCHEMA, response)), sanitize_context: None, + execution_codec_context: None, }) } @@ -3633,6 +3797,7 @@ fn invoke_request_payload_llm_context( .as_ref() .map(|response| json_envelope_infallible(JSON_SCHEMA, response)), sanitize_context: Some(context.into()), + execution_codec_context: None, }) } @@ -3943,6 +4108,18 @@ fn validate_registration_plan( "worker plugin '{plugin_id}' returned unspecified registration surface" ))); } + if registration.llm_execution_codec_context + && !matches!( + surface, + RegistrationSurface::LlmExecutionIntercept + | RegistrationSurface::LlmStreamExecutionIntercept + ) + { + return Err(PluginError::RegistrationFailed(format!( + "worker plugin '{plugin_id}' requested LLM execution codec context for incompatible surface {}", + surface.as_str_name() + ))); + } } let mut gate_names = std::collections::HashSet::new(); for gate in &response.conditional_middleware_guardrails { diff --git a/crates/core/tests/integration/worker_plugin_tests.rs b/crates/core/tests/integration/worker_plugin_tests.rs index c5a842d0f..532c71518 100644 --- a/crates/core/tests/integration/worker_plugin_tests.rs +++ b/crates/core/tests/integration/worker_plugin_tests.rs @@ -14,6 +14,7 @@ use nemo_relay::api::llm::{ LlmCallExecuteParams, LlmRequest, LlmStreamCallExecuteParams, llm_call_execute, llm_stream_call_execute, }; +use nemo_relay::api::runtime::LlmCodecIdentity; use nemo_relay::api::runtime::{LlmJsonStream, TASK_SCOPE_STACK, create_scope_stack}; use nemo_relay::api::scope::{ EmitMarkEventParams, PopScopeParams, PushScopeParams, ScopeType, event, pop_scope, push_scope, @@ -22,8 +23,10 @@ use nemo_relay::api::subscriber::{deregister_subscriber, flush_subscribers, regi use nemo_relay::api::tool::{ ToolCallExecuteParams, ToolExecutionResult, tool_call_execute, tool_request_intercepts, }; +use nemo_relay::codec::openai_chat::OpenAIChatCodec; use nemo_relay::codec::request::AnnotatedLlmRequest; -use nemo_relay::codec::traits::LlmCodec; +use nemo_relay::codec::response::AnnotatedLlmResponse; +use nemo_relay::codec::traits::{LlmCodec, LlmResponseCodec}; use nemo_relay::error::Result as FlowResult; use nemo_relay::observability::otel_logs::{OpenTelemetryLogConfig, OpenTelemetryLogSubscriber}; use nemo_relay::observability::otel_metrics::{ @@ -1581,28 +1584,9 @@ async fn python_worker_host_runtime_mark_and_mutated_request_round_trip() { let manifest_ref = PathBuf::from(env!("CARGO_MANIFEST_DIR")) .join("../../examples/python-grpc-worker-plugin/relay-plugin.toml"); let config = Map::from_iter([("tag".into(), json!("managed-environment"))]); - let activation = load_worker_plugins([WorkerPluginLoadSpec { - plugin_id: "examples.python_grpc_worker".into(), - manifest_ref: manifest_ref.to_string_lossy().into_owned(), - environment_ref: Some( - PathBuf::from(environment_ref) - .to_string_lossy() - .into_owned(), - ), - config: config.clone(), - }]) - .expect("managed Python worker should load"); - let mut cleanup = PythonWorkerCleanup::new(activation); - - let mut plugin_config = PluginConfig::default(); - plugin_config.components.push(PluginComponentSpec { - kind: "examples.python_grpc_worker".into(), - enabled: true, - config, - }); - test_initialize_plugin_host_exact(plugin_config) - .await - .expect("managed Python worker should initialize"); + let mut cleanup = + load_and_initialize_python_worker(&manifest_ref, &PathBuf::from(environment_ref), config) + .await; let events = Arc::new(Mutex::new(Vec::::new())); let captured = events.clone(); @@ -1694,6 +1678,147 @@ async fn python_worker_host_runtime_mark_and_mutated_request_round_trip() { drop(cleanup); } +#[tokio::test] +async fn python_worker_execution_codec_context_round_trips_host_codecs() { + let _guard = WORKER_PLUGIN_TEST_LOCK.lock().await; + let Some(environment_ref) = std::env::var_os("NEMO_RELAY_PYTHON_PLUGIN_TEST_ENVIRONMENT") + else { + eprintln!( + "skipping Python worker codec-context round-trip; \ + NEMO_RELAY_PYTHON_PLUGIN_TEST_ENVIRONMENT is unset" + ); + return; + }; + let (manifest_dir, manifest_ref) = write_python_codec_context_worker(); + let cleanup = load_and_initialize_python_worker( + &manifest_ref, + &PathBuf::from(environment_ref), + Map::new(), + ) + .await; + + struct Case { + name: &'static str, + identity_kind: &'static str, + identity_id: &'static str, + answer: &'static str, + request_codec: Arc, + response_codec: Arc, + } + + for case in [ + Case { + name: "builtin-openai-chat", + identity_kind: "builtin", + identity_id: "openai_chat", + answer: "provider answer", + request_codec: Arc::new(OpenAIChatCodec), + response_codec: Arc::new(OpenAIChatCodec), + }, + Case { + name: "runtime-openai-chat", + identity_kind: "runtime", + identity_id: "tests.openai_chat.v1", + answer: "runtime answer", + request_codec: Arc::new(RuntimeOpenAiChatCodec), + response_codec: Arc::new(RuntimeOpenAiChatCodec), + }, + // A third invocation proves the spawned worker remains responsive after + // exercising both directional codec capabilities. + Case { + name: "post-round-trip-health", + identity_kind: "builtin", + identity_id: "openai_chat", + answer: "healthy", + request_codec: Arc::new(OpenAIChatCodec), + response_codec: Arc::new(OpenAIChatCodec), + }, + ] { + let provider_only = json!({"lane": case.name, "preserve": true}); + let expected_provider_only = provider_only.clone(); + let answer = case.answer; + let response = llm_call_execute( + LlmCallExecuteParams::builder() + .name(format!("python-worker-codec-context-{}", case.name)) + .request(LlmRequest { + headers: Map::new(), + content: json!({ + "model": "caller-model", + "messages": [{"role": "user", "content": "hello"}], + "provider_only": provider_only, + }), + }) + .codec(case.request_codec) + .response_codec(case.response_codec) + .func(Arc::new(move |request| { + let expected_provider_only = expected_provider_only.clone(); + Box::pin(async move { + assert_eq!(request.content["model"], "worker-model"); + assert_eq!(request.content["provider_only"], expected_provider_only); + Ok(json!({ + "id": "chatcmpl-codec-context", + "object": "chat.completion", + "model": "provider-model", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": answer}, + "finish_reason": "stop", + }], + })) + }) + })) + .build(), + ) + .await + .unwrap_or_else(|error| panic!("{} codec-context call failed: {error}", case.name)); + + assert_eq!( + response["_codec_context_probe"], + json!({ + "request_kind": case.identity_kind, + "request_id": case.identity_id, + "response_kind": case.identity_kind, + "response_id": case.identity_id, + "decoded_model": "provider-model", + "decoded_message": case.answer, + }) + ); + } + + drop(cleanup); + drop(manifest_dir); +} + +struct RuntimeOpenAiChatCodec; + +impl LlmCodec for RuntimeOpenAiChatCodec { + fn codec_identity(&self) -> LlmCodecIdentity { + LlmCodecIdentity::Runtime("tests.openai_chat.v1".into()) + } + + fn decode(&self, request: &LlmRequest) -> FlowResult { + OpenAIChatCodec.decode(request) + } + + fn encode( + &self, + annotated: &AnnotatedLlmRequest, + original: &LlmRequest, + ) -> FlowResult { + OpenAIChatCodec.encode(annotated, original) + } +} + +impl LlmResponseCodec for RuntimeOpenAiChatCodec { + fn codec_identity(&self) -> LlmCodecIdentity { + LlmCodecIdentity::Runtime("tests.openai_chat.v1".into()) + } + + fn decode_response(&self, response: &Json) -> FlowResult { + OpenAIChatCodec.decode_response(response) + } +} + struct FixtureCodec; impl LlmCodec for FixtureCodec { @@ -1754,6 +1879,31 @@ impl PythonWorkerCleanup { } } +async fn load_and_initialize_python_worker( + manifest_ref: &Path, + environment_ref: &Path, + config: Map, +) -> PythonWorkerCleanup { + let activation = load_worker_plugins([WorkerPluginLoadSpec { + plugin_id: "examples.python_grpc_worker".into(), + manifest_ref: manifest_ref.to_string_lossy().into_owned(), + environment_ref: Some(environment_ref.to_string_lossy().into_owned()), + config: config.clone(), + }]) + .expect("managed Python worker should load"); + let cleanup = PythonWorkerCleanup::new(activation); + let mut plugin_config = PluginConfig::default(); + plugin_config.components.push(PluginComponentSpec { + kind: "examples.python_grpc_worker".into(), + enabled: true, + config, + }); + test_initialize_plugin_host_exact(plugin_config) + .await + .expect("managed Python worker should initialize"); + cleanup +} + impl Drop for PythonWorkerCleanup { fn drop(&mut self) { if let Some(subscriber_name) = self.subscriber_name.take() { @@ -1891,6 +2041,59 @@ entrypoint = {entrypoint} )) } +fn write_python_codec_context_worker() -> (TempDir, PathBuf) { + const WORKER: &str = r#" +from nemo_relay_plugin import WorkerPlugin, serve_plugin + + +class CodecContextProbe(WorkerPlugin): + plugin_id = "examples.python_grpc_worker" + + def register(self, ctx, config): + del config + + async def execute(_name, request, context, next_call): + if not context.available: + raise RuntimeError("execution codec context is unavailable") + if context.request_codec is None or context.response_codec is None: + raise RuntimeError("directional codec proxy is unavailable") + + annotated = await context.request_codec.decode(request) + annotated["model"] = "worker-model" + encoded = await context.request_codec.encode(annotated, request) + response = await next_call.call(encoded) + decoded = await context.response_codec.decode(response) + + result = dict(response) + result["_codec_context_probe"] = { + "request_kind": context.request_codec_identity.kind, + "request_id": context.request_codec_identity.id, + "response_kind": context.response_codec_identity.kind, + "response_id": context.response_codec_identity.id, + "decoded_model": decoded.get("model"), + "decoded_message": decoded.get("message"), + } + return result + + ctx.register_llm_execution_intercept_with_context("codec_context_probe", execute) + + +async def main(): + await serve_plugin(CodecContextProbe()) +"#; + + let relay = supported_relay_requirement(); + let (temp, manifest) = write_worker_manifest( + "examples.python_grpc_worker", + &relay, + "python", + "codec_context_probe:main", + ); + std::fs::write(temp.path().join("codec_context_probe.py"), WORKER) + .expect("Python codec-context worker fixture should be written"); + (temp, manifest) +} + fn write_manifest_text(contents: &str) -> (TempDir, PathBuf) { let temp = TempDir::new().expect("manifest tempdir should be created"); let manifest = temp.path().join("relay-plugin.toml"); diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index a2e9b99d3..550358444 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -12,8 +12,9 @@ use crate::api::optimization::{ LlmOptimizationRecorder, record_llm_optimization_contribution, scope_llm_optimization_recorder, }; use crate::api::runtime::{ - BuiltinLlmCodec, LlmCodecIdentity, LlmSanitizeRequestContext, LlmSanitizeResponseContext, - MiddlewareContinuationLease, NemoRelayContextState, + BuiltinLlmCodec, LlmCodecIdentity, LlmExecutionCodecContext, LlmExecutionNextFn, + LlmSanitizeRequestContext, LlmSanitizeResponseContext, MiddlewareContinuationLease, + NemoRelayContextState, }; use crate::api::tool::ToolExecutionResult; use crate::codec::openai_chat::OpenAIChatCodec; @@ -528,6 +529,7 @@ fn registration_plan_and_scope_type_helpers_validate_edges() { surface: RegistrationSurface::Subscriber as i32, priority: 0, break_chain: false, + llm_execution_codec_context: false, }], error: None, conditional_middleware_guardrails: Vec::new(), @@ -544,6 +546,7 @@ fn registration_plan_and_scope_type_helpers_validate_edges() { surface: 999, priority: 0, break_chain: false, + llm_execution_codec_context: false, }], error: None, conditional_middleware_guardrails: Vec::new(), @@ -564,6 +567,7 @@ fn registration_plan_and_scope_type_helpers_validate_edges() { surface: RegistrationSurface::Unspecified as i32, priority: 0, break_chain: false, + llm_execution_codec_context: false, }], error: None, conditional_middleware_guardrails: Vec::new(), @@ -576,6 +580,27 @@ fn registration_plan_and_scope_type_helpers_validate_edges() { .contains("unspecified registration surface") ); + let incompatible_codec_context = validate_registration_plan( + "fixture_worker", + &RegisterResponse { + registrations: vec![Registration { + local_name: "subscriber".into(), + surface: RegistrationSurface::Subscriber as i32, + priority: 0, + break_chain: false, + llm_execution_codec_context: true, + }], + error: None, + conditional_middleware_guardrails: Vec::new(), + }, + ) + .expect_err("codec context must be limited to LLM execution surfaces"); + assert!( + incompatible_codec_context + .to_string() + .contains("incompatible surface") + ); + let cases = [ (ProtoScopeType::Agent, crate::api::scope::ScopeType::Agent), ( @@ -932,6 +957,10 @@ async fn llm_worker_sanitizers_forward_codec_context_and_omission() { else { panic!("LLM sanitizer must receive an LLM invocation"); }; + assert!( + invocation.execution_codec_context.is_none(), + "sanitizers must not receive execution codec context" + ); let codec = match invocation.sanitize_context.as_ref() { Some(nemo_relay_worker_proto::v1::llm_invocation::SanitizeContext::RequestSanitizeContext(context)) => context.codec.as_ref(), Some(nemo_relay_worker_proto::v1::llm_invocation::SanitizeContext::ResponseSanitizeContext(context)) => context.codec.as_ref(), @@ -1066,6 +1095,10 @@ async fn llm_worker_codec_capabilities_are_active_only_during_sanitizer_invocati else { panic!("LLM sanitizer must receive an LLM invocation"); }; + assert!( + invocation.execution_codec_context.is_none(), + "sanitizers must not receive execution codec context" + ); let state = host_state .lock() .unwrap() @@ -1205,6 +1238,217 @@ async fn llm_worker_codec_capabilities_are_active_only_during_sanitizer_invocati ); } +#[tokio::test(flavor = "multi_thread")] +async fn llm_worker_execution_codec_context_is_opt_in_and_ephemeral() { + enable_operational_logs(); + let host_state = shared_worker_host_state(); + let seen = Arc::new(Mutex::new(None::)); + let (callback, _shutdown) = fake_callback_service({ + let host_state = Arc::clone(&host_state); + let seen = Arc::clone(&seen); + move |request| { + let registration_name = request.registration_name.clone(); + if registration_name == "legacy" { + let Some(invoke_request_payload::Payload::Llm(invocation)) = request.payload else { + panic!("LLM execution must receive an LLM invocation"); + }; + assert!(invocation.sanitize_context.is_none()); + assert!(invocation.execution_codec_context.is_none()); + } else { + let (request_id, response_id, invocation_id) = + execution_codec_capabilities(request); + let response_id = response_id.expect("response capability must be present"); + let state = host_state.lock().unwrap().clone().unwrap(); + state + .request_codec(&request_id, &invocation_id) + .expect("request capability resolves during callback"); + state + .response_codec(&response_id, &invocation_id) + .expect("response capability resolves during callback"); + *seen.lock().unwrap() = Some((request_id, Some(response_id), invocation_id)); + } + InvokeResponse { + result: Some(InvokeResult::Json(JsonResult { + value: Some(json_envelope(JSON_SCHEMA, &json!({"ok": true})).unwrap()), + error: None, + })), + } + } + }) + .await; + *host_state.lock().unwrap() = Some(callback.host_state.clone()); + + let next: LlmExecutionNextFn = Arc::new(|_| Box::pin(async { Ok(json!({"unused": true})) })); + callback + .invoke_llm_execution( + "legacy", + "model", + valid_llm_request(), + None, + Arc::clone(&next), + ) + .await + .unwrap(); + assert!(seen.lock().unwrap().is_none()); + + callback + .invoke_llm_execution( + "context", + "model", + valid_llm_request(), + Some(openai_execution_codec_context()), + next, + ) + .await + .unwrap(); + + let (request_id, response_id, invocation_id) = seen.lock().unwrap().take().unwrap(); + assert_request_codec_expired(&callback.host_state, &request_id, &invocation_id); + assert_response_codec_expired( + &callback.host_state, + response_id.as_deref().unwrap(), + &invocation_id, + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn cancelling_worker_execution_expires_context_and_continuation_state() { + enable_operational_logs(); + let (started_tx, started_rx) = oneshot::channel(); + let started_tx = Arc::new(Mutex::new(Some(started_tx))); + let (callback, _shutdown, mut cancel_rx) = fake_callback_service_with_handlers( + { + let started_tx = Arc::clone(&started_tx); + move |request| { + let started_tx = Arc::clone(&started_tx); + Box::pin(async move { + let (request_id, response_id, invocation_id) = + execution_codec_capabilities(request); + if let Some(started) = started_tx.lock().unwrap().take() { + let _ = started.send((request_id, response_id, invocation_id)); + } + std::future::pending::().await + }) + } + }, + |_| Box::pin(tokio_stream::empty()), + ) + .await; + + let callback_task = callback.clone(); + let task = tokio::spawn(async move { + callback_task + .invoke_llm_execution( + "cancel-context-execution", + "model", + valid_llm_request(), + Some(openai_execution_codec_context()), + Arc::new(|request| Box::pin(async move { Ok(request.content) })), + ) + .await + }); + let (request_id, response_id, invocation_id) = + tokio::time::timeout(std::time::Duration::from_secs(1), started_rx) + .await + .expect("worker execution must start") + .expect("worker execution must publish its capabilities"); + let response_id = response_id.expect("response capability must be present"); + callback + .host_state + .request_codec(&request_id, &invocation_id) + .expect("request capability must be active while the worker is pending"); + callback + .host_state + .response_codec(&response_id, &invocation_id) + .expect("response capability must be active while the worker is pending"); + + task.abort(); + let _ = task.await; + + let cancellation = tokio::time::timeout(std::time::Duration::from_secs(1), cancel_rx.recv()) + .await + .expect("caller cancellation must reach the worker") + .expect("cancellation channel remains open"); + assert_eq!(cancellation.invocation_id, invocation_id); + assert!(cancellation.reason.contains("host caller cancelled")); + assert_request_codec_expired(&callback.host_state, &request_id, &invocation_id); + assert_response_codec_expired(&callback.host_state, &response_id, &invocation_id); + assert!(callback.host_state.continuations.lock().unwrap().is_empty()); + assert!(callback.host_state.scope_stacks.lock().unwrap().is_empty()); +} + +#[tokio::test(flavor = "multi_thread")] +async fn llm_worker_stream_codec_context_is_request_only_and_expires_at_eof() { + enable_operational_logs(); + let host_state = shared_worker_host_state(); + let seen = Arc::new(Mutex::new(None::<(String, String)>)); + let (chunk_tx, chunk_rx) = mpsc::channel(1); + let chunk_rx = Arc::new(Mutex::new(Some(chunk_rx))); + let (callback, _shutdown) = fake_callback_service_with_stream( + |_| InvokeResponse { + result: Some(InvokeResult::Empty(EmptyResult {})), + }, + { + let host_state = Arc::clone(&host_state); + let seen = Arc::clone(&seen); + let chunk_rx = Arc::clone(&chunk_rx); + move |request| { + let (request_id, response_id, invocation_id) = + execution_codec_capabilities(request); + assert!( + response_id.is_none(), + "stream must not expose a response decoder" + ); + host_state + .lock() + .unwrap() + .clone() + .expect("host state") + .request_codec(&request_id, &invocation_id) + .expect("request capability resolves during stream invocation"); + *seen.lock().unwrap() = Some((request_id, invocation_id)); + let receiver = chunk_rx + .lock() + .unwrap() + .take() + .expect("test stream created once"); + Box::pin(tokio_stream::wrappers::ReceiverStream::new(receiver)) as FakeInvokeStream + } + }, + ) + .await; + *host_state.lock().unwrap() = Some(callback.host_state.clone()); + + let mut stream = callback + .invoke_llm_stream_execution( + "context-stream", + "model", + valid_llm_request(), + Some(openai_execution_codec_context()), + Arc::new(|_request| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), + ) + .await + .expect("host stream"); + let (request_id, invocation_id) = seen.lock().unwrap().clone().expect("stream context seen"); + callback + .host_state + .request_codec(&request_id, &invocation_id) + .expect("request capability remains active while the stream is open"); + + chunk_tx + .send(Ok(StreamChunk { + item: Some(StreamItem::Value( + json_envelope(JSON_SCHEMA, &json!({"done": true})).unwrap(), + )), + })) + .await + .expect("stream chunk accepted"); + drop(chunk_tx); + assert_eq!(stream.next().await.unwrap().unwrap(), json!({"done": true})); + assert!(stream.next().await.is_none()); + assert_request_codec_expired(&callback.host_state, &request_id, &invocation_id); +} + #[tokio::test(flavor = "multi_thread")] async fn cancelling_worker_sanitizer_expires_codec_capability() { enable_operational_logs(); @@ -1284,6 +1528,7 @@ async fn callback_stream_transport_error_surfaces_to_host_stream() { "stream_transport_error", "model", valid_llm_request(), + None, Arc::new(|_request| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), ) .await @@ -1337,6 +1582,7 @@ async fn callback_stream_stops_when_host_receiver_is_dropped() { "stream_receiver_drop", "model", valid_llm_request(), + None, Arc::new(|_request| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), ) .await @@ -1414,6 +1660,212 @@ async fn callback_timeout_sends_explicit_worker_cancellation() { assert!(cancellation.reason.contains("timed out")); } +#[tokio::test(start_paused = true)] +async fn continuation_bearing_worker_tool_callback_allows_slow_next() { + enable_operational_logs(); + let host_state = shared_worker_host_state(); + let (started_tx, started_rx) = oneshot::channel(); + let started_tx = Arc::new(Mutex::new(Some(started_tx))); + let (callback, _shutdown, mut cancel_rx) = fake_callback_service_with_handlers( + { + let host_state = Arc::clone(&host_state); + let started_tx = Arc::clone(&started_tx); + move |request| { + let host_state = Arc::clone(&host_state); + let started_tx = Arc::clone(&started_tx); + Box::pin(async move { + let continuation = continuation_for(&host_state, &request.continuation_id); + let Continuation::Tool { next, .. } = continuation else { + panic!("expected a tool continuation"); + }; + if let Some(started) = started_tx.lock().unwrap().take() { + let _ = started.send(()); + } + let response = next(json!({"input": "slow"})) + .await + .expect("slow downstream tool must complete"); + InvokeResponse { + result: Some(InvokeResult::ToolExecution(ToolExecutionInterceptResult { + outcome: Some(ProtoToolExecutionInterceptOutcome { + result: Some(JsonValue { + json: serde_json::to_vec(&response.result).unwrap(), + }), + annotation: None, + pending_marks: None, + }), + })), + } + }) + } + }, + |_| Box::pin(tokio_stream::empty()), + ) + .await; + *host_state.lock().unwrap() = Some(callback.host_state.clone()); + + let callback_task = callback.clone(); + let result = tokio::spawn(async move { + callback_task + .invoke_tool_execution( + "slow-tool-next", + "lookup", + json!({"input": "original"}), + Some("call-1"), + Arc::new(|value| { + Box::pin(async move { + tokio::time::sleep(WORKER_RPC_TIMEOUT + Duration::from_secs(1)).await; + Ok(ToolExecutionResult::new(json!({"slow": value}))) + }) + }), + ) + .await + }); + started_rx.await.expect("worker tool callback must start"); + tokio::time::advance(WORKER_RPC_TIMEOUT + Duration::from_secs(2)).await; + + assert_eq!( + result.await.unwrap().unwrap().result, + json!({"slow": {"input": "slow"}}) + ); + assert_no_worker_cancellation(&mut cancel_rx, "slow downstream tool execution"); +} + +#[tokio::test(start_paused = true)] +async fn continuation_bearing_worker_callback_allows_slow_next() { + enable_operational_logs(); + let host_state = shared_worker_host_state(); + let (started_tx, started_rx) = oneshot::channel(); + let started_tx = Arc::new(Mutex::new(Some(started_tx))); + let (callback, _shutdown, mut cancel_rx) = fake_callback_service_with_handlers( + { + let host_state = Arc::clone(&host_state); + let started_tx = Arc::clone(&started_tx); + move |request| { + let host_state = Arc::clone(&host_state); + let started_tx = Arc::clone(&started_tx); + Box::pin(async move { + let continuation = continuation_for(&host_state, &request.continuation_id); + let Continuation::Llm { next, .. } = continuation else { + panic!("expected an LLM continuation"); + }; + if let Some(started) = started_tx.lock().unwrap().take() { + let _ = started.send(()); + } + tokio::time::sleep(WORKER_RPC_TIMEOUT + Duration::from_secs(1)).await; + let response = next(valid_llm_request()) + .await + .expect("slow downstream provider must complete"); + InvokeResponse { + result: Some(InvokeResult::Json(JsonResult { + value: Some(json_envelope(JSON_SCHEMA, &response).unwrap()), + error: None, + })), + } + }) + } + }, + |_| Box::pin(tokio_stream::empty()), + ) + .await; + *host_state.lock().unwrap() = Some(callback.host_state.clone()); + + let callback_task = callback.clone(); + let result = tokio::spawn(async move { + callback_task + .invoke_llm_execution( + "slow-next", + "model", + valid_llm_request(), + None, + Arc::new(|_| Box::pin(async { Ok(json!({"slow": "completed"})) })), + ) + .await + }); + started_rx.await.expect("worker callback must start"); + tokio::time::advance(WORKER_RPC_TIMEOUT + Duration::from_secs(2)).await; + + assert_eq!(result.await.unwrap().unwrap(), json!({"slow": "completed"})); + assert_no_worker_cancellation(&mut cancel_rx, "slow downstream execution"); +} + +#[tokio::test(start_paused = true)] +async fn continuation_bearing_worker_stream_allows_slow_next() { + enable_operational_logs(); + let host_state = shared_worker_host_state(); + let (started_tx, started_rx) = oneshot::channel(); + let started_tx = Arc::new(Mutex::new(Some(started_tx))); + let (callback, _shutdown, mut cancel_rx) = fake_callback_service_with_async_stream_handler( + |_| { + Box::pin(async { + InvokeResponse { + result: Some(InvokeResult::Empty(EmptyResult {})), + } + }) + }, + { + let host_state = Arc::clone(&host_state); + let started_tx = Arc::clone(&started_tx); + move |request| { + let host_state = Arc::clone(&host_state); + let started_tx = Arc::clone(&started_tx); + Box::pin(async move { + let continuation = continuation_for(&host_state, &request.continuation_id); + let Continuation::LlmStream { next, .. } = continuation else { + panic!("expected an LLM stream continuation"); + }; + if let Some(started) = started_tx.lock().unwrap().take() { + let _ = started.send(()); + } + tokio::time::sleep(WORKER_RPC_TIMEOUT + Duration::from_secs(1)).await; + let stream = next(valid_llm_request()) + .await + .expect("slow downstream stream must open"); + Box::pin(stream.map(|item| { + item.map(|value| StreamChunk { + item: Some(StreamItem::Value( + json_envelope(JSON_SCHEMA, &value) + .expect("stream value must encode"), + )), + }) + .map_err(|error| Status::internal(error.to_string())) + })) as FakeInvokeStream + }) + } + }, + ) + .await; + *host_state.lock().unwrap() = Some(callback.host_state.clone()); + + let callback_task = callback.clone(); + let result = tokio::spawn(async move { + callback_task + .invoke_llm_stream_execution( + "slow-stream-next", + "model", + valid_llm_request(), + None, + Arc::new(|_| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( + json!({"slow_stream": "completed"}), + )]))) + }) + }), + ) + .await + }); + started_rx.await.expect("worker stream callback must start"); + tokio::time::advance(WORKER_RPC_TIMEOUT + Duration::from_secs(2)).await; + + let mut stream = result.await.unwrap().expect("slow worker stream must open"); + assert_eq!( + stream.next().await.unwrap().unwrap(), + json!({"slow_stream": "completed"}) + ); + assert!(stream.next().await.is_none()); + assert_no_worker_cancellation(&mut cancel_rx, "slow downstream stream setup"); +} + #[tokio::test(flavor = "multi_thread")] async fn dropping_callback_future_cancels_worker_and_cleans_host_state() { enable_operational_logs(); @@ -1660,57 +2112,75 @@ fn invocation_cleanup_releases_host_state_locks_before_unwinding() { #[tokio::test(flavor = "multi_thread")] async fn dropping_host_stream_sends_explicit_worker_cancellation() { enable_operational_logs(); - let (yield_tx, yield_rx) = oneshot::channel(); - let yield_rx = Arc::new(Mutex::new(Some(yield_rx))); - let (callback, _shutdown, mut cancel_rx) = fake_callback_service_with_handlers( - |_| { - Box::pin(async { - InvokeResponse { - result: Some(InvokeResult::Empty(EmptyResult {})), - } - }) - }, - { - let yield_rx = yield_rx.clone(); - move |_| { - let yield_rx = yield_rx - .lock() - .expect("yield lock") - .take() - .expect("stream should be created once"); - Box::pin(SignalChunkThenPendingStream { - yield_rx, - dropped: None, - yielded: false, - }) - } - }, - ) - .await; - let mut stream = callback - .invoke_llm_stream_execution( - "cancel_stream", - "model", - valid_llm_request(), - Arc::new(|_request| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), - ) - .await - .expect("host stream should be returned"); - yield_tx + let mut fixture = pending_worker_stream_with_codec_context("cancel_stream").await; + fixture + .yield_tx + .take() + .expect("yield signal sent once") .send(()) .expect("worker stream yield signal should be delivered"); - tokio::time::timeout(std::time::Duration::from_secs(1), stream.next()) + tokio::time::timeout( + std::time::Duration::from_secs(1), + fixture.stream.as_mut().unwrap().next(), + ) + .await + .expect("worker stream should yield before abandonment") + .expect("worker stream should yield before abandonment") + .expect("worker stream chunk should be valid"); + drop(fixture.stream.take()); + + assert_worker_stream_cancelled_and_cleaned(&mut fixture).await; +} + +#[tokio::test(flavor = "multi_thread")] +async fn closing_worker_stream_waits_for_cancellation_and_codec_cleanup() { + enable_operational_logs(); + let mut fixture = pending_worker_stream_with_codec_context("close_stream").await; + + fixture + .stream + .as_mut() + .unwrap() + .close() .await - .expect("worker stream should yield before abandonment") - .expect("worker stream should yield before abandonment") - .expect("worker stream chunk should be valid"); - drop(stream); + .expect("explicit close must wait for worker stream cleanup"); - let cancellation = tokio::time::timeout(std::time::Duration::from_secs(1), cancel_rx.recv()) + assert_request_codec_expired( + &fixture.callback.host_state, + &fixture.request_id, + &fixture.invocation_id, + ); + assert_worker_stream_cancelled_and_cleaned(&mut fixture).await; +} + +#[tokio::test] +async fn retrying_worker_stream_close_still_waits_after_the_first_future_is_cancelled() { + let (_stream_tx, stream_rx) = mpsc::channel(1); + let (completion_tx, completion_rx) = watch::channel(false); + let mut stream = WorkerForwardedLlmStream { + receiver: Some(tokio_stream::wrappers::ReceiverStream::new(stream_rx)), + completion: completion_rx, + }; + + // Creating close initiates shutdown by dropping the receiver. Cancelling + // that future must not consume the only observer of producer cleanup. + drop(Pin::new(&mut stream).close()); + + let mut retry = Pin::new(&mut stream).close(); + tokio::select! { + biased; + result = retry.as_mut() => { + panic!("retry completed before producer cleanup: {result:?}"); + } + _ = tokio::task::yield_now() => {} + } + + completion_tx.send_replace(true); + retry.await.expect("retry must observe producer cleanup"); + Pin::new(&mut stream) + .close() .await - .expect("host should cancel abandoned stream") - .expect("cancellation channel should remain open"); - assert!(cancellation.reason.contains("stopped consuming")); + .expect("completed close remains idempotent"); } #[tokio::test(flavor = "multi_thread")] @@ -2271,8 +2741,14 @@ async fn host_runtime_codec_capabilities_are_directional_authorized_and_ephemera let request_codec: Arc = codec.clone(); let response_codec: Arc = codec.clone(); let invocation_id = "sanitize-invocation"; - let request_capability = state.insert_request_codec(invocation_id, request_codec); - let response_capability = state.insert_response_codec(invocation_id, response_codec); + let request_capability_guard = state + .issue_request_codec(invocation_id, request_codec) + .unwrap(); + let response_capability_guard = state + .issue_response_codec(invocation_id, response_codec) + .unwrap(); + let request_capability = request_capability_guard.id().to_owned(); + let response_capability = response_capability_guard.id().to_owned(); let request = LlmRequest { headers: serde_json::Map::new(), content: json!({ @@ -2401,8 +2877,8 @@ async fn host_runtime_codec_capabilities_are_directional_authorized_and_ephemera .into_inner(); assert!(decoded.error.is_none()); - state.remove_codec(&request_capability); - state.remove_codec(&response_capability); + drop(request_capability_guard); + drop(response_capability_guard); let expired = service .decode_llm_codec_request(Request::new(LlmCodecDecodeRequest { activation_id: ACTIVATION_ID.into(), @@ -2490,6 +2966,97 @@ async fn host_runtime_service_reports_poisoned_internal_locks() { .await .expect_err("poisoned scope stack lock should fail"); assert_eq!(drop_error.code(), tonic::Code::Internal); + + let state = Arc::new(WorkerHostRuntimeState::new( + ACTIVATION_ID.into(), + AUTH_TOKEN.into(), + )); + poison_mutex({ + let state = state.clone(); + move || { + let _guard = state.codecs.lock().expect("codecs lock"); + panic!("poison codecs"); + } + }); + let insert_error = match state.issue_request_codec("invocation", Arc::new(OpenAIChatCodec)) { + Ok(_) => panic!("poisoned codec lock should reject capability insertion"), + Err(error) => error, + }; + assert!(matches!( + insert_error, + FlowError::Internal(message) if message.contains("codec lock poisoned") + )); + + let (callback, _shutdown) = fake_callback_service(|_| InvokeResponse { + result: Some(InvokeResult::Empty(EmptyResult {})), + }) + .await; + poison_mutex({ + let state = callback.host_state.clone(); + move || { + let _guard = state.codecs.lock().expect("codecs lock"); + panic!("poison callback codecs"); + } + }); + let codec = Arc::new(OpenAIChatCodec); + let sanitizer_error = callback + .invoke_llm_sanitize_request( + "poisoned-sanitizer-codec-context", + valid_llm_request(), + LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), + ) + .await + .expect_err("poisoned sanitizer codec setup must fail before invoking the worker"); + assert!(matches!( + sanitizer_error, + FlowError::Internal(message) if message.contains("codec lock poisoned") + )); + assert!( + callback + .host_state + .scope_stacks + .lock() + .expect("scope stack lock") + .is_empty(), + "failed sanitizer codec setup must remove its invocation scope stack" + ); + + let context = LlmExecutionCodecContext::new( + LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), + LlmSanitizeResponseContext::for_response_codec(Some(codec)), + ); + let error = callback + .invoke_llm_execution( + "poisoned-codec-context", + "model", + valid_llm_request(), + Some(context), + Arc::new(|request| Box::pin(async move { Ok(request.content) })), + ) + .await + .expect_err("poisoned codec setup must fail before invoking the worker"); + assert!(matches!( + error, + FlowError::Internal(message) if message.contains("codec lock poisoned") + )); + assert!( + callback + .host_state + .continuations + .lock() + .expect("continuation lock") + .is_empty(), + "failed codec setup must remove its continuation" + ); + assert!( + callback + .host_state + .scope_stacks + .lock() + .expect("scope stack lock") + .is_empty(), + "failed codec setup must remove its invocation scope stack" + ); } #[test] @@ -2862,6 +3429,99 @@ fn valid_llm_request() -> LlmRequest { } } +type SharedWorkerHostState = Arc>>>; +type ExecutionCodecCapabilities = (String, Option, String); + +fn shared_worker_host_state() -> SharedWorkerHostState { + Arc::new(Mutex::new(None)) +} + +fn openai_execution_codec_context() -> LlmExecutionCodecContext { + let codec = Arc::new(OpenAIChatCodec); + LlmExecutionCodecContext::new( + LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), + LlmSanitizeResponseContext::for_response_codec(Some(codec)), + ) +} + +fn execution_codec_capabilities(request: InvokeRequest) -> ExecutionCodecCapabilities { + let invocation_id = request.invocation_id; + let Some(invoke_request_payload::Payload::Llm(invocation)) = request.payload else { + panic!("LLM execution must receive an LLM invocation"); + }; + assert!(invocation.sanitize_context.is_none()); + let context = invocation + .execution_codec_context + .expect("opted-in execution must receive codec context"); + let request_context = context.request.expect("request codec context"); + let response_context = context.response.expect("response codec context"); + for identity in [request_context.codec, response_context.codec] { + let identity = identity.expect("codec identity"); + assert_eq!(identity.kind, LlmCodecKind::Builtin as i32); + assert_eq!(identity.id.as_deref(), Some("openai_chat")); + } + ( + request_context + .codec_capability_id + .expect("request capability must be present"), + response_context.codec_capability_id, + invocation_id, + ) +} + +fn assert_request_codec_expired( + state: &WorkerHostRuntimeState, + capability_id: &str, + invocation_id: &str, +) { + assert_eq!( + state + .request_codec(capability_id, invocation_id) + .err() + .expect("request codec capability must be expired") + .code(), + tonic::Code::NotFound + ); +} + +fn assert_response_codec_expired( + state: &WorkerHostRuntimeState, + capability_id: &str, + invocation_id: &str, +) { + assert_eq!( + state + .response_codec(capability_id, invocation_id) + .err() + .expect("response codec capability must be expired") + .code(), + tonic::Code::NotFound + ); +} + +fn continuation_for(state: &SharedWorkerHostState, continuation_id: &str) -> Continuation { + state + .lock() + .unwrap() + .as_ref() + .expect("host state") + .continuation(continuation_id) + .expect("continuation must remain active") +} + +fn assert_no_worker_cancellation( + cancel_rx: &mut mpsc::UnboundedReceiver, + operation: &str, +) { + assert!( + matches!( + cancel_rx.try_recv(), + Err(tokio::sync::mpsc::error::TryRecvError::Empty) + ), + "{operation} must not trigger worker cancellation" + ); +} + async fn fake_callback_service( invoke: impl Fn(InvokeRequest) -> InvokeResponse + Send + Sync + 'static, ) -> (WorkerPluginCallback, oneshot::Sender<()>) { @@ -2891,6 +3551,20 @@ async fn fake_callback_service_with_handlers( (callback, shutdown_tx, cancel_rx) } +async fn fake_callback_service_with_async_stream_handler( + invoke: impl Fn(InvokeRequest) -> FakeInvokeFuture + Send + Sync + 'static, + invoke_stream: impl Fn(InvokeRequest) -> FakeInvokeStreamFuture + Send + Sync + 'static, +) -> ( + WorkerPluginCallback, + oneshot::Sender<()>, + mpsc::UnboundedReceiver, +) { + let (client, shutdown_tx, cancel_rx, _register_calls) = + fake_worker_client_with_async_handlers(invoke, invoke_stream).await; + let (callback, shutdown_tx) = callback_for_client(client, shutdown_tx); + (callback, shutdown_tx, cancel_rx) +} + fn callback_for_client( client: PluginWorkerClient, shutdown_tx: oneshot::Sender<()>, @@ -2984,6 +3658,23 @@ async fn fake_worker_client_with_handlers( oneshot::Sender<()>, mpsc::UnboundedReceiver, Arc, +) { + let invoke_stream = Arc::new(invoke_stream); + fake_worker_client_with_async_handlers(invoke, move |request| { + let invoke_stream = Arc::clone(&invoke_stream); + Box::pin(async move { invoke_stream(request) }) + }) + .await +} + +async fn fake_worker_client_with_async_handlers( + invoke: impl Fn(InvokeRequest) -> FakeInvokeFuture + Send + Sync + 'static, + invoke_stream: impl Fn(InvokeRequest) -> FakeInvokeStreamFuture + Send + Sync + 'static, +) -> ( + PluginWorkerClient, + oneshot::Sender<()>, + mpsc::UnboundedReceiver, + Arc, ) { let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) .await @@ -3018,6 +3709,7 @@ fn registration(surface: RegistrationSurface, local_name: &str) -> Registration surface: surface as i32, priority: 0, break_chain: false, + llm_execution_codec_context: false, } } @@ -3034,7 +3726,7 @@ fn poison_mutex(f: impl FnOnce() + std::panic::UnwindSafe) { struct FakePluginWorker { invoke: Arc FakeInvokeFuture + Send + Sync>, - invoke_stream: Arc FakeInvokeStream + Send + Sync>, + invoke_stream: Arc FakeInvokeStreamFuture + Send + Sync>, cancel_tx: mpsc::UnboundedSender, register_calls: Arc, } @@ -3042,6 +3734,116 @@ struct FakePluginWorker { type FakeInvokeFuture = Pin + Send>>; type FakeInvokeStream = Pin> + Send>>; +type FakeInvokeStreamFuture = Pin + Send>>; + +struct WorkerStreamLifecycleFixture { + callback: WorkerPluginCallback, + stream: Option, + yield_tx: Option>, + cancel_rx: mpsc::UnboundedReceiver, + worker_stream_dropped_rx: Option>, + request_id: String, + invocation_id: String, + _shutdown: oneshot::Sender<()>, +} + +async fn pending_worker_stream_with_codec_context(name: &str) -> WorkerStreamLifecycleFixture { + let (context_tx, context_rx) = oneshot::channel(); + let context_tx = Arc::new(Mutex::new(Some(context_tx))); + let (yield_tx, yield_rx) = oneshot::channel(); + let yield_rx = Arc::new(Mutex::new(Some(yield_rx))); + let (worker_stream_dropped_tx, worker_stream_dropped_rx) = oneshot::channel(); + let worker_stream_dropped_tx = Arc::new(Mutex::new(Some(worker_stream_dropped_tx))); + let (callback, shutdown, cancel_rx) = fake_callback_service_with_handlers( + |_| { + Box::pin(async { + InvokeResponse { + result: Some(InvokeResult::Empty(EmptyResult {})), + } + }) + }, + { + let context_tx = Arc::clone(&context_tx); + let yield_rx = Arc::clone(&yield_rx); + let worker_stream_dropped_tx = Arc::clone(&worker_stream_dropped_tx); + move |request| { + let (request_id, response_id, invocation_id) = + execution_codec_capabilities(request); + assert!( + response_id.is_none(), + "stream must not expose a response decoder" + ); + context_tx + .lock() + .unwrap() + .take() + .expect("stream context sent once") + .send((request_id, invocation_id)) + .expect("stream context receiver remains open"); + Box::pin(SignalChunkThenPendingStream { + yield_rx: yield_rx + .lock() + .unwrap() + .take() + .expect("stream created once"), + dropped: worker_stream_dropped_tx.lock().unwrap().take(), + yielded: false, + }) as FakeInvokeStream + } + }, + ) + .await; + let stream = callback + .invoke_llm_stream_execution( + name, + "model", + valid_llm_request(), + Some(openai_execution_codec_context()), + Arc::new(|_| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), + ) + .await + .expect("host stream should be returned"); + let (request_id, invocation_id) = context_rx.await.expect("worker must publish codec context"); + callback + .host_state + .request_codec(&request_id, &invocation_id) + .expect("request capability remains active while the stream is open"); + WorkerStreamLifecycleFixture { + callback, + stream: Some(stream), + yield_tx: Some(yield_tx), + cancel_rx, + worker_stream_dropped_rx: Some(worker_stream_dropped_rx), + request_id, + invocation_id, + _shutdown: shutdown, + } +} + +async fn assert_worker_stream_cancelled_and_cleaned(fixture: &mut WorkerStreamLifecycleFixture) { + let cancellation = + tokio::time::timeout(std::time::Duration::from_secs(1), fixture.cancel_rx.recv()) + .await + .expect("host must cancel the stopped worker stream") + .expect("cancellation channel remains open"); + assert_eq!(cancellation.invocation_id, fixture.invocation_id); + assert!(cancellation.reason.contains("stopped consuming")); + tokio::time::timeout( + std::time::Duration::from_secs(1), + fixture + .worker_stream_dropped_rx + .take() + .expect("worker stream drop observed once"), + ) + .await + .expect("worker stream must be dropped") + .expect("worker stream drop signal must be delivered"); + assert_request_codec_expired( + &fixture.callback.host_state, + &fixture.request_id, + &fixture.invocation_id, + ); +} struct SignalChunkThenPendingStream { yield_rx: oneshot::Receiver<()>, @@ -3155,9 +3957,9 @@ impl PluginWorker for FakePluginWorker { &self, request: Request, ) -> std::result::Result, tonic::Status> { - Ok(tonic::Response::new((self.invoke_stream)( - request.into_inner(), - ))) + Ok(tonic::Response::new( + (self.invoke_stream)(request.into_inner()).await, + )) } async fn cancel_invocation( diff --git a/crates/core/tests/unit/llm_api_tests.rs b/crates/core/tests/unit/llm_api_tests.rs index 9ce1f02e7..37d405301 100644 --- a/crates/core/tests/unit/llm_api_tests.rs +++ b/crates/core/tests/unit/llm_api_tests.rs @@ -5,6 +5,8 @@ #![allow(clippy::await_holding_lock)] +#[cfg(feature = "worker-grpc")] +use std::collections::BTreeSet; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Barrier, Mutex}; @@ -20,11 +22,20 @@ use super::{ }; use crate::api::event::{Event, ScopeCategory}; use crate::api::optimization::finalize_optimization_summary; +#[cfg(feature = "worker-grpc")] +use crate::api::registry::{ + RuntimeRegistrationKind, deregister_conditional_middleware_guardrail, + deregister_llm_execution_intercept, register_conditional_middleware_guardrail, + register_contextual_llm_execution_intercept, register_llm_execution_intercept, + scope_register_llm_execution_intercept, +}; use crate::api::registry::{ deregister_llm_sanitize_request_guardrail, deregister_llm_sanitize_response_guardrail, register_llm_sanitize_request_guardrail, register_llm_sanitize_response_guardrail, }; use crate::api::runtime::{BuiltinLlmCodec, LlmCodecIdentity, LlmJsonStream}; +#[cfg(feature = "worker-grpc")] +use crate::api::runtime::{LlmExecutionCodecContext, LlmExecutionNextFn}; use crate::api::runtime::{ NemoRelayContextState, create_scope_stack, global_context, set_thread_scope_stack, }; @@ -117,6 +128,34 @@ fn multi_turn_annotation() -> Arc { Arc::new(OpenAIChatCodec.decode(&multi_turn_request()).unwrap()) } +#[cfg(feature = "worker-grpc")] +fn assert_openai_execution_context(context: &LlmExecutionCodecContext) { + assert_eq!( + context.request().codec(), + &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert!(context.request().resolve_codec().is_some()); + assert_eq!( + context.response().codec(), + &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert!(context.response().resolve_codec().is_some()); +} + +#[cfg(feature = "worker-grpc")] +async fn execute_openai_call(name: &str, func: LlmExecutionNextFn) -> crate::error::Result { + llm_call_execute( + LlmCallExecuteParams::builder() + .name(name) + .request(request()) + .func(func) + .codec(Arc::new(OpenAIChatCodec)) + .response_codec(Arc::new(OpenAIChatCodec)) + .build(), + ) + .await +} + struct ProjectionFailingCodec { projection_attempts: Arc, } @@ -245,6 +284,243 @@ fn response_sanitizer_context_preserves_all_codec_identity_states() { ); } +#[test] +#[cfg(feature = "worker-grpc")] +fn managed_execution_passes_codec_context_to_downstream_interceptor() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + let observations = Arc::new(Mutex::new(Vec::new())); + let captured = Arc::clone(&observations); + register_contextual_llm_execution_intercept( + "execution-codec-context-outer", + 1, + Arc::new(move |_name, request, context, next| { + let captured = Arc::clone(&captured); + Box::pin(async move { + assert_openai_execution_context(&context); + captured.lock().unwrap().push("before-next"); + let result = next(request).await; + assert_openai_execution_context(&context); + captured.lock().unwrap().push("after-next"); + result + }) + }), + ) + .unwrap(); + let captured = Arc::clone(&observations); + register_contextual_llm_execution_intercept( + "execution-codec-context-inner", + 2, + Arc::new(move |_name, request, context, next| { + let captured = Arc::clone(&captured); + Box::pin(async move { + assert_openai_execution_context(&context); + captured.lock().unwrap().push("inner"); + next(request).await + }) + }), + ) + .unwrap(); + + tokio::runtime::Runtime::new().unwrap().block_on(async { + execute_openai_call( + "execution-codec-context", + Arc::new(|_| Box::pin(async { Ok(json!({"ok": true})) })), + ) + .await + .unwrap(); + }); + + assert!(deregister_llm_execution_intercept("execution-codec-context-outer").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-context-inner").unwrap()); + let observations = observations.lock().unwrap(); + assert_eq!(*observations, ["before-next", "inner", "after-next"]); +} + +#[test] +#[cfg(feature = "worker-grpc")] +fn codec_context_keeps_legacy_scope_ordering_and_conditional_gating() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + let scope = push_scope( + PushScopeParams::builder() + .name("execution-codec-context-scope") + .scope_type(ScopeType::Custom) + .build(), + ) + .unwrap(); + let calls = Arc::new(Mutex::new(Vec::new())); + + let captured = Arc::clone(&calls); + register_llm_execution_intercept( + "execution-codec-legacy-global", + 10, + Arc::new(move |_name, request, next| { + let captured = Arc::clone(&captured); + Box::pin(async move { + captured.lock().unwrap().push("legacy-global-enter"); + let result = next(request).await; + captured.lock().unwrap().push("legacy-global-exit"); + result + }) + }), + ) + .unwrap(); + + register_contextual_llm_execution_intercept( + "execution-codec-gated-contextual", + 15, + Arc::new(move |_name, _request, _context, _next| { + Box::pin(async move { panic!("conditionally disabled interceptor must not execute") }) + }), + ) + .unwrap(); + let gate_kinds = BTreeSet::from([RuntimeRegistrationKind::LlmExecutionIntercept]); + register_conditional_middleware_guardrail( + "execution-codec-context-gate", + gate_kinds, + "execution-codec-gated-contextual", + Arc::new(|_, _| Some("disabled for regression test".into())), + ) + .unwrap(); + + let captured = Arc::clone(&calls); + register_contextual_llm_execution_intercept( + "execution-codec-contextual-global", + 20, + Arc::new(move |_name, request, context, next| { + let captured = Arc::clone(&captured); + Box::pin(async move { + assert_openai_execution_context(&context); + captured.lock().unwrap().push("contextual-global-enter"); + let result = next(request).await; + captured.lock().unwrap().push("contextual-global-exit"); + result + }) + }), + ) + .unwrap(); + + let captured = Arc::clone(&calls); + scope_register_llm_execution_intercept( + &scope.uuid, + "execution-codec-legacy-scope", + 30, + Arc::new(move |_name, request, next| { + let captured = Arc::clone(&captured); + Box::pin(async move { + captured.lock().unwrap().push("legacy-scope-enter"); + let result = next(request).await; + captured.lock().unwrap().push("legacy-scope-exit"); + result + }) + }), + ) + .unwrap(); + + let captured = Arc::clone(&calls); + let response = tokio::runtime::Runtime::new().unwrap().block_on(async { + execute_openai_call( + "execution-codec-context-mixed-chain", + Arc::new(move |_| { + let captured = Arc::clone(&captured); + Box::pin(async move { + captured.lock().unwrap().push("provider"); + Ok(json!({"ok": true})) + }) + }), + ) + .await + .unwrap() + }); + + assert_eq!(response, json!({"ok": true})); + assert!(deregister_conditional_middleware_guardrail("execution-codec-context-gate").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-legacy-global").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-gated-contextual").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-contextual-global").unwrap()); + pop_scope(PopScopeParams::builder().handle_uuid(&scope.uuid).build()).unwrap(); + + assert_eq!( + calls.lock().unwrap().as_slice(), + [ + "legacy-global-enter", + "contextual-global-enter", + "legacy-scope-enter", + "provider", + "legacy-scope-exit", + "contextual-global-exit", + "legacy-global-exit", + ] + ); +} + +#[test] +#[cfg(feature = "worker-grpc")] +fn execution_codec_context_does_not_follow_wire_format_mutation() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + register_llm_execution_intercept( + "execution-codec-change-wire-format", + 10, + Arc::new(move |_name, mut request, next| { + Box::pin(async move { + request.content = json!({ + "contents": [{ + "role": "user", + "parts": [{"text": "hello"}], + }], + }); + next(request).await + }) + }), + ) + .unwrap(); + + register_contextual_llm_execution_intercept( + "execution-codec-reject-stale-payload", + 20, + Arc::new(move |_name, request, context, _next| { + Box::pin(async move { + assert_openai_execution_context(&context); + let codec = context + .request() + .resolve_codec() + .expect("managed call must expose its selected request codec"); + codec.decode(&request).map(|_| json!({"unexpected": true})) + }) + }), + ) + .unwrap(); + + let provider_calls = Arc::new(AtomicUsize::new(0)); + let captured = Arc::clone(&provider_calls); + let error = tokio::runtime::Runtime::new() + .unwrap() + .block_on(async { + execute_openai_call( + "execution-codec-wire-format-invariant", + Arc::new(move |_| { + captured.fetch_add(1, Ordering::SeqCst); + Box::pin(async { Ok(json!({"unexpected": true})) }) + }), + ) + .await + }) + .expect_err("the selected OpenAI Chat codec must reject a Gemini wire payload"); + + assert!(deregister_llm_execution_intercept("execution-codec-change-wire-format").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-reject-stale-payload").unwrap()); + assert_eq!(provider_calls.load(Ordering::SeqCst), 0); + assert!(matches!(error, FlowError::InvalidArgument(_))); +} + impl LlmCodec for ProjectionFailingCodec { fn decode(&self, request: &LlmRequest) -> crate::error::Result { OpenAIChatCodec.decode(request) diff --git a/crates/worker-proto/build.rs b/crates/worker-proto/build.rs index 5fddbc74e..8bf3a8682 100644 --- a/crates/worker-proto/build.rs +++ b/crates/worker-proto/build.rs @@ -8,6 +8,7 @@ fn main() -> Result<(), Box> { let include = "proto"; let mut prost = prost_build::Config::new(); prost.protoc_executable(protoc_bin_vendored::protoc_bin_path()?); + prost.boxed(".nemo.relay.worker.v1.LlmInvocation.execution_codec_context"); tonic_prost_build::configure().compile_with_config(prost, &[proto], &[include])?; Ok(()) diff --git a/crates/worker-proto/proto/nemo/relay/worker/v1/plugin_worker.proto b/crates/worker-proto/proto/nemo/relay/worker/v1/plugin_worker.proto index 7aa990e31..b54dd0338 100644 --- a/crates/worker-proto/proto/nemo/relay/worker/v1/plugin_worker.proto +++ b/crates/worker-proto/proto/nemo/relay/worker/v1/plugin_worker.proto @@ -245,6 +245,9 @@ message Registration { int32 priority = 3; bool break_chain = 4; reserved 5; + // Requests invocation-scoped codec identities and supported codec operations + // for this execution interceptor. Older hosts ignore this field. + bool llm_execution_codec_context = 6; } message InvokeRequest { @@ -288,6 +291,14 @@ message LlmInvocation { LlmSanitizeRequestContext request_sanitize_context = 9; LlmSanitizeResponseContext response_sanitize_context = 10; } + // Present only for an execution interceptor that opted into codec context. + // Older workers ignore this additive field. + LlmExecutionCodecContext execution_codec_context = 11; +} + +message LlmExecutionCodecContext { + LlmSanitizeRequestContext request = 1; + LlmSanitizeResponseContext response = 2; } message LlmCodecIdentity { diff --git a/crates/worker-proto/tests/proto_tests.rs b/crates/worker-proto/tests/proto_tests.rs index 087abfeff..8a8189f76 100644 --- a/crates/worker-proto/tests/proto_tests.rs +++ b/crates/worker-proto/tests/proto_tests.rs @@ -6,9 +6,10 @@ use nemo_relay_worker_proto::v1::{ ConditionalMiddlewareGuardrailRegistration, ConditionalMiddlewareInvocation, EmitMarkRequest, GetRuntimeDiagnosticsRequest, GetRuntimeDiagnosticsResponse, HandshakeRequest, HealthRequest, - InvokeRequest, JsonEnvelope, JsonValue, RegisterConditionalMiddlewareGuardrailRequest, - RegistrationSurface, RuntimeDiagnostic, ScopeType, - ToolExecutionResult as ProtoToolExecutionResult, invoke_request, + InvokeRequest, JsonEnvelope, JsonValue, LlmCodecIdentity, LlmCodecKind, + LlmExecutionCodecContext, LlmInvocation, LlmSanitizeRequestContext, LlmSanitizeResponseContext, + RegisterConditionalMiddlewareGuardrailRequest, Registration, RegistrationSurface, + RuntimeDiagnostic, ScopeType, ToolExecutionResult as ProtoToolExecutionResult, invoke_request, }; use nemo_relay_worker_proto::{ WORKER_PROTOCOL_GRPC_V1, decode_json_envelope, decode_json_value, json_envelope, json_value, @@ -16,6 +17,24 @@ use nemo_relay_worker_proto::{ use prost::Message; use serde_json::json; +#[derive(Clone, PartialEq, Message)] +struct LegacyRegistration { + #[prost(string, tag = "1")] + local_name: String, + #[prost(int32, tag = "2")] + surface: i32, + #[prost(int32, tag = "3")] + priority: i32, + #[prost(bool, tag = "4")] + break_chain: bool, +} + +#[derive(Clone, PartialEq, Message)] +struct LegacyLlmInvocation { + #[prost(string, tag = "1")] + model_name: String, +} + #[test] fn worker_protocol_identifier_is_stable() { assert_eq!(WORKER_PROTOCOL_GRPC_V1, "grpc-v1"); @@ -148,6 +167,66 @@ fn request_field_numbers_are_stable() { ); } +#[test] +fn execution_codec_context_fields_are_additive_and_stable() { + let legacy_registration = LegacyRegistration { + local_name: "legacy".into(), + surface: RegistrationSurface::LlmExecutionIntercept as i32, + priority: 7, + break_chain: false, + }; + let decoded_by_new_host = Registration::decode(legacy_registration.encode_to_vec().as_slice()) + .expect("new host must decode a legacy registration"); + assert_eq!(decoded_by_new_host.local_name, "legacy"); + assert!(!decoded_by_new_host.llm_execution_codec_context); + + let contextual_registration = Registration { + llm_execution_codec_context: true, + ..Default::default() + }; + assert_eq!(contextual_registration.encode_to_vec(), b"\x30\x01"); + let decoded_by_legacy_host = + LegacyRegistration::decode(contextual_registration.encode_to_vec().as_slice()) + .expect("legacy host must ignore the additive registration field"); + assert_eq!(decoded_by_legacy_host, LegacyRegistration::default()); + + let legacy_invocation = LegacyLlmInvocation { + model_name: "legacy-model".into(), + ..Default::default() + }; + let decoded_by_new_worker = LlmInvocation::decode(legacy_invocation.encode_to_vec().as_slice()) + .expect("new worker must decode a legacy invocation"); + assert_eq!(decoded_by_new_worker.model_name, "legacy-model"); + assert!(decoded_by_new_worker.execution_codec_context.is_none()); + + let codec = LlmCodecIdentity { + kind: LlmCodecKind::Builtin as i32, + id: Some("openai_chat".into()), + }; + let invocation = LlmInvocation { + execution_codec_context: Some(Box::new(LlmExecutionCodecContext { + request: Some(LlmSanitizeRequestContext { + codec: Some(codec.clone()), + codec_capability_id: Some("request".into()), + }), + response: Some(LlmSanitizeResponseContext { + codec: Some(codec), + codec_capability_id: Some("response".into()), + }), + })), + ..Default::default() + }; + let encoded = invocation.encode_to_vec(); + assert_eq!(encoded.first(), Some(&0x5a)); // Field 11, length-delimited. + let decoded_by_legacy_worker = LegacyLlmInvocation::decode(encoded.as_slice()) + .expect("legacy worker must ignore the additive execution context"); + assert_eq!(decoded_by_legacy_worker, LegacyLlmInvocation::default()); + assert_eq!( + LlmInvocation::decode(encoded.as_slice()).unwrap(), + invocation + ); +} + #[test] fn runtime_diagnostics_messages_are_stable() { let request = GetRuntimeDiagnosticsRequest { diff --git a/crates/worker/src/lib.rs b/crates/worker/src/lib.rs index 953f384c3..2cad4082e 100644 --- a/crates/worker/src/lib.rs +++ b/crates/worker/src/lib.rs @@ -387,6 +387,68 @@ impl WorkerResponseCodec { .await } } + +/// Invocation-scoped codec identities and operations for an LLM execution interceptor. +/// +/// `is_available` is false when this SDK is running against a Relay host that +/// predates execution codec context. A supported host can still report no +/// active codec for either direction. +/// +/// This context identifies the codecs selected when Relay created the managed +/// invocation. Rewriting a request into another provider's wire format does +/// not change these identities; codec operations reject incompatible payloads. +#[derive(Clone)] +pub struct LlmExecutionContext { + available: bool, + request_codec_identity: LlmCodecIdentity, + response_codec_identity: LlmCodecIdentity, + request_codec: Option, + response_codec: Option, +} + +impl std::fmt::Debug for LlmExecutionContext { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("LlmExecutionContext") + .field("available", &self.available) + .field("request_codec_identity", &self.request_codec_identity) + .field("response_codec_identity", &self.response_codec_identity) + .finish_non_exhaustive() + } +} + +impl LlmExecutionContext { + /// Whether the Relay host supplied execution codec context. + #[must_use] + pub fn is_available(&self) -> bool { + self.available + } + + /// Identity of the active request codec, or `None` when no codec is active. + #[must_use] + pub fn request_codec_identity(&self) -> &LlmCodecIdentity { + &self.request_codec_identity + } + + /// Identity of the active response codec, or `None` when no codec is active. + #[must_use] + pub fn response_codec_identity(&self) -> &LlmCodecIdentity { + &self.response_codec_identity + } + + /// Invocation-scoped request codec proxy, when the host supplied one. + #[must_use] + pub fn request_codec(&self) -> Option { + self.request_codec.clone() + } + + /// Invocation-scoped response codec proxy, when the host supplied one. + #[must_use] + pub fn response_codec(&self) -> Option { + self.response_codec.clone() + } +} + type LlmConditionalFn = Arc BoxFutureResult> + Send + Sync>; type ConditionalMiddlewareFn = Arc< dyn Fn(BTreeSet, String) -> BoxFutureResult> @@ -402,9 +464,14 @@ type LlmRequestFn = Arc< + Send + Sync, >; -type LlmExecutionFn = Arc BoxFutureResult + Send + Sync>; -type LlmStreamExecutionFn = - Arc BoxFutureResult + Send + Sync>; +type LlmExecutionFn = Arc< + dyn Fn(&str, LlmRequest, LlmExecutionContext, LlmNext) -> BoxFutureResult + Send + Sync, +>; +type LlmStreamExecutionFn = Arc< + dyn Fn(&str, LlmRequest, LlmExecutionContext, LlmStreamNext) -> BoxFutureResult + + Send + + Sync, +>; #[derive(Default)] struct WorkerHandlers { @@ -832,7 +899,34 @@ impl PluginContext { ); self.handlers.llm_executions.insert( name.into(), - Arc::new(move |model, request, next| Box::pin(callback(model, request, next))), + Arc::new(move |model, request, _context, next| { + Box::pin(callback(model, request, next)) + }), + ); + } + + /// Registers an LLM execution intercept with invocation-scoped codec access. + pub fn register_llm_execution_intercept_with_context( + &mut self, + name: &str, + priority: i32, + callback: F, + ) where + F: Fn(&str, LlmRequest, LlmExecutionContext, LlmNext) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + self.push_registration_record( + name, + RegistrationSurface::LlmExecutionIntercept, + priority, + false, + true, + ); + self.handlers.llm_executions.insert( + name.into(), + Arc::new(move |model, request, context, next| { + Box::pin(callback(model, request, context, next)) + }), ); } @@ -859,7 +953,37 @@ impl PluginContext { ); self.handlers.llm_stream_executions.insert( name.into(), - Arc::new(move |model, request, next| Box::pin(callback(model, request, next))), + Arc::new(move |model, request, _context, next| { + Box::pin(callback(model, request, next)) + }), + ); + } + + /// Registers a streaming LLM execution intercept with request codec access. + /// + /// The context identifies the response codec but does not expose a response + /// decoder because stream chunks are not complete provider responses. + pub fn register_llm_stream_execution_intercept_with_context( + &mut self, + name: &str, + priority: i32, + callback: F, + ) where + F: Fn(&str, LlmRequest, LlmExecutionContext, LlmStreamNext) -> Fut + Send + Sync + 'static, + Fut: Future> + Send + 'static, + { + self.push_registration_record( + name, + RegistrationSurface::LlmStreamExecutionIntercept, + priority, + false, + true, + ); + self.handlers.llm_stream_executions.insert( + name.into(), + Arc::new(move |model, request, context, next| { + Box::pin(callback(model, request, context, next)) + }), ); } @@ -869,12 +993,24 @@ impl PluginContext { surface: RegistrationSurface, priority: i32, break_chain: bool, + ) { + self.push_registration_record(name, surface, priority, break_chain, false); + } + + fn push_registration_record( + &mut self, + name: &str, + surface: RegistrationSurface, + priority: i32, + break_chain: bool, + llm_execution_codec_context: bool, ) { self.handlers.registrations.push(Registration { local_name: name.into(), surface: surface as i32, priority, break_chain, + llm_execution_codec_context, }); } } @@ -2014,6 +2150,9 @@ impl PluginWorker for WorkerService { .cloned() .ok_or_else(|| Status::not_found("stream execution handler not registered"))?; let payload = llm_payload(request.payload).map_err(status_from_sdk)?; + let execution_context = payload + .execution_context(&self.runtime, &invocation_id) + .map_err(status_from_sdk)?; let request_value = required_json::(payload.request, "llm request").map_err(status_from_sdk)?; let next = LlmStreamNext { @@ -2028,7 +2167,7 @@ impl PluginWorker for WorkerService { TASK_SCOPE_CONTEXT .scope(open_scope.clone(), async { let future = with_thread_scope(&open_scope, || { - handler(&model_name, request_value, next) + handler(&model_name, request_value, execution_context, next) }); future.await }) @@ -2651,13 +2790,16 @@ impl WorkerService { scope: &Option, ) -> Result { let payload = llm_payload(request.payload)?; + let execution_context = payload.execution_context(&self.runtime, &request.invocation_id)?; let request_value = required_json::(payload.request, "llm request")?; let handler = self.llm_execution(&request.registration_name)?; let next = LlmNext { runtime: self.runtime.clone(), continuation_id: request.continuation_id, }; - let future = with_thread_scope(scope, || handler(&payload.model_name, request_value, next)); + let future = with_thread_scope(scope, || { + handler(&payload.model_name, request_value, execution_context, next) + }); Ok(json_response(future.await?)) } @@ -2844,9 +2986,55 @@ struct LlmPayload { annotated_request: Option, response: Option, sanitize_context: Option, + execution_codec_context: Option>, } impl LlmPayload { + fn execution_context( + &self, + runtime: &PluginRuntime, + invocation_id: &str, + ) -> Result { + let Some(context) = self.execution_codec_context.as_ref() else { + return Ok(LlmExecutionContext { + available: false, + request_codec_identity: LlmCodecIdentity::None, + response_codec_identity: LlmCodecIdentity::None, + request_codec: None, + response_codec: None, + }); + }; + let request = + require_execution_field(context.request.as_ref(), "request context is missing")?; + let response = + require_execution_field(context.response.as_ref(), "response context is missing")?; + let request_identity = + require_execution_field(request.codec.as_ref(), "request codec identity is missing")?; + let response_identity = require_execution_field( + response.codec.as_ref(), + "response codec identity is missing", + )?; + Ok(LlmExecutionContext { + available: true, + request_codec_identity: codec_identity_from_proto(Some(request_identity)), + response_codec_identity: codec_identity_from_proto(Some(response_identity)), + request_codec: request.codec_capability_id.as_ref().map(|capability_id| { + WorkerRequestCodec { + runtime: runtime.clone(), + capability_id: capability_id.clone(), + invocation_id: invocation_id.to_owned(), + } + }), + response_codec: response.codec_capability_id.as_ref().map(|capability_id| { + WorkerResponseCodec { + runtime: runtime.clone(), + capability_id: capability_id.clone(), + invocation_id: invocation_id.to_owned(), + } + }), + }) + } + fn sanitize_request_context(&self, invocation_id: &str) -> LlmSanitizeRequestContext { let codec = match self.sanitize_context.as_ref() { Some(nemo_relay_worker_proto::v1::llm_invocation::SanitizeContext::RequestSanitizeContext(context)) => context.codec.as_ref(), @@ -2929,6 +3117,12 @@ fn tool_payload( } } +fn require_execution_field(value: Option, detail: &str) -> Result { + value.ok_or_else(|| { + WorkerSdkError::InvalidInput(format!("malformed LLM execution codec context: {detail}")) + }) +} + fn llm_payload( payload: Option, ) -> Result { @@ -2939,6 +3133,7 @@ fn llm_payload( annotated_request: value.annotated_request, response: value.response, sanitize_context: value.sanitize_context, + execution_codec_context: value.execution_codec_context, }), _ => Err(WorkerSdkError::InvalidInput("expected llm payload".into())), } @@ -3581,3 +3776,6 @@ fn rustc_version_runtime() -> String { #[cfg(test)] #[path = "../tests/unit/codec_identity_tests.rs"] mod codec_identity_tests; +#[cfg(test)] +#[path = "../tests/unit/execution_context_tests.rs"] +mod execution_context_tests; diff --git a/crates/worker/tests/unit/execution_context_tests.rs b/crates/worker/tests/unit/execution_context_tests.rs new file mode 100644 index 000000000..0f9211f3c --- /dev/null +++ b/crates/worker/tests/unit/execution_context_tests.rs @@ -0,0 +1,98 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use super::*; + +fn disconnected_runtime() -> PluginRuntime { + PluginRuntime { + activation_id: "activation".into(), + auth_token: "token".into(), + host_endpoint: "http://127.0.0.1:1".into(), + host_channel: Arc::new(OnceCell::new()), + conditional_middleware_callbacks: Arc::new(Mutex::new(HashMap::new())), + } +} + +fn llm_payload( + execution_codec_context: Option>, +) -> LlmPayload { + LlmPayload { + model_name: "model".into(), + request: None, + annotated_request: None, + response: None, + sanitize_context: None, + execution_codec_context, + } +} + +#[test] +fn context_registration_is_opt_in_without_a_new_surface() { + let mut context = PluginContext::new(); + context.register_llm_execution_intercept("legacy", 7, |_, _, _| async { + Ok(serde_json::json!({"legacy": true})) + }); + context.register_llm_execution_intercept_with_context("context", 7, |_, _, _, _| async { + Ok(serde_json::json!({"context": true})) + }); + + let legacy = &context.handlers.registrations[0]; + let contextual = &context.handlers.registrations[1]; + assert_eq!( + legacy.surface, + RegistrationSurface::LlmExecutionIntercept as i32 + ); + assert_eq!(contextual.surface, legacy.surface); + assert_eq!(contextual.priority, legacy.priority); + assert!(!legacy.llm_execution_codec_context); + assert!(contextual.llm_execution_codec_context); +} + +#[test] +fn absent_execution_context_identifies_an_older_host() { + let payload = llm_payload(None); + + let context = payload + .execution_context(&disconnected_runtime(), "invocation") + .unwrap(); + assert!(!context.is_available()); + assert_eq!(context.request_codec_identity(), &LlmCodecIdentity::None); + assert_eq!(context.response_codec_identity(), &LlmCodecIdentity::None); + assert!(context.request_codec().is_none()); + assert!(context.response_codec().is_none()); +} + +#[test] +fn execution_context_from_a_new_host_preserves_identities_and_capabilities() { + let codec = nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: LlmCodecKind::Builtin as i32, + id: Some("openai_chat".into()), + }; + let payload = llm_payload(Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { + codec: Some(codec.clone()), + codec_capability_id: Some("request-capability".into()), + }), + response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { + codec: Some(codec), + codec_capability_id: Some("response-capability".into()), + }), + }, + ))); + + let context = payload + .execution_context(&disconnected_runtime(), "invocation") + .unwrap(); + assert!(context.is_available()); + assert_eq!( + context.request_codec_identity(), + &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert_eq!( + context.response_codec_identity(), + &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert!(context.request_codec().is_some()); + assert!(context.response_codec().is_some()); +} diff --git a/crates/worker/tests/worker_sdk_tests.rs b/crates/worker/tests/worker_sdk_tests.rs index ad6ee60a4..9128eff58 100644 --- a/crates/worker/tests/worker_sdk_tests.rs +++ b/crates/worker/tests/worker_sdk_tests.rs @@ -3286,6 +3286,7 @@ fn llm_invoke( annotated_request: annotated_request.map(json_env), response: response.map(json_env), sanitize_context: None, + execution_codec_context: None, }, )), } @@ -3310,6 +3311,7 @@ fn llm_invoke_without_request( annotated_request: None, response: None, sanitize_context: None, + execution_codec_context: None, }, )), } diff --git a/justfile b/justfile index d419197df..b7f4276a8 100644 --- a/justfile +++ b/justfile @@ -1152,6 +1152,12 @@ check-python-worker-proto: } assert pb.SUBSCRIBER == 1 assert pb.LLM_STREAM_EXECUTION_INTERCEPT == 25 + execution_context = pb.LlmInvocation.DESCRIPTOR.fields_by_name["execution_codec_context"] + assert execution_context.number == 11 + assert execution_context.containing_oneof is None + assert { + field.name for field in pb.LlmInvocation.DESCRIPTOR.oneofs_by_name["sanitize_context"].fields + } == {"request_sanitize_context", "response_sanitize_context"} tool_next = pb.DESCRIPTOR.services_by_name["RelayHostRuntime"].methods_by_name["ToolNext"] assert tool_next.output_type.full_name == "nemo.relay.worker.v1.ToolExecutionResultResponse" runtime_diagnostics = pb.DESCRIPTOR.services_by_name["RelayHostRuntime"].methods_by_name["GetRuntimeDiagnostics"] @@ -1566,7 +1572,7 @@ test-python-plugin-e2e: NEMO_RELAY_PYTHON_PLUGIN_TEST_ENVIRONMENT="$environment_ref" \ cargo nextest run --locked -p nemo-relay --features worker-grpc \ --test worker_plugin_integration \ - -E 'test(python_worker_host_runtime_mark_and_mutated_request_round_trip)' \ + -E 'test(python_worker_host_runtime_mark_and_mutated_request_round_trip) + test(python_worker_execution_codec_context_round_trips_host_codecs)' \ --no-capture \ --profile ci kill "$gateway_pid" 2>/dev/null || true diff --git a/python/plugin/src/nemo_relay_plugin/__init__.py b/python/plugin/src/nemo_relay_plugin/__init__.py index 1d255301a..b71d5ea44 100644 --- a/python/plugin/src/nemo_relay_plugin/__init__.py +++ b/python/plugin/src/nemo_relay_plugin/__init__.py @@ -29,6 +29,8 @@ LlmCodecIdentity: Typed discriminator for the active LLM codec. LlmSanitizeRequestContext: Per-call context supplied to an LLM request sanitizer. LlmSanitizeResponseContext: Per-call context supplied to an LLM response sanitizer. + LlmExecutionContext: Invocation-scoped codec context supplied to an LLM + execution intercept. WorkerRequestCodec: Invocation-scoped async proxy for an active request codec. WorkerResponseCodec: Invocation-scoped async proxy for an active response codec. AnnotatedLlmRequest: An annotated Relay LLM request represented as a JSON @@ -71,7 +73,11 @@ LlmConditionalCallback: LLM execution guardrail callback. LlmRequestCallback: LLM request intercept callback. LlmExecutionCallback: Unary LLM execution intercept callback. + LlmExecutionWithContextCallback: Unary LLM execution intercept callback + with codec context. LlmStreamExecutionCallback: Streaming LLM execution intercept callback. + LlmStreamExecutionWithContextCallback: Streaming LLM execution intercept + callback with request codec context. Public authoring types: WorkerPlugin: Base validation and registration contract for a plugin. @@ -101,6 +107,8 @@ LlmCodecIdentity, LlmConditionalCallback, LlmExecutionCallback, + LlmExecutionContext, + LlmExecutionWithContextCallback, LlmNext, LlmOptimizationContribution, LlmOptimizationDataSchema, @@ -117,6 +125,7 @@ LlmSanitizeResponseCallback, LlmSanitizeResponseContext, LlmStreamExecutionCallback, + LlmStreamExecutionWithContextCallback, LlmStreamNext, LogSeverity, MetricKind, @@ -164,6 +173,8 @@ "LlmConditionalCallback", "LlmCodecIdentity", "LlmExecutionCallback", + "LlmExecutionContext", + "LlmExecutionWithContextCallback", "LogSeverity", "MetricKind", "MetricMeasurement", @@ -185,6 +196,7 @@ "LlmSanitizeResponseCallback", "LlmStreamNext", "LlmStreamExecutionCallback", + "LlmStreamExecutionWithContextCallback", "PluginContext", "PluginRuntime", "RuntimeDiagnostic", diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index aeb0f159c..ba7f103fe 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -258,6 +258,22 @@ async def decode(self, response: Json) -> Json: return await self._runtime._decode_llm_codec_response(self._capability_id, self._invocation_id, response) +@dataclass(frozen=True) +class LlmExecutionContext: + """Invocation-scoped codec identities and operations for LLM execution middleware. + + The identities describe the codecs selected when Relay created the managed + invocation. Rewriting a request into another provider's wire format does + not select a new codec; incompatible codec operations fail. + """ + + available: bool + request_codec_identity: LlmCodecIdentity + response_codec_identity: LlmCodecIdentity + request_codec: WorkerRequestCodec | None = field(default=None, repr=False, compare=False) + response_codec: WorkerResponseCodec | None = field(default=None, repr=False, compare=False) + + def _llm_codec_identity(invocation: pb.LlmInvocation) -> LlmCodecIdentity: """Return the codec identity from a worker invocation.""" context = getattr(invocation, invocation.WhichOneof("sanitize_context") or "", None) @@ -282,6 +298,42 @@ def _llm_codec_capability(invocation: pb.LlmInvocation) -> str | None: return context.codec_capability_id if context is not None and context.HasField("codec_capability_id") else None +def _llm_execution_context( + invocation: pb.LlmInvocation, + runtime: "PluginRuntime", + invocation_id: str, +) -> LlmExecutionContext: + if not invocation.HasField("execution_codec_context"): + return LlmExecutionContext( + available=False, + request_codec_identity=LlmCodecIdentity("none"), + response_codec_identity=LlmCodecIdentity("none"), + ) + context = invocation.execution_codec_context + if not context.HasField("request") or not context.request.HasField("codec"): + raise WorkerSdkError("malformed LLM execution codec context: request codec identity is missing") + if not context.HasField("response") or not context.response.HasField("codec"): + raise WorkerSdkError("malformed LLM execution codec context: response codec identity is missing") + + request_id = context.request.codec_capability_id if context.request.HasField("codec_capability_id") else None + response_id = context.response.codec_capability_id if context.response.HasField("codec_capability_id") else None + request_identity = _codec_identity( + context.request.codec.kind, + context.request.codec.id if context.request.codec.HasField("id") else None, + ) + response_identity = _codec_identity( + context.response.codec.kind, + context.response.codec.id if context.response.codec.HasField("id") else None, + ) + return LlmExecutionContext( + available=True, + request_codec_identity=request_identity, + response_codec_identity=response_identity, + request_codec=(WorkerRequestCodec(runtime, request_id, invocation_id) if request_id else None), + response_codec=(WorkerResponseCodec(runtime, response_id, invocation_id) if response_id else None), + ) + + WORKER_PROTOCOL = "grpc-v1" JSON_SCHEMA = "nemo.relay.Json@1" DATA_SCHEMA_SCHEMA = "nemo.relay.DataSchema@1" @@ -1056,10 +1108,17 @@ def register(self, ctx: PluginContext, config: Json) -> None | Awaitable[None]: LlmRequestInterceptOutcome | Awaitable[LlmRequestInterceptOutcome], ] LlmExecutionCallback: TypeAlias = Callable[[str, LlmRequest, "LlmNext"], Json | Awaitable[Json]] +LlmExecutionWithContextCallback: TypeAlias = Callable[ + [str, LlmRequest, LlmExecutionContext, "LlmNext"], Json | Awaitable[Json] +] LlmStreamExecutionCallback: TypeAlias = Callable[ [str, LlmRequest, "LlmStreamNext"], Iterable[Json] | AsyncIterator[Json] | Awaitable[Iterable[Json] | AsyncIterator[Json]], ] +LlmStreamExecutionWithContextCallback: TypeAlias = Callable[ + [str, LlmRequest, LlmExecutionContext, "LlmStreamNext"], + Iterable[Json] | AsyncIterator[Json] | Awaitable[Iterable[Json] | AsyncIterator[Json]], +] @dataclass(slots=True) @@ -1081,8 +1140,8 @@ class _Handlers: llm_sanitize_responses: dict[str, LlmSanitizeResponseCallback] llm_conditionals: dict[str, LlmConditionalCallback] llm_requests: dict[str, LlmRequestCallback] - llm_executions: dict[str, LlmExecutionCallback] - llm_stream_executions: dict[str, LlmStreamExecutionCallback] + llm_executions: dict[str, LlmExecutionWithContextCallback] + llm_stream_executions: dict[str, LlmStreamExecutionWithContextCallback] @classmethod def empty(cls) -> _Handlers: @@ -1493,6 +1552,25 @@ def register_llm_execution_intercept( priority: Execution order. Lower values run first. """ self._push_registration(name, pb.LLM_EXECUTION_INTERCEPT, priority, False) + self._handlers.llm_executions[name] = lambda model, request, _context, next_call: callback( + model, request, next_call + ) + + def register_llm_execution_intercept_with_context( + self, + name: str, + callback: LlmExecutionWithContextCallback, + *, + priority: int = 0, + ) -> None: + """Register LLM execution middleware with invocation-scoped codecs.""" + self._push_registration( + name, + pb.LLM_EXECUTION_INTERCEPT, + priority, + False, + llm_execution_codec_context=True, + ) self._handlers.llm_executions[name] = callback def register_llm_stream_execution_intercept( @@ -1520,9 +1598,40 @@ def register_llm_stream_execution_intercept( error. """ self._push_registration(name, pb.LLM_STREAM_EXECUTION_INTERCEPT, priority, False) + self._handlers.llm_stream_executions[name] = lambda model, request, _context, next_call: callback( + model, request, next_call + ) + + def register_llm_stream_execution_intercept_with_context( + self, + name: str, + callback: LlmStreamExecutionWithContextCallback, + *, + priority: int = 0, + ) -> None: + """Register streaming middleware with request codec access. + + The context identifies the response codec but does not expose a + response decoder because stream chunks are not complete responses. + """ + self._push_registration( + name, + pb.LLM_STREAM_EXECUTION_INTERCEPT, + priority, + False, + llm_execution_codec_context=True, + ) self._handlers.llm_stream_executions[name] = callback - def _push_registration(self, name: str, surface: int, priority: int, break_chain: bool) -> None: + def _push_registration( + self, + name: str, + surface: int, + priority: int, + break_chain: bool, + *, + llm_execution_codec_context: bool = False, + ) -> None: if any( registration.local_name == name and registration.surface == surface for registration in self._handlers.registrations @@ -1534,6 +1643,7 @@ def _push_registration(self, name: str, surface: int, priority: int, break_chain surface=surface, priority=priority, break_chain=break_chain, + llm_execution_codec_context=llm_execution_codec_context, ) ) @@ -2478,9 +2588,10 @@ async def _produce_stream(self, request: Any, queue: asyncio.Queue[Any]) -> None handler = self._handler(self._handlers.llm_stream_executions, request.registration_name) payload = _require_payload(request, "llm") llm_request = _decode_required_envelope(payload.request, "llm request", LLM_REQUEST_SCHEMA) + execution_context = _llm_execution_context(payload, self._runtime, request.invocation_id) next_call = LlmStreamNext(self._runtime, request.continuation_id) with _bind_invocation_scope(request): - stream = await _maybe_await(handler(payload.model_name, llm_request, next_call)) + stream = await _maybe_await(handler(payload.model_name, llm_request, execution_context, next_call)) async for value in _as_async_iter(stream): await queue.put(pb.StreamChunk(value=_json_envelope(JSON_SCHEMA, value))) except asyncio.CancelledError: @@ -2653,6 +2764,7 @@ async def _invoke_llm_result(self, request: Any) -> Any: self._handler(self._handlers.llm_executions, request.registration_name)( payload.model_name, _decode_required_envelope(payload.request, "llm request", LLM_REQUEST_SCHEMA), + _llm_execution_context(payload, self._runtime, request.invocation_id), LlmNext(self._runtime, request.continuation_id), ) ) diff --git a/python/tests/plugin/test_public_api_docstrings.py b/python/tests/plugin/test_public_api_docstrings.py index 08b01c87d..fb6859568 100644 --- a/python/tests/plugin/test_public_api_docstrings.py +++ b/python/tests/plugin/test_public_api_docstrings.py @@ -36,7 +36,9 @@ "LlmConditionalCallback", "LlmRequestCallback", "LlmExecutionCallback", + "LlmExecutionWithContextCallback", "LlmStreamExecutionCallback", + "LlmStreamExecutionWithContextCallback", } diff --git a/python/tests/plugin/test_worker_sdk.py b/python/tests/plugin/test_worker_sdk.py index bc9269a12..492739ccc 100644 --- a/python/tests/plugin/test_worker_sdk.py +++ b/python/tests/plugin/test_worker_sdk.py @@ -736,6 +736,14 @@ def test_generated_proto_matches_worker_contract() -> None: "Shutdown", } assert pb.InvokeRequest.DESCRIPTOR.fields_by_name["auth_token"].number == 7 + assert pb.Registration.DESCRIPTOR.fields_by_name["llm_execution_codec_context"].number == 6 + execution_context = pb.LlmInvocation.DESCRIPTOR.fields_by_name["execution_codec_context"] + assert execution_context.number == 11 + assert execution_context.containing_oneof is None + assert {field.name for field in pb.LlmInvocation.DESCRIPTOR.oneofs_by_name["sanitize_context"].fields} == { + "request_sanitize_context", + "response_sanitize_context", + } assert pb.HealthRequest.DESCRIPTOR.fields_by_name["activation_id"].number == 1 assert pb.HealthRequest.DESCRIPTOR.fields_by_name["auth_token"].number == 2 assert pb.SUBSCRIBER == 1 @@ -1108,6 +1116,87 @@ def test_plugin_context_registers_llm_sanitizers_under_standard_names() -> None: ] +def test_execution_codec_context_is_opt_in_on_the_existing_surface() -> None: + context = PluginContext() + + async def legacy(_name: str, request: Json, next_call: Any) -> Json: + return await next_call.call(request) + + async def contextual(_name: str, request: Json, _context: Any, next_call: Any) -> Json: + return await next_call.call(request) + + context.register_llm_execution_intercept("legacy", legacy, priority=7) + context.register_llm_execution_intercept_with_context("contextual", contextual, priority=7) + + legacy_registration, contextual_registration = context._handlers.registrations + assert legacy_registration.surface == pb.LLM_EXECUTION_INTERCEPT + assert contextual_registration.surface == legacy_registration.surface + assert contextual_registration.priority == legacy_registration.priority + assert not legacy_registration.llm_execution_codec_context + assert contextual_registration.llm_execution_codec_context + + +async def test_contextual_execution_callback_handles_old_and_new_hosts() -> None: + seen: list[bool] = [] + + class ContextualExecutionPlugin(WorkerPlugin): + plugin_id = "tests.contextual_execution" + + def register(self, ctx: PluginContext, config: Json) -> None: + del config + + async def execution(name: str, request: Json, context: Any, next_call: Any) -> Json: + del name + seen.append(context.available) + if context.available: + assert context.request_codec_identity == plugin_api.LlmCodecIdentity("builtin", "openai_chat") + assert context.response_codec_identity == plugin_api.LlmCodecIdentity("builtin", "openai_chat") + assert context.request_codec is not None + assert context.response_codec is not None + await context.request_codec.decode(request) + result = await next_call.call(request) + if context.available: + await context.response_codec.decode(result) + return result + + ctx.register_llm_execution_intercept_with_context("execution", execution) + + host = RecordingHostStub() + service = _service(ContextualExecutionPlugin(), host) + register = await _register(service) + assert register.registrations[0].llm_execution_codec_context + + new_context = pb.LlmExecutionCodecContext( + request=pb.LlmSanitizeRequestContext( + codec=pb.LlmCodecIdentity(kind=pb.LLM_CODEC_KIND_BUILTIN, id="openai_chat"), + codec_capability_id="request-capability", + ), + response=pb.LlmSanitizeResponseContext( + codec=pb.LlmCodecIdentity(kind=pb.LLM_CODEC_KIND_BUILTIN, id="openai_chat"), + codec_capability_id="response-capability", + ), + ) + for execution_context in [None, new_context]: + payload = _llm_payload(request={"content": {"model": "gpt-test"}}) + if execution_context is not None: + payload.execution_codec_context.CopyFrom(execution_context) + result = await _invoke_json_async( + service, + "execution", + pb.LLM_EXECUTION_INTERCEPT, + payload=payload, + ) + assert "next_llm" in result + + assert seen == [False, True] + codec_capabilities = [ + request.codec_capability_id + for request in host.requests + if isinstance(request, (pb.LlmCodecDecodeRequest, pb.LlmCodecDecodeResponse)) + ] + assert codec_capabilities == ["request-capability", "response-capability"] + + async def test_llm_sanitizers_receive_codec_context_and_can_omit_payloads() -> None: seen: list[tuple[str, LlmSanitizeRequestContext | LlmSanitizeResponseContext]] = [] From 48b634ee6668bbe2eca1524456e034fb86c5f549 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 22 Sep 2026 10:29:53 -0400 Subject: [PATCH 02/22] feat!: pass codec context to LLM execution intercepts Signed-off-by: Alex Fournier --- crates/adaptive/src/acg_component.rs | 4 +- .../adaptive/src/response_cache/intercept.rs | 4 +- crates/adaptive/src/runtime/features.rs | 8 +- .../integration/runtime_integration_tests.rs | 4 +- .../tests/unit/acg_component_tests.rs | 30 +- .../tests/unit/runtime_features_tests.rs | 24 +- .../coverage/daemon/worker_managed_tests.rs | 2 +- .../cli/tests/coverage/shared/server_tests.rs | 12 +- crates/core/src/api/llm.rs | 10 +- crates/core/src/api/registry.rs | 76 +- crates/core/src/api/runtime.rs | 5 +- crates/core/src/api/runtime/callbacks.rs | 17 +- .../src/api/runtime/llm_execution_context.rs | 255 ++++--- crates/core/src/api/runtime/state.rs | 68 +- crates/core/src/context/registries.rs | 11 +- crates/core/src/plugin.rs | 44 -- crates/core/src/plugin/dynamic.rs | 31 + crates/core/src/plugin/dynamic/native.rs | 551 ++++++++++++--- crates/core/src/plugin/dynamic/worker.rs | 134 ++-- .../src/plugins/nemo_guardrails/python.rs | 19 +- .../src/plugins/nemo_guardrails/remote.rs | 11 +- .../tests/fixtures/native_plugin/src/lib.rs | 87 ++- .../tests/fixtures/worker_plugin/src/main.rs | 4 +- .../tests/integration/api_surface_tests.rs | 8 +- .../tests/integration/middleware_tests.rs | 34 +- .../tests/integration/native_plugin_tests.rs | 93 ++- .../tests/integration/worker_plugin_tests.rs | 50 +- .../core/tests/unit/dynamic_worker_tests.rs | 137 ++-- crates/core/tests/unit/llm_api_tests.rs | 586 ++++++++++++++-- crates/core/tests/unit/native_plugin_tests.rs | 430 +++++++++--- crates/core/tests/unit/plugin_tests.rs | 18 +- crates/ffi/nemo_relay.h | 23 +- crates/ffi/src/callable.rs | 151 +++- crates/ffi/tests/integration/api_tests.rs | 7 +- .../tests/integration/callable_extra_tests.rs | 23 +- crates/ffi/tests/unit/api/registry_tests.rs | 10 + crates/ffi/tests/unit/api_tests.rs | 7 +- crates/ffi/tests/unit/callable_tests.rs | 192 ++++- crates/node/plugin.d.ts | 9 +- crates/node/root-types.d.ts | 8 + crates/node/src/api/mod.rs | 26 +- crates/node/src/callable.rs | 100 ++- crates/node/src/promise_call.rs | 40 +- crates/node/tests/adaptive_tests.mjs | 6 +- crates/node/tests/llm_tests.mjs | 261 +++++-- crates/node/tests/scope_local_tests.mjs | 80 ++- crates/node/tests/typed_tests.mjs | 2 +- crates/plugin/README.md | 16 +- crates/plugin/src/async_sdk.rs | 327 +++++++-- crates/plugin/src/lib.rs | 395 ++++++++++- crates/plugin/tests/typed_callbacks.rs | 660 +++++++++++++++--- crates/python/src/py_api/mod.rs | 8 +- crates/python/src/py_callable.rs | 77 +- crates/python/src/py_types/core.rs | 31 +- crates/python/src/py_types/mod.rs | 1 + .../python/tests/coverage/coverage_tests.rs | 75 +- .../tests/coverage/py_api_coverage_tests.rs | 4 +- .../coverage/py_callable_coverage_tests.rs | 16 +- .../coverage/py_plugin_coverage_tests.rs | 12 +- crates/worker-proto/README.md | 8 + .../nemo/relay/worker/v1/plugin_worker.proto | 8 +- crates/worker-proto/tests/proto_tests.rs | 40 +- crates/worker/README.md | 10 + crates/worker/src/lib.rs | 180 ++--- .../tests/unit/execution_context_tests.rs | 118 +++- crates/worker/tests/worker_sdk_tests.rs | 55 +- docs/about-nemo-relay/release-notes/index.mdx | 16 + docs/build-plugins/about.mdx | 13 +- .../language-binding/register-behavior.mdx | 26 +- docs/build-plugins/native/about.mdx | 11 +- .../native/build-and-package.mdx | 2 +- .../native/native-abi-reference.mdx | 33 +- docs/build-plugins/native/wrap-execution.mdx | 6 +- .../package-discoverable-plugins.mdx | 12 +- docs/build-plugins/workers/about.mdx | 7 + .../workers/grpc-v1-protocol.mdx | 22 +- .../workers/middleware-and-continuations.mdx | 30 +- docs/build-plugins/workers/python.mdx | 6 +- docs/build-plugins/workers/rust.mdx | 6 +- docs/reference/migration-guides.mdx | 34 + .../language-binding-plugin/node/main.mjs | 19 +- .../node/test-plugin.mjs | 13 +- .../language-binding-plugin/python/main.py | 3 + .../language-binding-plugin/rust/src/lib.rs | 4 +- examples/python-grpc-worker-plugin/README.md | 7 +- .../worker.py | 3 +- .../relay-plugin.toml | 4 +- .../tests/test_worker.py | 28 +- examples/rust-grpc-worker-plugin/README.md | 5 + .../rust-grpc-worker-plugin/relay-plugin.toml | 2 +- examples/rust-grpc-worker-plugin/src/lib.rs | 4 +- .../rust-grpc-worker-plugin/tests/config.rs | 2 +- .../tests/lifecycle.rs | 2 +- examples/rust-native-plugin/README.md | 5 + examples/rust-native-plugin/relay-plugin.toml | 2 +- examples/rust-native-plugin/src/execution.rs | 4 +- .../rust-native-plugin/tests/lifecycle.rs | 2 +- go/nemo_relay/adaptive_plugin_test.go | 8 +- go/nemo_relay/callbacks.go | 41 +- go/nemo_relay/intercepts/intercepts_test.go | 8 +- go/nemo_relay/llm_test.go | 129 +++- go/nemo_relay/nemo_relay.go | 5 +- go/nemo_relay/plugin.go | 5 +- go/nemo_relay/scope_local_test.go | 4 +- python/nemo_relay/__init__.py | 18 +- python/nemo_relay/__init__.pyi | 16 +- python/nemo_relay/_native.pyi | 17 +- python/nemo_relay/intercepts.py | 14 +- python/nemo_relay/scope_local.py | 12 +- .../plugin/src/nemo_relay_plugin/__init__.py | 13 +- python/plugin/src/nemo_relay_plugin/_api.py | 130 ++-- .../plugin/test_public_api_docstrings.py | 2 - python/tests/plugin/test_worker_sdk.py | 149 ++-- python/tests/test_adaptive.py | 8 +- python/tests/test_llm.py | 212 +++++- python/tests/test_scope_local.py | 12 +- 116 files changed, 5196 insertions(+), 1693 deletions(-) diff --git a/crates/adaptive/src/acg_component.rs b/crates/adaptive/src/acg_component.rs index cea5bfc20..b5f192ea7 100644 --- a/crates/adaptive/src/acg_component.rs +++ b/crates/adaptive/src/acg_component.rs @@ -666,7 +666,7 @@ pub(crate) fn create_acg_llm_execution_intercept( plugin: Arc, ) -> LlmExecutionFn { Arc::new( - move |_name: &str, request: LlmRequest, next: LlmExecutionNextFn| { + move |_name: &str, request: LlmRequest, _context, next: LlmExecutionNextFn| { let cache = hot_cache.clone(); let agent_id = agent_id.clone(); let provider = provider.clone(); @@ -689,7 +689,7 @@ pub(crate) fn create_acg_llm_stream_execution_intercept( plugin: Arc, ) -> LlmStreamExecutionFn { Arc::new( - move |_name: &str, request: LlmRequest, next: LlmStreamExecutionNextFn| { + move |_name: &str, request: LlmRequest, _context, next: LlmStreamExecutionNextFn| { let cache = hot_cache.clone(); let agent_id = agent_id.clone(); let provider = provider.clone(); diff --git a/crates/adaptive/src/response_cache/intercept.rs b/crates/adaptive/src/response_cache/intercept.rs index 01d876bfe..5172d8be2 100644 --- a/crates/adaptive/src/response_cache/intercept.rs +++ b/crates/adaptive/src/response_cache/intercept.rs @@ -128,7 +128,7 @@ pub(crate) fn make_intercept( config: Arc, ) -> LlmExecutionFn { Arc::new( - move |provider: &str, request: LlmRequest, next: LlmExecutionNextFn| { + move |provider: &str, request: LlmRequest, _context, next: LlmExecutionNextFn| { let store = Arc::clone(&store); let config = Arc::clone(&config); let provider = provider.to_string(); @@ -150,7 +150,7 @@ pub(crate) fn make_stream_intercept( config: Arc, ) -> LlmStreamExecutionFn { Arc::new( - move |provider: &str, request: LlmRequest, next: LlmStreamExecutionNextFn| { + move |provider: &str, request: LlmRequest, _context, next: LlmStreamExecutionNextFn| { let store = Arc::clone(&store); let config = Arc::clone(&config); let provider = provider.to_string(); diff --git a/crates/adaptive/src/runtime/features.rs b/crates/adaptive/src/runtime/features.rs index 4814add67..e361050cc 100644 --- a/crates/adaptive/src/runtime/features.rs +++ b/crates/adaptive/src/runtime/features.rs @@ -806,7 +806,7 @@ impl AdaptiveFeature for AcgFeature { ctx.register_llm_execution_intercept( &self.execution_name, self.priority, - Arc::new(move |name, request, next| { + Arc::new(move |name, request, context, next| { let execution_intercept = execution_intercept.clone(); let bound_scopes = bound_scopes.clone(); let name = name.to_string(); @@ -818,7 +818,7 @@ impl AdaptiveFeature for AcgFeature { if has_bound_scopes { return next(request).await; } - execution_intercept(&name, request, next).await + execution_intercept(&name, request, context, next).await }) }), )?; @@ -832,7 +832,7 @@ impl AdaptiveFeature for AcgFeature { ctx.register_llm_stream_execution_intercept( &self.stream_name, self.priority, - Arc::new(move |name, request, next| { + Arc::new(move |name, request, context, next| { let stream_intercept = stream_intercept.clone(); let bound_scopes = bound_scopes.clone(); let name = name.to_string(); @@ -844,7 +844,7 @@ impl AdaptiveFeature for AcgFeature { if has_bound_scopes { return next(request).await; } - stream_intercept(&name, request, next).await + stream_intercept(&name, request, context, next).await }) }), ) diff --git a/crates/adaptive/tests/integration/runtime_integration_tests.rs b/crates/adaptive/tests/integration/runtime_integration_tests.rs index 1c92455fc..6e364b4cb 100644 --- a/crates/adaptive/tests/integration/runtime_integration_tests.rs +++ b/crates/adaptive/tests/integration/runtime_integration_tests.rs @@ -753,7 +753,7 @@ impl Plugin for HeaderPlugin { ctx.register_llm_execution_intercept( "llm_exec_plugin", priority, - Arc::new(|_name, request, next| { + Arc::new(|_name, request, _context, next| { Box::pin(async move { let mut response = next(request).await?; if let Json::Object(ref mut map) = response { @@ -766,7 +766,7 @@ impl Plugin for HeaderPlugin { ctx.register_llm_stream_execution_intercept( "llm_stream_exec_plugin", priority, - Arc::new(|_name, request, next| { + Arc::new(|_name, request, _context, next| { Box::pin(async move { let mut stream = next(request).await?; let mut chunks = Vec::new(); diff --git a/crates/adaptive/tests/unit/acg_component_tests.rs b/crates/adaptive/tests/unit/acg_component_tests.rs index dd833b47a..3aa637603 100644 --- a/crates/adaptive/tests/unit/acg_component_tests.rs +++ b/crates/adaptive/tests/unit/acg_component_tests.rs @@ -17,7 +17,9 @@ use crate::config::AcgComponentConfig; use crate::storage::memory::InMemoryBackend; use crate::storage::traits::StorageBackendDyn; use nemo_relay::api::llm::LlmRequest; -use nemo_relay::api::runtime::{LlmExecutionNextFn, LlmJsonStream, LlmStreamExecutionNextFn}; +use nemo_relay::api::runtime::{ + LlmExecutionContext, LlmExecutionNextFn, LlmJsonStream, LlmStreamExecutionNextFn, +}; use nemo_relay::codec::request::{AnnotatedLlmRequest, Message, MessageContent}; use serde_json::{Value, json}; use tokio_stream::StreamExt; @@ -698,7 +700,7 @@ async fn acg_component_stream_execution_intercept_rewrites_streaming_requests() }) }); - let mut stream = intercept("anthropic", request, next) + let mut stream = intercept("anthropic", request, LlmExecutionContext::default(), next) .await .expect("stream intercept should succeed"); let first = stream @@ -1302,7 +1304,7 @@ async fn acg_component_execution_intercept_rewrites_non_streaming_requests() { ); let next: LlmExecutionNextFn = Arc::new(|req| Box::pin(async move { Ok(req.content) })); - let result = intercept("anthropic", request, next) + let result = intercept("anthropic", request, LlmExecutionContext::default(), next) .await .expect("execution intercept should succeed"); @@ -1515,9 +1517,14 @@ async fn acg_component_execution_intercept_passes_original_request_when_translat ); let next: LlmExecutionNextFn = Arc::new(|req| Box::pin(async move { Ok(req.content) })); - let result = intercept("anthropic", invalid_request.clone(), next) - .await - .expect("execution intercept should succeed"); + let result = intercept( + "anthropic", + invalid_request.clone(), + LlmExecutionContext::default(), + next, + ) + .await + .expect("execution intercept should succeed"); assert_eq!(result, invalid_request.content); } @@ -1544,9 +1551,14 @@ async fn acg_component_stream_execution_intercept_passes_original_request_when_t }) }); - let mut stream = intercept("anthropic", invalid_request.clone(), next) - .await - .expect("stream intercept should pass through"); + let mut stream = intercept( + "anthropic", + invalid_request.clone(), + LlmExecutionContext::default(), + next, + ) + .await + .expect("stream intercept should pass through"); let first = stream .next() .await diff --git a/crates/adaptive/tests/unit/runtime_features_tests.rs b/crates/adaptive/tests/unit/runtime_features_tests.rs index 34dd15e13..30ff246d8 100644 --- a/crates/adaptive/tests/unit/runtime_features_tests.rs +++ b/crates/adaptive/tests/unit/runtime_features_tests.rs @@ -211,7 +211,7 @@ fn assert_llm_execution_intercept_registered(name: &str) { register_llm_execution_intercept( name, i32::MAX, - Arc::new(|_name, request, next| next(request)), + Arc::new(|_name, request, _context, next| next(request)), ), name, ); @@ -221,7 +221,7 @@ fn assert_llm_execution_intercept_absent(name: &str) { register_llm_execution_intercept( name, i32::MAX, - Arc::new(|_name, request, next| next(request)), + Arc::new(|_name, request, _context, next| next(request)), ) .unwrap(); deregister_llm_execution_intercept(name).unwrap(); @@ -232,7 +232,7 @@ fn assert_llm_stream_execution_intercept_registered(name: &str) { register_llm_stream_execution_intercept( name, i32::MAX, - Arc::new(|_name, request, next| next(request)), + Arc::new(|_name, request, _context, next| next(request)), ), name, ); @@ -242,7 +242,7 @@ fn assert_llm_stream_execution_intercept_absent(name: &str) { register_llm_stream_execution_intercept( name, i32::MAX, - Arc::new(|_name, request, next| next(request)), + Arc::new(|_name, request, _context, next| next(request)), ) .unwrap(); deregister_llm_stream_execution_intercept(name).unwrap(); @@ -807,13 +807,13 @@ async fn registration_context_registers_all_supported_callback_types() { ctx.register_llm_execution_intercept( "adaptive_test_execution", 6, - Arc::new(|_name, request, _next| Box::pin(async move { Ok(request.content) })), + Arc::new(|_name, request, _context, _next| Box::pin(async move { Ok(request.content) })), ) .unwrap(); ctx.register_llm_stream_execution_intercept( "adaptive_test_stream", 7, - Arc::new(|_name, request, _next| { + Arc::new(|_name, request, _context, _next| { Box::pin(async move { Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( request.content @@ -1017,7 +1017,7 @@ async fn acg_feature_reports_execution_registration_conflicts() { register_llm_execution_intercept( &execution_name, 1, - Arc::new(|_name, request, next| next(request)), + Arc::new(|_name, request, _context, next| next(request)), ) .unwrap(); @@ -1146,8 +1146,12 @@ async fn response_cache_feature_cleans_up_when_llm_registration_conflicts() { Uuid::now_v7(), ); let name = feature.name.clone(); - register_llm_execution_intercept(&name, 1, Arc::new(|_name, request, next| next(request))) - .unwrap(); + register_llm_execution_intercept( + &name, + 1, + Arc::new(|_name, request, _context, next| next(request)), + ) + .unwrap(); let error = { let mut ctx = RegistrationContext::new(&mut runtime); @@ -1180,7 +1184,7 @@ async fn response_cache_feature_cleans_up_when_stream_registration_conflicts() { register_llm_stream_execution_intercept( &stream_name, 1, - Arc::new(|_name, request, next| next(request)), + Arc::new(|_name, request, _context, next| next(request)), ) .unwrap(); diff --git a/crates/cli/tests/coverage/daemon/worker_managed_tests.rs b/crates/cli/tests/coverage/daemon/worker_managed_tests.rs index ce5a9fecd..b0bca14f8 100644 --- a/crates/cli/tests/coverage/daemon/worker_managed_tests.rs +++ b/crates/cli/tests/coverage/daemon/worker_managed_tests.rs @@ -1479,7 +1479,7 @@ async fn managed_runtime_rejects_response_mutating_execution_middleware() { register_llm_execution_intercept( INTERCEPT, 1, - Arc::new(|_name, _request, _next| Box::pin(async { Ok(json!({})) })), + Arc::new(|_name, _request, _context, _next| Box::pin(async { Ok(json!({})) })), ) .expect("register execution middleware"); diff --git a/crates/cli/tests/coverage/shared/server_tests.rs b/crates/cli/tests/coverage/shared/server_tests.rs index 04bbacc92..535b289fd 100644 --- a/crates/cli/tests/coverage/shared/server_tests.rs +++ b/crates/cli/tests/coverage/shared/server_tests.rs @@ -4964,7 +4964,7 @@ async fn gateway_concurrent_next_uses_canonical_selected_buffered_response() { register_llm_execution_intercept( INTERCEPT_NAME, 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { let release_losing = release_from_intercept.clone(); Box::pin(async move { if request.content["messages"][0]["content"].as_str() != Some(MARKER) { @@ -5052,7 +5052,7 @@ async fn gateway_concurrent_next_uses_canonical_selected_streaming_response() { register_llm_stream_execution_intercept( INTERCEPT_NAME, 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { let release_losing = release_from_intercept.clone(); Box::pin(async move { if request.content["messages"][0]["content"].as_str() != Some(MARKER) { @@ -5135,7 +5135,7 @@ async fn gateway_concurrent_next_relays_selected_buffered_failure() { register_llm_execution_intercept( INTERCEPT_NAME, 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { let release_losing = release_from_intercept.clone(); Box::pin(async move { if request.content["messages"][0]["content"].as_str() != Some(MARKER) { @@ -5214,7 +5214,7 @@ async fn gateway_concurrent_next_relays_selected_streaming_failure() { register_llm_stream_execution_intercept( INTERCEPT_NAME, 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { let release_losing = release_from_intercept.clone(); Box::pin(async move { if request.content["messages"][0]["content"].as_str() != Some(MARKER) { @@ -5514,7 +5514,7 @@ async fn gateway_surfaces_post_upstream_intercept_rejection_instead_of_relaying_ register_llm_execution_intercept( INTERCEPT_NAME, 1, - Arc::new(|_name, request, next| { + Arc::new(|_name, request, _context, next| { Box::pin(async move { let marked = request.content["messages"][0]["content"].as_str() == Some(MARKER); let response = next(request).await?; @@ -6165,7 +6165,7 @@ async fn model_call_policy_still_applies_on_a_named_upstream() { register_llm_execution_intercept( INTERCEPT, 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { Box::pin(async move { if request.content["messages"][0]["content"].as_str() == Some("blocked by policy") { return Err(nemo_relay::error::FlowError::GuardrailRejected( diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index d5d8a88b2..76b673e21 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -29,9 +29,9 @@ use crate::api::runtime::subscriber_dispatcher::{ dispatch_sanitized_event, dispatch_transformed_event, register_pending_publication, }; use crate::api::runtime::{ - EventSubscriberFn, LlmCollectorFn, LlmExecutionCodecContext, LlmExecutionNextFn, - LlmFinalizerFn, LlmJsonStream, LlmSanitizeRequestContext, LlmSanitizeResponseContext, - LlmStreamExecutionNextFn, MiddlewareContinuationContext, + EventSubscriberFn, LlmCollectorFn, LlmExecutionContext, LlmExecutionNextFn, LlmFinalizerFn, + LlmJsonStream, LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn, + MiddlewareContinuationContext, }; use crate::api::runtime::{ScopeStackHandle, capture_trace_context, current_scope_stack}; use crate::api::scope::event; @@ -1739,7 +1739,7 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { ); let execution_name = name.clone(); let event_uuid = handle.uuid; - let execution_context = LlmExecutionCodecContext::for_codecs(request_codec, &response_codec); + let execution_context = LlmExecutionContext::for_unary_codecs(request_codec, &response_codec); let execution = with_active_event_trace_context( event_uuid, Some(active_trace_context), @@ -1972,7 +1972,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu let execution_name = name.clone(); let event_uuid = handle.uuid; let stream_started_at = Instant::now(); - let execution_context = LlmExecutionCodecContext::for_codecs(request_codec, &response_codec); + let execution_context = LlmExecutionContext::for_streaming_codec(request_codec); let execution = with_active_event_trace_context( event_uuid, Some(active_trace_context), diff --git a/crates/core/src/api/registry.rs b/crates/core/src/api/registry.rs index 5d8e16f56..10f627e95 100644 --- a/crates/core/src/api/registry.rs +++ b/crates/core/src/api/registry.rs @@ -8,10 +8,7 @@ use crate::api::runtime::{ ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, LlmConditionalFn, LlmExecutionFn, LlmRequestInterceptFn, LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionFn, ToolConditionalFn, ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn, - adapt_llm_execution_fn, adapt_llm_stream_execution_fn, }; -#[cfg(feature = "worker-grpc")] -use crate::api::runtime::{ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn}; use crate::api::runtime::{current_scope_stack, global_context}; use crate::api::shared::ensure_runtime_owner; use crate::error::{FlowError, Result}; @@ -335,15 +332,6 @@ pub(crate) type Intercept = RegistryRecord>; /// A priority-ordered execution intercept registration record. pub(crate) type ExecutionIntercept = RegistryRecord; -macro_rules! adapt_execution_callable { - ($callable:expr) => { - $callable - }; - ($callable:expr, $adapter:path) => { - $adapter($callable) - }; -} - macro_rules! global_guardrail_registry_api { ( $(#[$register_meta:meta])* @@ -475,7 +463,6 @@ macro_rules! global_execution_registry_api { $deregister_name:ident, $field:ident, $fn_type:ty - $(, $adapter:path)? ) => { $(#[$register_meta])* /// @@ -498,11 +485,7 @@ macro_rules! global_execution_registry_api { .map_err(|error| FlowError::Internal(error.to_string()))?; state .$field - .register(ExecutionIntercept::new( - name, - priority, - adapt_execution_callable!(callable $(, $adapter)?), - )) + .register(ExecutionIntercept::new(name, priority, callable)) .map_err(FlowError::AlreadyExists) } @@ -677,7 +660,6 @@ macro_rules! scope_execution_registry_api { $deregister_name:ident, $field:ident, $fn_type:ty - $(, $adapter:path)? ) => { $(#[$register_meta])* /// @@ -708,11 +690,7 @@ macro_rules! scope_execution_registry_api { .ok_or_else(|| FlowError::NotFound(format!("scope {scope_uuid} not found")))?; registries .$field - .register(ExecutionIntercept::new( - name, - priority, - adapt_execution_callable!(callable $(, $adapter)?), - )) + .register(ExecutionIntercept::new(name, priority, callable)) .map_err(FlowError::AlreadyExists) } @@ -898,8 +876,7 @@ global_execution_registry_api!( /// Deregister a global LLM execution intercept. deregister_llm_execution_intercept, llm_execution_intercepts, - LlmExecutionFn, - adapt_llm_execution_fn + LlmExecutionFn ); global_execution_registry_api!( /// Register a global streaming LLM execution intercept. @@ -909,48 +886,9 @@ global_execution_registry_api!( /// Deregister a global streaming LLM execution intercept. deregister_llm_stream_execution_intercept, llm_stream_execution_intercepts, - LlmStreamExecutionFn, - adapt_llm_stream_execution_fn + LlmStreamExecutionFn ); -/// Register a global non-streaming LLM execution intercept that receives -/// Relay's private invocation context. -#[cfg(feature = "worker-grpc")] -pub(crate) fn register_contextual_llm_execution_intercept( - name: &str, - priority: i32, - callable: ContextualLlmExecutionFn, -) -> Result<()> { - ensure_runtime_owner()?; - let context = global_context(); - let mut state = context - .write() - .map_err(|error| FlowError::Internal(error.to_string()))?; - state - .llm_execution_intercepts - .register(ExecutionIntercept::new(name, priority, callable)) - .map_err(FlowError::AlreadyExists) -} - -/// Register a global streaming LLM execution intercept that receives Relay's -/// private invocation context. -#[cfg(feature = "worker-grpc")] -pub(crate) fn register_contextual_llm_stream_execution_intercept( - name: &str, - priority: i32, - callable: ContextualLlmStreamExecutionFn, -) -> Result<()> { - ensure_runtime_owner()?; - let context = global_context(); - let mut state = context - .write() - .map_err(|error| FlowError::Internal(error.to_string()))?; - state - .llm_stream_execution_intercepts - .register(ExecutionIntercept::new(name, priority, callable)) - .map_err(FlowError::AlreadyExists) -} - scope_guardrail_registry_api!( /// Register a scope-local mark event sanitizer. scope_register_mark_sanitize_guardrail, @@ -1106,8 +1044,7 @@ scope_execution_registry_api!( /// Deregister a scope-local LLM execution intercept. scope_deregister_llm_execution_intercept, llm_execution_intercepts, - LlmExecutionFn, - adapt_llm_execution_fn + LlmExecutionFn ); scope_execution_registry_api!( /// Register a scope-local streaming LLM execution intercept. @@ -1117,6 +1054,5 @@ scope_execution_registry_api!( /// Deregister a scope-local streaming LLM execution intercept. scope_deregister_llm_stream_execution_intercept, llm_stream_execution_intercepts, - LlmStreamExecutionFn, - adapt_llm_stream_execution_fn + LlmStreamExecutionFn ); diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index 8bdeb8de6..a86713a57 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -25,10 +25,7 @@ pub use continuation_context::MiddlewareContinuationContext; #[cfg(test)] pub(crate) use continuation_context::MiddlewareContinuationLease; pub use global::global_context; -pub(crate) use llm_execution_context::{ - ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn, LlmExecutionCodecContext, - adapt_llm_execution_fn, adapt_llm_stream_execution_fn, -}; +pub use llm_execution_context::LlmExecutionContext; pub(crate) use scope_stack::capture_trace_context; pub use scope_stack::{ PropagationContext, ScopeStack, ScopeStackHandle, TASK_SCOPE_STACK, ThreadScopeStackBinding, diff --git a/crates/core/src/api/runtime/callbacks.rs b/crates/core/src/api/runtime/callbacks.rs index a583b325c..8248fdf3e 100644 --- a/crates/core/src/api/runtime/callbacks.rs +++ b/crates/core/src/api/runtime/callbacks.rs @@ -493,12 +493,14 @@ pub type LlmExecutionNextFn = /// Wrap or replace non-streaming LLM execution. /// /// A non-streaming execution intercept receives the logical provider name, the -/// current request, and the continuation representing the rest of the chain. +/// current request, the invocation codec context, and the continuation +/// representing the rest of the chain. /// /// # Parameters /// - First argument: Logical provider or model family name. /// - Second argument: Current LLM request. -/// - Third argument: Continuation for the remaining execution chain. +/// - Third argument: Request and unary-response codec context. +/// - Fourth argument: Continuation for the remaining execution chain. /// /// # Returns /// A future resolving to the provider response JSON. @@ -510,6 +512,7 @@ pub type LlmExecutionFn = Arc< dyn Fn( &str, LlmRequest, + super::llm_execution_context::LlmExecutionContext, LlmExecutionNextFn, ) -> Pin> + Send>> + Send @@ -628,9 +631,7 @@ pub type LlmFinalizerFn = Box Json + Send>; /// # Returns /// A shared reference to a scope-local streaming execution registry. pub(crate) type LlmStreamExecutionRegistryRef<'a> = &'a crate::registry::SortedRegistry< - crate::api::registry::ExecutionIntercept< - super::llm_execution_context::ContextualLlmStreamExecutionFn, - >, + crate::api::registry::ExecutionIntercept, >; /// Slice of scope-local streaming execution registries. /// @@ -670,11 +671,14 @@ pub type LlmStreamExecutionNextFn = Arc< /// /// A streaming execution intercept can observe or modify the request before /// invoking the continuation, and it can also replace the returned stream. +/// Its execution context exposes the request codec but deliberately omits a +/// response codec because Relay codecs decode complete responses, not chunks. /// /// # Parameters /// - First argument: Logical provider or model family name. /// - Second argument: Current LLM request. -/// - Third argument: Continuation for the remaining streaming execution chain. +/// - Third argument: Request codec context with no response direction. +/// - Fourth argument: Continuation for the remaining streaming execution chain. /// /// # Returns /// A future resolving to a JSON chunk stream. @@ -686,6 +690,7 @@ pub type LlmStreamExecutionFn = Arc< dyn Fn( &str, LlmRequest, + super::llm_execution_context::LlmExecutionContext, LlmStreamExecutionNextFn, ) -> Pin> + Send>> + Send diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index 5c5a730d6..e7c52235a 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -1,113 +1,200 @@ // SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -//! Invocation-scoped codec context for internal LLM execution adapters. +//! Invocation-scoped codec context for LLM execution intercepts. -use std::future::Future; -use std::pin::Pin; use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use super::callbacks::{LlmSanitizeRequestContext, LlmSanitizeResponseContext}; use crate::api::llm::LlmRequest; +use crate::codec::request::AnnotatedLlmRequest; +use crate::codec::response::AnnotatedLlmResponse; use crate::codec::traits::{LlmCodec, LlmResponseCodec}; -use crate::error::Result; +use crate::error::{FlowError, Result}; use crate::json::Json; -use super::callbacks::{ - LlmExecutionFn, LlmExecutionNextFn, LlmJsonStream, LlmStreamExecutionFn, - LlmStreamExecutionNextFn, -}; -#[cfg(feature = "worker-grpc")] -use super::callbacks::{LlmSanitizeRequestContext, LlmSanitizeResponseContext}; +const INACTIVE_EXECUTION_CODEC_ERROR: &str = "LLM execution codec capability is no longer active"; + +#[derive(Debug)] +struct ExecutionCodecGate { + active: AtomicBool, +} + +impl ExecutionCodecGate { + fn new() -> Self { + Self { + active: AtomicBool::new(true), + } + } + + fn ensure_active(&self) -> Result<()> { + if self.active.load(Ordering::Acquire) { + Ok(()) + } else { + Err(FlowError::InvalidArgument( + INACTIVE_EXECUTION_CODEC_ERROR.into(), + )) + } + } + + fn revoke(&self) { + self.active.store(false, Ordering::Release); + } +} + +/// Revokes codec capabilities issued to one execution-intercept invocation. +pub(crate) struct LlmExecutionCodecLeaseGuard { + gate: Arc, +} + +impl Drop for LlmExecutionCodecLeaseGuard { + fn drop(&mut self) { + self.gate.revoke(); + } +} + +struct RevocableRequestCodec { + codec: Arc, + gate: Arc, +} + +impl LlmCodec for RevocableRequestCodec { + fn codec_identity(&self) -> super::LlmCodecIdentity { + self.codec.codec_identity() + } + + fn decode(&self, request: &LlmRequest) -> Result { + self.gate.ensure_active()?; + self.codec.decode(request) + } + + fn encode(&self, annotated: &AnnotatedLlmRequest, original: &LlmRequest) -> Result { + self.gate.ensure_active()?; + self.codec.encode(annotated, original) + } +} -/// Active request and response codecs for one managed LLM execution. +struct RevocableResponseCodec { + codec: Arc, + gate: Arc, +} + +impl LlmResponseCodec for RevocableResponseCodec { + fn codec_identity(&self) -> super::LlmCodecIdentity { + self.codec.codec_identity() + } + + fn decode_response(&self, response: &Json) -> Result { + self.gate.ensure_active()?; + self.codec.decode_response(response) + } +} + +/// Active request and response codec context for one managed LLM execution. /// -/// Public execution-interceptor callbacks remain unchanged. Relay adapts them -/// into one private context-aware callback shape so language and process -/// bridges receive the active codecs explicitly. +/// The request direction is always present and distinguishes an invocation +/// with no request codec from an invocation with a built-in, runtime, or opaque +/// codec. Unary execution also carries a response direction. Streaming +/// execution deliberately leaves [`Self::response_codec`] unavailable because +/// Relay's response codecs operate on complete provider responses rather than +/// individual stream chunks. /// -/// The context describes the codecs selected when the managed invocation was -/// created. Execution interceptors may rewrite payloads within that codec's -/// contract, but changing the provider wire format does not select a new codec. -/// A subsequent decode or encode will reject an incompatible payload rather -/// than silently infer another codec. -#[derive(Clone)] -pub(crate) struct LlmExecutionCodecContext { - #[cfg(feature = "worker-grpc")] - request: LlmSanitizeRequestContext, - #[cfg(feature = "worker-grpc")] - response: LlmSanitizeResponseContext, +/// The codecs are fixed when the managed invocation is created. Rewriting a +/// payload does not select another codec; decoding or encoding an incompatible +/// wire representation fails rather than inferring a different format. +/// Resolved codec capabilities are valid only for the callback that received +/// this context. Unary capabilities expire when that callback settles; +/// streaming request capabilities remain valid until its returned stream ends +/// or closes. Retained capabilities return [`FlowError::InvalidArgument`] +/// after expiry. +#[derive(Clone, Debug, Default)] +pub struct LlmExecutionContext { + request_codec: LlmSanitizeRequestContext, + response_codec: Option, } -impl LlmExecutionCodecContext { - #[cfg(feature = "worker-grpc")] - pub(crate) fn new( - request: LlmSanitizeRequestContext, - response: LlmSanitizeResponseContext, +impl LlmExecutionContext { + /// Construct an execution context from its directional codec contexts. + #[must_use] + pub fn new( + request_codec: LlmSanitizeRequestContext, + response_codec: Option, ) -> Self { - Self { request, response } + Self { + request_codec, + response_codec, + } } - pub(crate) fn for_codecs( + /// Construct the context for a unary managed execution. + pub(crate) fn for_unary_codecs( request_codec: Option>, response_codec: &Option>, ) -> Self { - #[cfg(feature = "worker-grpc")] - { - Self::new( - LlmSanitizeRequestContext::for_request_codec(request_codec), - LlmSanitizeResponseContext::for_response_codec(response_codec.clone()), - ) - } - #[cfg(not(feature = "worker-grpc"))] - { - let _ = (request_codec, response_codec); - Self {} - } + Self::new( + LlmSanitizeRequestContext::for_request_codec(request_codec), + Some(LlmSanitizeResponseContext::for_response_codec( + response_codec.clone(), + )), + ) } - #[cfg(feature = "worker-grpc")] - pub(crate) fn request(&self) -> &LlmSanitizeRequestContext { - &self.request + /// Construct the context for a streaming managed execution. + pub(crate) fn for_streaming_codec(request_codec: Option>) -> Self { + Self::new( + LlmSanitizeRequestContext::for_request_codec(request_codec), + None, + ) } - #[cfg(feature = "worker-grpc")] - pub(crate) fn response(&self) -> &LlmSanitizeResponseContext { - &self.response + /// Issue revocable codec facades for one execution-intercept invocation. + /// + /// The source context retains Relay's selected codecs, but callbacks only + /// receive the facades created here. Dropping the returned guard makes all + /// retained facade clones fail without exposing the underlying codec. + pub(crate) fn lease(&self) -> (Self, LlmExecutionCodecLeaseGuard) { + let gate = Arc::new(ExecutionCodecGate::new()); + let request_codec = match self.request_codec.resolve_codec() { + Some(codec) => LlmSanitizeRequestContext::for_request_codec(Some(Arc::new( + RevocableRequestCodec { + codec, + gate: Arc::clone(&gate), + }, + ))), + None => LlmSanitizeRequestContext::with_identity(self.request_codec.codec().clone()), + }; + let response_codec = + self.response_codec + .as_ref() + .map(|context| match context.resolve_codec() { + Some(codec) => LlmSanitizeResponseContext::for_response_codec(Some(Arc::new( + RevocableResponseCodec { + codec, + gate: Arc::clone(&gate), + }, + ))), + None => LlmSanitizeResponseContext::with_identity(context.codec().clone()), + }); + + ( + Self::new(request_codec, response_codec), + LlmExecutionCodecLeaseGuard { gate }, + ) } -} -/// Private non-streaming execution callback used by Relay's registries. -pub(crate) type ContextualLlmExecutionFn = Arc< - dyn Fn( - &str, - LlmRequest, - LlmExecutionCodecContext, - LlmExecutionNextFn, - ) -> Pin> + Send>> - + Send - + Sync, ->; - -/// Private streaming execution callback used by Relay's registries. -pub(crate) type ContextualLlmStreamExecutionFn = Arc< - dyn Fn( - &str, - LlmRequest, - LlmExecutionCodecContext, - LlmStreamExecutionNextFn, - ) -> Pin> + Send>> - + Send - + Sync, ->; - -/// Adapt the stable public callback into Relay's private context-aware shape. -pub(crate) fn adapt_llm_execution_fn(callback: LlmExecutionFn) -> ContextualLlmExecutionFn { - Arc::new(move |name, request, _context, next| callback(name, request, next)) -} + /// Return the request-direction codec identity and revocable capability. + #[must_use] + pub fn request_codec(&self) -> &LlmSanitizeRequestContext { + &self.request_codec + } -/// Adapt the stable public stream callback into Relay's private context-aware shape. -pub(crate) fn adapt_llm_stream_execution_fn( - callback: LlmStreamExecutionFn, -) -> ContextualLlmStreamExecutionFn { - Arc::new(move |name, request, _context, next| callback(name, request, next)) + /// Return the unary response-direction codec identity and revocable capability. + /// + /// Streaming execution returns `None` because Relay does not expose a + /// completed-response codec for individual stream chunks. + #[must_use] + pub fn response_codec(&self) -> Option<&LlmSanitizeResponseContext> { + self.response_codec.as_ref() + } } diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index 4825a45af..fa5d6b284 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -42,10 +42,9 @@ use crate::api::runtime::callbacks::{ use crate::api::runtime::continuation_context::{ MiddlewareContinuationContext, MiddlewareContinuationGuard, MiddlewareContinuationLease, }; +use crate::api::runtime::llm_execution_context::LlmExecutionCodecLeaseGuard; use crate::api::runtime::subscriber_dispatcher; -use crate::api::runtime::{ - ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn, LlmExecutionCodecContext, -}; +use crate::api::runtime::{LlmExecutionContext, LlmExecutionFn, LlmStreamExecutionFn}; use crate::api::scope::{CreateScopeHandleParams, EndScopeHandleParams, ScopeHandle, ScopeType}; use crate::api::shared::snapshot_event_sanitizers; use crate::api::tool::ToolHandle; @@ -70,6 +69,11 @@ struct ContinuationGuardedLlmStream { guard: Option, } +struct CodecGuardedLlmStream { + inner: LlmJsonStream, + guard: Option, +} + struct ContextualizedLlmStream { inner: LlmJsonStream, context: MiddlewareContinuationContext, @@ -159,6 +163,45 @@ fn guard_stream_continuation( }) } +impl Stream for CodecGuardedLlmStream { + type Item = crate::error::Result; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + let result = Pin::new(&mut this.inner).poll_next(cx); + if matches!(&result, Poll::Ready(None) | Poll::Ready(Some(Err(_)))) { + this.guard.take(); + } + result + } +} + +impl LlmStreamInner for CodecGuardedLlmStream { + fn terminalize(self: Pin<&mut Self>) { + let this = self.get_mut(); + this.guard.take(); + this.inner.terminalize(); + } + + fn close( + self: Pin<&mut Self>, + ) -> Pin> + Send + '_>> { + Box::pin(async move { + let this = self.get_mut(); + let result = this.inner.close().await; + this.guard.take(); + result + }) + } +} + +fn guard_stream_codec(stream: LlmJsonStream, guard: LlmExecutionCodecLeaseGuard) -> LlmJsonStream { + LlmJsonStream::from_closeable(CodecGuardedLlmStream { + inner: stream, + guard: Some(guard), + }) +} + struct GuardrailScopeCompletion<'a> { handle: Option, subscribers: &'a [EventSubscriberFn], @@ -246,11 +289,10 @@ pub struct NemoRelayContextState { /// Global LLM request intercepts that can rewrite or annotate requests. pub(crate) llm_request_intercepts: SortedRegistry>, /// Global non-streaming LLM execution intercepts that wrap callback execution. - pub(crate) llm_execution_intercepts: - SortedRegistry>, + pub(crate) llm_execution_intercepts: SortedRegistry>, /// Global streaming LLM execution intercepts that wrap stream-producing callbacks. pub(crate) llm_stream_execution_intercepts: - SortedRegistry>, + SortedRegistry>, /// Global lifecycle subscribers notified after runtime events are emitted. pub(crate) event_subscribers: HashMap, /// Whether LLM start events retain complete sanitized request payloads. @@ -1813,8 +1855,8 @@ impl NemoRelayContextState { &self, name: &str, default_fn: LlmExecutionNextFn, - scope_locals: &[&SortedRegistry>], - execution_context: LlmExecutionCodecContext, + scope_locals: &[&SortedRegistry>], + execution_context: LlmExecutionContext, ) -> LlmExecutionNextFn { let matching = merge_execution_intercept_callables( &self.llm_execution_intercepts, @@ -1833,6 +1875,7 @@ impl NemoRelayContextState { let current_name = current_name.clone(); let current_context = current_context.clone(); Box::pin(async move { + let (current_context, codec_guard) = current_context.lease(); let (continuation, continuation_guard) = MiddlewareContinuationLease::capture(); let raw_next: LlmExecutionNextFn = Arc::new(move |request| { let invocation = continuation.begin(); @@ -1843,6 +1886,7 @@ impl NemoRelayContextState { }); let result = callable(¤t_name, request, current_context, raw_next).await; drop(continuation_guard); + drop(codec_guard); result }) }); @@ -1870,7 +1914,7 @@ impl NemoRelayContextState { name: &str, default_fn: LlmStreamExecutionNextFn, scope_locals: LlmStreamExecutionRegistryRefs<'_>, - execution_context: LlmExecutionCodecContext, + execution_context: LlmExecutionContext, ) -> LlmStreamExecutionNextFn { let matching = merge_execution_intercept_callables( &self.llm_stream_execution_intercepts, @@ -1889,6 +1933,7 @@ impl NemoRelayContextState { let current_name = current_name.clone(); let current_context = current_context.clone(); Box::pin(async move { + let (current_context, codec_guard) = current_context.lease(); let (continuation, continuation_guard) = MiddlewareContinuationLease::capture(); let raw_next: LlmStreamExecutionNextFn = Arc::new(move |request| { let invocation = continuation.begin(); @@ -1901,7 +1946,10 @@ impl NemoRelayContextState { }) }); let result = callable(¤t_name, request, current_context, raw_next).await; - result.map(|stream| guard_stream_continuation(stream, continuation_guard)) + result.map(|stream| { + let stream = guard_stream_continuation(stream, continuation_guard); + guard_stream_codec(stream, codec_guard) + }) }) }); } diff --git a/crates/core/src/context/registries.rs b/crates/core/src/context/registries.rs index 5934fdf65..32991a24d 100644 --- a/crates/core/src/context/registries.rs +++ b/crates/core/src/context/registries.rs @@ -14,9 +14,9 @@ use crate::api::registry::{ runtime_registration_is_enabled, }; use crate::api::runtime::{ - ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn, EventSanitizeFn, EventSubscriberFn, - LlmConditionalFn, LlmRequestInterceptFn, LlmSanitizeRequestFn, LlmSanitizeResponseFn, - ToolConditionalFn, ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn, + EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionFn, LlmRequestInterceptFn, + LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionFn, ToolConditionalFn, + ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn, }; use crate::registry::SortedRegistry; @@ -55,11 +55,10 @@ pub(crate) struct ScopeLocalRegistries { /// LLM request intercepts that can rewrite or annotate requests. pub(crate) llm_request_intercepts: SortedRegistry>, /// Non-streaming LLM execution intercepts that wrap callback execution. - pub(crate) llm_execution_intercepts: - SortedRegistry>, + pub(crate) llm_execution_intercepts: SortedRegistry>, /// Streaming LLM execution intercepts that wrap stream-producing callbacks. pub(crate) llm_stream_execution_intercepts: - SortedRegistry>, + SortedRegistry>, /// Scope-local lifecycle subscribers visible while the owning scope is active. pub(crate) event_subscribers: HashMap, } diff --git a/crates/core/src/plugin.rs b/crates/core/src/plugin.rs index cc2f85e94..9b80e694d 100644 --- a/crates/core/src/plugin.rs +++ b/crates/core/src/plugin.rs @@ -41,18 +41,12 @@ use crate::api::registry::{ register_tool_request_intercept, register_tool_sanitize_request_guardrail, register_tool_sanitize_response_guardrail, }; -#[cfg(feature = "worker-grpc")] -use crate::api::registry::{ - register_contextual_llm_execution_intercept, register_contextual_llm_stream_execution_intercept, -}; use crate::api::runtime::{ ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionFn, LlmRequestInterceptFn, LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionFn, ToolConditionalFn, ToolExecutionFn, ToolInterceptFn, ToolSanitizeFn, }; -#[cfg(feature = "worker-grpc")] -use crate::api::runtime::{ContextualLlmExecutionFn, ContextualLlmStreamExecutionFn}; use crate::api::subscriber::{deregister_subscriber, register_subscriber}; pub use nemo_relay_types::plugin::{ConfigDiagnostic, DiagnosticLevel}; @@ -892,25 +886,6 @@ impl PluginRegistrationContext { ) } - /// Registers an internal context-aware LLM execution intercept and records - /// its rollback closure. - #[cfg(feature = "worker-grpc")] - pub(crate) fn register_contextual_llm_execution_intercept( - &mut self, - name: &str, - priority: i32, - callback: ContextualLlmExecutionFn, - ) -> Result<()> { - self.register_execution_intercept( - name, - priority, - callback, - "llm execution intercept", - register_contextual_llm_execution_intercept, - deregister_llm_execution_intercept, - ) - } - /// Registers an LLM stream execution intercept and records its rollback closure. pub fn register_llm_stream_execution_intercept( &mut self, @@ -928,25 +903,6 @@ impl PluginRegistrationContext { ) } - /// Registers an internal context-aware streaming LLM execution intercept - /// and records its rollback closure. - #[cfg(feature = "worker-grpc")] - pub(crate) fn register_contextual_llm_stream_execution_intercept( - &mut self, - name: &str, - priority: i32, - callback: ContextualLlmStreamExecutionFn, - ) -> Result<()> { - self.register_execution_intercept( - name, - priority, - callback, - "llm stream execution intercept", - register_contextual_llm_stream_execution_intercept, - deregister_llm_stream_execution_intercept, - ) - } - fn register_execution_intercept( &mut self, name: &str, diff --git a/crates/core/src/plugin/dynamic.rs b/crates/core/src/plugin/dynamic.rs index 13da05f65..a6b496073 100644 --- a/crates/core/src/plugin/dynamic.rs +++ b/crates/core/src/plugin/dynamic.rs @@ -144,6 +144,37 @@ pub(super) fn validate_tool_execution_context_compatibility( Ok(()) } +#[cfg(feature = "worker-grpc")] +pub(super) fn validate_llm_execution_context_compatibility( + relay: &str, + plugin_kind: &str, +) -> crate::plugin::Result<()> { + let requirement = VersionReq::parse(relay).map_err(|error| { + PluginError::InvalidConfig(format!("invalid compat.relay version requirement: {error}")) + })?; + if version_requirement_matches_minor(&requirement, 0, 9) { + return Err(PluginError::InvalidConfig(format!( + "dynamic plugin '{plugin_kind}' registers an LLM execution intercept and must declare compat.relay = \">=0.10,<1.0\" or another range that excludes Relay 0.9" + ))); + } + Ok(()) +} + +pub(super) fn validate_native_abi_compatibility( + relay: &str, + plugin_kind: &str, +) -> crate::plugin::Result<()> { + let requirement = VersionReq::parse(relay).map_err(|error| { + PluginError::InvalidConfig(format!("invalid compat.relay version requirement: {error}")) + })?; + if version_requirement_matches_minor(&requirement, 0, 9) { + return Err(PluginError::InvalidConfig(format!( + "dynamic native plugin '{plugin_kind}' uses native ABI v7 and must declare compat.relay = \">=0.10,<1.0\" or another range that excludes Relay 0.9" + ))); + } + Ok(()) +} + fn version_requirement_matches_minor(requirement: &VersionReq, major: u64, minor: u64) -> bool { let first_version = Version::new(major, minor, 0); let mut next_minor = first_version.clone(); diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 0ace3703e..1fe534178 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -29,8 +29,8 @@ use crate::api::registry::{ }; use crate::api::runtime::{ ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, - LlmCodecIdentity, LlmConditionalFn, LlmExecutionFn, LlmExecutionNextFn, LlmJsonStream, - LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, + LlmCodecIdentity, LlmConditionalFn, LlmExecutionContext, LlmExecutionFn, LlmExecutionNextFn, + LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, LlmStreamExecutionNextFn, MiddlewareContinuationContext, ToolConditionalFn, ToolExecutionContext, ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, @@ -57,17 +57,20 @@ use chrono::{DateTime, Utc}; use libloading::{Library, Symbol}; use nemo_relay_plugin::{ NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, + NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT, NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, NemoRelayNativeAsyncCallbackState, - NemoRelayNativeAsyncCompletion, NemoRelayNativeAsyncLlmStreamOpenCb, - NemoRelayNativeAsyncLlmStreamPullCb, NemoRelayNativeAsyncMiddlewareCb, - NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, NemoRelayNativeAsyncNextResultCb, - NemoRelayNativeAsyncNextStreamCb, NemoRelayNativeAsyncStream, + NemoRelayNativeAsyncCompletion, NemoRelayNativeAsyncLlmExecutionCb, + NemoRelayNativeAsyncLlmStreamOpenCb, NemoRelayNativeAsyncLlmStreamPullCb, + NemoRelayNativeAsyncMiddlewareCb, NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, + NemoRelayNativeAsyncNextResultCb, NemoRelayNativeAsyncNextStreamCb, NemoRelayNativeAsyncStream, NemoRelayNativeAsyncStreamMiddlewareCb, NemoRelayNativeConditionalMiddlewareCb, NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, - NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeLlmAsyncStream, - NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, + NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, + NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, + NemoRelayNativeLlmExecutionRequestContext, NemoRelayNativeLlmExecutionResponseContext, NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, @@ -89,7 +92,7 @@ use super::{ DynamicPluginKind, DynamicPluginManifest, DynamicPluginManifestLoad, DynamicPluginTeardownOutcome, deregister_tracked_registrations_checked, validate_annotated_request_consumer_compatibility, validate_dynamic_plugin_relay_compatibility, - validate_tool_execution_context_compatibility, + validate_native_abi_compatibility, validate_tool_execution_context_compatibility, }; /// Native plugin load request derived from host dynamic-plugin state. @@ -357,6 +360,7 @@ fn load_one_native_plugin( .as_deref() .expect("validated native manifest must declare compat.relay") .to_string(); + validate_native_abi_compatibility(&relay_compat, &spec.plugin_id)?; if manifest.compat.native_api.as_deref().map(str::trim) != Some("1") { return Err(PluginError::InvalidConfig(format!( "dynamic plugin '{}' declares unsupported compat.native_api '{}'; expected 1", @@ -410,30 +414,11 @@ fn load_one_native_plugin( library_path.display() )) })?; - let mut status = entry(native_host_api(), &mut plugin); - // Older SDKs reject newer tables. Negotiate from the current v6 table - // through separately frozen v5, v4, v3, and v2 tables so their struct sizes and - // function pointers do not change as the current ABI grows. - if status == NemoRelayStatus::InvalidArg { - drop_native_plugin_descriptor(&mut plugin); - status = entry(native_host_api_v5(), &mut plugin); - } - if status == NemoRelayStatus::InvalidArg { - drop_native_plugin_descriptor(&mut plugin); - status = entry(native_host_api_v4(), &mut plugin); - } - if status == NemoRelayStatus::InvalidArg { - drop_native_plugin_descriptor(&mut plugin); - status = entry(native_host_api_v3(), &mut plugin); - } - if status == NemoRelayStatus::InvalidArg { - drop_native_plugin_descriptor(&mut plugin); - status = entry(native_host_api_v2(), &mut plugin); - } + let status = entry(native_host_api(), &mut plugin); if status != NemoRelayStatus::Ok { drop_native_plugin_descriptor(&mut plugin); return Err(PluginError::RegistrationFailed(format!( - "native plugin entry symbol '{symbol}' failed: {}", + "native plugin entry symbol '{symbol}' rejected native ABI {NEMO_RELAY_NATIVE_ABI_VERSION}; rebuild the plugin against this Relay release: {}", native_last_error_message().unwrap_or_else(|| format!("{status:?}")) ))); } @@ -593,6 +578,96 @@ struct NativeHostString(Vec); struct NativeHostLlmRequestCodec(Arc); struct NativeHostLlmResponseCodec(Arc); +/// Borrows the host codec handles and owns the native strings exposed to one +/// native LLM execution callback. +struct NativeLlmExecutionContextBridge<'a> { + request_codec: Option<&'a NativeHostLlmRequestCodec>, + request_kind: NemoRelayNativeLlmCodecKind, + request_id: Option, + response_codec: Option<&'a NativeHostLlmResponseCodec>, + response_kind: Option, + response_id: Option, +} + +impl<'a> NativeLlmExecutionContextBridge<'a> { + fn new( + context: &LlmExecutionContext, + request_codec: Option<&'a NativeHostLlmRequestCodec>, + response_codec: Option<&'a NativeHostLlmResponseCodec>, + ) -> FlowResult { + let (request_kind, request_id) = + native_llm_codec_identity(context.request_codec().codec())?; + let request_id = request_id.map(|value| value as usize); + + let (response_kind, response_id) = if let Some(response) = context.response_codec() { + let (kind, id) = match native_llm_codec_identity(response.codec()) { + Ok(value) => value, + Err(error) => { + if let Some(request_id) = request_id { + unsafe { native_string_free(request_id as *mut NemoRelayNativeString) }; + } + return Err(error); + } + }; + (Some(kind), id.map(|value| value as usize)) + } else { + (None, None) + }; + + Ok(Self { + request_codec, + request_kind, + request_id, + response_codec, + response_kind, + response_id, + }) + } + + fn with_native_context( + &self, + callback: impl FnOnce(NemoRelayNativeLlmExecutionContext) -> T, + ) -> T { + let request_codec = NemoRelayNativeLlmExecutionRequestContext { + codec_kind: self.request_kind, + codec_id: self + .request_id + .map_or(ptr::null(), |value| value as *const NemoRelayNativeString), + codec: self + .request_codec + .map_or(ptr::null(), |value| std::ptr::from_ref(value).cast()), + }; + let response_codec = + self.response_kind + .map(|codec_kind| NemoRelayNativeLlmExecutionResponseContext { + codec_kind, + codec_id: self + .response_id + .map_or(ptr::null(), |value| value as *const NemoRelayNativeString), + codec: self + .response_codec + .map_or(ptr::null(), |value| std::ptr::from_ref(value).cast()), + }); + callback(NemoRelayNativeLlmExecutionContext { + request_codec, + response_codec: response_codec + .as_ref() + .map_or(ptr::null(), std::ptr::from_ref), + }) + } +} + +impl Drop for NativeLlmExecutionContextBridge<'_> { + fn drop(&mut self) { + if let Some(request_id) = self.request_id.take() { + unsafe { native_string_free(request_id as *mut NemoRelayNativeString) }; + } + if let Some(response_id) = self.response_id.take() { + unsafe { native_string_free(response_id as *mut NemoRelayNativeString) }; + } + } +} + struct NativeHostScopeHandle(ScopeHandle); struct NativeHostScopeStack(ScopeStackHandle); @@ -872,25 +947,41 @@ 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_v7) + .v6 + .v5 + .v4 + .v3 + .v1 as *const NemoRelayNativeHostApiV1 +} + +#[cfg(test)] +fn native_host_api_v6() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); &HOST_API.get_or_init(build_native_host_api_v6).v5.v4.v3.v1 as *const NemoRelayNativeHostApiV1 } +#[cfg(test)] fn native_host_api_v5() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); &HOST_API.get_or_init(build_native_host_api_v5).v4.v3.v1 as *const NemoRelayNativeHostApiV1 } +#[cfg(test)] fn native_host_api_v4() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); &HOST_API.get_or_init(build_native_host_api_v4).v3.v1 as *const NemoRelayNativeHostApiV1 } +#[cfg(test)] fn native_host_api_v3() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); &HOST_API.get_or_init(build_native_host_api_v3).v1 as *const NemoRelayNativeHostApiV1 } +#[cfg(test)] fn native_host_api_v2() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); HOST_API.get_or_init(build_native_host_api_v2) as *const _ @@ -1029,7 +1120,7 @@ fn build_native_host_api_v5() -> NemoRelayNativeHostApiV5 { fn build_native_host_api_v6() -> NemoRelayNativeHostApiV6 { let mut v5 = build_native_host_api_v5(); - v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION; + v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING; v5.v4.v3.v1.struct_size = std::mem::size_of::(); NemoRelayNativeHostApiV6 { v5, @@ -1037,6 +1128,20 @@ fn build_native_host_api_v6() -> NemoRelayNativeHostApiV6 { } } +fn build_native_host_api_v7() -> NemoRelayNativeHostApiV7 { + let mut v6 = build_native_host_api_v6(); + v6.v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT; + v6.v5.v4.v3.v1.struct_size = std::mem::size_of::(); + NemoRelayNativeHostApiV7 { + v6, + plugin_context_register_async_llm_execution_intercept: + native_plugin_context_register_async_llm_execution_intercept, + async_stream_retain: native_async_stream_retain, + async_stream_llm_request_codec_decode: native_async_stream_llm_request_codec_decode, + async_stream_llm_request_codec_encode: native_async_stream_llm_request_codec_encode, + } +} + fn read_native_string(value: *const NemoRelayNativeString) -> crate::plugin::Result { if value.is_null() { return Ok(String::new()); @@ -1718,9 +1823,10 @@ fn make_user_data( const NATIVE_ASYNC_STREAM_CHANNEL_CAPACITY: usize = 64; -enum NativeAsyncCodecCapability { - Request(Arc), - Response(Arc), +#[derive(Default)] +struct NativeAsyncCodecCapabilities { + request: Option>, + response: Option>, } struct NativeAsyncCompletion { @@ -1729,7 +1835,8 @@ struct NativeAsyncCompletion { next_invoked: AtomicBool, next_abort: Mutex>, continuation_aborts: Mutex>, - codec: Option, + request_codec: Option, + response_codec: Option, #[cfg(test)] before_settlement_lock: Option>, // A pending native callback can continue running after its completion @@ -1920,6 +2027,7 @@ struct NativeAsyncStream { backpressured: AtomicBool, downstream_aborts: Mutex>, settlement: Mutex<()>, + request_codec: Option, #[cfg(test)] before_settlement_lock: Option>, _callback_user_data: Option>, @@ -2031,12 +2139,20 @@ impl Drop for NativeAsyncStreamReceiver { } } -async fn invoke_native_async_callback( - cb: NemoRelayNativeAsyncMiddlewareCb, +enum NativeAsyncCallback { + Middleware(NemoRelayNativeAsyncMiddlewareCb), + LlmExecution { + callback: NemoRelayNativeAsyncLlmExecutionCb, + context: LlmExecutionContext, + }, +} + +async fn invoke_native_async_callback_inner( + callback: NativeAsyncCallback, user_data: Arc, invocation: Json, next: Option, - codec: Option, + codecs: NativeAsyncCodecCapabilities, ) -> FlowResult { let runtime = if next.is_some() { Some(tokio::runtime::Handle::try_current().map_err(|error| { @@ -2057,11 +2173,30 @@ async fn invoke_native_async_callback( next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec, + request_codec: codecs.request.map(NativeHostLlmRequestCodec), + response_codec: codecs.response.map(NativeHostLlmResponseCodec), #[cfg(test)] before_settlement_lock: None, _callback_user_data: Some(user_data.clone()), }); + let native_context = match &callback { + NativeAsyncCallback::LlmExecution { context, .. } => { + Some(NativeLlmExecutionContextBridge::new( + context, + completion.request_codec.as_ref(), + completion.response_codec.as_ref(), + )) + } + NativeAsyncCallback::Middleware(_) => None, + } + .transpose(); + let native_context = match native_context { + Ok(context) => context, + Err(error) => { + unsafe { native_string_free(invocation as *mut NemoRelayNativeString) }; + return Err(error); + } + }; let mut wait = NativeAsyncWait { completion: Arc::clone(&completion), receiver, @@ -2085,17 +2220,34 @@ async fn invoke_native_async_callback( // SDK can capture it before moving the future to its own executor. let previous_thread_stack = capture_thread_scope_stack(); sync_thread_scope_stack(current_scope_stack()); - let state = catch_unwind(AssertUnwindSafe(|| 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, - ) + let state = catch_unwind(AssertUnwindSafe(|| match callback { + NativeAsyncCallback::Middleware(callback) => unsafe { + callback( + user_data.ptr, + invocation as *const NemoRelayNativeString, + next_ref + .map(|next| next as *const NemoRelayNativeAsyncNext) + .unwrap_or(ptr::null()), + completion_ref as *const NemoRelayNativeAsyncCompletion, + ) + }, + NativeAsyncCallback::LlmExecution { callback, .. } => native_context + .as_ref() + .expect("LLM execution callbacks always build a native context") + .with_native_context(|context| unsafe { + callback( + user_data.ptr, + invocation as *const NemoRelayNativeString, + std::ptr::from_ref(&context), + next_ref + .map(|next| next as *const NemoRelayNativeAsyncNext) + .unwrap_or(ptr::null()), + completion_ref as *const NemoRelayNativeAsyncCompletion, + ) + }), })); restore_thread_scope_stack(previous_thread_stack); + drop(native_context); let state = match state { Ok(state) => state, Err(_) => { @@ -2142,6 +2294,41 @@ async fn invoke_native_async_callback( wait.receive().await } +async fn invoke_native_async_callback( + callback: NemoRelayNativeAsyncMiddlewareCb, + user_data: Arc, + invocation: Json, + next: Option, + codecs: NativeAsyncCodecCapabilities, +) -> FlowResult { + invoke_native_async_callback_inner( + NativeAsyncCallback::Middleware(callback), + user_data, + invocation, + next, + codecs, + ) + .await +} + +async fn invoke_native_async_llm_execution_callback( + callback: NemoRelayNativeAsyncLlmExecutionCb, + user_data: Arc, + invocation: Json, + next: LlmExecutionNextFn, + codecs: NativeAsyncCodecCapabilities, + context: LlmExecutionContext, +) -> FlowResult { + invoke_native_async_callback_inner( + NativeAsyncCallback::LlmExecution { callback, context }, + user_data, + invocation, + Some(NativeAsyncNextInner::Llm(next)), + codecs, + ) + .await +} + unsafe extern "C" fn native_async_completion_resolve_json( completion: *const NemoRelayNativeAsyncCompletion, value_json: *const NemoRelayNativeString, @@ -2268,12 +2455,11 @@ unsafe extern "C" fn native_async_completion_llm_request_codec_decode( Ok(completion) => completion, Err(status) => return status, }; - let Some(NativeAsyncCodecCapability::Request(codec)) = &completion.codec else { + let Some(codec) = &completion.request_codec else { set_native_last_error("async completion has no request codec capability"); return NemoRelayStatus::InvalidArg; }; - let codec = NativeHostLlmRequestCodec(Arc::clone(codec)); - unsafe { native_llm_request_codec_decode(std::ptr::from_ref(&codec).cast(), request_json, out) } + unsafe { native_llm_request_codec_decode(std::ptr::from_ref(codec).cast(), request_json, out) } } unsafe extern "C" fn native_async_completion_llm_request_codec_encode( @@ -2291,14 +2477,13 @@ unsafe extern "C" fn native_async_completion_llm_request_codec_encode( Ok(completion) => completion, Err(status) => return status, }; - let Some(NativeAsyncCodecCapability::Request(codec)) = &completion.codec else { + let Some(codec) = &completion.request_codec else { set_native_last_error("async completion has no request codec capability"); return NemoRelayStatus::InvalidArg; }; - let codec = NativeHostLlmRequestCodec(Arc::clone(codec)); unsafe { native_llm_request_codec_encode( - std::ptr::from_ref(&codec).cast(), + std::ptr::from_ref(codec).cast(), annotated_json, original_json, out, @@ -2320,13 +2505,12 @@ unsafe extern "C" fn native_async_completion_llm_response_codec_decode( Ok(completion) => completion, Err(status) => return status, }; - let Some(NativeAsyncCodecCapability::Response(codec)) = &completion.codec else { + let Some(codec) = &completion.response_codec else { set_native_last_error("async completion has no response codec capability"); return NemoRelayStatus::InvalidArg; }; - let codec = NativeHostLlmResponseCodec(Arc::clone(codec)); unsafe { - native_llm_response_codec_decode(std::ptr::from_ref(&codec).cast(), response_json, out) + native_llm_response_codec_decode(std::ptr::from_ref(codec).cast(), response_json, out) } } @@ -2566,6 +2750,74 @@ unsafe extern "C" fn native_async_stream_is_backpressured( .is_some_and(|stream| stream.backpressured.load(Ordering::Acquire)) } +unsafe extern "C" fn native_async_stream_retain( + stream: *const NemoRelayNativeAsyncStream, +) -> NemoRelayStatus { + if stream.is_null() { + return NemoRelayStatus::NullPointer; + } + unsafe { Arc::increment_strong_count(stream.cast::()) }; + NemoRelayStatus::Ok +} + +fn with_active_async_stream_request_codec( + stream: *const NemoRelayNativeAsyncStream, + operation: impl FnOnce(&NativeHostLlmRequestCodec) -> NemoRelayStatus, +) -> NemoRelayStatus { + let Some(stream) = (unsafe { (stream as *const NativeAsyncStream).as_ref() }) else { + return NemoRelayStatus::NullPointer; + }; + let _settlement = stream + .settlement + .lock() + .unwrap_or_else(|error| error.into_inner()); + if stream.cancelled.load(Ordering::Acquire) || stream.settled.load(Ordering::Acquire) { + set_native_last_error("native async stream codec capability is expired"); + return NemoRelayStatus::InvalidArg; + } + let Some(codec) = &stream.request_codec else { + set_native_last_error("native async stream has no request codec capability"); + return NemoRelayStatus::InvalidArg; + }; + operation(codec) +} + +unsafe extern "C" fn native_async_stream_llm_request_codec_decode( + stream: *const NemoRelayNativeAsyncStream, + request_json: *const NemoRelayNativeString, + out: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + if out.is_null() { + set_native_last_error("request codec decode output is null"); + return NemoRelayStatus::NullPointer; + } + unsafe { *out = ptr::null_mut() }; + with_active_async_stream_request_codec(stream, |codec| unsafe { + native_llm_request_codec_decode(std::ptr::from_ref(codec).cast(), request_json, out) + }) +} + +unsafe extern "C" fn native_async_stream_llm_request_codec_encode( + stream: *const NemoRelayNativeAsyncStream, + annotated_json: *const NemoRelayNativeString, + original_json: *const NemoRelayNativeString, + out: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + if out.is_null() { + set_native_last_error("request codec encode output is null"); + return NemoRelayStatus::NullPointer; + } + unsafe { *out = ptr::null_mut() }; + with_active_async_stream_request_codec(stream, |codec| unsafe { + native_llm_request_codec_encode( + std::ptr::from_ref(codec).cast(), + annotated_json, + original_json, + out, + ) + }) +} + unsafe extern "C" fn native_async_stream_release(stream: *const NemoRelayNativeAsyncStream) { if !stream.is_null() { let stream = unsafe { Arc::from_raw(stream as *const NativeAsyncStream) }; @@ -3456,7 +3708,7 @@ fn wrap_native_async_tool_json( user_data, serde_json::json!({"name": name, "value": value}), None, - None, + NativeAsyncCodecCapabilities::default(), ) .await?; Ok(value) @@ -3479,7 +3731,7 @@ fn wrap_native_async_tool_conditional( user_data, serde_json::json!({"name": name, "value": value}), None, - None, + NativeAsyncCodecCapabilities::default(), ) .await? { @@ -3508,7 +3760,7 @@ fn wrap_native_async_llm_conditional( user_data, serde_json::json!({"request": request}), None, - None, + NativeAsyncCodecCapabilities::default(), ) .await? { @@ -3532,16 +3784,17 @@ fn wrap_native_async_llm_sanitize_request( Arc::new(move |request, context| { let user_data = user_data.clone(); let codec = native_async_codec_identity(context.codec()); - let capability = context - .resolve_codec() - .map(NativeAsyncCodecCapability::Request); + let capabilities = NativeAsyncCodecCapabilities { + request: context.resolve_codec(), + response: None, + }; Box::pin(async move { let value = invoke_native_async_callback( cb, user_data, serde_json::json!({"request": request, "context": codec}), None, - capability, + capabilities, ) .await?; if value.is_null() { @@ -3565,16 +3818,17 @@ fn wrap_native_async_llm_sanitize_response( Arc::new(move |response, context| { let user_data = user_data.clone(); let codec = native_async_codec_identity(context.codec()); - let capability = context - .resolve_codec() - .map(NativeAsyncCodecCapability::Response); + let capabilities = NativeAsyncCodecCapabilities { + request: None, + response: context.resolve_codec(), + }; Box::pin(async move { let value = invoke_native_async_callback( cb, user_data, serde_json::json!({"response": response, "context": codec}), None, - capability, + capabilities, ) .await?; Ok((!value.is_null()).then_some(value)) @@ -3619,7 +3873,7 @@ fn wrap_native_async_llm_request_intercept( "annotated": annotated, }), None, - None, + NativeAsyncCodecCapabilities::default(), ) .await?, ) @@ -3648,7 +3902,7 @@ fn wrap_native_async_event_sanitize( user_data, serde_json::json!({"event": event, "fields": fields}), None, - None, + NativeAsyncCodecCapabilities::default(), ) .await?, ) @@ -3675,7 +3929,7 @@ fn wrap_native_async_event_metadata_injector( user_data, serde_json::json!({"event": event}), None, - None, + NativeAsyncCodecCapabilities::default(), ) .await?, ) @@ -3710,7 +3964,7 @@ fn wrap_native_async_tool_execution( user_data, invocation, Some(NativeAsyncNextInner::Tool(next)), - None, + NativeAsyncCodecCapabilities::default(), ) .await?; deserialize_native_tool_outcome(outcome).map_err(|error| { @@ -3722,21 +3976,28 @@ fn wrap_native_async_tool_execution( fn wrap_native_async_llm_execution( instance: Arc, - cb: NemoRelayNativeAsyncMiddlewareCb, + cb: NemoRelayNativeAsyncLlmExecutionCb, 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| { + Arc::new(move |name, request, context, next| { let user_data = user_data.clone(); let name = name.to_owned(); + let capabilities = NativeAsyncCodecCapabilities { + request: context.request_codec().resolve_codec(), + response: context + .response_codec() + .and_then(LlmSanitizeResponseContext::resolve_codec), + }; Box::pin(async move { - invoke_native_async_callback( + invoke_native_async_llm_execution_callback( cb, user_data, serde_json::json!({"name": name, "request": request}), - Some(NativeAsyncNextInner::Llm(next)), - None, + next, + capabilities, + context, ) .await }) @@ -3757,9 +4018,13 @@ fn wrap_native_incremental_llm_stream_execution_with_user_data( cb: NemoRelayNativeAsyncStreamMiddlewareCb, user_data: Arc, ) -> LlmStreamExecutionFn { - Arc::new(move |name, request, next| { + Arc::new(move |name, request, context, next| { let user_data = user_data.clone(); let name = name.to_owned(); + let request_codec = context + .request_codec() + .resolve_codec() + .map(NativeHostLlmRequestCodec); Box::pin(async move { let (sender, receiver) = tokio::sync::mpsc::channel(NATIVE_ASYNC_STREAM_CHANNEL_CAPACITY); @@ -3770,10 +4035,16 @@ fn wrap_native_incremental_llm_stream_execution_with_user_data( backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec, #[cfg(test)] before_settlement_lock: None, _callback_user_data: Some(user_data.clone()), }); + let native_context = NativeLlmExecutionContextBridge::new( + &context, + stream.request_codec.as_ref(), + None, + )?; let output = NativeAsyncStreamReceiver { receiver, stream: Arc::clone(&stream), @@ -3801,12 +4072,15 @@ fn wrap_native_incremental_llm_stream_execution_with_user_data( let previous_thread_stack = capture_thread_scope_stack(); sync_thread_scope_stack(current_scope_stack()); let state = catch_unwind(AssertUnwindSafe(|| unsafe { - cb( - user_data.ptr, - invocation, - next_ref as *const NemoRelayNativeAsyncNext, - stream_ref as *const NemoRelayNativeAsyncStream, - ) + native_context.with_native_context(|context| { + cb( + user_data.ptr, + invocation, + std::ptr::from_ref(&context), + next_ref as *const NemoRelayNativeAsyncNext, + stream_ref as *const NemoRelayNativeAsyncStream, + ) + }) })); restore_thread_scope_stack(previous_thread_stack); unsafe { native_string_free(invocation) }; @@ -3867,6 +4141,37 @@ unsafe extern "C" fn native_plugin_context_register_async_stream_middleware( } } +unsafe extern "C" fn native_plugin_context_register_async_llm_execution_intercept( + ctx: *mut NemoRelayNativePluginContext, + name: *const NemoRelayNativeString, + priority: i32, + cb: NemoRelayNativeAsyncLlmExecutionCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> NemoRelayStatus { + clear_native_last_error(); + let user_data_guard = NativeCallbackUserDataGuard::new(user_data, free_fn); + 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 (user_data, free_fn) = user_data_guard.transfer(); + let context = unsafe { &mut *host_ctx.ctx }; + match context.register_llm_execution_intercept( + &name, + priority, + wrap_native_async_llm_execution(instance, cb, user_data, free_fn), + ) { + Ok(()) => NemoRelayStatus::Ok, + Err(error) => status_from_plugin_error(error), + } +} + unsafe extern "C" fn native_plugin_context_register_async_middleware( ctx: *mut NemoRelayNativePluginContext, kind: u32, @@ -3897,9 +4202,13 @@ unsafe extern "C" fn native_plugin_context_register_async_middleware( return NemoRelayStatus::InvalidArg; } }; - if kind == NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept { + if matches!( + kind, + NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept + | NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept + ) { set_native_last_error( - "completion-based LLM stream middleware is unsupported; use plugin_context_register_async_stream_middleware", + "LLM execution middleware requires its dedicated ABI-v7 registration function", ); return NemoRelayStatus::InvalidArg; } @@ -3978,14 +4287,9 @@ unsafe extern "C" fn native_plugin_context_register_async_middleware( 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 => { - unreachable!("completion-based stream middleware was rejected before registration") + NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept + | NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept => { + unreachable!("LLM execution middleware was rejected before registration") } NemoRelayNativeAsyncMiddlewareKind::MarkSanitize => context .register_mark_sanitize_guardrail( @@ -4766,6 +5070,7 @@ unsafe extern "C" fn native_plugin_context_register_llm_execution_intercept( free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { clear_native_last_error(); + let user_data_guard = NativeCallbackUserDataGuard::new(user_data, free_fn); let host_ctx = match host_ctx_mut(ctx) { Ok(ctx) => ctx, Err(status) => return status, @@ -4776,6 +5081,7 @@ unsafe extern "C" fn native_plugin_context_register_llm_execution_intercept( Ok(name) => name, Err(status) => return status, }; + let (user_data, free_fn) = user_data_guard.transfer(); match ctx.register_llm_execution_intercept( &name, priority, @@ -4795,6 +5101,7 @@ unsafe extern "C" fn native_plugin_context_register_llm_stream_execution_interce free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { clear_native_last_error(); + let user_data_guard = NativeCallbackUserDataGuard::new(user_data, free_fn); let host_ctx = match host_ctx_mut(ctx) { Ok(ctx) => ctx, Err(status) => return status, @@ -4805,6 +5112,7 @@ unsafe extern "C" fn native_plugin_context_register_llm_stream_execution_interce Ok(name) => name, Err(status) => return status, }; + let (user_data, free_fn) = user_data_guard.transfer(); match ctx.register_llm_stream_execution_intercept( &name, priority, @@ -5529,10 +5837,12 @@ fn wrap_llm_execution_fn( free_fn: NemoRelayNativeFreeFn, ) -> LlmExecutionFn { let user_data = make_user_data(instance, user_data, free_fn); - Arc::new(move |name, request, next| { + Arc::new(move |name, request, context, next| { let name = name.to_owned(); let user_data = user_data.clone(); - Box::pin(async move { call_llm_execution_callback(cb, &user_data, &name, &request, next) }) + Box::pin(async move { + call_llm_execution_callback(cb, &user_data, &name, &request, &context, next) + }) }) } @@ -5541,9 +5851,23 @@ fn call_llm_execution_callback( user_data: &NativeCallbackUserData, name: &str, request: &LlmRequest, + context: &LlmExecutionContext, next: LlmExecutionNextFn, ) -> FlowResult { clear_native_last_error(); + let request_codec = context + .request_codec() + .resolve_codec() + .map(NativeHostLlmRequestCodec); + let response_codec = context + .response_codec() + .and_then(LlmSanitizeResponseContext::resolve_codec) + .map(NativeHostLlmResponseCodec); + let native_context = NativeLlmExecutionContextBridge::new( + context, + request_codec.as_ref(), + response_codec.as_ref(), + )?; 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) @@ -5552,16 +5876,18 @@ fn call_llm_execution_callback( .ok_or_else(|| FlowError::Internal("failed to allocate native LLM request".into()))?; let next_ctx = Box::into_raw(Box::new(next)) as *mut c_void; let mut out = ptr::null_mut(); - let status = unsafe { + let status = native_context.with_native_context(|context| unsafe { cb( user_data.ptr, name_string, request_string, + context, native_llm_next, next_ctx, &mut out, ) - }; + }); + drop(native_context); unsafe { drop(Box::from_raw(next_ctx as *mut LlmExecutionNextFn)); native_string_free(name_string); @@ -5611,12 +5937,12 @@ fn wrap_llm_stream_execution_fn( free_fn: NemoRelayNativeFreeFn, ) -> LlmStreamExecutionFn { let user_data = make_user_data(instance, user_data, free_fn); - Arc::new(move |name, request, next| { + Arc::new(move |name, request, context, next| { let name = name.to_owned(); let user_data = user_data.clone(); - Box::pin( - async move { call_llm_stream_execution_callback(cb, user_data, &name, &request, next) }, - ) + Box::pin(async move { + call_llm_stream_execution_callback(cb, user_data, &name, &request, &context, next) + }) }) } @@ -5625,9 +5951,17 @@ fn call_llm_stream_execution_callback( user_data: Arc, name: &str, request: &LlmRequest, + context: &LlmExecutionContext, next: LlmStreamExecutionNextFn, ) -> FlowResult { clear_native_last_error(); + let request_codec = context + .request_codec() + .resolve_codec() + .map(NativeHostLlmRequestCodec) + .map(Box::new); + let native_context = + NativeLlmExecutionContextBridge::new(context, request_codec.as_deref(), None)?; 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) @@ -5636,16 +5970,18 @@ fn call_llm_stream_execution_callback( .ok_or_else(|| FlowError::Internal("failed to allocate native LLM request".into()))?; let next_ctx = NativeStreamNextContext::new(Box::into_raw(Box::new(next)) as *mut c_void); let mut out = NemoRelayNativeLlmStreamV1::default(); - let status = unsafe { + let status = native_context.with_native_context(|context| unsafe { cb( user_data.ptr, name_string, request_string, + context, native_llm_stream_next, next_ctx.ptr, &mut out, ) - }; + }); + drop(native_context); unsafe { native_string_free(name_string); native_string_free(request_string); @@ -5657,7 +5993,7 @@ fn call_llm_stream_execution_callback( "native LLM stream execution failed", )); } - native_stream_to_relay_stream(out, Some(next_ctx), Some(user_data)) + native_stream_to_relay_stream(out, Some(next_ctx), Some(user_data), request_codec) } unsafe extern "C" fn native_llm_stream_next( @@ -5698,6 +6034,7 @@ struct NativeRelayLlmStream { finished: bool, _next_ctx: Option, _callback_user_data: Option>, + _request_codec: Option>, } unsafe impl Send for NativeRelayLlmStream {} @@ -5707,6 +6044,7 @@ impl NativeRelayLlmStream { raw: NemoRelayNativeLlmStreamV1, next_ctx: Option, callback_user_data: Option>, + request_codec: Option>, ) -> FlowResult { if raw.struct_size != std::mem::size_of::() { let struct_size = raw.struct_size; @@ -5727,6 +6065,7 @@ impl NativeRelayLlmStream { finished: false, _next_ctx: next_ctx, _callback_user_data: callback_user_data, + _request_codec: request_codec, }) } @@ -5805,11 +6144,13 @@ fn native_stream_to_relay_stream( raw: NemoRelayNativeLlmStreamV1, next_ctx: Option, callback_user_data: Option>, + request_codec: Option>, ) -> FlowResult { Ok(LlmJsonStream::new(NativeRelayLlmStream::from_raw( raw, next_ctx, callback_user_data, + request_codec, )?)) } diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index ae65a49a3..aeaf86879 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -83,7 +83,7 @@ use crate::api::runtime::subscriber_dispatcher::{ PublicationBuffer, capture_nested_publication_buffer, with_nested_publication_buffer, }; use crate::api::runtime::{ - EventMetadataInjectorFn, EventSanitizeFn, LlmCodecIdentity, LlmExecutionCodecContext, + EventMetadataInjectorFn, EventSanitizeFn, LlmCodecIdentity, LlmExecutionContext, LlmExecutionNextFn, LlmJsonStream, LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn, LlmStreamInner, MiddlewareContinuationContext, ToolExecutionContext, ToolExecutionNextFn, current_scope_stack, with_scope_stack, @@ -106,7 +106,7 @@ use super::{ DynamicPluginKind, DynamicPluginManifest, DynamicPluginManifestLoad, DynamicPluginTeardownOutcome, WorkerRuntime, deregister_tracked_registrations_checked, validate_annotated_request_consumer_compatibility, validate_dynamic_plugin_relay_compatibility, - validate_tool_execution_context_compatibility, + validate_llm_execution_context_compatibility, validate_tool_execution_context_compatibility, }; const JSON_SCHEMA: &str = "nemo.relay.Json@1"; @@ -1133,6 +1133,17 @@ impl WorkerPluginInstance { }) { validate_tool_execution_context_compatibility(&self.relay_compat, &self.plugin_kind)?; } + if registrations.iter().any(|registration| { + RegistrationSurface::try_from(registration.surface).is_ok_and(|surface| { + matches!( + surface, + RegistrationSurface::LlmExecutionIntercept + | RegistrationSurface::LlmStreamExecutionIntercept + ) + }) + }) { + validate_llm_execution_context_compatibility(&self.relay_compat, &self.plugin_kind)?; + } let initial_gates = register.conditional_middleware_guardrails; for gate in initial_gates { let kinds = gate @@ -1413,7 +1424,6 @@ impl WorkerPluginInstance { ) -> crate::plugin::Result<()> { let name = registration.local_name.as_str(); let priority = registration.priority; - let include_codec_context = registration.llm_execution_codec_context; let instance = Arc::new(self.clone_for_callback()); let callback_name = name.to_owned(); match surface { @@ -1476,29 +1486,28 @@ impl WorkerPluginInstance { }) }), ), - RegistrationSurface::LlmExecutionIntercept => ctx - .register_contextual_llm_execution_intercept( - name, - priority, - Arc::new(move |model_name, request, context, next| { - let instance = instance.clone(); - let callback_name = callback_name.clone(); - let model_name = model_name.to_owned(); - Box::pin(async move { - instance - .invoke_llm_execution( - &callback_name, - &model_name, - request, - include_codec_context.then_some(context), - next, - ) - .await - }) - }), - ), + RegistrationSurface::LlmExecutionIntercept => ctx.register_llm_execution_intercept( + name, + priority, + Arc::new(move |model_name, request, context, next| { + let instance = instance.clone(); + let callback_name = callback_name.clone(); + let model_name = model_name.to_owned(); + Box::pin(async move { + instance + .invoke_llm_execution( + &callback_name, + &model_name, + request, + context, + next, + ) + .await + }) + }), + ), RegistrationSurface::LlmStreamExecutionIntercept => ctx - .register_contextual_llm_stream_execution_intercept( + .register_llm_stream_execution_intercept( name, priority, Arc::new(move |model_name, request, context, next| { @@ -1511,7 +1520,7 @@ impl WorkerPluginInstance { &callback_name, &model_name, request, - include_codec_context.then_some(context), + context, next, ) .await @@ -2038,7 +2047,7 @@ impl WorkerPluginCallback { registration_name: &str, model_name: &str, request: LlmRequest, - execution_context: Option, + execution_context: LlmExecutionContext, next: LlmExecutionNextFn, ) -> FlowResult { let continuation_id = self @@ -2055,13 +2064,9 @@ impl WorkerPluginCallback { None, )), ); - let codec_capabilities = execution_context - .as_ref() - .map(|context| self.attach_llm_execution_codec_context(&mut invoke, context, true)) - .transpose(); - let _codec_capabilities = self - .cleanup_after_setup_error(&invoke, codec_capabilities)? - .unwrap_or_default(); + let codec_capabilities = + self.attach_llm_execution_codec_context(&mut invoke, &execution_context); + let _codec_capabilities = self.cleanup_after_setup_error(&invoke, codec_capabilities)?; json_from_invoke_response(self.invoke_async(invoke).await?) } @@ -2070,7 +2075,7 @@ impl WorkerPluginCallback { registration_name: &str, model_name: &str, request: LlmRequest, - execution_context: Option, + execution_context: LlmExecutionContext, next: LlmStreamExecutionNextFn, ) -> FlowResult { let continuation_id = self @@ -2087,13 +2092,9 @@ impl WorkerPluginCallback { None, )), ); - let codec_capabilities = execution_context - .as_ref() - .map(|context| self.attach_llm_execution_codec_context(&mut invoke, context, false)) - .transpose(); - let codec_capabilities = self - .cleanup_after_setup_error(&invoke, codec_capabilities)? - .unwrap_or_default(); + let codec_capabilities = + self.attach_llm_execution_codec_context(&mut invoke, &execution_context); + let codec_capabilities = self.cleanup_after_setup_error(&invoke, codec_capabilities)?; let mut client = self.client.clone(); let mut guard = WorkerInvocationGuard::new(self, &invoke); let (tx, rx) = mpsc::channel(16); @@ -2172,16 +2173,15 @@ impl WorkerPluginCallback { fn attach_llm_execution_codec_context( &self, invoke: &mut InvokeRequest, - context: &LlmExecutionCodecContext, - include_response_capability: bool, + context: &LlmExecutionContext, ) -> FlowResult> { let mut guards = Vec::with_capacity(2); let mut request = ProtoLlmSanitizeRequestContext { - codec: Some(codec_identity_to_proto(context.request().codec())), + codec: Some(codec_identity_to_proto(context.request_codec().codec())), codec_capability_id: None, }; - if let Some(codec) = context.request().resolve_codec() { + if let Some(codec) = context.request_codec().resolve_codec() { let capability = self .host_state .issue_request_codec(&invoke.invocation_id, codec)?; @@ -2189,24 +2189,30 @@ impl WorkerPluginCallback { guards.push(capability); } - let mut response = ProtoLlmSanitizeResponseContext { - codec: Some(codec_identity_to_proto(context.response().codec())), - codec_capability_id: None, - }; - if include_response_capability && let Some(codec) = context.response().resolve_codec() { - let capability = self - .host_state - .issue_response_codec(&invoke.invocation_id, codec)?; - response.codec_capability_id = Some(capability.id().into()); - guards.push(capability); - } + let response = context + .response_codec() + .map(|response_context| -> FlowResult<_> { + let mut response = ProtoLlmSanitizeResponseContext { + codec: Some(codec_identity_to_proto(response_context.codec())), + codec_capability_id: None, + }; + if let Some(codec) = response_context.resolve_codec() { + let capability = self + .host_state + .issue_response_codec(&invoke.invocation_id, codec)?; + response.codec_capability_id = Some(capability.id().into()); + guards.push(capability); + } + Ok(response) + }) + .transpose()?; let Some(invoke_request_payload::Payload::Llm(llm)) = invoke.payload.as_mut() else { unreachable!("LLM execution invocation must have an LLM payload"); }; llm.execution_codec_context = Some(Box::new(ProtoLlmExecutionCodecContext { request: Some(request), - response: Some(response), + response, })); Ok(guards) } @@ -4108,18 +4114,6 @@ fn validate_registration_plan( "worker plugin '{plugin_id}' returned unspecified registration surface" ))); } - if registration.llm_execution_codec_context - && !matches!( - surface, - RegistrationSurface::LlmExecutionIntercept - | RegistrationSurface::LlmStreamExecutionIntercept - ) - { - return Err(PluginError::RegistrationFailed(format!( - "worker plugin '{plugin_id}' requested LLM execution codec context for incompatible surface {}", - surface.as_str_name() - ))); - } } let mut gate_names = std::collections::HashSet::new(); for gate in &response.conditional_middleware_guardrails { diff --git a/crates/core/src/plugins/nemo_guardrails/python.rs b/crates/core/src/plugins/nemo_guardrails/python.rs index b50f6f490..11a4180f4 100644 --- a/crates/core/src/plugins/nemo_guardrails/python.rs +++ b/crates/core/src/plugins/nemo_guardrails/python.rs @@ -54,7 +54,7 @@ pub(super) fn register_local_backend( let llm_runtime = Arc::clone(&runtime); let enable_input = config.input; let enable_output = config.output; - let llm_execution: LlmExecutionFn = Arc::new(move |_name, request, next| { + let llm_execution: LlmExecutionFn = Arc::new(move |_name, request, _context, next| { let runtime = Arc::clone(&llm_runtime); Box::pin(async move { runtime @@ -71,14 +71,15 @@ pub(super) fn register_local_backend( let stream_runtime = Arc::clone(&runtime); let enable_input = config.input; let enable_output = config.output; - let llm_stream_execution: LlmStreamExecutionFn = Arc::new(move |_name, request, next| { - let runtime = Arc::clone(&stream_runtime); - Box::pin(async move { - runtime - .execute_llm_stream(request, next, enable_input, enable_output) - .await - }) - }); + let llm_stream_execution: LlmStreamExecutionFn = + Arc::new(move |_name, request, _context, next| { + let runtime = Arc::clone(&stream_runtime); + Box::pin(async move { + runtime + .execute_llm_stream(request, next, enable_input, enable_output) + .await + }) + }); ctx.register_llm_stream_execution_intercept( "nemo_guardrails_local_stream", config.priority, diff --git a/crates/core/src/plugins/nemo_guardrails/remote.rs b/crates/core/src/plugins/nemo_guardrails/remote.rs index bf1f713e3..d349cd151 100644 --- a/crates/core/src/plugins/nemo_guardrails/remote.rs +++ b/crates/core/src/plugins/nemo_guardrails/remote.rs @@ -979,17 +979,18 @@ pub(super) fn register_remote_backend( if config.input || config.output { let llm_execution_runtime = Arc::clone(&runtime); - let llm_execution: LlmExecutionFn = Arc::new(move |_name, request, _next| { + let llm_execution: LlmExecutionFn = Arc::new(move |_name, request, _context, _next| { let runtime = Arc::clone(&llm_execution_runtime); Box::pin(async move { runtime.execute(request, false).await }) }); ctx.register_llm_execution_intercept("llm_remote_backend", config.priority, llm_execution)?; let llm_stream_runtime = Arc::clone(&runtime); - let llm_stream_execution: LlmStreamExecutionFn = Arc::new(move |_name, request, _next| { - let runtime = Arc::clone(&llm_stream_runtime); - Box::pin(async move { runtime.execute_stream(request).await }) - }); + let llm_stream_execution: LlmStreamExecutionFn = + Arc::new(move |_name, request, _context, _next| { + let runtime = Arc::clone(&llm_stream_runtime); + Box::pin(async move { runtime.execute_stream(request).await }) + }); ctx.register_llm_stream_execution_intercept( "llm_stream_remote_backend", config.priority, diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index fda6a7907..fd1385191 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -12,14 +12,16 @@ use nemo_relay_plugin::{ Json, LlmJsonAsyncStream, LlmRequest, LlmRequestInterceptOutcome, MetricKind, MetricMeasurement, MetricValueType, NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, - NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, NativeExecutorConfig, NativePlugin, - NemoRelayNativeAsyncCallbackState, - NemoRelayNativeAsyncMiddlewareCb, NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, - NemoRelayNativeAsyncStream, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, - NemoRelayNativeHostApiV4, NemoRelayNativeHostApiV5, NemoRelayNativePluginContext, - NemoRelayNativePluginV1, NemoRelayNativeString, NemoRelayNativeToolNextFn, NemoRelayStatus, - PendingMarkSpec, PluginContext, PluginRuntime, RuntimeRegistrationKind, ScopeCategory, - ScopeType, ToolExecutionInterceptOutcome, + NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, + NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, NativeExecutorConfig, NativePlugin, + NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncMiddlewareCb, + NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, NemoRelayNativeAsyncStream, + NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, + NemoRelayNativeLlmExecutionContext, NemoRelayNativePluginContext, NemoRelayNativePluginV1, + NemoRelayNativeString, NemoRelayNativeToolNextFn, NemoRelayStatus, PendingMarkSpec, + PluginContext, PluginRuntime, RuntimeRegistrationKind, ScopeCategory, ScopeType, + ToolExecutionInterceptOutcome, }; use serde_json::{Map, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; @@ -328,7 +330,7 @@ impl NativePlugin for FixtureNativePlugin { ctx.register_llm_execution_intercept( "fixture_llm_execution", 0, - |_name, request, next| async move { + |_name, request, _context, next| async move { let response = next .call(mark_llm_request( request, @@ -341,7 +343,7 @@ impl NativePlugin for FixtureNativePlugin { ctx.register_llm_stream_execution_intercept( "fixture_llm_stream_execution", 0, - |_name, request, next| async move { + |_name, request, _context, next| async move { let stream = next .call(mark_llm_request( request, @@ -528,7 +530,7 @@ pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v2( } } -/// Raw ABI-v5 entry used to verify the current table. +/// Raw ABI-v5 entry used to verify that stale native binaries are rejected. #[unsafe(no_mangle)] pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v5( host: *const NemoRelayNativeHostApiV1, @@ -538,13 +540,30 @@ pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v5( fixture_compat_entry( host, out, - NEMO_RELAY_NATIVE_ABI_VERSION, + NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, std::mem::size_of::(), b"fixture_native_v5", ) } } +/// Raw ABI-v6 entry used to verify that the immediately stale callback layout is rejected. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v6( + host: *const NemoRelayNativeHostApiV1, + out: *mut NemoRelayNativePluginV1, +) -> NemoRelayStatus { + unsafe { + fixture_compat_entry( + host, + out, + NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, + std::mem::size_of::(), + b"fixture_native_v6", + ) + } +} + /// Raw ABI-v4 entry used to verify fallback for plugins built with the previous SDK. #[unsafe(no_mangle)] pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v4( @@ -943,7 +962,7 @@ unsafe extern "C" fn raw_register_event_sanitize_errors( } struct FixtureAsyncPlugin { - host: Option>, + host: Option>, } impl NativePlugin for FixtureAsyncPlugin { @@ -957,25 +976,25 @@ impl NativePlugin for FixtureAsyncPlugin { ctx: &mut PluginContext<'_>, ) -> nemo_relay_plugin::Result<()> { let host = ctx.host_api(); - if host.abi_version < 3 - || host.struct_size < std::mem::size_of::() + if host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION + || host.struct_size < std::mem::size_of::() { - return Err("fixture async plugin requires ABI v3".into()); + return Err("fixture async plugin requires ABI v7".into()); } self.host = Some(Box::new(unsafe { - *(host as *const _ as *const NemoRelayNativeHostApiV3) + *(host as *const _ as *const NemoRelayNativeHostApiV7) })); let user_data = self .host .as_deref() - .map(|host| (host as *const NemoRelayNativeHostApiV3).cast_mut().cast()) + .map(|host| (host as *const NemoRelayNativeHostApiV7).cast_mut().cast()) .expect("fixture async host was initialized"); let registrations: [( NemoRelayNativeAsyncMiddlewareKind, &str, NemoRelayNativeAsyncMiddlewareCb, - ); 13] = [ + ); 12] = [ ( NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest, "fixture_async_tool_sanitize_request", @@ -1021,11 +1040,6 @@ impl NativePlugin for FixtureAsyncPlugin { "fixture_async_llm_request", raw_async_passthrough_callback, ), - ( - NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept, - "fixture_async_llm_execution", - raw_async_tool_execution_callback, - ), ( NemoRelayNativeAsyncMiddlewareKind::MarkSanitize, "fixture_async_mark", @@ -1058,6 +1072,20 @@ impl NativePlugin for FixtureAsyncPlugin { return Err(format!("async registration failed: {status:?}")); } } + let status = unsafe { + ctx.register_async_llm_execution_intercept_raw( + "fixture_async_llm_execution", + 0, + raw_async_llm_execution_callback, + user_data, + None, + ) + }; + if status != NemoRelayStatus::Ok { + return Err(format!( + "async LLM execution registration failed: {status:?}" + )); + } let status = unsafe { ctx.register_async_stream_middleware_raw( "fixture_async_llm_stream", @@ -1115,6 +1143,7 @@ unsafe extern "C" fn raw_async_stream_forward( unsafe extern "C" fn raw_async_stream_callback( user_data: *mut c_void, invocation_json: *const NemoRelayNativeString, + _context: *const NemoRelayNativeLlmExecutionContext, next: *const NemoRelayNativeAsyncNext, stream: *const NemoRelayNativeAsyncStream, ) -> u32 { @@ -1161,6 +1190,16 @@ unsafe extern "C" fn raw_async_stream_callback( } } +unsafe extern "C" fn raw_async_llm_execution_callback( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + _context: *const NemoRelayNativeLlmExecutionContext, + next: *const NemoRelayNativeAsyncNext, + completion: *const nemo_relay_plugin::NemoRelayNativeAsyncCompletion, +) -> u32 { + unsafe { raw_async_tool_execution_callback(user_data, invocation_json, next, completion) } +} + unsafe extern "C" fn raw_async_allow_callback( user_data: *mut c_void, _invocation_json: *const NemoRelayNativeString, diff --git a/crates/core/tests/fixtures/worker_plugin/src/main.rs b/crates/core/tests/fixtures/worker_plugin/src/main.rs index 9089b515d..0de5ae310 100644 --- a/crates/core/tests/fixtures/worker_plugin/src/main.rs +++ b/crates/core/tests/fixtures/worker_plugin/src/main.rs @@ -375,7 +375,7 @@ fn register_fixture_llm_hooks( ctx.register_llm_execution_intercept( "fixture_llm_execution", 0, - |_name, request, next: LlmNext| async move { + |_name, request, _context, next: LlmNext| async move { let response = next .call(mark_llm_request( request, @@ -388,7 +388,7 @@ fn register_fixture_llm_hooks( ctx.register_llm_stream_execution_intercept( "fixture_llm_stream_execution", 0, - move |_name, request, next: LlmStreamNext| async move { + move |_name, request, _context, next: LlmStreamNext| async move { if llm_stream_open_error { return Err(WorkerSdkError::Callback( "fixture LLM stream open error requested".into(), diff --git a/crates/core/tests/integration/api_surface_tests.rs b/crates/core/tests/integration/api_surface_tests.rs index 6233468fc..0d2b80c06 100644 --- a/crates/core/tests/integration/api_surface_tests.rs +++ b/crates/core/tests/integration/api_surface_tests.rs @@ -1212,7 +1212,7 @@ fn assert_global_llm_registry() { register_llm_execution_intercept( "llm-execution", 1, - Arc::new(|_name, request, _next| Box::pin(async move { Ok(request.content) })), + Arc::new(|_name, request, _context, _next| Box::pin(async move { Ok(request.content) })), ) .unwrap(); assert!(deregister_llm_execution_intercept("llm-execution").unwrap()); @@ -1220,7 +1220,7 @@ fn assert_global_llm_registry() { register_llm_stream_execution_intercept( "llm-stream", 1, - Arc::new(|_name, request, _next| { + Arc::new(|_name, request, _context, _next| { Box::pin(async move { Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( request.content @@ -1479,7 +1479,7 @@ fn assert_scope_llm_registry(scope_uuid: &uuid::Uuid) { scope_uuid, "llm-execution", 1, - Arc::new(|_name, request, _next| Box::pin(async move { Ok(request.content) })), + Arc::new(|_name, request, _context, _next| Box::pin(async move { Ok(request.content) })), ) .unwrap(); assert!(scope_deregister_llm_execution_intercept(scope_uuid, "llm-execution").unwrap()); @@ -1488,7 +1488,7 @@ fn assert_scope_llm_registry(scope_uuid: &uuid::Uuid) { scope_uuid, "llm-stream", 1, - Arc::new(|_name, request, _next| { + Arc::new(|_name, request, _context, _next| { Box::pin(async move { Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( request.content diff --git a/crates/core/tests/integration/middleware_tests.rs b/crates/core/tests/integration/middleware_tests.rs index 58bd5c4e3..df0ca2404 100644 --- a/crates/core/tests/integration/middleware_tests.rs +++ b/crates/core/tests/integration/middleware_tests.rs @@ -2228,7 +2228,7 @@ async fn execution_next_is_revoked_after_each_interceptor_settles() { register_llm_execution_intercept( "late_llm_next", 1, - Arc::new(move |_name, _request, next| { + Arc::new(move |_name, _request, _context, next| { *captured_llm_next.lock().unwrap() = Some(next); ready_result(Ok(json!({"source": "llm-intercept"}))) }), @@ -2268,7 +2268,7 @@ async fn execution_next_is_revoked_after_each_interceptor_settles() { register_llm_stream_execution_intercept( "late_llm_stream_next", 1, - Arc::new(move |_name, _request, next| { + Arc::new(move |_name, _request, _context, next| { *captured_stream_next.lock().unwrap() = Some(next); Box::pin(async { Ok(LlmJsonStream::new(futures::stream::iter(vec![Ok( @@ -2332,7 +2332,7 @@ async fn stream_next_is_revoked_when_the_managed_stream_terminalizes_with_an_err register_llm_stream_execution_intercept( "upstream_error_stream_next", 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { *captured_upstream_error_next.lock().unwrap() = Some(next.clone()); next(request) }), @@ -2377,7 +2377,7 @@ async fn stream_next_is_revoked_when_the_managed_stream_terminalizes_with_an_err register_llm_stream_execution_intercept( "collector_error_stream_next", 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { *captured_collector_error_next.lock().unwrap() = Some(next.clone()); next(request) }), @@ -2592,7 +2592,7 @@ async fn stream_next_preserves_each_invocation_scope_while_polling() { register_llm_stream_execution_intercept( "scoped_stream_next", 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { let first_stack = first_stack.clone(); let second_stack = second_stack.clone(); Box::pin(async move { @@ -2670,7 +2670,7 @@ async fn stream_next_remains_active_during_interceptor_stream_close() { register_llm_stream_execution_intercept( "close_calls_stream_next", 1, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { Box::pin(async move { Ok(LlmJsonStream::from_closeable(CloseCallsStreamNext { next: Some(next), @@ -3029,7 +3029,7 @@ async fn dropping_pending_llm_execution_closes_the_managed_lifecycle() { register_llm_execution_intercept( "pending_llm_execution", 1, - Arc::new(move |_name, _request, _next| { + Arc::new(move |_name, _request, _context, _next| { if let Some(sender) = entered_tx.lock().unwrap().take() { let _ = sender.send(()); } @@ -4496,7 +4496,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { register_llm_execution_intercept( "lock_global_llm_execution", 1, - Arc::new(move |_, request, next| { + Arc::new(move |_, request, _context, next| { record_middleware_callback(&tracked, "llm_execution_global"); assert_middleware_callback_locks_are_free(); Box::pin(async move { next(request).await }) @@ -4508,7 +4508,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { &scope.uuid, "lock_scope_llm_execution", 2, - Arc::new(move |_, request, next| { + Arc::new(move |_, request, _context, next| { record_middleware_callback(&tracked, "llm_execution_scope"); assert_middleware_callback_locks_are_free(); Box::pin(async move { next(request).await }) @@ -4519,7 +4519,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { register_llm_stream_execution_intercept( "lock_global_llm_stream_execution", 1, - Arc::new(move |_, request, next| { + Arc::new(move |_, request, _context, next| { record_middleware_callback(&tracked, "llm_stream_execution_global"); assert_middleware_callback_locks_are_free(); Box::pin(async move { next(request).await }) @@ -4531,7 +4531,7 @@ async fn test_llm_middleware_callbacks_run_without_registry_or_scope_locks() { &scope.uuid, "lock_scope_llm_stream_execution", 2, - Arc::new(move |_, request, next| { + Arc::new(move |_, request, _context, next| { record_middleware_callback(&tracked, "llm_stream_execution_scope"); assert_middleware_callback_locks_are_free(); Box::pin(async move { next(request).await }) @@ -5883,7 +5883,7 @@ async fn test_managed_llm_materializes_optimization_mark_and_end_summary() { register_llm_execution_intercept( "optimization_execution_contributor", 1, - Arc::new(|_name, request, next| { + Arc::new(|_name, request, _context, next| { let contribution = LlmOptimizationContribution::new("test.execution", "test_execution_kind"); assert!(record_llm_optimization_contribution(contribution)); @@ -5986,7 +5986,7 @@ async fn execution_optimization_mark_keeps_decision_commit_timestamp_order() { register_llm_execution_intercept( "optimization_timestamp_contributor", 1, - Arc::new(|_name, request, next| { + Arc::new(|_name, request, _context, next| { Box::pin(async move { event( EmitMarkEventParams::builder() @@ -6096,7 +6096,7 @@ async fn test_stream_optimization_mark_uses_the_llm_captured_sanitizer_scope() { register_llm_stream_execution_intercept( "stream_optimization_sanitizer_contributor", 1, - Arc::new(|_name, request, next| { + Arc::new(|_name, request, _context, next| { let mut contribution = LlmOptimizationContribution::new("test.stream", "stream_test"); contribution.payload_schema = Some(DataSchema { name: "test.stream_evidence".to_string(), @@ -6332,7 +6332,7 @@ async fn test_llm_execution_intercept_chain() { register_llm_execution_intercept( "llm_exec_1", 1, - Arc::new(move |_name, req, next| { + Arc::new(move |_name, req, _context, next| { let o = o1.clone(); Box::pin(async move { o.lock().unwrap().push("intercept_before".into()); @@ -6413,7 +6413,7 @@ async fn test_llm_start_emits_before_short_circuit_execution_intercept() { register_llm_execution_intercept( "llm_short_circuit_exec", 1, - Arc::new(move |_name, mut req, _next| { + Arc::new(move |_name, mut req, _context, _next| { Box::pin(async move { req.content .as_object_mut() @@ -6507,7 +6507,7 @@ async fn test_llm_stream_start_emits_before_short_circuit_execution_intercept() register_llm_stream_execution_intercept( "llm_stream_short_circuit_exec", 1, - Arc::new(move |_name, mut req, _next| { + Arc::new(move |_name, mut req, _context, _next| { Box::pin(async move { req.content .as_object_mut() diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index 656d91c07..84c6ffffc 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -965,47 +965,17 @@ async fn native_tool_execution_rejects_null_malformed_and_error_outcomes() { activation.clear(); } -#[tokio::test] -async fn native_api_one_preserves_results_after_abi_v2_negotiation() { - let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; +#[test] +fn native_api_one_does_not_admit_a_stale_abi_v2_binary() { + let _guard = NATIVE_PLUGIN_TEST_LOCK.blocking_lock(); let fixture = build_fixture_plugin(); let manifest_ref = write_manifest_with_symbol(&fixture, "nemo_relay_fixture_abi_v2_api1"); - let activation = load_native_plugins([load_spec("fixture_native", &manifest_ref)]) - .expect("native API 1 plugin should negotiate the ABI v2 host table"); - let mut cleanup = NativePluginTestCleanup::new(); - - let mut plugin_config = PluginConfig::default(); - plugin_config.components.push(PluginComponentSpec { - kind: "fixture_native".into(), - enabled: true, - config: Map::new(), - }); - test_initialize_plugin_host_exact(plugin_config) - .await - .expect("ABI v2 native API 1 fixture should initialize"); - cleanup.mark_plugin_configuration_active(); - - let result = tool_call_execute( - ToolCallExecuteParams::builder() - .name("fixture-abi-v2-api1") - .args(json!({"input": true})) - .func(Arc::new(|args| { - Box::pin(async move { - Ok(ToolExecutionResult::annotated( - args, - json!({"source": "provider"}), - )) - }) - })) - .build(), - ) - .await - .expect("ABI v2 callback should use the canonical native API 1 result contract"); - assert_eq!(result.result, json!({"input": true})); - assert_eq!(result.annotation, Some(json!({"source": "provider"}))); - - drop(cleanup); - activation.clear(); + let error = expect_native_load_error_from_specs( + [load_spec("fixture_native", &manifest_ref)], + "a native_api=1 manifest must not make an ABI-v2 binary compatible", + ); + assert!(error.contains("rejected native ABI 7"), "{error}"); + assert!(error.contains("rebuild the plugin"), "{error}"); } #[tokio::test] @@ -1128,6 +1098,28 @@ async fn native_loader_rejects_manifest_that_admits_pre_zero_eight_relay() { ); } +#[tokio::test] +async fn native_abi_v7_rejects_manifest_that_admits_relay_zero_nine() { + let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; + let fixture = build_fixture_plugin(); + let manifest_ref = write_manifest_text(ManifestOptions { + manifest_dir: fixture.manifest_dir.path(), + plugin_id: "fixture_native", + relay: ">=0.9,<1.0", + library: &fixture.library_path.to_string_lossy(), + symbol: "nemo_relay_fixture_native_plugin", + integrity: None, + }); + let error = expect_native_load_error_from_specs( + [load_spec("fixture_native", &manifest_ref)], + "an ABI-v7 native plugin must exclude Relay 0.9", + ); + assert!( + error.contains("uses native ABI v7") && error.contains("excludes Relay 0.9"), + "{error}" + ); +} + #[test] fn native_loader_resolves_manifest_directory_and_relative_library_paths() { let _guard = NATIVE_PLUGIN_TEST_LOCK.blocking_lock(); @@ -1158,7 +1150,7 @@ fn native_loader_resolves_manifest_directory_and_relative_library_paths() { } #[test] -fn native_loader_falls_back_to_abi_v3_plugins() { +fn native_loader_rejects_abi_v3_plugins() { let _guard = NATIVE_PLUGIN_TEST_LOCK.blocking_lock(); let fixture = build_fixture_plugin(); let manifest_ref = write_manifest_text(ManifestOptions { @@ -1170,17 +1162,21 @@ fn native_loader_falls_back_to_abi_v3_plugins() { integrity: None, }); - let activation = load_native_plugins([load_spec("fixture_native_v3", &manifest_ref)]) - .expect("ABI-v3 fixture should load through compatibility fallback"); - activation.clear(); + let error = expect_native_load_error_from_specs( + [load_spec("fixture_native_v3", &manifest_ref)], + "ABI-v3 plugins must be rebuilt for ABI v7", + ); + assert!(error.contains("rejected native ABI 7"), "{error}"); + assert!(error.contains("rebuild the plugin"), "{error}"); } #[test] -fn native_loader_supports_current_v5_frozen_v4_and_legacy_v2_plugins() { +fn native_loader_rejects_v2_v4_v5_and_v6_plugins() { let _guard = NATIVE_PLUGIN_TEST_LOCK.blocking_lock(); let fixture = build_fixture_plugin(); for (plugin_id, symbol) in [ + ("fixture_native_v6", "nemo_relay_fixture_native_plugin_v6"), ("fixture_native_v5", "nemo_relay_fixture_native_plugin_v5"), ("fixture_native_v4", "nemo_relay_fixture_native_plugin_v4"), ("fixture_native_v2", "nemo_relay_fixture_native_plugin_v2"), @@ -1194,9 +1190,12 @@ fn native_loader_supports_current_v5_frozen_v4_and_legacy_v2_plugins() { integrity: None, }); - let activation = load_native_plugins([load_spec(plugin_id, &manifest_ref)]) - .unwrap_or_else(|error| panic!("ABI compatibility fixture should load: {error}")); - activation.clear(); + let error = expect_native_load_error_from_specs( + [load_spec(plugin_id, &manifest_ref)], + "stale native plugins must be rebuilt for ABI v7", + ); + assert!(error.contains("rejected native ABI 7"), "{error}"); + assert!(error.contains("rebuild the plugin"), "{error}"); } } diff --git a/crates/core/tests/integration/worker_plugin_tests.rs b/crates/core/tests/integration/worker_plugin_tests.rs index 532c71518..f29b2d62d 100644 --- a/crates/core/tests/integration/worker_plugin_tests.rs +++ b/crates/core/tests/integration/worker_plugin_tests.rs @@ -1455,6 +1455,28 @@ fn worker_loader_rejects_manifest_that_admits_pre_zero_eight_relay() { ); } +#[test] +fn worker_llm_execution_context_requires_zero_ten_compatibility() { + let _guard = WORKER_PLUGIN_TEST_LOCK.blocking_lock(); + let fixture = build_fixture_worker(); + let (_manifest_dir, manifest_ref) = + write_manifest_with_relay(fixture.binary_path(), ">=0.9,<1.0"); + + let error = match load_worker_plugins([WorkerPluginLoadSpec { + plugin_id: "fixture_worker".into(), + manifest_ref: manifest_ref.to_string_lossy().into_owned(), + environment_ref: None, + config: Map::new(), + }]) { + Ok(activation) => { + activation.clear(); + panic!("an execution-context worker must exclude Relay 0.9"); + } + Err(error) => error.to_string(), + }; + assert!(error.contains("excludes Relay 0.9"), "{error}"); +} + #[test] fn invalid_worker_relay_requirement_reports_parse_error() { let _guard = WORKER_PLUGIN_TEST_LOCK.blocking_lock(); @@ -2053,29 +2075,31 @@ class CodecContextProbe(WorkerPlugin): del config async def execute(_name, request, context, next_call): - if not context.available: - raise RuntimeError("execution codec context is unavailable") - if context.request_codec is None or context.response_codec is None: - raise RuntimeError("directional codec proxy is unavailable") - - annotated = await context.request_codec.decode(request) + if context.response_codec is None: + raise RuntimeError("unary response codec context is unavailable") + request_codec = context.request_codec.resolve_codec() + response_codec = context.response_codec.resolve_codec() + if request_codec is None or response_codec is None: + raise RuntimeError("directional codec capability is unavailable") + + annotated = await request_codec.decode(request) annotated["model"] = "worker-model" - encoded = await context.request_codec.encode(annotated, request) + encoded = await request_codec.encode(annotated, request) response = await next_call.call(encoded) - decoded = await context.response_codec.decode(response) + decoded = await response_codec.decode(response) result = dict(response) result["_codec_context_probe"] = { - "request_kind": context.request_codec_identity.kind, - "request_id": context.request_codec_identity.id, - "response_kind": context.response_codec_identity.kind, - "response_id": context.response_codec_identity.id, + "request_kind": context.request_codec.codec.kind, + "request_id": context.request_codec.codec.id, + "response_kind": context.response_codec.codec.kind, + "response_id": context.response_codec.codec.id, "decoded_model": decoded.get("model"), "decoded_message": decoded.get("message"), } return result - ctx.register_llm_execution_intercept_with_context("codec_context_probe", execute) + ctx.register_llm_execution_intercept("codec_context_probe", execute) async def main(): diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index 550358444..b661b5991 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -12,7 +12,7 @@ use crate::api::optimization::{ LlmOptimizationRecorder, record_llm_optimization_contribution, scope_llm_optimization_recorder, }; use crate::api::runtime::{ - BuiltinLlmCodec, LlmCodecIdentity, LlmExecutionCodecContext, LlmExecutionNextFn, + BuiltinLlmCodec, LlmCodecIdentity, LlmExecutionContext, LlmExecutionNextFn, LlmSanitizeRequestContext, LlmSanitizeResponseContext, MiddlewareContinuationLease, NemoRelayContextState, }; @@ -529,7 +529,6 @@ fn registration_plan_and_scope_type_helpers_validate_edges() { surface: RegistrationSurface::Subscriber as i32, priority: 0, break_chain: false, - llm_execution_codec_context: false, }], error: None, conditional_middleware_guardrails: Vec::new(), @@ -546,7 +545,6 @@ fn registration_plan_and_scope_type_helpers_validate_edges() { surface: 999, priority: 0, break_chain: false, - llm_execution_codec_context: false, }], error: None, conditional_middleware_guardrails: Vec::new(), @@ -567,7 +565,6 @@ fn registration_plan_and_scope_type_helpers_validate_edges() { surface: RegistrationSurface::Unspecified as i32, priority: 0, break_chain: false, - llm_execution_codec_context: false, }], error: None, conditional_middleware_guardrails: Vec::new(), @@ -580,27 +577,6 @@ fn registration_plan_and_scope_type_helpers_validate_edges() { .contains("unspecified registration surface") ); - let incompatible_codec_context = validate_registration_plan( - "fixture_worker", - &RegisterResponse { - registrations: vec![Registration { - local_name: "subscriber".into(), - surface: RegistrationSurface::Subscriber as i32, - priority: 0, - break_chain: false, - llm_execution_codec_context: true, - }], - error: None, - conditional_middleware_guardrails: Vec::new(), - }, - ) - .expect_err("codec context must be limited to LLM execution surfaces"); - assert!( - incompatible_codec_context - .to_string() - .contains("incompatible surface") - ); - let cases = [ (ProtoScopeType::Agent, crate::api::scope::ScopeType::Agent), ( @@ -1239,7 +1215,7 @@ async fn llm_worker_codec_capabilities_are_active_only_during_sanitizer_invocati } #[tokio::test(flavor = "multi_thread")] -async fn llm_worker_execution_codec_context_is_opt_in_and_ephemeral() { +async fn llm_worker_execution_codec_context_is_required_and_ephemeral() { enable_operational_logs(); let host_state = shared_worker_host_state(); let seen = Arc::new(Mutex::new(None::)); @@ -1247,26 +1223,16 @@ async fn llm_worker_execution_codec_context_is_opt_in_and_ephemeral() { let host_state = Arc::clone(&host_state); let seen = Arc::clone(&seen); move |request| { - let registration_name = request.registration_name.clone(); - if registration_name == "legacy" { - let Some(invoke_request_payload::Payload::Llm(invocation)) = request.payload else { - panic!("LLM execution must receive an LLM invocation"); - }; - assert!(invocation.sanitize_context.is_none()); - assert!(invocation.execution_codec_context.is_none()); - } else { - let (request_id, response_id, invocation_id) = - execution_codec_capabilities(request); - let response_id = response_id.expect("response capability must be present"); - let state = host_state.lock().unwrap().clone().unwrap(); - state - .request_codec(&request_id, &invocation_id) - .expect("request capability resolves during callback"); - state - .response_codec(&response_id, &invocation_id) - .expect("response capability resolves during callback"); - *seen.lock().unwrap() = Some((request_id, Some(response_id), invocation_id)); - } + let (request_id, response_id, invocation_id) = execution_codec_capabilities(request); + let response_id = response_id.expect("response capability must be present"); + let state = host_state.lock().unwrap().clone().unwrap(); + state + .request_codec(&request_id, &invocation_id) + .expect("request capability resolves during callback"); + state + .response_codec(&response_id, &invocation_id) + .expect("response capability resolves during callback"); + *seen.lock().unwrap() = Some((request_id, Some(response_id), invocation_id)); InvokeResponse { result: Some(InvokeResult::Json(JsonResult { value: Some(json_envelope(JSON_SCHEMA, &json!({"ok": true})).unwrap()), @@ -1279,24 +1245,12 @@ async fn llm_worker_execution_codec_context_is_opt_in_and_ephemeral() { *host_state.lock().unwrap() = Some(callback.host_state.clone()); let next: LlmExecutionNextFn = Arc::new(|_| Box::pin(async { Ok(json!({"unused": true})) })); - callback - .invoke_llm_execution( - "legacy", - "model", - valid_llm_request(), - None, - Arc::clone(&next), - ) - .await - .unwrap(); - assert!(seen.lock().unwrap().is_none()); - callback .invoke_llm_execution( "context", "model", valid_llm_request(), - Some(openai_execution_codec_context()), + openai_execution_codec_context(), next, ) .await @@ -1342,7 +1296,7 @@ async fn cancelling_worker_execution_expires_context_and_continuation_state() { "cancel-context-execution", "model", valid_llm_request(), - Some(openai_execution_codec_context()), + openai_execution_codec_context(), Arc::new(|request| Box::pin(async move { Ok(request.content) })), ) .await @@ -1424,7 +1378,7 @@ async fn llm_worker_stream_codec_context_is_request_only_and_expires_at_eof() { "context-stream", "model", valid_llm_request(), - Some(openai_execution_codec_context()), + openai_stream_execution_codec_context(), Arc::new(|_request| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), ) .await @@ -1528,7 +1482,7 @@ async fn callback_stream_transport_error_surfaces_to_host_stream() { "stream_transport_error", "model", valid_llm_request(), - None, + openai_stream_execution_codec_context(), Arc::new(|_request| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), ) .await @@ -1582,7 +1536,7 @@ async fn callback_stream_stops_when_host_receiver_is_dropped() { "stream_receiver_drop", "model", valid_llm_request(), - None, + openai_stream_execution_codec_context(), Arc::new(|_request| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), ) .await @@ -1776,7 +1730,7 @@ async fn continuation_bearing_worker_callback_allows_slow_next() { "slow-next", "model", valid_llm_request(), - None, + openai_execution_codec_context(), Arc::new(|_| Box::pin(async { Ok(json!({"slow": "completed"})) })), ) .await @@ -1843,7 +1797,7 @@ async fn continuation_bearing_worker_stream_allows_slow_next() { "slow-stream-next", "model", valid_llm_request(), - None, + openai_stream_execution_codec_context(), Arc::new(|_| { Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( @@ -3021,16 +2975,16 @@ async fn host_runtime_service_reports_poisoned_internal_locks() { "failed sanitizer codec setup must remove its invocation scope stack" ); - let context = LlmExecutionCodecContext::new( + let context = LlmExecutionContext::new( LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), - LlmSanitizeResponseContext::for_response_codec(Some(codec)), + Some(LlmSanitizeResponseContext::for_response_codec(Some(codec))), ); let error = callback .invoke_llm_execution( "poisoned-codec-context", "model", valid_llm_request(), - Some(context), + context, Arc::new(|request| Box::pin(async move { Ok(request.content) })), ) .await @@ -3436,15 +3390,23 @@ fn shared_worker_host_state() -> SharedWorkerHostState { Arc::new(Mutex::new(None)) } -fn openai_execution_codec_context() -> LlmExecutionCodecContext { +fn openai_execution_codec_context() -> LlmExecutionContext { let codec = Arc::new(OpenAIChatCodec); - LlmExecutionCodecContext::new( + LlmExecutionContext::new( LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), - LlmSanitizeResponseContext::for_response_codec(Some(codec)), + Some(LlmSanitizeResponseContext::for_response_codec(Some(codec))), + ) +} + +fn openai_stream_execution_codec_context() -> LlmExecutionContext { + LlmExecutionContext::new( + LlmSanitizeRequestContext::for_request_codec(Some(Arc::new(OpenAIChatCodec))), + None, ) } fn execution_codec_capabilities(request: InvokeRequest) -> ExecutionCodecCapabilities { + let surface = RegistrationSurface::try_from(request.surface).expect("registration surface"); let invocation_id = request.invocation_id; let Some(invoke_request_payload::Payload::Llm(invocation)) = request.payload else { panic!("LLM execution must receive an LLM invocation"); @@ -3452,19 +3414,33 @@ fn execution_codec_capabilities(request: InvokeRequest) -> ExecutionCodecCapabil assert!(invocation.sanitize_context.is_none()); let context = invocation .execution_codec_context - .expect("opted-in execution must receive codec context"); + .expect("execution must receive codec context"); let request_context = context.request.expect("request codec context"); - let response_context = context.response.expect("response codec context"); - for identity in [request_context.codec, response_context.codec] { - let identity = identity.expect("codec identity"); - assert_eq!(identity.kind, LlmCodecKind::Builtin as i32); - assert_eq!(identity.id.as_deref(), Some("openai_chat")); - } + let request_identity = request_context.codec.expect("request codec identity"); + assert_eq!(request_identity.kind, LlmCodecKind::Builtin as i32); + assert_eq!(request_identity.id.as_deref(), Some("openai_chat")); + let response_capability_id = match surface { + RegistrationSurface::LlmExecutionIntercept => { + let response_context = context.response.expect("unary response codec context"); + let response_identity = response_context.codec.expect("response codec identity"); + assert_eq!(response_identity.kind, LlmCodecKind::Builtin as i32); + assert_eq!(response_identity.id.as_deref(), Some("openai_chat")); + response_context.codec_capability_id + } + RegistrationSurface::LlmStreamExecutionIntercept => { + assert!( + context.response.is_none(), + "stream execution must not receive response codec context" + ); + None + } + other => panic!("unexpected execution surface: {other:?}"), + }; ( request_context .codec_capability_id .expect("request capability must be present"), - response_context.codec_capability_id, + response_capability_id, invocation_id, ) } @@ -3709,7 +3685,6 @@ fn registration(surface: RegistrationSurface, local_name: &str) -> Registration surface: surface as i32, priority: 0, break_chain: false, - llm_execution_codec_context: false, } } @@ -3798,7 +3773,7 @@ async fn pending_worker_stream_with_codec_context(name: &str) -> WorkerStreamLif name, "model", valid_llm_request(), - Some(openai_execution_codec_context()), + openai_stream_execution_codec_context(), Arc::new(|_| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), ) .await diff --git a/crates/core/tests/unit/llm_api_tests.rs b/crates/core/tests/unit/llm_api_tests.rs index 37d405301..a68ceb179 100644 --- a/crates/core/tests/unit/llm_api_tests.rs +++ b/crates/core/tests/unit/llm_api_tests.rs @@ -5,7 +5,6 @@ #![allow(clippy::await_holding_lock)] -#[cfg(feature = "worker-grpc")] use std::collections::BTreeSet; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Barrier, Mutex}; @@ -22,20 +21,20 @@ use super::{ }; use crate::api::event::{Event, ScopeCategory}; use crate::api::optimization::finalize_optimization_summary; -#[cfg(feature = "worker-grpc")] use crate::api::registry::{ RuntimeRegistrationKind, deregister_conditional_middleware_guardrail, - deregister_llm_execution_intercept, register_conditional_middleware_guardrail, - register_contextual_llm_execution_intercept, register_llm_execution_intercept, - scope_register_llm_execution_intercept, + deregister_llm_execution_intercept, deregister_llm_stream_execution_intercept, + register_conditional_middleware_guardrail, register_llm_execution_intercept, + register_llm_stream_execution_intercept, scope_register_llm_execution_intercept, }; use crate::api::registry::{ deregister_llm_sanitize_request_guardrail, deregister_llm_sanitize_response_guardrail, register_llm_sanitize_request_guardrail, register_llm_sanitize_response_guardrail, }; -use crate::api::runtime::{BuiltinLlmCodec, LlmCodecIdentity, LlmJsonStream}; -#[cfg(feature = "worker-grpc")] -use crate::api::runtime::{LlmExecutionCodecContext, LlmExecutionNextFn}; +use crate::api::runtime::{ + BuiltinLlmCodec, LlmCodecIdentity, LlmExecutionContext, LlmExecutionNextFn, LlmJsonStream, + LlmStreamExecutionNextFn, +}; use crate::api::runtime::{ NemoRelayContextState, create_scope_stack, global_context, set_thread_scope_stack, }; @@ -128,21 +127,22 @@ fn multi_turn_annotation() -> Arc { Arc::new(OpenAIChatCodec.decode(&multi_turn_request()).unwrap()) } -#[cfg(feature = "worker-grpc")] -fn assert_openai_execution_context(context: &LlmExecutionCodecContext) { +fn assert_openai_execution_context(context: &LlmExecutionContext) { assert_eq!( - context.request().codec(), + context.request_codec().codec(), &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) ); - assert!(context.request().resolve_codec().is_some()); + assert!(context.request_codec().resolve_codec().is_some()); + let response = context + .response_codec() + .expect("unary execution must expose its response-codec direction"); assert_eq!( - context.response().codec(), + response.codec(), &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) ); - assert!(context.response().resolve_codec().is_some()); + assert!(response.resolve_codec().is_some()); } -#[cfg(feature = "worker-grpc")] async fn execute_openai_call(name: &str, func: LlmExecutionNextFn) -> crate::error::Result { llm_call_execute( LlmCallExecuteParams::builder() @@ -156,6 +156,34 @@ async fn execute_openai_call(name: &str, func: LlmExecutionNextFn) -> crate::err .await } +async fn execute_openai_stream( + name: &str, + func: LlmStreamExecutionNextFn, +) -> crate::error::Result { + llm_stream_call_execute( + LlmStreamCallExecuteParams::builder() + .name(name) + .request(request()) + .func(func) + .codec(Arc::new(OpenAIChatCodec)) + .response_codec(Arc::new(OpenAIChatCodec)) + .collector(Box::new(|_chunk| Ok(()))) + .finalizer(Box::new(|| json!({"ok": true}))) + .build(), + ) + .await +} + +fn assert_execution_codec_inactive(error: FlowError) { + match error { + FlowError::InvalidArgument(message) => assert_eq!( + message, + "LLM execution codec capability is no longer active" + ), + error => panic!("expected an inactive execution codec error, got {error}"), + } +} + struct ProjectionFailingCodec { projection_attempts: Arc, } @@ -285,7 +313,6 @@ fn response_sanitizer_context_preserves_all_codec_identity_states() { } #[test] -#[cfg(feature = "worker-grpc")] fn managed_execution_passes_codec_context_to_downstream_interceptor() { let _guard = lock_global_runtime(); reset_global(); @@ -293,7 +320,7 @@ fn managed_execution_passes_codec_context_to_downstream_interceptor() { let observations = Arc::new(Mutex::new(Vec::new())); let captured = Arc::clone(&observations); - register_contextual_llm_execution_intercept( + register_llm_execution_intercept( "execution-codec-context-outer", 1, Arc::new(move |_name, request, context, next| { @@ -310,7 +337,7 @@ fn managed_execution_passes_codec_context_to_downstream_interceptor() { ) .unwrap(); let captured = Arc::clone(&observations); - register_contextual_llm_execution_intercept( + register_llm_execution_intercept( "execution-codec-context-inner", 2, Arc::new(move |_name, request, context, next| { @@ -340,8 +367,331 @@ fn managed_execution_passes_codec_context_to_downstream_interceptor() { } #[test] -#[cfg(feature = "worker-grpc")] -fn codec_context_keeps_legacy_scope_ordering_and_conditional_gating() { +fn managed_execution_codec_context_decodes_encodes_and_decodes_response() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + register_llm_execution_intercept( + "execution-codec-round-trip", + 1, + Arc::new(move |_name, request, context, next| { + Box::pin(async move { + let request_codec = context + .request_codec() + .resolve_codec() + .expect("managed execution must expose the resolved request codec"); + let mut annotated = request_codec.decode(&request)?; + let Some(Message::User { content, .. }) = annotated.messages.last_mut() else { + panic!("fixture request must end with a user message"); + }; + *content = MessageContent::Text("rewritten by interceptor".into()); + let request = request_codec.encode(&annotated, &request)?; + + let response = next(request).await?; + let response_codec = context + .response_codec() + .expect("unary execution must expose the response direction") + .resolve_codec() + .expect("managed execution must expose the resolved response codec"); + let annotated_response = response_codec.decode_response(&response)?; + assert_eq!( + annotated_response.message, + Some(MessageContent::Text("accepted".into())) + ); + Ok(response) + }) + }), + ) + .unwrap(); + + let original = LlmRequest { + headers: serde_json::Map::new(), + content: json!({ + "model": "demo", + "messages": [{"role": "user", "content": "original"}], + "provider_only": {"preserved": true} + }), + }; + let response = tokio::runtime::Runtime::new().unwrap().block_on(async { + llm_call_execute( + LlmCallExecuteParams::builder() + .name("execution-codec-round-trip") + .request(original) + .func(Arc::new(|request| { + Box::pin(async move { + assert_eq!( + request.content["messages"][0]["content"], + json!("rewritten by interceptor") + ); + assert_eq!(request.content["provider_only"], json!({"preserved": true})); + Ok(json!({ + "id": "chatcmpl-test", + "object": "chat.completion", + "model": "demo", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "accepted"}, + "finish_reason": "stop" + }] + })) + }) + })) + .codec(Arc::new(OpenAIChatCodec)) + .response_codec(Arc::new(OpenAIChatCodec)) + .build(), + ) + .await + .unwrap() + }); + + assert_eq!(response["choices"][0]["message"]["content"], "accepted"); + assert!(deregister_llm_execution_intercept("execution-codec-round-trip").unwrap()); +} + +#[test] +fn unary_execution_codec_facades_expire_after_interceptor_settlement() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + let retained_request = Arc::new(Mutex::new(None::>)); + let retained_response = Arc::new(Mutex::new(None::>)); + let request_capture = Arc::clone(&retained_request); + let response_capture = Arc::clone(&retained_response); + register_llm_execution_intercept( + "execution-codec-expiry", + 1, + Arc::new(move |_name, request, context, next| { + let request_capture = Arc::clone(&request_capture); + let response_capture = Arc::clone(&response_capture); + Box::pin(async move { + let request_codec = context + .request_codec() + .resolve_codec() + .expect("request codec must be available during execution"); + request_codec.decode(&request)?; + *request_capture.lock().unwrap() = Some(request_codec); + + let response = next(request).await?; + let response_codec = context + .response_codec() + .and_then(|context| context.resolve_codec()) + .expect("response codec must be available during unary execution"); + response_codec.decode_response(&response)?; + *response_capture.lock().unwrap() = Some(response_codec); + Ok(response) + }) + }), + ) + .unwrap(); + + let response = json!({ + "id": "chatcmpl-expiry", + "object": "chat.completion", + "model": "demo", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "accepted"}, + "finish_reason": "stop" + }] + }); + let provider_response = response.clone(); + tokio::runtime::Runtime::new().unwrap().block_on(async { + let actual = execute_openai_call( + "execution-codec-expiry", + Arc::new(move |_| { + let response = provider_response.clone(); + Box::pin(async move { Ok(response) }) + }), + ) + .await + .unwrap(); + assert_eq!(actual, response); + }); + + let request_codec = retained_request.lock().unwrap().clone().unwrap(); + assert_execution_codec_inactive(request_codec.decode(&request()).unwrap_err()); + let annotated = OpenAIChatCodec.decode(&request()).unwrap(); + assert_execution_codec_inactive(request_codec.encode(&annotated, &request()).unwrap_err()); + let response_codec = retained_response.lock().unwrap().clone().unwrap(); + assert_execution_codec_inactive(response_codec.decode_response(&response).unwrap_err()); + + assert!(deregister_llm_execution_intercept("execution-codec-expiry").unwrap()); +} + +#[test] +fn nested_execution_interceptors_have_independent_codec_leases() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + let retained_outer = Arc::new(Mutex::new(None::>)); + let retained_inner = Arc::new(Mutex::new(None::>)); + + let outer_capture = Arc::clone(&retained_outer); + let inner_observation = Arc::clone(&retained_inner); + register_llm_execution_intercept( + "execution-codec-independent-outer", + 1, + Arc::new(move |_name, request, context, next| { + let outer_capture = Arc::clone(&outer_capture); + let inner_observation = Arc::clone(&inner_observation); + Box::pin(async move { + let outer_codec = context.request_codec().resolve_codec().unwrap(); + *outer_capture.lock().unwrap() = Some(Arc::clone(&outer_codec)); + let original = request.clone(); + let response = next(request).await?; + + let inner_codec = inner_observation.lock().unwrap().clone().unwrap(); + assert_execution_codec_inactive(inner_codec.decode(&original).unwrap_err()); + outer_codec.decode(&original)?; + Ok(response) + }) + }), + ) + .unwrap(); + + let inner_capture = Arc::clone(&retained_inner); + register_llm_execution_intercept( + "execution-codec-independent-inner", + 2, + Arc::new(move |_name, request, context, next| { + let inner_capture = Arc::clone(&inner_capture); + Box::pin(async move { + *inner_capture.lock().unwrap() = context.request_codec().resolve_codec(); + next(request).await + }) + }), + ) + .unwrap(); + + tokio::runtime::Runtime::new().unwrap().block_on(async { + execute_openai_call( + "execution-codec-independent-leases", + Arc::new(|_| Box::pin(async { Ok(json!({"ok": true})) })), + ) + .await + .unwrap(); + }); + + let outer_codec = retained_outer.lock().unwrap().clone().unwrap(); + assert_execution_codec_inactive(outer_codec.decode(&request()).unwrap_err()); + assert!(deregister_llm_execution_intercept("execution-codec-independent-outer").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-independent-inner").unwrap()); +} + +#[test] +fn streaming_execution_codec_facade_lives_with_returned_stream() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + let retained = Arc::new(Mutex::new(Vec::>::new())); + let captured = Arc::clone(&retained); + register_llm_stream_execution_intercept( + "execution-codec-stream-lifetime", + 1, + Arc::new(move |_name, request, context, next| { + let captured = Arc::clone(&captured); + Box::pin(async move { + let codec = context.request_codec().resolve_codec().unwrap(); + codec.decode(&request)?; + captured.lock().unwrap().push(codec); + next(request).await + }) + }), + ) + .unwrap(); + + tokio::runtime::Runtime::new().unwrap().block_on(async { + let mut completed = execute_openai_stream( + "execution-codec-stream-lifetime", + Arc::new(|_| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok(json!({ + "chunk": true + }))]))) + }) + }), + ) + .await + .unwrap(); + let completed_codec = retained.lock().unwrap()[0].clone(); + completed_codec.decode(&request()).unwrap(); + assert!(completed.next().await.unwrap().is_ok()); + completed_codec.decode(&request()).unwrap(); + assert!(completed.next().await.is_none()); + assert_execution_codec_inactive(completed_codec.decode(&request()).unwrap_err()); + + let mut closed = execute_openai_stream( + "execution-codec-stream-lifetime", + Arc::new(|_| { + Box::pin(async { + Ok(LlmJsonStream::new(futures_util::stream::pending::< + crate::error::Result, + >())) + }) + }), + ) + .await + .unwrap(); + let closed_codec = retained.lock().unwrap()[1].clone(); + closed_codec.decode(&request()).unwrap(); + closed.close().await.unwrap(); + assert_execution_codec_inactive(closed_codec.decode(&request()).unwrap_err()); + }); + + assert!(deregister_llm_stream_execution_intercept("execution-codec-stream-lifetime").unwrap()); +} + +#[test] +fn dropping_unconsumed_stream_expires_execution_codec_facade() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + let retained = Arc::new(Mutex::new(None::>)); + let captured = Arc::clone(&retained); + register_llm_stream_execution_intercept( + "execution-codec-stream-drop", + 1, + Arc::new(move |_name, request, context, next| { + let captured = Arc::clone(&captured); + Box::pin(async move { + *captured.lock().unwrap() = context.request_codec().resolve_codec(); + next(request).await + }) + }), + ) + .unwrap(); + + tokio::runtime::Runtime::new().unwrap().block_on(async { + let stream = execute_openai_stream( + "execution-codec-stream-drop", + Arc::new(|_| { + Box::pin(async { + Ok(LlmJsonStream::new(futures_util::stream::pending::< + crate::error::Result, + >())) + }) + }), + ) + .await + .unwrap(); + let codec = retained.lock().unwrap().clone().unwrap(); + codec.decode(&request()).unwrap(); + + drop(stream); + + assert_execution_codec_inactive(codec.decode(&request()).unwrap_err()); + }); + + assert!(deregister_llm_stream_execution_intercept("execution-codec-stream-drop").unwrap()); +} + +#[test] +fn codec_context_preserves_scope_ordering_and_conditional_gating() { let _guard = lock_global_runtime(); reset_global(); set_thread_scope_stack(create_scope_stack()); @@ -357,22 +707,23 @@ fn codec_context_keeps_legacy_scope_ordering_and_conditional_gating() { let captured = Arc::clone(&calls); register_llm_execution_intercept( - "execution-codec-legacy-global", + "execution-codec-global-first", 10, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, context, next| { let captured = Arc::clone(&captured); Box::pin(async move { - captured.lock().unwrap().push("legacy-global-enter"); + assert_openai_execution_context(&context); + captured.lock().unwrap().push("global-first-enter"); let result = next(request).await; - captured.lock().unwrap().push("legacy-global-exit"); + captured.lock().unwrap().push("global-first-exit"); result }) }), ) .unwrap(); - register_contextual_llm_execution_intercept( - "execution-codec-gated-contextual", + register_llm_execution_intercept( + "execution-codec-gated", 15, Arc::new(move |_name, _request, _context, _next| { Box::pin(async move { panic!("conditionally disabled interceptor must not execute") }) @@ -383,22 +734,22 @@ fn codec_context_keeps_legacy_scope_ordering_and_conditional_gating() { register_conditional_middleware_guardrail( "execution-codec-context-gate", gate_kinds, - "execution-codec-gated-contextual", + "execution-codec-gated", Arc::new(|_, _| Some("disabled for regression test".into())), ) .unwrap(); let captured = Arc::clone(&calls); - register_contextual_llm_execution_intercept( - "execution-codec-contextual-global", + register_llm_execution_intercept( + "execution-codec-global-second", 20, Arc::new(move |_name, request, context, next| { let captured = Arc::clone(&captured); Box::pin(async move { assert_openai_execution_context(&context); - captured.lock().unwrap().push("contextual-global-enter"); + captured.lock().unwrap().push("global-second-enter"); let result = next(request).await; - captured.lock().unwrap().push("contextual-global-exit"); + captured.lock().unwrap().push("global-second-exit"); result }) }), @@ -408,14 +759,15 @@ fn codec_context_keeps_legacy_scope_ordering_and_conditional_gating() { let captured = Arc::clone(&calls); scope_register_llm_execution_intercept( &scope.uuid, - "execution-codec-legacy-scope", + "execution-codec-scope", 30, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, context, next| { let captured = Arc::clone(&captured); Box::pin(async move { - captured.lock().unwrap().push("legacy-scope-enter"); + assert_openai_execution_context(&context); + captured.lock().unwrap().push("scope-enter"); let result = next(request).await; - captured.lock().unwrap().push("legacy-scope-exit"); + captured.lock().unwrap().push("scope-exit"); result }) }), @@ -440,27 +792,116 @@ fn codec_context_keeps_legacy_scope_ordering_and_conditional_gating() { assert_eq!(response, json!({"ok": true})); assert!(deregister_conditional_middleware_guardrail("execution-codec-context-gate").unwrap()); - assert!(deregister_llm_execution_intercept("execution-codec-legacy-global").unwrap()); - assert!(deregister_llm_execution_intercept("execution-codec-gated-contextual").unwrap()); - assert!(deregister_llm_execution_intercept("execution-codec-contextual-global").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-global-first").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-gated").unwrap()); + assert!(deregister_llm_execution_intercept("execution-codec-global-second").unwrap()); pop_scope(PopScopeParams::builder().handle_uuid(&scope.uuid).build()).unwrap(); assert_eq!( calls.lock().unwrap().as_slice(), [ - "legacy-global-enter", - "contextual-global-enter", - "legacy-scope-enter", + "global-first-enter", + "global-second-enter", + "scope-enter", "provider", - "legacy-scope-exit", - "contextual-global-exit", - "legacy-global-exit", + "scope-exit", + "global-second-exit", + "global-first-exit", ] ); } #[test] -#[cfg(feature = "worker-grpc")] +fn managed_execution_distinguishes_absent_codecs_from_streaming_response_unavailability() { + let _guard = lock_global_runtime(); + reset_global(); + set_thread_scope_stack(create_scope_stack()); + + register_llm_execution_intercept( + "execution-codec-absent", + 1, + Arc::new(move |_name, request, context, next| { + Box::pin(async move { + assert_eq!(context.request_codec().codec(), &LlmCodecIdentity::None); + assert!(context.request_codec().resolve_codec().is_none()); + let response = context + .response_codec() + .expect("unary execution exposes an absent response codec explicitly"); + assert_eq!(response.codec(), &LlmCodecIdentity::None); + assert!(response.resolve_codec().is_none()); + next(request).await + }) + }), + ) + .unwrap(); + + tokio::runtime::Runtime::new().unwrap().block_on(async { + llm_call_execute( + LlmCallExecuteParams::builder() + .name("execution-codec-absent") + .request(request()) + .func(Arc::new(|_| Box::pin(async { Ok(json!({"ok": true})) }))) + .build(), + ) + .await + .unwrap(); + }); + + assert!(deregister_llm_execution_intercept("execution-codec-absent").unwrap()); + + register_llm_stream_execution_intercept( + "execution-codec-stream-request-only", + 1, + Arc::new(move |_name, request, context, next| { + Box::pin(async move { + assert_eq!( + context.request_codec().codec(), + &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert!(context.request_codec().resolve_codec().is_some()); + assert!( + context.response_codec().is_none(), + "streaming execution must not expose a completed-response codec" + ); + next(request).await + }) + }), + ) + .unwrap(); + + tokio::runtime::Runtime::new().unwrap().block_on(async { + let mut stream = llm_stream_call_execute( + LlmStreamCallExecuteParams::builder() + .name("execution-codec-stream-request-only") + .request(request()) + .func(Arc::new(|_| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok(json!({ + "chunk": true + }))]))) + }) + })) + .codec(Arc::new(OpenAIChatCodec)) + .response_codec(Arc::new(OpenAIChatCodec)) + .collector(Box::new(|_chunk| Ok(()))) + .finalizer(Box::new(|| json!({"ok": true}))) + .build(), + ) + .await + .unwrap(); + assert_eq!( + stream.next().await.unwrap().unwrap(), + json!({"chunk": true}) + ); + assert!(stream.next().await.is_none()); + }); + + assert!( + deregister_llm_stream_execution_intercept("execution-codec-stream-request-only").unwrap() + ); +} + +#[test] fn execution_codec_context_does_not_follow_wire_format_mutation() { let _guard = lock_global_runtime(); reset_global(); @@ -469,7 +910,7 @@ fn execution_codec_context_does_not_follow_wire_format_mutation() { register_llm_execution_intercept( "execution-codec-change-wire-format", 10, - Arc::new(move |_name, mut request, next| { + Arc::new(move |_name, mut request, _context, next| { Box::pin(async move { request.content = json!({ "contents": [{ @@ -483,14 +924,14 @@ fn execution_codec_context_does_not_follow_wire_format_mutation() { ) .unwrap(); - register_contextual_llm_execution_intercept( + register_llm_execution_intercept( "execution-codec-reject-stale-payload", 20, Arc::new(move |_name, request, context, _next| { Box::pin(async move { assert_openai_execution_context(&context); let codec = context - .request() + .request_codec() .resolve_codec() .expect("managed call must expose its selected request codec"); codec.decode(&request).map(|_| json!({"unexpected": true})) @@ -549,6 +990,55 @@ impl LlmResponseCodec for ProjectionFailingCodec { } } +#[test] +fn execution_context_preserves_runtime_and_opaque_codec_identities() { + let runtime_request: Arc = Arc::new(RuntimeIdentityCodec); + let runtime_response: Arc = Arc::new(RuntimeIdentityCodec); + let runtime_context = + LlmExecutionContext::for_unary_codecs(Some(runtime_request), &Some(runtime_response)); + assert_eq!( + runtime_context.request_codec().codec(), + &LlmCodecIdentity::Runtime("com.example.chat.v1".into()) + ); + assert_eq!( + runtime_context.response_codec().unwrap().codec(), + &LlmCodecIdentity::Runtime("com.example.chat.v1".into()) + ); + assert!(runtime_context.request_codec().resolve_codec().is_some()); + assert!( + runtime_context + .response_codec() + .unwrap() + .resolve_codec() + .is_some() + ); + + let opaque_request: Arc = Arc::new(ProjectionFailingCodec { + projection_attempts: Arc::new(AtomicUsize::new(0)), + }); + let opaque_response: Arc = Arc::new(ProjectionFailingCodec { + projection_attempts: Arc::new(AtomicUsize::new(0)), + }); + let opaque_context = + LlmExecutionContext::for_unary_codecs(Some(opaque_request), &Some(opaque_response)); + assert_eq!( + opaque_context.request_codec().codec(), + &LlmCodecIdentity::Opaque + ); + assert_eq!( + opaque_context.response_codec().unwrap().codec(), + &LlmCodecIdentity::Opaque + ); + assert!(opaque_context.request_codec().resolve_codec().is_some()); + assert!( + opaque_context + .response_codec() + .unwrap() + .resolve_codec() + .is_some() + ); +} + fn emit_compaction() { event( EmitMarkEventParams::builder() diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index ad70be1dd..3aac06e8d 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -136,7 +136,8 @@ fn native_async_release_defers_library_guard_drop_to_host_reaper() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: Some(callback_user_data), }); @@ -633,57 +634,57 @@ fn assert_native_digest_edges() { fn assert_native_host_api_versions() { let current = native_host_api(); + let frozen_v6 = native_host_api_v6(); let frozen_v5 = native_host_api_v5(); let frozen_v4 = native_host_api_v4(); let frozen_v3 = native_host_api_v3(); let legacy = native_host_api_v2(); - assert!(!current.is_null()); - assert!(!frozen_v5.is_null()); - assert!(!frozen_v4.is_null()); - assert!(!frozen_v3.is_null()); - assert!(!legacy.is_null()); - assert_eq!( - unsafe { (*current).abi_version }, - NEMO_RELAY_NATIVE_ABI_VERSION - ); - assert_eq!(unsafe { (*frozen_v3).abi_version }, 3); - assert_eq!( - unsafe { (*frozen_v5).abi_version }, - NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT - ); - assert_eq!( - unsafe { (*frozen_v4).abi_version }, - NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL - ); - assert_eq!( - unsafe { (*legacy).abi_version }, - NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY - ); - assert_eq!( - unsafe { (*current).struct_size }, - std::mem::size_of::() - ); - assert_eq!( - unsafe { (*frozen_v5).struct_size }, - std::mem::size_of::() - ); - assert_eq!( - unsafe { (*frozen_v4).struct_size }, - std::mem::size_of::() - ); - assert_eq!( - unsafe { (*frozen_v3).struct_size }, - std::mem::size_of::() - ); - assert_eq!( - unsafe { (*legacy).struct_size }, - std::mem::size_of::() - ); + assert_native_host_api_descriptor( + current, + NEMO_RELAY_NATIVE_ABI_VERSION, + std::mem::size_of::(), + ); + assert_native_host_api_descriptor( + frozen_v6, + NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, + std::mem::size_of::(), + ); + assert_native_host_api_descriptor( + frozen_v5, + NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, + std::mem::size_of::(), + ); + assert_native_host_api_descriptor( + frozen_v4, + NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, + std::mem::size_of::(), + ); + assert_native_host_api_descriptor( + frozen_v3, + 3, + std::mem::size_of::(), + ); + assert_native_host_api_descriptor( + legacy, + NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, + std::mem::size_of::(), + ); + assert_native_host_api_v7_layout(); assert_native_host_api_v6_layout(); assert_native_host_api_v5_layout(); assert_native_host_api_v4_layout(); } +fn assert_native_host_api_descriptor( + host: *const NemoRelayNativeHostApiV1, + expected_version: u32, + expected_size: usize, +) { + assert!(!host.is_null()); + assert_eq!(unsafe { (*host).abi_version }, expected_version); + assert_eq!(unsafe { (*host).struct_size }, expected_size); +} + fn assert_native_host_api_v4_layout() { #[cfg(target_pointer_width = "64")] { @@ -803,6 +804,71 @@ fn assert_native_host_api_v6_layout() { } } +fn assert_native_host_api_v7_layout() { + #[cfg(target_pointer_width = "64")] + { + assert_eq!(std::mem::align_of::(), 8); + assert_eq!(std::mem::size_of::(), 648); + assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV7, v6), 0); + assert_eq!( + std::mem::offset_of!( + NemoRelayNativeHostApiV7, + plugin_context_register_async_llm_execution_intercept + ), + 616 + ); + assert_eq!( + std::mem::offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), + 624 + ); + assert_eq!( + std::mem::offset_of!( + NemoRelayNativeHostApiV7, + async_stream_llm_request_codec_decode + ), + 632 + ); + assert_eq!( + std::mem::offset_of!( + NemoRelayNativeHostApiV7, + async_stream_llm_request_codec_encode + ), + 640 + ); + } + #[cfg(target_pointer_width = "32")] + { + assert_eq!(std::mem::align_of::(), 4); + assert_eq!(std::mem::size_of::(), 320); + assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV7, v6), 0); + assert_eq!( + std::mem::offset_of!( + NemoRelayNativeHostApiV7, + plugin_context_register_async_llm_execution_intercept + ), + 304 + ); + assert_eq!( + std::mem::offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), + 308 + ); + assert_eq!( + std::mem::offset_of!( + NemoRelayNativeHostApiV7, + async_stream_llm_request_codec_decode + ), + 312 + ); + assert_eq!( + std::mem::offset_of!( + NemoRelayNativeHostApiV7, + async_stream_llm_request_codec_encode + ), + 316 + ); + } +} + #[tokio::test] async fn native_async_wait_and_rejection_cover_dropped_and_aborted_continuations() { let (sender, receiver) = tokio::sync::oneshot::channel::>(); @@ -814,7 +880,8 @@ async fn native_async_wait_and_rejection_cover_dropped_and_aborted_continuations next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }), @@ -836,7 +903,8 @@ async fn native_async_wait_and_rejection_cover_dropped_and_aborted_continuations next_invoked: AtomicBool::new(true), next_abort: Mutex::new(Some(abort)), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -871,6 +939,7 @@ fn native_stream_callback_guard_covers_terminal_drop_modes() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -930,6 +999,7 @@ async fn native_async_stream_forwarding_reports_conversion_and_stream_errors() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }) @@ -1144,7 +1214,8 @@ fn accepted_native_callbacks_settle_when_cancelled_before_first_poll() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1190,6 +1261,7 @@ fn accepted_native_callbacks_settle_when_cancelled_before_first_poll() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1281,6 +1353,7 @@ fn native_async_stream_entrypoints_cover_closed_full_and_settled_channels() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1308,6 +1381,7 @@ fn native_async_stream_entrypoints_cover_closed_full_and_settled_channels() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1333,6 +1407,7 @@ fn native_async_stream_entrypoints_cover_closed_full_and_settled_channels() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1356,6 +1431,7 @@ fn native_async_stream_entrypoints_cover_closed_full_and_settled_channels() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1432,6 +1508,7 @@ async fn native_async_stream_next_entrypoint_validates_handle_kind_and_request() backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1529,6 +1606,7 @@ unsafe extern "C" fn invoke_native_next_then_return_state( unsafe extern "C" fn invoke_native_stream_next_then_return_state( user_data: *mut c_void, invocation_json: *const NemoRelayNativeString, + _context: *const NemoRelayNativeLlmExecutionContext, next: *const NemoRelayNativeAsyncNext, stream: *const NemoRelayNativeAsyncStream, ) -> u32 { @@ -1573,6 +1651,7 @@ unsafe extern "C" fn invoke_native_stream_next_then_return_state( unsafe extern "C" fn invoke_detached_next_and_finish_replacement_stream( user_data: *mut c_void, invocation_json: *const NemoRelayNativeString, + _context: *const NemoRelayNativeLlmExecutionContext, next: *const NemoRelayNativeAsyncNext, stream: *const NemoRelayNativeAsyncStream, ) -> u32 { @@ -1807,7 +1886,7 @@ fn assert_native_json_output_and_host_api() { 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::() ); } @@ -1846,7 +1925,8 @@ fn native_async_next_abi_runs_tool_llm_and_stream_continuations() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1885,7 +1965,8 @@ fn native_async_next_abi_runs_tool_llm_and_stream_continuations() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -1946,7 +2027,8 @@ fn native_async_next_reports_a_revoked_continuation_without_calling_the_provider next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2206,7 +2288,8 @@ fn owned_native_result_continuation_is_aborted_when_completion_is_cancelled() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2406,7 +2489,8 @@ fn native_async_next_preserves_runtime_context_for_unary_and_stream_continuation next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2465,6 +2549,7 @@ fn native_async_next_preserves_runtime_context_for_unary_and_stream_continuation backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2652,7 +2737,8 @@ fn native_async_next_panics_settle_unary_and_stream_errors() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2697,6 +2783,7 @@ fn native_async_next_panics_settle_unary_and_stream_errors() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2773,7 +2860,8 @@ fn native_async_next_is_permanently_one_shot() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2826,7 +2914,8 @@ fn cancelled_native_async_next_does_not_start_unary_or_stream_continuations() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2859,6 +2948,7 @@ fn cancelled_native_async_next_does_not_start_unary_or_stream_continuations() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2922,7 +3012,8 @@ fn malformed_llm_next_does_not_consume_the_completion() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -2972,6 +3063,7 @@ fn native_async_stream_next_supports_repeated_concurrent_calls() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3089,6 +3181,7 @@ fn native_async_stream_settlement_rejects_late_next_and_aborts_in_flight_next() backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3210,6 +3303,7 @@ fn native_async_stream_next_stops_callbacks_after_false() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3284,6 +3378,7 @@ fn native_async_stream_in_flight_cancellation_releases_callback_state() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3366,6 +3461,7 @@ fn native_async_stream_cancellation_before_first_poll_releases_callback_state() backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3442,7 +3538,8 @@ fn native_async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlemen next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3479,7 +3576,8 @@ fn native_async_completion_abi_rejects_invalid_duplicate_and_cancelled_settlemen next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3507,7 +3605,8 @@ fn completed_native_async_wait_is_not_marked_cancelled() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3563,7 +3662,8 @@ fn native_async_completion_cancellation_wins_resolve_and_reject_settlement_races next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: Some(Arc::clone(&settlement_checkpoint)), _callback_user_data: None, }); @@ -3646,7 +3746,8 @@ fn cancelling_completion_aborts_pending_native_next() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: None, + request_codec: None, + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -3745,7 +3846,7 @@ fn native_async_callback_contract_errors_abort_an_invoked_next() { }) } }))), - None, + NativeAsyncCodecCapabilities::default(), )) .unwrap_err(); @@ -3820,6 +3921,7 @@ fn native_async_stream_contract_errors_abort_an_invoked_next() { headers: Map::new(), content: Json::Null, }, + LlmExecutionContext::default(), next, )); let error = match result { @@ -3879,6 +3981,7 @@ fn native_replacement_stream_does_not_wait_for_a_detached_pending_next() { headers: Map::new(), content: Json::Null, }, + LlmExecutionContext::default(), next, )) .expect("replacement stream must not wait for detached downstream construction"); @@ -3924,6 +4027,7 @@ fn native_async_stream_settlement_cannot_succeed_after_cancellation() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: Some(Arc::clone(&settlement_checkpoint)), _callback_user_data: None, }); @@ -3993,6 +4097,7 @@ fn native_async_stream_push_is_bounded_retryable_and_incremental() { backpressured: AtomicBool::new(false), downstream_aborts: Mutex::new(HashMap::new()), settlement: Mutex::new(()), + request_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -4722,6 +4827,7 @@ unsafe extern "C" fn noop_llm_execution( _user_data: *mut c_void, _name: *const NemoRelayNativeString, _request_json: *const NemoRelayNativeString, + _context: NemoRelayNativeLlmExecutionContext, _next_fn: NemoRelayNativeLlmNextFn, _next_ctx: *mut c_void, _out_json: *mut *mut NemoRelayNativeString, @@ -4733,6 +4839,7 @@ unsafe extern "C" fn noop_llm_stream_execution( _user_data: *mut c_void, _name: *const NemoRelayNativeString, _request_json: *const NemoRelayNativeString, + _context: NemoRelayNativeLlmExecutionContext, _next_fn: NemoRelayNativeLlmStreamNextFn, _next_ctx: *mut c_void, _out_stream: *mut NemoRelayNativeLlmStreamV1, @@ -5649,7 +5756,7 @@ fn native_codec_operations_report_json_and_codec_failures() { } #[test] -fn native_v4_completion_scoped_codecs_enforce_direction_and_expiration() { +fn native_completion_scoped_codecs_enforce_direction_and_expiration() { let (sender, _receiver) = tokio::sync::oneshot::channel(); let completion = Arc::new(NativeAsyncCompletion { sender: Mutex::new(Some(sender)), @@ -5657,9 +5764,10 @@ fn native_v4_completion_scoped_codecs_enforce_direction_and_expiration() { next_invoked: AtomicBool::new(false), next_abort: Mutex::new(None), continuation_aborts: Mutex::new(HashMap::new()), - codec: Some(NativeAsyncCodecCapability::Request( - Arc::new(OpenAIChatCodec) as Arc, + request_codec: Some(NativeHostLlmRequestCodec( + Arc::new(OpenAIChatCodec) as Arc )), + response_codec: None, before_settlement_lock: None, _callback_user_data: None, }); @@ -5743,6 +5851,81 @@ fn native_v4_completion_scoped_codecs_enforce_direction_and_expiration() { } } +#[test] +fn native_stream_scoped_codec_expires_after_stream_settlement() { + let (sender, _receiver) = tokio::sync::mpsc::channel(1); + let stream = Arc::new(NativeAsyncStream { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + settled: AtomicBool::new(false), + backpressured: AtomicBool::new(false), + downstream_aborts: Mutex::new(HashMap::new()), + settlement: Mutex::new(()), + request_codec: Some(NativeHostLlmRequestCodec( + Arc::new(OpenAIChatCodec) as Arc + )), + before_settlement_lock: None, + _callback_user_data: None, + }); + let stream_ref = Arc::into_raw(Arc::clone(&stream)) as *const NemoRelayNativeAsyncStream; + assert_eq!( + unsafe { native_async_stream_retain(stream_ref) }, + NemoRelayStatus::Ok + ); + let request = LlmRequest { + headers: Map::new(), + content: json!({ + "model": "gpt-test", + "messages": [{"role": "user", "content": "secret"}] + }), + }; + let request_json = native_string(&serde_json::to_string(&request).unwrap()); + let mut output = ptr::null_mut(); + assert_eq!( + unsafe { + native_async_stream_llm_request_codec_decode(stream_ref, request_json, &mut output) + }, + NemoRelayStatus::Ok + ); + let annotated: AnnotatedLlmRequest = + serde_json::from_str(&read_native_string(output).unwrap()).unwrap(); + unsafe { native_string_free(output) }; + let annotated_json = native_string(&serde_json::to_string(&annotated).unwrap()); + assert_eq!( + unsafe { + native_async_stream_llm_request_codec_encode( + stream_ref, + annotated_json, + request_json, + &mut output, + ) + }, + NemoRelayStatus::Ok + ); + unsafe { native_string_free(output) }; + + assert_eq!( + unsafe { native_async_stream_finish(stream_ref) }, + NemoRelayStatus::Ok + ); + let sentinel = native_string("expired"); + output = sentinel; + assert_eq!( + unsafe { + native_async_stream_llm_request_codec_decode(stream_ref, request_json, &mut output) + }, + NemoRelayStatus::InvalidArg + ); + assert!(output.is_null()); + unsafe { + native_string_free(sentinel); + native_string_free(request_json); + native_string_free(annotated_json); + native_async_stream_release(stream_ref); + native_async_stream_release(stream_ref); + } +} + #[test] fn native_codec_operations_contain_codec_panics() { let request_codec = @@ -6070,6 +6253,7 @@ unsafe extern "C" fn llm_execution_error( _user_data: *mut c_void, _name: *const NemoRelayNativeString, _request_json: *const NemoRelayNativeString, + _context: NemoRelayNativeLlmExecutionContext, _next_fn: NemoRelayNativeLlmNextFn, _next_ctx: *mut c_void, out_json: *mut *mut NemoRelayNativeString, @@ -6084,6 +6268,7 @@ unsafe extern "C" fn llm_stream_execution_error( _user_data: *mut c_void, _name: *const NemoRelayNativeString, _request_json: *const NemoRelayNativeString, + _context: NemoRelayNativeLlmExecutionContext, _next_fn: NemoRelayNativeLlmStreamNextFn, _next_ctx: *mut c_void, out_stream: *mut NemoRelayNativeLlmStreamV1, @@ -6363,11 +6548,16 @@ async fn native_callback_wrappers_release_error_outputs_and_preserve_reasons() { None, ); assert!( - llm_execution("model", request.clone(), llm_next(Ok(Json::Null))) - .await - .unwrap_err() - .to_string() - .contains("LLM execution failed") + llm_execution( + "model", + request.clone(), + LlmExecutionContext::default(), + llm_next(Ok(Json::Null)), + ) + .await + .unwrap_err() + .to_string() + .contains("LLM execution failed") ); let stream_next: LlmStreamExecutionNextFn = @@ -6375,12 +6565,17 @@ async fn native_callback_wrappers_release_error_outputs_and_preserve_reasons() { let llm_stream_execution = wrap_llm_stream_execution_fn(instance, llm_stream_execution_error, ptr::null_mut(), None); assert!( - llm_stream_execution("model", request, stream_next) - .await - .err() - .expect("native stream callback should fail") - .to_string() - .contains("LLM stream execution failed") + llm_stream_execution( + "model", + request, + LlmExecutionContext::default(), + stream_next, + ) + .await + .err() + .expect("native stream callback should fail") + .to_string() + .contains("LLM stream execution failed") ); } @@ -6546,6 +6741,76 @@ fn native_llm_sanitize_context_preserves_all_codec_identity_states() { } } +#[test] +fn native_execution_context_is_directional_and_streaming_omits_response_codec() { + let request_codec: Arc = Arc::new(OpenAIChatCodec); + let response_codec: Arc = Arc::new(OpenAIChatCodec); + let unary = LlmExecutionContext::new( + LlmSanitizeRequestContext::for_request_codec(Some(request_codec.clone())), + Some(LlmSanitizeResponseContext::for_response_codec(Some( + response_codec.clone(), + ))), + ); + let native_request_codec = NativeHostLlmRequestCodec(request_codec.clone()); + let native_response_codec = NativeHostLlmResponseCodec(response_codec); + let bridge = NativeLlmExecutionContextBridge::new( + &unary, + Some(&native_request_codec), + Some(&native_response_codec), + ) + .unwrap(); + bridge.with_native_context(|context| { + assert_eq!( + context.request_codec.codec_kind, + NemoRelayNativeLlmCodecKind::BuiltIn + ); + assert_eq!( + read_native_string(context.request_codec.codec_id).unwrap(), + "openai_chat" + ); + assert!(!context.request_codec.codec.is_null()); + assert!(!context.response_codec.is_null()); + let response = unsafe { &*context.response_codec }; + assert_eq!(response.codec_kind, NemoRelayNativeLlmCodecKind::BuiltIn); + assert_eq!( + read_native_string(response.codec_id).unwrap(), + "openai_chat" + ); + assert!(!response.codec.is_null()); + }); + + let streaming = LlmExecutionContext::new( + LlmSanitizeRequestContext::for_request_codec(Some(request_codec)), + None, + ); + let bridge = + NativeLlmExecutionContextBridge::new(&streaming, Some(&native_request_codec), None) + .unwrap(); + bridge.with_native_context(|context| { + assert!(!context.request_codec.codec.is_null()); + assert!(context.response_codec.is_null()); + }); + + let absent = LlmExecutionContext::new( + LlmSanitizeRequestContext::for_request_codec(None), + Some(LlmSanitizeResponseContext::for_response_codec(None)), + ); + let bridge = NativeLlmExecutionContextBridge::new(&absent, None, None).unwrap(); + bridge.with_native_context(|context| { + assert_eq!( + context.request_codec.codec_kind, + NemoRelayNativeLlmCodecKind::None + ); + assert!(context.request_codec.codec_id.is_null()); + assert!(context.request_codec.codec.is_null()); + assert!(!context.response_codec.is_null()); + let response = unsafe { &*context.response_codec }; + assert_eq!(response.codec_kind, NemoRelayNativeLlmCodecKind::None); + assert!(response.codec_id.is_null()); + assert!(response.codec.is_null()); + }); +} + #[test] fn native_async_llm_sanitize_context_uses_stable_codec_envelope() { assert_eq!( @@ -7082,7 +7347,7 @@ async fn native_stream_adapter_covers_chunks_end_errors_and_cancellation() { NativeStreamItem::Json(json!({"chunk": 1})), NativeStreamItem::End, ]); - let mut stream = native_stream_to_relay_stream(raw, None, None).unwrap(); + let mut stream = native_stream_to_relay_stream(raw, None, None, None).unwrap(); assert_eq!(stream.next().await.unwrap().unwrap(), json!({"chunk": 1})); assert!(stream.next().await.is_none()); assert_eq!(cancel_count.load(Ordering::SeqCst), 0); @@ -7091,7 +7356,7 @@ async fn native_stream_adapter_covers_chunks_end_errors_and_cancellation() { let (raw, cancel_count, drop_count) = test_native_stream([NativeStreamItem::Json(json!({ "chunk": 2 }))]); - let stream = native_stream_to_relay_stream(raw, None, None).unwrap(); + let stream = native_stream_to_relay_stream(raw, None, None, None).unwrap(); drop(stream); assert_eq!(cancel_count.load(Ordering::SeqCst), 1); assert_eq!(drop_count.load(Ordering::SeqCst), 1); @@ -7103,7 +7368,7 @@ async fn native_stream_adapter_covers_chunks_end_errors_and_cancellation() { NativeStreamItem::ErrorWithJson(NemoRelayStatus::InvalidArg), ] { let (raw, _, drop_count) = test_native_stream([item]); - let mut stream = native_stream_to_relay_stream(raw, None, None).unwrap(); + let mut stream = native_stream_to_relay_stream(raw, None, None, None).unwrap(); assert!(stream.next().await.unwrap().is_err()); assert!(stream.next().await.is_none()); assert_eq!(drop_count.load(Ordering::SeqCst), 1); @@ -7111,16 +7376,16 @@ async fn native_stream_adapter_covers_chunks_end_errors_and_cancellation() { let (mut raw, _, drop_count) = test_native_stream([]); raw.struct_size = 0; - assert!(NativeRelayLlmStream::from_raw(raw, None, None).is_err()); + assert!(NativeRelayLlmStream::from_raw(raw, None, None, None).is_err()); assert_eq!(drop_count.load(Ordering::SeqCst), 1); let (mut raw, _, drop_count) = test_native_stream([]); raw.next = None; - assert!(NativeRelayLlmStream::from_raw(raw, None, None).is_err()); + assert!(NativeRelayLlmStream::from_raw(raw, None, None, None).is_err()); assert_eq!(drop_count.load(Ordering::SeqCst), 1); let (raw, _, drop_count) = test_native_stream([NativeStreamItem::EndWithJson]); - let mut stream = native_stream_to_relay_stream(raw, None, None).unwrap(); + let mut stream = native_stream_to_relay_stream(raw, None, None, None).unwrap(); assert!(stream.next().await.is_none()); assert_eq!(drop_count.load(Ordering::SeqCst), 1); @@ -7129,6 +7394,7 @@ async fn native_stream_adapter_covers_chunks_end_errors_and_cancellation() { finished: false, _next_ctx: None, _callback_user_data: None, + _request_codec: None, }; assert!(invalid.next().await.unwrap().is_err()); assert!(invalid.next().await.is_none()); diff --git a/crates/core/tests/unit/plugin_tests.rs b/crates/core/tests/unit/plugin_tests.rs index c3b21f310..40f4bba24 100644 --- a/crates/core/tests/unit/plugin_tests.rs +++ b/crates/core/tests/unit/plugin_tests.rs @@ -1377,13 +1377,13 @@ fn test_plugin_registration_context_covers_all_registration_helpers() { ctx.register_llm_execution_intercept( "llm-exec", 1, - Arc::new(|_name, request, _next| Box::pin(async move { Ok(request.content) })), + Arc::new(|_name, request, _context, _next| Box::pin(async move { Ok(request.content) })), ) .unwrap(); ctx.register_llm_stream_execution_intercept( "llm-stream", 1, - Arc::new(|_name, request, _next| { + Arc::new(|_name, request, _context, _next| { Box::pin(async move { Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( request.content @@ -2344,14 +2344,16 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { ctx.register_llm_execution_intercept( "llm-exec", 1, - Arc::new(|_name, request, _next| Box::pin(async move { Ok(request.content) })), + Arc::new(|_name, request, _context, _next| Box::pin(async move { Ok(request.content) })), ) .unwrap(); expect_registration_failed( ctx.register_llm_execution_intercept( "llm-exec", 1, - Arc::new(|_name, request, _next| Box::pin(async move { Ok(request.content) })), + Arc::new(|_name, request, _context, _next| { + Box::pin(async move { Ok(request.content) }) + }), ), "llm execution intercept:", ); @@ -2359,7 +2361,7 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { ctx.register_llm_stream_execution_intercept( "llm-stream", 1, - Arc::new(|_name, request, _next| { + Arc::new(|_name, request, _context, _next| { Box::pin(async move { Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( request.content @@ -2372,7 +2374,7 @@ fn test_plugin_registration_context_maps_duplicate_registration_errors() { ctx.register_llm_stream_execution_intercept( "llm-stream", 1, - Arc::new(|_name, request, _next| { + Arc::new(|_name, request, _context, _next| { Box::pin(async move { Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( request.content @@ -2494,13 +2496,13 @@ fn test_plugin_registration_context_maps_deregistration_errors() { ctx.register_llm_execution_intercept( "llm-exec", 1, - Arc::new(|_name, request, _next| Box::pin(async move { Ok(request.content) })), + Arc::new(|_name, request, _context, _next| Box::pin(async move { Ok(request.content) })), ) .unwrap(); ctx.register_llm_stream_execution_intercept( "llm-stream", 1, - Arc::new(|_name, request, _next| { + Arc::new(|_name, request, _context, _next| { Box::pin(async move { Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( request.content diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 13b5bcd10..5b5368bd9 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -381,6 +381,24 @@ typedef NemoRelayStatus (*NemoRelayLlmRequestInterceptCb)(void *user_data, const char *annotated_json, char **out_outcome_json); +/** + * Directional codec context supplied to an LLM execution intercept. + * + * `request_codec` is always present. `response_codec` is non-null for unary + * execution and null for streaming execution, where Relay has no completed + * response to decode. + */ +typedef struct NemoRelayLlmExecutionContext { + /** + * Active request codec identity and capability. + */ + struct NemoRelayLlmSanitizeRequestContext request_codec; + /** + * Active unary-response codec context, or null for streaming execution. + */ + const struct NemoRelayLlmSanitizeResponseContext *response_codec; +} NemoRelayLlmExecutionContext; + /** * Runtime-provided "next" callback for LLM execution middleware chain. * Takes a native JSON C string, returns a response JSON C string. @@ -393,10 +411,13 @@ typedef char *(*NemoRelayLlmExecNextFn)(const char *native_json, void *next_ctx) /** * Callback for LLM execution intercepts with middleware chain support. - * Receives native JSON C string plus a `next` callback and its context. + * Receives the managed LLM call name, native JSON C string, execution context, + * plus a `next` callback and its context. */ typedef char *(*NemoRelayLlmExecInterceptCb)(void *user_data, + const char *name, const char *native_json, + struct NemoRelayLlmExecutionContext context, NemoRelayLlmExecNextFn next_fn, void *next_ctx); diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 640359f39..4af0e5a0e 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -25,10 +25,11 @@ use std::sync::Arc; use libc::c_char; use nemo_relay::api::runtime::{ EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, - LlmConditionalFn, LlmExecutionNextFn, LlmJsonStream, LlmRequestInterceptFn, - LlmSanitizeRequestContext, LlmSanitizeRequestFn, LlmSanitizeResponseContext, - LlmSanitizeResponseFn, LlmStreamExecutionNextFn, ToolConditionalFn, ToolExecutionContext, - ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + LlmConditionalFn, LlmExecutionContext, LlmExecutionFn, LlmExecutionNextFn, LlmJsonStream, + LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, + LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, + LlmStreamExecutionNextFn, ToolConditionalFn, ToolExecutionContext, ToolExecutionFn, + ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, }; use serde_json::Value as Json; use tokio_stream::StreamExt; @@ -179,6 +180,7 @@ pub enum NemoRelayLlmSanitizeCodecKind { /// Codec identity supplied to an LLM sanitizer. `codec_id` is null for /// `None` and `Opaque`, and is valid only for the duration of the callback. #[repr(C)] +#[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmSanitizeRequestContext { /// Kind of active codec identity. pub codec_kind: NemoRelayLlmSanitizeCodecKind, @@ -190,6 +192,7 @@ pub struct NemoRelayLlmSanitizeRequestContext { /// Directional codec context supplied to an LLM response sanitizer. #[repr(C)] +#[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmSanitizeResponseContext { /// Kind of active codec identity. pub codec_kind: NemoRelayLlmSanitizeCodecKind, @@ -199,6 +202,20 @@ pub struct NemoRelayLlmSanitizeResponseContext { pub codec: *const crate::types::FfiLlmSanitizeResponseCodec, } +/// Directional codec context supplied to an LLM execution intercept. +/// +/// `request_codec` is always present. `response_codec` is non-null for unary +/// execution and null for streaming execution, where Relay has no completed +/// response to decode. +#[repr(C)] +#[derive(Debug, Clone, Copy)] +pub struct NemoRelayLlmExecutionContext { + /// Active request codec identity and capability. + pub request_codec: NemoRelayLlmSanitizeRequestContext, + /// Active unary-response codec context, or null for streaming execution. + pub response_codec: *const NemoRelayLlmSanitizeResponseContext, +} + /// LLM request sanitizer. It receives the request first and its codec context /// second. Return null to omit the observability payload. The request is /// borrowed, but returning that same pointer is supported as a pass-through. @@ -241,10 +258,13 @@ pub type NemoRelayLlmExecNextFn = unsafe extern "C" fn(native_json: *const c_char, next_ctx: *mut libc::c_void) -> *mut c_char; /// Callback for LLM execution intercepts with middleware chain support. -/// Receives native JSON C string plus a `next` callback and its context. +/// Receives the managed LLM call name, native JSON C string, execution context, +/// plus a `next` callback and its context. pub type NemoRelayLlmExecInterceptCb = unsafe extern "C" fn( user_data: *mut libc::c_void, + name: *const c_char, native_json: *const c_char, + context: NemoRelayLlmExecutionContext, next_fn: NemoRelayLlmExecNextFn, next_ctx: *mut libc::c_void, ) -> *mut c_char; @@ -588,24 +608,20 @@ async fn call_tool_exec_intercept_cb( } } -/// Wrap a C LLM execution intercept callback into an `Arc ...>`. +/// Wrap a C LLM execution intercept callback. pub fn wrap_llm_exec_intercept_fn( cb: NemoRelayLlmExecInterceptCb, user_data: *mut libc::c_void, free_fn: NemoRelayFreeFn, -) -> Arc< - dyn Fn( - &str, - LlmRequest, - LlmExecutionNextFn, - ) -> Pin> + Send>> - + Send - + Sync, -> { +) -> LlmExecutionFn { let ud = make_user_data(user_data, free_fn); Arc::new( - move |_name: &str, request: LlmRequest, next: LlmExecutionNextFn| { + move |name: &str, + request: LlmRequest, + context: LlmExecutionContext, + next: LlmExecutionNextFn| { let ud = ud.clone(); + let c_name = CString::new(name).unwrap_or_default(); Box::pin(async move { let next_box = Box::new(next); let next_ctx = Box::into_raw(next_box) as *mut libc::c_void; @@ -645,7 +661,24 @@ pub fn wrap_llm_exec_intercept_fn( let request_json = serde_json::to_value(&request).unwrap_or(Json::Null); let c_request = json_to_c_string(&request_json); clear_last_error(); - let result_ptr = unsafe { cb(ud.ptr, c_request, llm_next_trampoline, next_ctx) }; + let result_ptr = + match with_ffi_llm_execution_context(&context, |ffi_context| unsafe { + cb( + ud.ptr, + c_name.as_ptr(), + c_request, + ffi_context, + llm_next_trampoline, + next_ctx, + ) + }) { + Ok(result) => result, + Err(error) => { + unsafe { drop(Box::from_raw(next_ctx as *mut LlmExecutionNextFn)) }; + unsafe { nemo_relay_string_free_internal(c_request) }; + return Err(error); + } + }; unsafe { drop(Box::from_raw(next_ctx as *mut LlmExecutionNextFn)) }; unsafe { nemo_relay_string_free_internal(c_request) }; let result = @@ -664,19 +697,15 @@ pub fn wrap_llm_stream_exec_intercept_fn( cb: NemoRelayLlmExecInterceptCb, user_data: *mut libc::c_void, free_fn: NemoRelayFreeFn, -) -> Arc< - dyn Fn( - &str, - LlmRequest, - LlmStreamExecutionNextFn, - ) -> Pin> + Send>> - + Send - + Sync, -> { +) -> LlmStreamExecutionFn { let ud = make_user_data(user_data, free_fn); Arc::new( - move |_name: &str, request: LlmRequest, next: LlmStreamExecutionNextFn| { + move |name: &str, + request: LlmRequest, + context: LlmExecutionContext, + next: LlmStreamExecutionNextFn| { let ud = ud.clone(); + let c_name = CString::new(name).unwrap_or_default(); Box::pin(async move { let next_box = Box::new(next); let next_ctx = Box::into_raw(next_box) as *mut libc::c_void; @@ -722,7 +751,25 @@ pub fn wrap_llm_stream_exec_intercept_fn( let c_request = json_to_c_string(&request_json); clear_last_error(); let result_ptr = - unsafe { cb(ud.ptr, c_request, llm_stream_next_trampoline, next_ctx) }; + match with_ffi_llm_execution_context(&context, |ffi_context| unsafe { + cb( + ud.ptr, + c_name.as_ptr(), + c_request, + ffi_context, + llm_stream_next_trampoline, + next_ctx, + ) + }) { + Ok(result) => result, + Err(error) => { + unsafe { + drop(Box::from_raw(next_ctx as *mut LlmStreamExecutionNextFn)) + }; + unsafe { nemo_relay_string_free_internal(c_request) }; + return Err(error); + } + }; unsafe { drop(Box::from_raw(next_ctx as *mut LlmStreamExecutionNextFn)) }; unsafe { nemo_relay_string_free_internal(c_request) }; let result = json_result_from_ptr( @@ -943,6 +990,54 @@ fn ffi_codec_identity( }) } +fn with_ffi_llm_execution_context( + context: &LlmExecutionContext, + callback: impl FnOnce(NemoRelayLlmExecutionContext) -> T, +) -> Result { + let request_context = context.request_codec(); + let (request_kind, request_id) = ffi_codec_identity(request_context.codec())?; + let request_codec = request_context + .resolve_codec() + .map(crate::types::FfiLlmSanitizeRequestCodec); + let response_context = context.response_codec(); + let (response_kind, response_id, response_codec) = match response_context { + Some(context) => { + let (kind, id) = ffi_codec_identity(context.codec())?; + let codec = context + .resolve_codec() + .map(crate::types::FfiLlmSanitizeResponseCodec); + (Some(kind), id, codec) + } + None => (None, None, None), + }; + let request = NemoRelayLlmSanitizeRequestContext { + codec_kind: request_kind, + codec_id: request_id + .as_ref() + .map_or(std::ptr::null(), |name| name.as_ptr()), + codec: request_codec + .as_ref() + .map_or(std::ptr::null(), std::ptr::from_ref), + }; + let response = response_kind.map(|codec_kind| NemoRelayLlmSanitizeResponseContext { + codec_kind, + codec_id: response_id + .as_ref() + .map_or(std::ptr::null(), |name| name.as_ptr()), + codec: response_codec + .as_ref() + .map_or(std::ptr::null(), std::ptr::from_ref), + }); + // Identity strings and opaque codec wrappers stay in this stack frame for + // the complete callback. Foreign code must not retain any of their pointers. + Ok(callback(NemoRelayLlmExecutionContext { + request_codec: request, + response_codec: response + .as_ref() + .map_or(std::ptr::null(), std::ptr::from_ref), + })) +} + /// Wrap a C LLM conditional callback into a Rust closure. pub fn wrap_llm_conditional_fn( cb: NemoRelayLlmConditionalCb, diff --git a/crates/ffi/tests/integration/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index e4b1ce552..5175da0b9 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -15,8 +15,9 @@ use serde_json::{Value as Json, json}; use uuid::Uuid; use nemo_relay_ffi::callable::{ - NemoRelayLlmExecNextFn, NemoRelayLlmSanitizeCodecKind, NemoRelayLlmSanitizeRequestContext, - NemoRelayLlmSanitizeResponseContext, NemoRelayToolExecNextFn, + NemoRelayLlmExecNextFn, NemoRelayLlmExecutionContext, NemoRelayLlmSanitizeCodecKind, + NemoRelayLlmSanitizeRequestContext, NemoRelayLlmSanitizeResponseContext, + NemoRelayToolExecNextFn, }; use nemo_relay_ffi::convert::nemo_relay_string_free; use nemo_relay_ffi::error::{NemoRelayStatus, nemo_relay_last_error, set_last_error}; @@ -558,7 +559,9 @@ unsafe extern "C" fn codec_encode_cb( unsafe extern "C" fn llm_exec_intercept_cb( _user_data: *mut libc::c_void, + _name: *const c_char, native_json: *const c_char, + _context: NemoRelayLlmExecutionContext, next_fn: NemoRelayLlmExecNextFn, next_ctx: *mut libc::c_void, ) -> *mut c_char { diff --git a/crates/ffi/tests/integration/callable_extra_tests.rs b/crates/ffi/tests/integration/callable_extra_tests.rs index f3b079fc2..d044bfd6b 100644 --- a/crates/ffi/tests/integration/callable_extra_tests.rs +++ b/crates/ffi/tests/integration/callable_extra_tests.rs @@ -31,7 +31,9 @@ unsafe extern "C" fn tool_exec_intercept_null_next_cb( unsafe extern "C" fn llm_exec_intercept_null_next_cb( _user_data: *mut libc::c_void, + _name: *const c_char, _native_json: *const c_char, + _context: NemoRelayLlmExecutionContext, next_fn: NemoRelayLlmExecNextFn, next_ctx: *mut libc::c_void, ) -> *mut c_char { @@ -207,7 +209,12 @@ fn test_callable_extra_trampoline_and_helper_paths() { }) }); let llm_err = runtime - .block_on(llm_intercept("llm", make_request(), llm_next)) + .block_on(llm_intercept( + "llm", + make_request(), + nemo_relay::api::runtime::LlmExecutionContext::default(), + llm_next, + )) .unwrap_err(); assert!(llm_err.to_string().contains("llm next failed")); @@ -222,7 +229,12 @@ fn test_callable_extra_trampoline_and_helper_paths() { }) }); let mut empty_stream = runtime - .block_on(llm_stream_intercept("llm", make_request(), empty_next)) + .block_on(llm_stream_intercept( + "llm", + make_request(), + nemo_relay::api::runtime::LlmExecutionContext::default(), + empty_next, + )) .unwrap(); let empty_item = runtime .block_on(async { empty_stream.next().await }) @@ -233,7 +245,12 @@ fn test_callable_extra_trampoline_and_helper_paths() { let err_next: LlmStreamExecutionNextFn = Arc::new(|_request| { Box::pin(async move { Err(FlowError::Internal("stream next failed".into())) }) }); - let stream_err = match runtime.block_on(llm_stream_intercept("llm", make_request(), err_next)) { + let stream_err = match runtime.block_on(llm_stream_intercept( + "llm", + make_request(), + nemo_relay::api::runtime::LlmExecutionContext::default(), + err_next, + )) { Ok(_) => panic!("expected llm stream intercept error"), Err(err) => err, }; diff --git a/crates/ffi/tests/unit/api/registry_tests.rs b/crates/ffi/tests/unit/api/registry_tests.rs index 77aa8705a..911b8e40b 100644 --- a/crates/ffi/tests/unit/api/registry_tests.rs +++ b/crates/ffi/tests/unit/api/registry_tests.rs @@ -2656,10 +2656,20 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { }) ); + let intercept_name = cstring("ffi-intercept"); let request = cstring(r#"{"headers":{},"content":{"model":"ffi-model","messages":[]}}"#); let llm_intercept_json = take_string(llm_exec_intercept_cb( ptr::null_mut(), + intercept_name.as_ptr(), request.as_ptr(), + NemoRelayLlmExecutionContext { + request_codec: NemoRelayLlmSanitizeRequestContext { + codec_kind: NemoRelayLlmSanitizeCodecKind::None, + codec_id: ptr::null(), + codec: ptr::null(), + }, + response_codec: ptr::null(), + }, llm_next_passthrough, ptr::null_mut(), )) diff --git a/crates/ffi/tests/unit/api_tests.rs b/crates/ffi/tests/unit/api_tests.rs index 76df77156..83b8665dc 100644 --- a/crates/ffi/tests/unit/api_tests.rs +++ b/crates/ffi/tests/unit/api_tests.rs @@ -16,8 +16,9 @@ use serde_json::{Value as Json, json}; use uuid::Uuid; use crate::callable::{ - NemoRelayLlmExecNextFn, NemoRelayLlmSanitizeCodecKind, NemoRelayLlmSanitizeRequestContext, - NemoRelayLlmSanitizeResponseContext, NemoRelayToolExecNextFn, + NemoRelayLlmExecNextFn, NemoRelayLlmExecutionContext, NemoRelayLlmSanitizeCodecKind, + NemoRelayLlmSanitizeRequestContext, NemoRelayLlmSanitizeResponseContext, + NemoRelayToolExecNextFn, }; use crate::convert::nemo_relay_string_free; use crate::error::{NemoRelayStatus, nemo_relay_last_error}; @@ -656,7 +657,9 @@ unsafe extern "C" fn codec_encode_cb( unsafe extern "C" fn llm_exec_intercept_cb( _user_data: *mut libc::c_void, + _name: *const c_char, native_json: *const c_char, + _context: NemoRelayLlmExecutionContext, next_fn: NemoRelayLlmExecNextFn, next_ctx: *mut libc::c_void, ) -> *mut c_char { diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index abf6b75eb..6dd25b6b3 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -259,10 +259,13 @@ unsafe extern "C" fn llm_exec_error_cb( unsafe extern "C" fn llm_exec_intercept_cb( _user_data: *mut libc::c_void, + name: *const c_char, native_json: *const c_char, + _context: NemoRelayLlmExecutionContext, next_fn: NemoRelayLlmExecNextFn, next_ctx: *mut libc::c_void, ) -> *mut c_char { + assert_eq!(unsafe { CStr::from_ptr(name) }.to_str().unwrap(), "llm"); let result_ptr = unsafe { next_fn(native_json, next_ctx) }; if result_ptr.is_null() { return std::ptr::null_mut(); @@ -276,10 +279,13 @@ unsafe extern "C" fn llm_exec_intercept_cb( unsafe extern "C" fn llm_exec_short_circuit_cb( _user_data: *mut libc::c_void, + name: *const c_char, native_json: *const c_char, + _context: NemoRelayLlmExecutionContext, _next_fn: NemoRelayLlmExecNextFn, _next_ctx: *mut libc::c_void, ) -> *mut c_char { + assert_eq!(unsafe { CStr::from_ptr(name) }.to_str().unwrap(), "llm"); let request: Json = serde_json::from_str(unsafe { CStr::from_ptr(native_json) }.to_str().unwrap()).unwrap(); let response = json!({ @@ -289,6 +295,129 @@ unsafe extern "C" fn llm_exec_short_circuit_cb( CString::new(response.to_string()).unwrap().into_raw() } +struct OpaqueExecutionCodec; + +impl nemo_relay::codec::traits::LlmCodec for OpaqueExecutionCodec { + fn decode( + &self, + request: &LlmRequest, + ) -> nemo_relay::error::Result { + Ok(nemo_relay::codec::request::AnnotatedLlmRequest { + model: request.content["model"].as_str().map(str::to_owned), + ..Default::default() + }) + } + + fn encode( + &self, + annotated: &nemo_relay::codec::request::AnnotatedLlmRequest, + original: &LlmRequest, + ) -> nemo_relay::error::Result { + let mut request = original.clone(); + request.content["model"] = json!(annotated.model); + request.content["encoded"] = json!(true); + Ok(request) + } +} + +impl nemo_relay::codec::traits::LlmResponseCodec for OpaqueExecutionCodec { + fn decode_response( + &self, + response: &Json, + ) -> nemo_relay::error::Result { + Ok(nemo_relay::codec::response::AnnotatedLlmResponse { + model: response["model"].as_str().map(str::to_owned), + ..Default::default() + }) + } +} + +unsafe extern "C" fn llm_exec_absent_context_cb( + _user_data: *mut libc::c_void, + name: *const c_char, + native_json: *const c_char, + context: NemoRelayLlmExecutionContext, + next_fn: NemoRelayLlmExecNextFn, + next_ctx: *mut libc::c_void, +) -> *mut c_char { + assert_eq!( + unsafe { CStr::from_ptr(name) }.to_str().unwrap(), + "ffi-none" + ); + assert_eq!( + context.request_codec.codec_kind, + NemoRelayLlmSanitizeCodecKind::None + ); + assert!(context.request_codec.codec_id.is_null()); + assert!(context.request_codec.codec.is_null()); + assert!(!context.response_codec.is_null()); + let response = unsafe { &*context.response_codec }; + assert_eq!(response.codec_kind, NemoRelayLlmSanitizeCodecKind::None); + assert!(response.codec_id.is_null()); + assert!(response.codec.is_null()); + unsafe { next_fn(native_json, next_ctx) } +} + +unsafe extern "C" fn llm_exec_opaque_context_cb( + _user_data: *mut libc::c_void, + name: *const c_char, + native_json: *const c_char, + context: NemoRelayLlmExecutionContext, + next_fn: NemoRelayLlmExecNextFn, + next_ctx: *mut libc::c_void, +) -> *mut c_char { + assert_eq!( + unsafe { CStr::from_ptr(name) }.to_str().unwrap(), + "ffi-opaque" + ); + assert_eq!( + context.request_codec.codec_kind, + NemoRelayLlmSanitizeCodecKind::Opaque + ); + assert!(context.request_codec.codec_id.is_null()); + assert!(!context.request_codec.codec.is_null()); + + let request: LlmRequest = + serde_json::from_str(unsafe { CStr::from_ptr(native_json) }.to_str().unwrap()).unwrap(); + let request = FfiLLMRequest(request); + let annotated = unsafe { + crate::api::nemo_relay_llm_sanitize_request_codec_decode( + context.request_codec.codec, + std::ptr::from_ref(&request), + ) + }; + assert!(!annotated.is_null()); + let encoded = unsafe { + crate::api::nemo_relay_llm_sanitize_request_codec_encode( + context.request_codec.codec, + annotated, + std::ptr::from_ref(&request), + ) + }; + unsafe { nemo_relay_string_free_internal(annotated) }; + assert!(!encoded.is_null()); + let encoded_json = + CString::new(serde_json::to_string(&unsafe { &*encoded }.0).unwrap()).unwrap(); + let result = unsafe { next_fn(encoded_json.as_ptr(), next_ctx) }; + unsafe { drop(Box::from_raw(encoded)) }; + assert!(!result.is_null()); + + assert!(!context.response_codec.is_null()); + let response = unsafe { &*context.response_codec }; + assert_eq!(response.codec_kind, NemoRelayLlmSanitizeCodecKind::Opaque); + assert!(response.codec_id.is_null()); + assert!(!response.codec.is_null()); + let decoded_ptr = unsafe { + crate::api::nemo_relay_llm_sanitize_response_codec_decode(response.codec, result) + }; + assert!(!decoded_ptr.is_null()); + let decoded: Json = + serde_json::from_str(unsafe { CStr::from_ptr(decoded_ptr) }.to_str().unwrap()).unwrap(); + unsafe { nemo_relay_string_free_internal(decoded_ptr) }; + assert_eq!(decoded["model"], json!("test-model")); + result +} + static COLLECTED_COUNT: AtomicUsize = AtomicUsize::new(0); unsafe extern "C" fn collector_cb(_chunk: *const c_char) { @@ -648,9 +777,62 @@ fn assert_llm_exec_callbacks(runtime: &tokio::runtime::Runtime) { let next: LlmExecutionNextFn = Arc::new(|request| Box::pin(async move { Ok(json!({"model": request.content["model"]})) })); let intercepted = runtime - .block_on(intercept("llm", make_request(), next)) + .block_on(intercept( + "llm", + make_request(), + nemo_relay::api::runtime::LlmExecutionContext::default(), + next, + )) .unwrap(); assert_eq!(intercepted["intercepted"], json!(true)); + + let absent_intercept = + wrap_llm_exec_intercept_fn(llm_exec_absent_context_cb, std::ptr::null_mut(), None); + let absent_next: LlmExecutionNextFn = + Arc::new(|request| Box::pin(async move { Ok(json!({"model": request.content["model"]})) })); + let absent = runtime + .block_on(absent_intercept( + "ffi-none", + make_request(), + nemo_relay::api::runtime::LlmExecutionContext::new( + Default::default(), + Some(Default::default()), + ), + absent_next, + )) + .unwrap(); + assert_eq!(absent["model"], json!("test-model")); + + let opaque_intercept = + wrap_llm_exec_intercept_fn(llm_exec_opaque_context_cb, std::ptr::null_mut(), None); + let opaque_next: LlmExecutionNextFn = Arc::new(|request| { + Box::pin(async move { + assert_eq!(request.content["encoded"], json!(true)); + Ok(json!({"model": request.content["model"]})) + }) + }); + let request_codec: Arc = + Arc::new(OpaqueExecutionCodec); + let response_codec: Arc = + Arc::new(OpaqueExecutionCodec); + let opaque = runtime + .block_on(opaque_intercept( + "ffi-opaque", + make_request(), + nemo_relay::api::runtime::LlmExecutionContext::new( + nemo_relay::api::runtime::LlmSanitizeRequestContext::for_request_codec(Some( + request_codec, + )), + Some( + nemo_relay::api::runtime::LlmSanitizeResponseContext::for_response_codec(Some( + response_codec, + )), + ), + ), + opaque_next, + )) + .unwrap(); + assert_eq!(opaque["model"], json!("test-model")); } fn assert_llm_stream_callbacks(runtime: &tokio::runtime::Runtime) { @@ -669,7 +851,12 @@ fn assert_llm_stream_callbacks(runtime: &tokio::runtime::Runtime) { }) }); let mut intercepted_stream = runtime - .block_on(stream_intercept("llm", make_request(), next_stream)) + .block_on(stream_intercept( + "llm", + make_request(), + nemo_relay::api::runtime::LlmExecutionContext::default(), + next_stream, + )) .unwrap(); let first = runtime.block_on(async { intercepted_stream.next().await.unwrap().unwrap() }); assert_eq!(first["intercepted"], json!(true)); @@ -689,6 +876,7 @@ fn assert_llm_stream_callbacks(runtime: &tokio::runtime::Runtime) { .block_on(stream_intercept_with_next( "llm", make_request(), + nemo_relay::api::runtime::LlmExecutionContext::default(), next_stream, )) .unwrap(); diff --git a/crates/node/plugin.d.ts b/crates/node/plugin.d.ts index c363d8506..86d0bb26f 100644 --- a/crates/node/plugin.d.ts +++ b/crates/node/plugin.d.ts @@ -7,6 +7,7 @@ import type { EventMetadata, EventSanitizeFields, Json, + LlmExecutionContext, LlmRequestInterceptOutcome, LlmSanitizeRequestContext, LlmSanitizeResponseContext, @@ -21,6 +22,7 @@ export type { EventMetadataScalar, EventMetadataValue, LlmCodecIdentity, + LlmExecutionContext, LlmOptimizationContribution, LlmOptimizationDataSchema, LlmOptimizationModel, @@ -238,7 +240,11 @@ export interface PluginContext { registerLlmExecutionIntercept( name: string, priority: number, - callback: (request: Json, next: (request: Json) => Json | Promise) => Json | Promise, + callback: ( + request: Json, + context: LlmExecutionContext, + next: (request: Json) => Json | Promise, + ) => Json | Promise, ): void; /** * Register an LLM streaming execution intercept for this component. @@ -251,6 +257,7 @@ export interface PluginContext { priority: number, callback: ( request: Json, + context: LlmExecutionContext, next: (request: Json) => Promise>, ) => AsyncIterable | Promise>, ): void; diff --git a/crates/node/root-types.d.ts b/crates/node/root-types.d.ts index 278589649..ef4f88edc 100644 --- a/crates/node/root-types.d.ts +++ b/crates/node/root-types.d.ts @@ -25,6 +25,14 @@ export interface LlmSanitizeResponseContext { resolveCodec(): import('./typed').LlmResponseCodec | null; } +/** Codec capabilities for one managed LLM execution intercept invocation. */ +export interface LlmExecutionContext { + /** Request codec identity plus optional decode and encode capability. */ + requestCodec: LlmSanitizeRequestContext; + /** Unary response codec identity plus optional decode capability; `null` for streaming execution. */ + responseCodec: LlmSanitizeResponseContext | null; +} + /** Schema tag attached to an opaque optimization contribution payload. */ export interface LlmOptimizationDataSchema { name: string; diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index adb25766d..6e9f19386 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -4267,9 +4267,9 @@ pub fn deregister_llm_request_intercept(name: String) -> Result { /// Register an LLM execution intercept following the middleware chain pattern. /// -/// The `callable` receives the request and a `next` function. Call `next(request)` to -/// invoke the next intercept or original implementation; skip calling `next` to -/// short-circuit the chain. `next` may be called repeatedly or concurrently while +/// The `callable` receives the request, codec context, and a `next` function. Call +/// `next(request)` to invoke the next intercept or original implementation; skip calling +/// `next` to short-circuit the chain. `next` may be called repeatedly or concurrently while /// `callable` is pending; each call receives an isolated scope-stack branch, and /// unfinished or later calls reject after `callable` settles. #[napi] @@ -4278,7 +4278,7 @@ pub fn register_llm_execution_intercept( name: String, priority: i32, #[napi( - ts_arg_type = "(request: Json, next: (request: Json) => Json | Promise) => Json | Promise" + ts_arg_type = "(request: Json, context: LlmExecutionContext, next: (request: Json) => Json | Promise) => Json | Promise" )] callable: JsFunction, ) -> Result<()> { @@ -4306,7 +4306,9 @@ pub fn deregister_llm_execution_intercept(name: String) -> Result { /// Register a streaming LLM execution intercept following the middleware chain pattern. /// -/// The `callable` receives the request and a `next` function. Call `next(request)` to +/// The `callable` receives the request, request-codec context, and a `next` function. The +/// response codec is `null` because streaming execution has no complete-response codec. +/// Call `next(request)` to /// invoke the next intercept or original streaming implementation; in Node the /// returned promise resolves to a lazy `AsyncIterable`. Return it directly or wrap it /// `next` to short-circuit the chain. `next` may be called repeatedly or concurrently @@ -4319,7 +4321,7 @@ pub fn register_llm_stream_execution_intercept( name: String, priority: i32, #[napi( - ts_arg_type = "(request: Json, next: (request: Json) => Promise>) => AsyncIterable | Promise>" + ts_arg_type = "(request: Json, context: LlmExecutionContext, next: (request: Json) => Promise>) => AsyncIterable | Promise>" )] callable: JsFunction, ) -> Result<()> { @@ -4919,8 +4921,8 @@ pub fn scope_deregister_llm_request_intercept(scope_uuid: String, name: String) /// Register a scope-local LLM execution intercept following the middleware chain pattern. /// -/// The `callable` receives the request and a `next` function. Call `next(request)` to -/// invoke the next intercept or original implementation; skip calling `next` to +/// The `callable` receives the request, codec context, and a `next` function. Call +/// `next(request)` to invoke the next intercept or original implementation; skip calling `next` to /// short-circuit the chain. `next` may be called repeatedly or concurrently while /// `callable` is pending; each call receives an isolated scope-stack branch, and /// unfinished or later calls reject after `callable` settles. @@ -4931,7 +4933,7 @@ pub fn scope_register_llm_execution_intercept( name: String, priority: i32, #[napi( - ts_arg_type = "(request: Json, next: (request: Json) => Json | Promise) => Json | Promise" + ts_arg_type = "(request: Json, context: LlmExecutionContext, next: (request: Json) => Json | Promise) => Json | Promise" )] callable: JsFunction, ) -> Result<()> { @@ -4964,7 +4966,9 @@ pub fn scope_deregister_llm_execution_intercept(scope_uuid: String, name: String /// Register a scope-local streaming LLM execution intercept following the middleware chain pattern. /// -/// The `callable` receives the request and a `next` function. Call `next(request)` to +/// The `callable` receives the request, request-codec context, and a `next` function. The +/// response codec is `null` because streaming execution has no complete-response codec. +/// Call `next(request)` to /// invoke the next intercept or original streaming implementation; in Node the /// returned promise resolves to a lazy `AsyncIterable`. Return it directly or wrap it /// `next` to short-circuit the chain. `next` may be called repeatedly or concurrently @@ -4978,7 +4982,7 @@ pub fn scope_register_llm_stream_execution_intercept( name: String, priority: i32, #[napi( - ts_arg_type = "(request: Json, next: (request: Json) => Promise>) => AsyncIterable | Promise>" + ts_arg_type = "(request: Json, context: LlmExecutionContext, next: (request: Json) => Promise>) => AsyncIterable | Promise>" )] callable: JsFunction, ) -> Result<()> { diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index 018994484..1f0d05fe8 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -23,10 +23,11 @@ use napi::{Env, JsFunction, JsObject, JsUnknown, NapiRaw, NapiValue}; use napi_derive::napi; use nemo_relay::api::runtime::{ EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, - LlmConditionalFn, LlmExecutionNextFn, LlmJsonStream, LlmRequestInterceptFn, - LlmSanitizeRequestContext, LlmSanitizeRequestFn, LlmSanitizeResponseContext, - LlmSanitizeResponseFn, LlmStreamExecutionNextFn, ToolConditionalFn, ToolExecutionContext, - ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + LlmConditionalFn, LlmExecutionContext, LlmExecutionFn, LlmExecutionNextFn, + LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, + LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, + LlmStreamExecutionNextFn, ToolConditionalFn, ToolExecutionContext, ToolExecutionNextFn, + ToolInterceptFn, ToolSanitizeFn, }; use serde::{Deserialize, Serialize}; use serde_json::Value as Json; @@ -1162,6 +1163,45 @@ pub(crate) fn js_llm_sanitize_response_context_to_napi( Ok(js_object_to_unknown(env, object)) } +fn js_llm_execution_context_args( + request: LlmRequest, + context: LlmExecutionContext, +) -> crate::promise_call::Arg0Builder { + Box::new(move |env| { + let mut args = env.create_array_with_length(2)?; + let request = serde_json::to_value(request) + .map_err(|error| napi::Error::from_reason(error.to_string()))?; + args.set_element(0, json_to_js_unknown(env, request)?)?; + + let mut execution = env.create_object()?; + execution.set_named_property( + "requestCodec", + js_llm_sanitize_request_context_to_napi( + env, + js_llm_sanitize_request_context(context.request_codec()), + )?, + )?; + let response = match context.response_codec() { + Some(context) => js_llm_sanitize_response_context_to_napi( + env, + js_llm_sanitize_response_context(context), + )?, + None => { + let null = env.get_null()?; + unsafe { JsUnknown::from_raw_unchecked(env.raw(), null.raw()) } + } + }; + execution.set_named_property("responseCodec", response)?; + args.set_element(1, js_object_to_unknown(env, execution))?; + Ok(js_object_to_unknown(env, args)) + }) +} + +fn json_to_js_unknown(env: &Env, value: Json) -> napi::Result { + let raw = unsafe { Json::to_napi_value(env.raw(), value) }?; + Ok(unsafe { JsUnknown::from_raw_unchecked(env.raw(), raw) }) +} + /// Wrap a JS function for LLM conditional guardrails: `(request: object) => string | null`. pub fn wrap_js_llm_conditional_fn( func: ThreadsafeFunction, @@ -1636,26 +1676,18 @@ async fn call_js_tool_exec_intercept( }) } -/// Wrap a JS function `(request, next) => result` for LLM execution intercept. +/// Wrap a JS function `(request, context, next) => result` for LLM execution intercept. /// /// The JS callback receives the `LlmRequest` serialized as a plain JSON object /// and a real `next(request)` function that returns a Promise for the downstream /// result. -pub fn wrap_js_llm_exec_intercept_fn( - func: Arc, -) -> Arc< - dyn Fn( - &str, - LlmRequest, - LlmExecutionNextFn, - ) -> Pin> + Send>> - + Send - + Sync, -> { +pub fn wrap_js_llm_exec_intercept_fn(func: Arc) -> LlmExecutionFn { Arc::new( - move |_name: &str, request: LlmRequest, next: LlmExecutionNextFn| { + move |_name: &str, + request: LlmRequest, + context: LlmExecutionContext, + next: LlmExecutionNextFn| { let func = func.clone(); - let req_json = serde_json::to_value(&request).unwrap_or(Json::Null); let next_json: JsonNextFn = Arc::new(move |next_request_json| { let next = next.clone(); Box::pin(async move { @@ -1666,32 +1698,28 @@ pub fn wrap_js_llm_exec_intercept_fn( next(next_request).await }) }); - Box::pin(async move { func.call_with_json_next(req_json, next_json).await }) + let args = js_llm_execution_context_args(request, context); + Box::pin(async move { + func.call_spread_with_arg0_and_json_next(args, next_json) + .await + }) }, ) } -/// Wrap a JS function `(request, next) => result` for LLM stream execution intercept. +/// Wrap a JS function `(request, context, next) => result` for LLM stream execution intercept. /// /// The JS callback receives the `LlmRequest` serialized as a plain JSON object /// and a real `next(request)` function whose Promise resolves to a lazy /// async iterable. The callback can return it directly or wrap it with an async /// generator without materializing the downstream stream. -pub fn wrap_js_llm_stream_exec_intercept_fn( - func: Arc, -) -> Arc< - dyn Fn( - &str, - LlmRequest, - LlmStreamExecutionNextFn, - ) -> Pin> + Send>> - + Send - + Sync, -> { +pub fn wrap_js_llm_stream_exec_intercept_fn(func: Arc) -> LlmStreamExecutionFn { Arc::new( - move |_name: &str, request: LlmRequest, next: LlmStreamExecutionNextFn| { + move |_name: &str, + request: LlmRequest, + context: LlmExecutionContext, + next: LlmStreamExecutionNextFn| { let func = func.clone(); - let req_json = serde_json::to_value(&request).unwrap_or(Json::Null); let next_stream: JsonStreamNextFn = Arc::new(move |next_request_json| { let next = next.clone(); Box::pin(async move { @@ -1702,7 +1730,11 @@ pub fn wrap_js_llm_stream_exec_intercept_fn( next(next_request).await }) }); - Box::pin(async move { func.call_with_stream_next(req_json, next_stream).await }) + let args = js_llm_execution_context_args(request, context); + Box::pin(async move { + func.call_spread_with_arg0_and_stream_next(args, next_stream) + .await + }) }, ) } diff --git a/crates/node/src/promise_call.rs b/crates/node/src/promise_call.rs index 06d6c4131..b4b244eac 100644 --- a/crates/node/src/promise_call.rs +++ b/crates/node/src/promise_call.rs @@ -731,12 +731,48 @@ impl PromiseAwareFn { .await } + /// Call a spread JavaScript callback with builder-constructed arguments + /// followed by a middleware-style `next(arg)` callback. + pub async fn call_spread_with_arg0_and_json_next( + &self, + build_arg0: Arg0Builder, + next: JsonNextFn, + ) -> FlowResult { + self.call_inner( + PrimaryArg::Build(build_arg0), + CallMode::SPREAD, + Some(NextFn::Json(next)), + ) + .await + } + /// Call the JS function with a middleware-style `next(arg)` callback that /// resolves to a lazy downstream stream. pub async fn call_with_stream_next( &self, args: Json, next: JsonStreamNextFn, + ) -> FlowResult { + self.call_with_stream_next_inner(PrimaryArg::Json(args), false, next) + .await + } + + /// Call a spread JavaScript callback with builder-constructed arguments + /// followed by a middleware-style streaming `next(arg)` callback. + pub async fn call_spread_with_arg0_and_stream_next( + &self, + build_arg0: Arg0Builder, + next: JsonStreamNextFn, + ) -> FlowResult { + self.call_with_stream_next_inner(PrimaryArg::Build(build_arg0), true, next) + .await + } + + async fn call_with_stream_next_inner( + &self, + arg0: PrimaryArg, + spread: bool, + next: JsonStreamNextFn, ) -> FlowResult { let (ready_sender, ready_receiver) = tokio::sync::oneshot::channel(); let (chunk_sender, chunk_receiver) = @@ -759,8 +795,8 @@ impl PromiseAwareFn { .ok_or_else(closed_tsfn_error)?; let status = tsfn.call( Ok(CallArgs { - arg0: PrimaryArg::Json(args), - spread: false, + arg0, + spread, next: Some(NextFn::Stream(next)), publication: false, publication_context_id: publication_callback_context_id(), diff --git a/crates/node/tests/adaptive_tests.mjs b/crates/node/tests/adaptive_tests.mjs index bd23dfbe7..4aaf8d3af 100644 --- a/crates/node/tests/adaptive_tests.mjs +++ b/crates/node/tests/adaptive_tests.mjs @@ -198,14 +198,16 @@ describe('core plugins', () => { ...args, nodeToolPlugin: `priority:${pluginConfig.priority}`, })); - context.registerLlmExecutionIntercept('llmExec', 17, async (request, next) => { + context.registerLlmExecutionIntercept('llmExec', 17, async (request, _context, next) => { const result = await next(request); return { ...result, nodeLlmPlugin: `priority:${pluginConfig.priority}`, }; }); - context.registerLlmStreamExecutionIntercept('llmStreamExec', 17, async (request, next) => next(request)); + context.registerLlmStreamExecutionIntercept('llmStreamExec', 17, async (request, _context, next) => + next(request), + ); }, }); diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index 040ae0405..6ac75ede2 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -1144,7 +1144,7 @@ describe('LLM intercepts', () => { const observed = []; const explicitRoot = '018f13f0-7c1a-7a80-8000-000000000799'; registerSubscriber('node_llm_exec_propagation_parent', (event) => events.push(event)); - registerLlmExecutionIntercept('node_llm_exec_propagation_parent', 10, async (request, next) => { + registerLlmExecutionIntercept('node_llm_exec_propagation_parent', 10, async (request, _context, next) => { const before = lib.capturePropagationContext(); const rooted = lib.capturePropagationContextWithRoot(explicitRoot); assert.equal(rooted.rootUuid, explicitRoot); @@ -1265,7 +1265,7 @@ describe('LLM intercepts', () => { const events = []; const observed = []; registerSubscriber('node_llm_exec_propagated_trace_root', (event) => events.push(event)); - registerLlmExecutionIntercept('node_llm_exec_propagated_trace_root', 10, async (request, next) => { + registerLlmExecutionIntercept('node_llm_exec_propagated_trace_root', 10, async (request, _context, next) => { observed.push(lib.captureTraceparent()); return next(request); }); @@ -1477,12 +1477,12 @@ describe('LLM intercepts', () => { }); it('execution intercept', () => { - registerLlmExecutionIntercept('node_llm_exec_int', 10, async (native, next) => next(native)); + registerLlmExecutionIntercept('node_llm_exec_int', 10, async (native, _context, next) => next(native)); deregisterLlmExecutionIntercept('node_llm_exec_int'); }); it('stream execution intercept', () => { - registerLlmStreamExecutionIntercept('node_llm_stream_exec', 10, async (native, next) => next(native)); + registerLlmStreamExecutionIntercept('node_llm_stream_exec', 10, async (native, _context, next) => next(native)); deregisterLlmStreamExecutionIntercept('node_llm_stream_exec'); }); @@ -1607,7 +1607,7 @@ describe('LLM intercepts', () => { }); it('execution intercept composes with next', async () => { - registerLlmExecutionIntercept('node_llm_exec_repl', 10, async (native, next) => { + registerLlmExecutionIntercept('node_llm_exec_repl', 10, async (native, _context, next) => { native.content.intercepted = true; const result = await next(native); return { @@ -1633,6 +1633,149 @@ describe('LLM intercepts', () => { deregisterLlmExecutionIntercept('node_llm_exec_repl'); }); + it('execution intercept receives directional codec context', async () => { + const codec = new lib.OpenAIChatCodec(); + const response = { + id: 'chatcmpl-execution-context', + model: 'test-model', + choices: [{ index: 0, message: { role: 'assistant', content: 'ok' }, finish_reason: 'stop' }], + }; + let observed = false; + registerLlmExecutionIntercept('node_llm_execution_context', 10, async (request, context, next) => { + assert.deepEqual(context.requestCodec.codec, { kind: 'opaque' }); + const requestCodec = context.requestCodec.resolveCodec(); + assert.notEqual(requestCodec, null); + assert.equal(requestCodec.decode(request).model, 'test-model'); + + assert.notEqual(context.responseCodec, null); + assert.deepEqual(context.responseCodec.codec, { kind: 'opaque' }); + const result = await next(request); + const responseCodec = context.responseCodec.resolveCodec(); + assert.notEqual(responseCodec, null); + assert.equal(responseCodec.decodeResponse(result).model, 'test-model'); + observed = true; + return result; + }); + try { + const result = await llmCallExecute( + 'node_llm_execution_context', + makeNative(), + () => response, + null, + null, + null, + null, + null, + codec.decode.bind(codec), + ({ annotated, original }) => codec.encode(annotated, original), + codec.decodeResponse.bind(codec), + ); + assert.deepEqual(result, response); + assert.equal(observed, true); + } finally { + deregisterLlmExecutionIntercept('node_llm_execution_context'); + } + }); + + it('execution intercept reports absent codecs without capabilities', async () => { + let observed = false; + registerLlmExecutionIntercept('node_llm_execution_context_absent', 10, async (request, context, next) => { + assert.deepEqual(context.requestCodec.codec, { kind: 'none' }); + assert.equal(context.requestCodec.resolveCodec(), null); + assert.notEqual(context.responseCodec, null); + assert.deepEqual(context.responseCodec.codec, { kind: 'none' }); + assert.equal(context.responseCodec.resolveCodec(), null); + observed = true; + return next(request); + }); + try { + const response = await llmCallExecute( + 'node_llm_execution_context_absent', + makeNative(), + () => ({ ok: true }), + null, + null, + null, + null, + null, + ); + assert.deepEqual(response, { ok: true }); + assert.equal(observed, true); + } finally { + deregisterLlmExecutionIntercept('node_llm_execution_context_absent'); + } + }); + + it('execution codec capability expires when the callback settles', async () => { + const codec = new lib.OpenAIChatCodec(); + let retainedCodec; + registerLlmExecutionIntercept('node_llm_execution_codec_expiry', 10, async (request, context, next) => { + retainedCodec = context.requestCodec.resolveCodec(); + assert.notEqual(retainedCodec, null); + return next(request); + }); + try { + await llmCallExecute( + 'node_llm_execution_codec_expiry', + makeNative(), + () => ({ ok: true }), + null, + null, + null, + null, + null, + codec.decode.bind(codec), + ({ annotated, original }) => codec.encode(annotated, original), + ); + } finally { + deregisterLlmExecutionIntercept('node_llm_execution_codec_expiry'); + } + + assert.notEqual(retainedCodec, undefined); + assert.throws(() => retainedCodec.decode(makeNative()), /LLM execution codec capability is no longer active/i); + }); + + it('stream execution context omits the response codec', async () => { + const codec = new lib.OpenAIChatCodec(); + let observed = false; + let retainedCodec; + registerLlmStreamExecutionIntercept('node_llm_stream_execution_context', 10, async (request, context, next) => { + assert.deepEqual(context.requestCodec.codec, { kind: 'opaque' }); + retainedCodec = context.requestCodec.resolveCodec(); + assert.notEqual(retainedCodec, null); + assert.equal(context.responseCodec, null); + observed = true; + return next(request); + }); + try { + const stream = await llmStreamCallExecute( + 'node_llm_stream_execution_context', + makeNative(), + (wrapper) => { + lib.pushStreamChunk(wrapper.__nemo_relay_stream_id, { token: 'ok' }); + lib.endStream(wrapper.__nemo_relay_stream_id); + }, + null, + () => ({}), + null, + null, + null, + null, + null, + codec.decode.bind(codec), + ({ annotated, original }) => codec.encode(annotated, original), + codec.decodeResponse.bind(codec), + ); + assert.equal(retainedCodec.decode(makeNative()).model, 'test-model'); + assert.deepEqual(await stream.next(), { token: 'ok' }); + assert.equal(await stream.next(), null); + assert.throws(() => retainedCodec.decode(makeNative()), /LLM execution codec capability is no longer active/i); + assert.equal(observed, true); + } finally { + deregisterLlmStreamExecutionIntercept('node_llm_stream_execution_context'); + } + }); + it('execution intercept rejects a detached next call after settlement', async () => { let releaseLateNext; const lateGate = new Promise((resolve) => { @@ -1640,7 +1783,7 @@ describe('LLM intercepts', () => { }); let lateNext; let providerCalls = 0; - registerLlmExecutionIntercept('node_llm_exec_late_next', 10, async (native, next) => { + registerLlmExecutionIntercept('node_llm_exec_late_next', 10, async (native, _context, next) => { lateNext = lateGate.then(() => next(native)); return { source: 'intercept' }; }); @@ -1675,7 +1818,7 @@ describe('LLM intercepts', () => { }); let downstream; let providerSideEffects = 0; - registerLlmExecutionIntercept('node_llm_exec_abort_started_provider', 10, async (native, next) => { + registerLlmExecutionIntercept('node_llm_exec_abort_started_provider', 10, async (native, _context, next) => { downstream = next(native); downstream.catch(() => undefined); await started; @@ -1722,7 +1865,7 @@ describe('LLM intercepts', () => { }); it('execution intercept rejects invalid next request payloads', async () => { - registerLlmExecutionIntercept('node_llm_exec_invalid_next', 10, async (_native, next) => { + registerLlmExecutionIntercept('node_llm_exec_invalid_next', 10, async (_native, _context, next) => { return next({ headers: 1, content: { @@ -1793,7 +1936,7 @@ describe('LLM intercepts', () => { }); it('execution intercept rejects non-JSON next arguments without aborting Node', async () => { - registerLlmExecutionIntercept('node_llm_exec_bigint_next', 10, async (_native, next) => next(1n)); + registerLlmExecutionIntercept('node_llm_exec_bigint_next', 10, async (_native, _context, next) => next(1n)); try { await assert.rejects( () => llmCallExecute('bigint_next_llm', makeNative(), () => ({ ok: true })), @@ -1805,7 +1948,7 @@ describe('LLM intercepts', () => { }); it('stream execution intercept composes with next', async () => { - registerLlmStreamExecutionIntercept('node_llm_stream_exec_repl', 10, async (native, next) => { + registerLlmStreamExecutionIntercept('node_llm_stream_exec_repl', 10, async (native, _context, next) => { native.content.intercepted = true; const downstream = await next(native); return (async function* () { @@ -1853,7 +1996,7 @@ describe('LLM intercepts', () => { }); let lateNext; let providerCalls = 0; - registerLlmStreamExecutionIntercept('node_llm_stream_late_next', 10, async (native, next) => { + registerLlmStreamExecutionIntercept('node_llm_stream_late_next', 10, async (native, _context, next) => { lateNext = lateGate.then(() => next(native)); return (async function* () { yield { source: 'intercept' }; @@ -1879,7 +2022,7 @@ describe('LLM intercepts', () => { const invocationStack = lib.createScopeStack(); const invocationScope = lib.withScopeStack(invocationStack, () => lib.getHandle().uuid); let retainedNext; - registerLlmStreamExecutionIntercept('node_llm_stream_retained_next_scope', 10, async (native, next) => { + registerLlmStreamExecutionIntercept('node_llm_stream_retained_next_scope', 10, async (native, _context, next) => { retainedNext = () => next(native); return (async function* () { for (let index = 0; index < 64; index += 1) { @@ -1909,7 +2052,9 @@ describe('LLM intercepts', () => { const completion = new Promise((resolve) => { releaseCompletion = resolve; }); - registerLlmStreamExecutionIntercept('node_llm_stream_incremental', 10, async (request, next) => next(request)); + registerLlmStreamExecutionIntercept('node_llm_stream_incremental', 10, async (request, _context, next) => + next(request), + ); try { const stream = await llmStreamCallExecute('incremental_stream_llm', makeNative(), (wrapper) => { lib.pushStreamChunk(wrapper.__nemo_relay_stream_id, { token: 'first' }); @@ -1956,7 +2101,7 @@ describe('LLM intercepts', () => { it('stream execution intercept closes a transformed downstream stream early', async () => { let providerEnded = false; - registerLlmStreamExecutionIntercept('node_llm_stream_early_close', 10, async (request, next) => { + registerLlmStreamExecutionIntercept('node_llm_stream_early_close', 10, async (request, _context, next) => { const downstream = await next(request); return (async function* () { yield* downstream; @@ -2007,18 +2152,21 @@ describe('LLM intercepts', () => { it('stream execution intercept close waits for iterator cleanup', async () => { let cleaned = false; - registerLlmStreamExecutionIntercept('node_llm_stream_await_cleanup', 10, async (_request, _next, signal) => - (async function* () { - try { - yield { token: 'first' }; - if (!signal.aborted) { - await new Promise((resolve) => signal.addEventListener('abort', resolve, { once: true })); + registerLlmStreamExecutionIntercept( + 'node_llm_stream_await_cleanup', + 10, + async (_request, _context, _next, signal) => + (async function* () { + try { + yield { token: 'first' }; + if (!signal.aborted) { + await new Promise((resolve) => signal.addEventListener('abort', resolve, { once: true })); + } + } finally { + await new Promise((resolve) => setTimeout(resolve, 10)); + cleaned = true; } - } finally { - await new Promise((resolve) => setTimeout(resolve, 10)); - cleaned = true; - } - })(), + })(), ); try { const stream = await llmStreamCallExecute('await_cleanup_stream_llm', makeNative(), () => {}); @@ -2126,26 +2274,30 @@ describe('LLM intercepts', () => { const secondStack = lib.createScopeStack(); const firstScope = lib.withScopeStack(firstStack, () => lib.getHandle().uuid); const secondScope = lib.withScopeStack(secondStack, () => lib.getHandle().uuid); - registerLlmStreamExecutionIntercept('node_llm_stream_next_scope_replacements', 10, async (native, next) => { - const [first, second] = await Promise.all([ - lib.withScopeStack(firstStack, () => - next({ - ...native, - content: { ...native.content, branch: 'first' }, - }), - ), - lib.withScopeStack(secondStack, () => - next({ - ...native, - content: { ...native.content, branch: 'second' }, - }), - ), - ]); - return (async function* () { - yield* first; - yield* second; - })(); - }); + registerLlmStreamExecutionIntercept( + 'node_llm_stream_next_scope_replacements', + 10, + async (native, _context, next) => { + const [first, second] = await Promise.all([ + lib.withScopeStack(firstStack, () => + next({ + ...native, + content: { ...native.content, branch: 'first' }, + }), + ), + lib.withScopeStack(secondStack, () => + next({ + ...native, + content: { ...native.content, branch: 'second' }, + }), + ), + ]); + return (async function* () { + yield* first; + yield* second; + })(); + }, + ); try { const stream = await llmStreamCallExecute('scoped_next_stream_llm', makeNative(), (wrapper) => { lib.pushStreamChunk(wrapper.__nemo_relay_stream_id, { @@ -2172,7 +2324,7 @@ describe('LLM intercepts', () => { }); it('stream execution intercept rejects non-JSON next arguments without aborting Node', async () => { - registerLlmStreamExecutionIntercept('node_llm_stream_bigint_next', 10, async (_native, next) => next(1n)); + registerLlmStreamExecutionIntercept('node_llm_stream_bigint_next', 10, async (_native, _context, next) => next(1n)); try { await assert.rejects( () => @@ -2196,14 +2348,14 @@ describe('LLM intercepts', () => { releaseBlocker = resolve; }); - registerLlmStreamExecutionIntercept('node_llm_stream_snapshot_target', 100, async (request, next) => { + registerLlmStreamExecutionIntercept('node_llm_stream_snapshot_target', 100, async (request, _context, next) => { const downstream = await next(request); return (async function* () { yield* downstream; yield { snapshotted: true }; })(); }); - registerLlmStreamExecutionIntercept('node_llm_stream_snapshot_blocker', -100, async (request, next) => { + registerLlmStreamExecutionIntercept('node_llm_stream_snapshot_blocker', -100, async (request, _context, next) => { blockerEntered(); await release; return next(request); @@ -2255,7 +2407,7 @@ describe('LLM intercepts', () => { headers: {}, content: { messages: [], model: 'test-model' }, }; - lib.registerLlmStreamExecutionIntercept('process-exit-stream', 10, async (value, next) => next(value)); + lib.registerLlmStreamExecutionIntercept('process-exit-stream', 10, async (value, _context, next) => next(value)); const stream = await lib.llmStreamCallExecute( 'process-exit-llm', request, @@ -2285,7 +2437,7 @@ describe('LLM intercepts', () => { }); it('stream execution intercept rejects invalid next request payloads', async () => { - registerLlmStreamExecutionIntercept('node_llm_stream_invalid_next', 10, async (_native, next) => { + registerLlmStreamExecutionIntercept('node_llm_stream_invalid_next', 10, async (_native, _context, next) => { return next({ headers: 1, content: { @@ -2441,6 +2593,15 @@ describe('LLM intercepts', () => { 2, 'global and scope-local context intercept declarations must expose the tool execution context', ); + assert.equal( + declarations.split('context: LlmExecutionContext').length - 1, + 4, + 'global and scope-local unary and streaming LLM intercept declarations must expose the execution codec context', + ); + assert.match( + declarations, + /export interface LlmExecutionContext \{[\s\S]*?requestCodec: LlmSanitizeRequestContext[\s\S]*?responseCodec: LlmSanitizeResponseContext \| null/, + ); assert.doesNotMatch(declarations, /registerToolExecutionInterceptV2|scopeRegisterToolExecutionInterceptV2/); }); diff --git a/crates/node/tests/scope_local_tests.mjs b/crates/node/tests/scope_local_tests.mjs index 46ddf0278..f65bbb1ad 100644 --- a/crates/node/tests/scope_local_tests.mjs +++ b/crates/node/tests/scope_local_tests.mjs @@ -569,7 +569,7 @@ describe('Scope-local auto-cleanup on scope pop', () => { it('scope-local llm execution intercept is cleaned up when scope is popped', async () => { const scope = pushScope('sl_cleanup_llm_exec', ScopeType.Agent, null, null); - scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_cleanup_llm_exec_int', 10, async (request, next) => { + scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_cleanup_llm_exec_int', 10, async (request, _context, next) => { const updated = { ...request, content: { @@ -606,7 +606,7 @@ describe('Scope-local auto-cleanup on scope pop', () => { scope.uuid, 'sl_cleanup_llm_stream_exec_int', 10, - async (request, next) => { + async (request, _context, next) => { const updated = { ...request, content: { @@ -925,7 +925,7 @@ describe('Priority merge of global and scope-local middleware', () => { }); it('scope-local llm execution intercept and global intercept merge', async () => { - lib.registerLlmExecutionIntercept('sl_llm_merge_global_exec', 5, async (request, next) => { + lib.registerLlmExecutionIntercept('sl_llm_merge_global_exec', 5, async (request, _context, next) => { const result = await next({ ...request, content: { @@ -940,7 +940,7 @@ describe('Priority merge of global and scope-local middleware', () => { }); const scope = pushScope('sl_llm_merge_exec_scope', ScopeType.Agent, null, null); - scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_llm_merge_local_exec', 15, async (request, next) => { + scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_llm_merge_local_exec', 15, async (request, _context, next) => { const result = await next({ ...request, content: { @@ -1059,7 +1059,9 @@ describe('Scope-local LLM intercepts', () => { it('register and deregister scope-local llm execution intercept', () => { const scope = pushScope('sl_llm_exec_int_scope', ScopeType.Agent, null, null); - scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_llm_exec_int', 10, async (request, next) => next(request)); + scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_llm_exec_int', 10, async (request, _context, next) => + next(request), + ); const removed = scopeDeregisterLlmExecutionIntercept(scope.uuid, 'sl_llm_exec_int'); assert.equal(removed, true); popScope(scope); @@ -1067,7 +1069,7 @@ describe('Scope-local LLM intercepts', () => { it('scope-local llm execution intercept composes with next', async () => { const scope = pushScope('sl_llm_exec_compose_scope', ScopeType.Agent, null, null); - scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_llm_exec_compose', 10, async (request, next) => { + scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_llm_exec_compose', 10, async (request, _context, next) => { const result = await next({ ...request, content: { @@ -1106,7 +1108,7 @@ describe('Scope-local LLM intercepts', () => { it('scope-local llm execution intercept rejects invalid next request payloads', async () => { const scope = pushScope('sl_llm_exec_invalid_scope', ScopeType.Agent, null, null); - scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_llm_exec_invalid', 10, async (_request, next) => { + scopeRegisterLlmExecutionIntercept(scope.uuid, 'sl_llm_exec_invalid', 10, async (_request, _context, next) => { const downstream = await next({ headers: 1, content: { @@ -1144,7 +1146,7 @@ describe('Scope-local LLM intercepts', () => { it('register and deregister scope-local llm stream execution intercept', () => { const scope = pushScope('sl_llm_stream_int_scope', ScopeType.Agent, null, null); - scopeRegisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_int', 10, async (request, next) => + scopeRegisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_int', 10, async (request, _context, next) => next(request), ); const removed = scopeDeregisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_int'); @@ -1154,19 +1156,24 @@ describe('Scope-local LLM intercepts', () => { it('scope-local llm stream execution intercept composes with next', async () => { const scope = pushScope('sl_llm_stream_compose_scope', ScopeType.Agent, null, null); - scopeRegisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_compose', 10, async (request, next) => { - const downstream = await next({ - ...request, - content: { - ...request.content, - touchedByScopeStream: true, - }, - }); - return (async function* () { - yield* downstream; - yield { wrappedByScopeStream: true }; - })(); - }); + scopeRegisterLlmStreamExecutionIntercept( + scope.uuid, + 'sl_llm_stream_compose', + 10, + async (request, _context, next) => { + const downstream = await next({ + ...request, + content: { + ...request.content, + touchedByScopeStream: true, + }, + }); + return (async function* () { + yield* downstream; + yield { wrappedByScopeStream: true }; + })(); + }, + ); try { const stream = await llmStreamCallExecute( @@ -1205,14 +1212,19 @@ describe('Scope-local LLM intercepts', () => { it('scope-local llm stream execution intercept rejects invalid next request payloads', async () => { const scope = pushScope('sl_llm_stream_invalid_scope', ScopeType.Agent, null, null); - scopeRegisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_invalid', 10, async (_request, next) => { - return next({ - headers: 1, - content: { - model: 'broken', - }, - }); - }); + scopeRegisterLlmStreamExecutionIntercept( + scope.uuid, + 'sl_llm_stream_invalid', + 10, + async (_request, _context, next) => { + return next({ + headers: 1, + content: { + model: 'broken', + }, + }); + }, + ); try { await assert.rejects( @@ -1244,11 +1256,11 @@ describe('Scope-local LLM intercepts', () => { it('duplicate scope-local llm stream execution intercept fails', () => { const scope = pushScope('sl_llm_stream_dup_scope', ScopeType.Agent, null, null); - scopeRegisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_dup', 10, async (request, next) => + scopeRegisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_dup', 10, async (request, _context, next) => next(request), ); assert.throws(() => { - scopeRegisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_dup', 20, async (request, next) => + scopeRegisterLlmStreamExecutionIntercept(scope.uuid, 'sl_llm_stream_dup', 20, async (request, _context, next) => next(request), ); }); @@ -1402,10 +1414,12 @@ describe('Scope-local subscriber receives events', () => { })), () => scopeDeregisterLlmRequestIntercept('not-a-uuid', 'bad_llm_int'), () => - scopeRegisterLlmExecutionIntercept('not-a-uuid', 'bad_llm_exec', 10, async (request, next) => next(request)), + scopeRegisterLlmExecutionIntercept('not-a-uuid', 'bad_llm_exec', 10, async (request, _context, next) => + next(request), + ), () => scopeDeregisterLlmExecutionIntercept('not-a-uuid', 'bad_llm_exec'), () => - scopeRegisterLlmStreamExecutionIntercept('not-a-uuid', 'bad_llm_stream', 10, async (request, next) => + scopeRegisterLlmStreamExecutionIntercept('not-a-uuid', 'bad_llm_stream', 10, async (request, _context, next) => next(request), ), () => scopeDeregisterLlmStreamExecutionIntercept('not-a-uuid', 'bad_llm_stream'), diff --git a/crates/node/tests/typed_tests.mjs b/crates/node/tests/typed_tests.mjs index e55451076..78c1c30da 100644 --- a/crates/node/tests/typed_tests.mjs +++ b/crates/node/tests/typed_tests.mjs @@ -618,7 +618,7 @@ describe('typedLlmExecute', () => { }); let downstream; let providerSideEffects = 0; - registerLlmExecutionIntercept('typed_llm_abort_started_provider', 10, async (request, next) => { + registerLlmExecutionIntercept('typed_llm_abort_started_provider', 10, async (request, _context, next) => { downstream = next(request); downstream.catch(() => undefined); await started; diff --git a/crates/plugin/README.md b/crates/plugin/README.md index 2c2691a29..81da76a94 100644 --- a/crates/plugin/README.md +++ b/crates/plugin/README.md @@ -32,7 +32,7 @@ the dynamic-library boundary on the stable C-compatible ABI. | `PluginContext` | Installs component-owned subscribers, guardrails, intercepts, continuations, and streams. | | `PluginRuntime` | Emits marks and manages Relay-owned scopes and scope stacks through typed host helpers. | | `nemo_relay_plugin!` | Exports the one versioned native entry point used by the loader. | -| Native ABI v5 | Keeps C-compatible host and plugin tables behind the safe Rust interface while the host retains frozen v4, v3, and v2 tables for previously compiled plugins. | +| Native ABI v7 | Keeps C-compatible host and plugin tables behind the safe Rust interface. ABI v7 adds directional codec context to LLM execution callbacks and intentionally rejects plugins compiled with an older callback layout. | | Typed async middleware | Drives guardrails, sanitizers, and intercepts on a per-component SDK-owned Tokio executor. Subscribers and raw ABI registrations remain synchronous. | | Async continuations and streams | `ToolNext`, `LlmNext`, and `LlmStreamNext` support repeated or concurrent downstream calls. Streaming LLM continuations use a pull-based host handle. | | Tool results | `ToolNext` returns `ToolExecutionResult`, which keeps an application result and optional annotation together. | @@ -84,7 +84,8 @@ Build the `cdylib`, describe its entry symbol and compatibility in a `relay-plugin.toml` manifest, then register it through the Relay CLI. Refer to the complete example for platform-specific artifact and manifest setup. -Typed async plugins require `compat.relay = ">=0.8.0,<1.0"`. Relay creates one +Native plugins built with the 0.10 SDK require +`compat.relay = ">=0.10.0,<1.0"`. Relay creates one SDK-owned Tokio executor for each configured plugin component. It defaults to two workers: enough for modest concurrent async I/O without broadly oversubscribing the host. Increase the count only when measured I/O concurrency @@ -100,8 +101,15 @@ Relay 0.9 advances the C host-table ABI to v5 for `ToolExecutionContext`; the ho the frozen v4 table for previously compiled plugins. Plugins that register a context-aware tool execution intercept must rebuild and set `compat.relay = ">=0.9.0,<1.0"`; the required registration is unavailable in the v4 host -table. Typed async plugins that do not use this registration may retain -`compat.relay = ">=0.8.0,<1.0"`. +table. Under Relay 0.9, typed async plugins that did not use this registration +could retain `compat.relay = ">=0.8.0,<1.0"`. + +Relay 0.10 advances the internal table to ABI v7 and makes +`LlmExecutionContext` part of every unary and streaming LLM execution callback. +Because this changes callback layouts, the 0.10 host rejects every native plugin +compiled against an older table. Rebuild the plugin with the 0.10 SDK and set +`compat.relay = ">=0.10.0,<1.0"`. The authored manifest contract remains +`compat.native_api = "1"`. Set a plugin-wide default in Rust, then let the component's TOML configuration override it: diff --git a/crates/plugin/src/async_sdk.rs b/crates/plugin/src/async_sdk.rs index 332ae714b..ec416c35f 100644 --- a/crates/plugin/src/async_sdk.rs +++ b/crates/plugin/src/async_sdk.rs @@ -153,6 +153,12 @@ struct HostV4(NemoRelayNativeHostApiV4); unsafe impl Send for HostV4 {} unsafe impl Sync for HostV4 {} +#[derive(Clone, Copy)] +struct HostV7(NemoRelayNativeHostApiV7); + +unsafe impl Send for HostV7 {} +unsafe impl Sync for HostV7 {} + struct Completion { host: HostV4, raw: *const NemoRelayNativeAsyncCompletion, @@ -235,6 +241,77 @@ impl CompletionRef { }; Ok(LlmSanitizeResponseContext { codec, resolved }) } + + fn execution_request_context( + self, + codec: LlmCodecIdentity, + resolved: bool, + ) -> Result> { + let resolved = if resolved { + let status = unsafe { (self.host.0.async_completion_retain)(self.raw) }; + status_result(status, "retain native async completion capability")?; + Some(LlmExecutionRequestCodec { + owner: LlmExecutionRequestCodecOwner::Completion { + host: self.host.0, + completion: self.raw, + }, + _lifetime: PhantomData, + }) + } else { + None + }; + Ok(LlmExecutionRequestContext { codec, resolved }) + } + + fn execution_response_context( + self, + codec: LlmCodecIdentity, + resolved: bool, + ) -> Result> { + let resolved = if resolved { + let status = unsafe { (self.host.0.async_completion_retain)(self.raw) }; + status_result(status, "retain native async completion capability")?; + Some(LlmExecutionResponseCodec { + host: self.host.0, + completion: self.raw, + _lifetime: PhantomData, + }) + } else { + None + }; + Ok(LlmExecutionResponseContext { codec, resolved }) + } +} + +#[derive(Clone, Copy)] +struct StreamRef { + host: HostV7, + raw: *const NemoRelayNativeAsyncStream, +} + +unsafe impl Send for StreamRef {} + +impl StreamRef { + fn execution_request_context( + self, + codec: LlmCodecIdentity, + resolved: bool, + ) -> Result> { + let resolved = if resolved { + let status = unsafe { (self.host.0.async_stream_retain)(self.raw) }; + status_result(status, "retain native async stream capability")?; + Some(LlmExecutionRequestCodec { + owner: LlmExecutionRequestCodecOwner::Stream { + host: self.host.0, + stream: self.raw, + }, + _lifetime: PhantomData, + }) + } else { + None + }; + Ok(LlmExecutionRequestContext { codec, resolved }) + } } struct NextInner { @@ -534,8 +611,14 @@ unsafe extern "C" fn unary_next_callback( } type UnaryFuture = Pin> + Send>>; -type UnaryAdapter = - dyn Fn(Json, Option>, CompletionRef) -> UnaryFuture + Send + Sync; +type UnaryAdapter = dyn Fn( + Json, + Option>, + Option>, + CompletionRef, + ) -> UnaryFuture + + Send + + Sync; struct UnaryCallbackState { host: HostV4, @@ -554,6 +637,26 @@ unsafe extern "C" fn unary_trampoline( invocation_json: *const NemoRelayNativeString, next: *const NemoRelayNativeAsyncNext, completion: *const NemoRelayNativeAsyncCompletion, +) -> u32 { + unsafe { unary_trampoline_impl(user_data, invocation_json, ptr::null(), next, completion) } +} + +unsafe extern "C" fn unary_llm_execution_trampoline( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + context: *const NemoRelayNativeLlmExecutionContext, + next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, +) -> u32 { + unsafe { unary_trampoline_impl(user_data, invocation_json, context, next, completion) } +} + +unsafe fn unary_trampoline_impl( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + context: *const NemoRelayNativeLlmExecutionContext, + next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, ) -> u32 { let state = unsafe { &*user_data.cast::() }; let completion = Completion { @@ -572,10 +675,13 @@ unsafe extern "C" fn unary_trampoline( }); let invocation = read_json_value(&state.host.0.v3.v1, invocation_json, "async invocation") .map_err(|status| format!("invalid async invocation: {status:?}")); + let context = (!context.is_null()) + .then(|| llm_execution_context_from_completion(completion_ref, unsafe { &*context })) + .transpose(); let binding = ScopePollBinding::capture(state.host.0.v3.v1); - let future = catch_unwind(AssertUnwindSafe(|| match invocation { - Ok(invocation) => (state.adapter)(invocation, next, completion_ref), - Err(error) => Box::pin(async move { Err(error) }) as UnaryFuture, + let future = catch_unwind(AssertUnwindSafe(|| match (invocation, context) { + (Ok(invocation), Ok(context)) => (state.adapter)(invocation, context, next, completion_ref), + (Err(error), _) | (_, Err(error)) => Box::pin(async move { Err(error) }) as UnaryFuture, })); if let Err(error) = state.executor.ensure_started() { completion.reject(&error); @@ -850,10 +956,11 @@ struct EventMetadataInvocation { } type StreamFuture = Pin> + Send>>; -type StreamAdapter = dyn Fn(Json, LlmStreamNext) -> StreamFuture + Send + Sync; +type StreamAdapter = + dyn Fn(Json, LlmExecutionContext<'static>, LlmStreamNext) -> StreamFuture + Send + Sync; struct StreamCallbackState { - host: HostV4, + host: HostV7, executor: Arc, adapter: Box, } @@ -865,7 +972,7 @@ unsafe extern "C" fn drop_stream_callback(user_data: *mut c_void) { } struct OutputStream { - host: HostV4, + host: HostV7, raw: *const NemoRelayNativeAsyncStream, } @@ -874,18 +981,19 @@ unsafe impl Sync for OutputStream {} impl OutputStream { fn cancelled(&self) -> bool { - unsafe { (self.host.0.v3.async_stream_is_cancelled)(self.raw) } + unsafe { (self.host.0.v6.v5.v4.v3.async_stream_is_cancelled)(self.raw) } } async fn push(&self, value: &Json) -> Result<()> { - let value = HostString::from_json(&self.host.0.v3.v1, value) + let value = HostString::from_json(&self.host.0.v6.v5.v4.v3.v1, value) .ok_or_else(|| "failed to serialize native stream chunk".to_string())?; loop { if self.cancelled() { return Err("native stream consumer cancelled".into()); } - let status = - unsafe { (self.host.0.v3.async_stream_push_json)(self.raw, value.as_ptr()) }; + let status = unsafe { + (self.host.0.v6.v5.v4.v3.async_stream_push_json)(self.raw, value.as_ptr()) + }; match status { NemoRelayStatus::Ok => return Ok(()), NemoRelayStatus::Backpressured => { @@ -898,19 +1006,20 @@ impl OutputStream { fn finish(&self) -> Result<()> { status_result( - unsafe { (self.host.0.v3.async_stream_finish)(self.raw) }, + unsafe { (self.host.0.v6.v5.v4.v3.async_stream_finish)(self.raw) }, "finish native stream", ) } async fn reject(&self, error: &str) { - if let Some(error) = HostString::new(&self.host.0.v3.v1, error) { + if let Some(error) = HostString::new(&self.host.0.v6.v5.v4.v3.v1, error) { loop { if self.cancelled() { break; } - let status = - unsafe { (self.host.0.v3.async_stream_reject)(self.raw, error.as_ptr()) }; + let status = unsafe { + (self.host.0.v6.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()) + }; match status { NemoRelayStatus::Backpressured => { tokio::time::sleep(CANCELLATION_POLL_INTERVAL).await; @@ -922,9 +1031,9 @@ impl OutputStream { } fn reject_once(&self, error: &str) { - if let Some(error) = HostString::new(&self.host.0.v3.v1, error) { + if let Some(error) = HostString::new(&self.host.0.v6.v5.v4.v3.v1, error) { unsafe { - (self.host.0.v3.async_stream_reject)(self.raw, error.as_ptr()); + (self.host.0.v6.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()); } } } @@ -932,13 +1041,14 @@ impl OutputStream { impl Drop for OutputStream { fn drop(&mut self) { - unsafe { (self.host.0.v3.async_stream_release)(self.raw) }; + unsafe { (self.host.0.v6.v5.v4.v3.async_stream_release)(self.raw) }; } } unsafe extern "C" fn stream_trampoline( user_data: *mut c_void, invocation_json: *const NemoRelayNativeString, + context: *const NemoRelayNativeLlmExecutionContext, next: *const NemoRelayNativeAsyncNext, stream: *const NemoRelayNativeAsyncStream, ) -> u32 { @@ -952,24 +1062,39 @@ unsafe extern "C" fn stream_trampoline( return NemoRelayNativeAsyncCallbackState::Pending as u32; } let next = LlmStreamNext(Arc::new(NextInner { - host: state.host, + host: HostV4(state.host.0.v6.v5.v4), raw: next, })); - let invocation = read_json_value(&state.host.0.v3.v1, invocation_json, "stream invocation") - .map_err(|status| format!("invalid native stream invocation: {status:?}")); - let bindings = ScopePollBinding::capture(state.host.0.v3.v1).and_then(|future| { - ScopePollBinding::capture(state.host.0.v3.v1).map(|stream| (future, stream)) + let invocation = read_json_value( + &state.host.0.v6.v5.v4.v3.v1, + invocation_json, + "stream invocation", + ) + .map_err(|status| format!("invalid native stream invocation: {status:?}")); + let context = unsafe { context.as_ref() } + .ok_or_else(|| "native LLM stream execution context was null".to_string()) + .and_then(|context| { + llm_stream_execution_context_from_native( + StreamRef { + host: state.host, + raw: stream, + }, + context, + ) + }); + let bindings = ScopePollBinding::capture(state.host.0.v6.v5.v4.v3.v1).and_then(|future| { + ScopePollBinding::capture(state.host.0.v6.v5.v4.v3.v1).map(|stream| (future, stream)) }); - let future = catch_unwind(AssertUnwindSafe(|| match invocation { - Ok(invocation) => (state.adapter)(invocation, next), - Err(error) => Box::pin(async move { Err(error) }) as StreamFuture, + let future = catch_unwind(AssertUnwindSafe(|| match (invocation, context) { + (Ok(invocation), Ok(context)) => (state.adapter)(invocation, context, next), + (Err(error), _) | (_, Err(error)) => Box::pin(async move { Err(error) }) as StreamFuture, })); let future = future.unwrap_or_else(|_| { Box::pin(async move { Err("typed native stream callback panicked".into()) }) }); if let Err(error) = state.executor.ensure_started() { output.reject_once(&error); - set_last_error(&state.host.0.v3.v1, &error); + set_last_error(&state.host.0.v6.v5.v4.v3.v1, &error); return NemoRelayNativeAsyncCallbackState::Pending as u32; } let task = async move { @@ -1031,7 +1156,7 @@ unsafe extern "C" fn stream_trampoline( } }; if let Err(error) = state.executor.spawn(task) { - set_last_error(&state.host.0.v3.v1, &error); + set_last_error(&state.host.0.v6.v5.v4.v3.v1, &error); } NemoRelayNativeAsyncCallbackState::Pending as u32 } @@ -1074,6 +1199,68 @@ impl CodecIdentityInvocation { } } +fn execution_codec_identity( + host: &NemoRelayNativeHostApiV1, + kind: NemoRelayNativeLlmCodecKind, + codec_id: *const NemoRelayNativeString, +) -> Result { + let id = (!codec_id.is_null()) + .then(|| read_host_string(host, codec_id).map_err(|_| "invalid LLM codec ID".to_string())) + .transpose()?; + match (kind, id) { + (NemoRelayNativeLlmCodecKind::None, _) => Ok(LlmCodecIdentity::None), + (NemoRelayNativeLlmCodecKind::Opaque, _) => Ok(LlmCodecIdentity::Opaque), + (NemoRelayNativeLlmCodecKind::BuiltIn, Some(id)) => BuiltinLlmCodec::from_id(&id) + .map(LlmCodecIdentity::BuiltIn) + .ok_or_else(|| format!("unknown built-in LLM codec: {id}")), + (NemoRelayNativeLlmCodecKind::Runtime, Some(id)) => Ok(LlmCodecIdentity::Runtime(id)), + (kind, None) => Err(format!("missing LLM codec ID for {kind:?}")), + } +} + +fn llm_execution_context_from_completion( + completion: CompletionRef, + context: &NemoRelayNativeLlmExecutionContext, +) -> Result> { + let request = context.request_codec; + let host = &completion.host.0.v3.v1; + let request_codec = completion.execution_request_context( + execution_codec_identity(host, request.codec_kind, request.codec_id)?, + !request.codec.is_null(), + )?; + let response_codec = unsafe { context.response_codec.as_ref() } + .map(|response| { + completion.execution_response_context( + execution_codec_identity(host, response.codec_kind, response.codec_id)?, + !response.codec.is_null(), + ) + }) + .transpose()?; + Ok(LlmExecutionContext { + request_codec, + response_codec, + }) +} + +fn llm_stream_execution_context_from_native( + stream: StreamRef, + context: &NemoRelayNativeLlmExecutionContext, +) -> Result> { + if !context.response_codec.is_null() { + return Err("native LLM stream execution context exposed a response codec".into()); + } + let request = context.request_codec; + let host = &stream.host.0.v6.v5.v4.v3.v1; + let request_codec = stream.execution_request_context( + execution_codec_identity(host, request.codec_kind, request.codec_id)?, + !request.codec.is_null(), + )?; + Ok(LlmExecutionContext { + request_codec, + response_codec: None, + }) +} + impl PluginContext<'_> { fn host_v4(&self) -> Result { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL @@ -1086,6 +1273,17 @@ impl PluginContext<'_> { })) } + fn host_v7(&self) -> Result { + if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT + || self.host.struct_size < std::mem::size_of::() + { + return Err("typed LLM execution middleware requires Relay ABI v7".into()); + } + Ok(HostV7(unsafe { + *(self.host as *const _ as *const NemoRelayNativeHostApiV7) + })) + } + fn register_unary_adapter( &mut self, kind: NemoRelayNativeAsyncMiddlewareKind, @@ -1121,6 +1319,33 @@ impl PluginContext<'_> { } } + fn register_llm_execution_adapter( + &mut self, + name: &str, + priority: i32, + adapter: Box, + ) -> Result<()> { + let state = Box::into_raw(Box::new(UnaryCallbackState { + host: self.host_v4()?, + executor: Arc::clone(&self.executor), + adapter, + })); + let status = unsafe { + self.register_async_llm_execution_intercept_raw( + name, + priority, + unary_llm_execution_trampoline, + state.cast(), + Some(drop_unary_callback), + ) + }; + if status == NemoRelayStatus::Ok { + Ok(()) + } else { + Err(status_message(self.host, status, "LLM execution intercept")) + } + } + fn register_event_adapter( &mut self, kind: NemoRelayNativeAsyncMiddlewareKind, @@ -1138,7 +1363,7 @@ impl PluginContext<'_> { name, priority, false, - Box::new(move |value, _, _| { + Box::new(move |value, _, _, _| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: EventInvocation = serde_json::from_value(value) @@ -1169,7 +1394,7 @@ impl PluginContext<'_> { name, priority, false, - Box::new(move |value, _, _| { + Box::new(move |value, _, _, _| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: EventMetadataInvocation = serde_json::from_value(value) @@ -1258,7 +1483,7 @@ impl PluginContext<'_> { name, priority, break_chain, - Box::new(move |value, _, _| { + Box::new(move |value, _, _, _| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: NameValueInvocation = @@ -1326,7 +1551,7 @@ impl PluginContext<'_> { name, priority, false, - Box::new(move |value, _, _| { + Box::new(move |value, _, _, _| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: NameValueInvocation = @@ -1376,7 +1601,7 @@ impl PluginContext<'_> { name, priority, false, - Box::new(move |value, next, _| { + Box::new(move |value, _, next, _| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: ToolExecutionInvocation = @@ -1413,7 +1638,7 @@ impl PluginContext<'_> { name, priority, false, - Box::new(move |value, _, completion| { + Box::new(move |value, _, _, completion| { let callback = Arc::clone(&callback); Box::pin(async move { #[derive(Deserialize)] @@ -1448,7 +1673,7 @@ impl PluginContext<'_> { name, priority, false, - Box::new(move |value, _, completion| { + Box::new(move |value, _, _, completion| { let callback = Arc::clone(&callback); Box::pin(async move { #[derive(Deserialize)] @@ -1483,7 +1708,7 @@ impl PluginContext<'_> { name, priority, false, - Box::new(move |value, _, _| { + Box::new(move |value, _, _, _| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: RequestInvocation = @@ -1513,7 +1738,7 @@ impl PluginContext<'_> { name, priority, break_chain, - Box::new(move |value, _, _| { + Box::new(move |value, _, _, _| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: LlmRequestInterceptInvocation = @@ -1535,24 +1760,27 @@ impl PluginContext<'_> { callback: F, ) -> Result<()> where - F: Fn(String, LlmRequest, LlmNext) -> Fut + Send + Sync + 'static, + F: Fn(String, LlmRequest, LlmExecutionContext<'static>, LlmNext) -> Fut + + Send + + Sync + + 'static, Fut: Future> + Send + 'static, { let callback = Arc::new(callback); - self.register_unary_adapter( - NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept, + self.register_llm_execution_adapter( name, priority, - false, - Box::new(move |value, next, _| { + Box::new(move |value, context, next, _| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: NameRequestInvocation = serde_json::from_value(value).map_err(|error| error.to_string())?; + let context = context + .ok_or_else(|| "native LLM execution context was null".to_string())?; let next = LlmNext( next.ok_or_else(|| "LLM execution continuation was null".to_string())?, ); - callback(invocation.name, invocation.request, next).await + callback(invocation.name, invocation.request, context, next).await }) }), ) @@ -1566,19 +1794,22 @@ impl PluginContext<'_> { callback: F, ) -> Result<()> where - F: Fn(String, LlmRequest, LlmStreamNext) -> Fut + Send + Sync + 'static, + F: Fn(String, LlmRequest, LlmExecutionContext<'static>, LlmStreamNext) -> Fut + + Send + + Sync + + 'static, Fut: Future> + Send + 'static, { let callback = Arc::new(callback); let state = Box::into_raw(Box::new(StreamCallbackState { - host: self.host_v4()?, + host: self.host_v7()?, executor: Arc::clone(&self.executor), - adapter: Box::new(move |value, next| { + adapter: Box::new(move |value, context, next| { let callback = Arc::clone(&callback); Box::pin(async move { let invocation: NameRequestInvocation = serde_json::from_value(value).map_err(|error| error.to_string())?; - callback(invocation.name, invocation.request, next).await + callback(invocation.name, invocation.request, context, next).await }) }), })); diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index 369e60f75..c10f74983 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -50,10 +50,14 @@ use serde_json::Map; /// Native plugin ABI version supported by this crate. /// -/// Version 6 adds host-routed operational logging for native plugins. -/// Hosts retain frozen version-4, version-3, and version-2 tables for -/// already-built plugins that target those layouts. -pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 6; +/// Version 7 makes LLM execution intercept callbacks context-aware. +/// +/// This is an intentional callback-layout break. Native plugins must rebuild +/// against this Relay release even though their authored `native_api` +/// compatibility label remains `1`. +pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 7; +/// ABI version that introduced uniform LLM execution codec context. +pub const NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT: u32 = 7; /// ABI version that introduced host-routed operational logging. pub const NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING: u32 = 6; /// ABI version that introduced context-aware raw tool execution intercepts. @@ -86,6 +90,44 @@ pub struct LlmSanitizeResponseContext<'a> { // completion capability; callback-scoped native contexts never construct it. unsafe impl Send for LlmSanitizeResponseContext<'_> {} +/// Per-call codec context delivered to an LLM execution intercept. +/// +/// Request codec context is always available. Unary execution also supplies a +/// response codec context, while streaming execution leaves it unavailable +/// until Relay has a completed-response streaming codec contract. +pub struct LlmExecutionContext<'a> { + request_codec: LlmExecutionRequestContext<'a>, + response_codec: Option>, +} + +/// Request codec context for one LLM execution intercept invocation. +pub struct LlmExecutionRequestContext<'a> { + /// Identity of the active request codec. + pub codec: LlmCodecIdentity, + resolved: Option>, +} + +/// Response codec context for one unary LLM execution intercept invocation. +pub struct LlmExecutionResponseContext<'a> { + /// Identity of the active response codec. + pub codec: LlmCodecIdentity, + resolved: Option>, +} + +impl<'a> LlmExecutionContext<'a> { + /// Return the active request codec context. + #[must_use] + pub fn request_codec(&self) -> &LlmExecutionRequestContext<'a> { + &self.request_codec + } + + /// Return the unary response codec context, or `None` for streaming execution. + #[must_use] + pub fn response_codec(&self) -> Option<&LlmExecutionResponseContext<'a>> { + self.response_codec.as_ref() + } +} + /// Status codes returned by stable native ABI functions. #[repr(i32)] #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -180,6 +222,40 @@ pub struct NemoRelayNativeLlmSanitizeResponseContext { pub codec: *const NemoRelayNativeLlmResponseCodec, } +/// Request codec context passed to a native LLM execution intercept. +#[repr(C)] +#[derive(Debug, Clone, Copy)] +pub struct NemoRelayNativeLlmExecutionRequestContext { + /// Discriminator for the active request codec. + pub codec_kind: NemoRelayNativeLlmCodecKind, + /// Optional borrowed built-in or runtime codec identifier. + pub codec_id: *const NemoRelayNativeString, + /// Borrowed request codec capability, or null when no codec is active. + pub codec: *const NemoRelayNativeLlmRequestCodec, +} + +/// Response codec context passed to a native unary LLM execution intercept. +#[repr(C)] +#[derive(Debug, Clone, Copy)] +pub struct NemoRelayNativeLlmExecutionResponseContext { + /// Discriminator for the active response codec. + pub codec_kind: NemoRelayNativeLlmCodecKind, + /// Optional borrowed built-in or runtime codec identifier. + pub codec_id: *const NemoRelayNativeString, + /// Borrowed response codec capability, or null when no codec is active. + pub codec: *const NemoRelayNativeLlmResponseCodec, +} + +/// Codec context passed to native LLM execution intercept callbacks. +#[repr(C)] +#[derive(Debug, Clone, Copy)] +pub struct NemoRelayNativeLlmExecutionContext { + /// Request codec context, always present. + pub request_codec: NemoRelayNativeLlmExecutionRequestContext, + /// Unary response codec context, or null for streaming execution. + pub response_codec: *const NemoRelayNativeLlmExecutionResponseContext, +} + /// Safe completion-backed request codec facade for typed native plugins. pub struct LlmSanitizeRequestCodec<'a> { async_host: NemoRelayNativeHostApiV4, @@ -273,6 +349,160 @@ impl LlmSanitizeResponseCodec<'_> { } } +enum LlmExecutionRequestCodecOwner { + Completion { + host: NemoRelayNativeHostApiV4, + completion: *const NemoRelayNativeAsyncCompletion, + }, + Stream { + host: NemoRelayNativeHostApiV7, + stream: *const NemoRelayNativeAsyncStream, + }, +} + +/// Invocation-lifetime request codec facade for an LLM execution intercept. +pub struct LlmExecutionRequestCodec<'a> { + owner: LlmExecutionRequestCodecOwner, + _lifetime: PhantomData<&'a NemoRelayNativeLlmRequestCodec>, +} + +// SAFETY: construction retains the completion or stream that owns the codec. +// Host operations are thread-safe and reject calls after that owner settles. +unsafe impl Send for LlmExecutionRequestCodec<'_> {} +unsafe impl Sync for LlmExecutionRequestCodec<'_> {} + +impl Drop for LlmExecutionRequestCodec<'_> { + fn drop(&mut self) { + match self.owner { + LlmExecutionRequestCodecOwner::Completion { host, completion } => unsafe { + (host.v3.async_completion_release)(completion) + }, + LlmExecutionRequestCodecOwner::Stream { host, stream } => unsafe { + (host.v6.v5.v4.v3.async_stream_release)(stream) + }, + } + } +} + +impl LlmExecutionRequestCodec<'_> { + /// Decode an opaque request into Relay's normalized request model. + pub fn decode(&self, request: &LlmRequest) -> Result { + match self.owner { + LlmExecutionRequestCodecOwner::Completion { host, completion } => { + native_codec_call(&host.v3.v1, |out| unsafe { + let request = HostString::from_json(&host.v3.v1, request) + .ok_or_else(|| "failed to serialize LLM request".to_string())?; + let status = (host.async_completion_llm_request_codec_decode)( + completion, + request.as_ptr(), + out, + ); + codec_status(&host.v3.v1, status) + }) + } + LlmExecutionRequestCodecOwner::Stream { host, stream } => { + native_codec_call(&host.v6.v5.v4.v3.v1, |out| unsafe { + let request = HostString::from_json(&host.v6.v5.v4.v3.v1, request) + .ok_or_else(|| "failed to serialize LLM request".to_string())?; + let status = + (host.async_stream_llm_request_codec_decode)(stream, request.as_ptr(), out); + codec_status(&host.v6.v5.v4.v3.v1, status) + }) + } + } + } + + /// Encode normalized changes onto the original opaque request. + pub fn encode( + &self, + annotated: &AnnotatedLlmRequest, + original: &LlmRequest, + ) -> Result { + match self.owner { + LlmExecutionRequestCodecOwner::Completion { host, completion } => { + native_codec_call(&host.v3.v1, |out| unsafe { + let annotated = HostString::from_json(&host.v3.v1, annotated) + .ok_or_else(|| "failed to serialize annotated request".to_string())?; + let original = HostString::from_json(&host.v3.v1, original) + .ok_or_else(|| "failed to serialize original request".to_string())?; + let status = (host.async_completion_llm_request_codec_encode)( + completion, + annotated.as_ptr(), + original.as_ptr(), + out, + ); + codec_status(&host.v3.v1, status) + }) + } + LlmExecutionRequestCodecOwner::Stream { host, stream } => { + native_codec_call(&host.v6.v5.v4.v3.v1, |out| unsafe { + let annotated = HostString::from_json(&host.v6.v5.v4.v3.v1, annotated) + .ok_or_else(|| "failed to serialize annotated request".to_string())?; + let original = HostString::from_json(&host.v6.v5.v4.v3.v1, original) + .ok_or_else(|| "failed to serialize original request".to_string())?; + let status = (host.async_stream_llm_request_codec_encode)( + stream, + annotated.as_ptr(), + original.as_ptr(), + out, + ); + codec_status(&host.v6.v5.v4.v3.v1, status) + }) + } + } + } +} + +/// Invocation-lifetime response codec facade for a unary LLM execution intercept. +pub struct LlmExecutionResponseCodec<'a> { + host: NemoRelayNativeHostApiV4, + completion: *const NemoRelayNativeAsyncCompletion, + _lifetime: PhantomData<&'a NemoRelayNativeLlmResponseCodec>, +} + +// SAFETY: construction retains the completion that owns the codec. Host +// operations are thread-safe and reject calls after that completion settles. +unsafe impl Send for LlmExecutionResponseCodec<'_> {} +unsafe impl Sync for LlmExecutionResponseCodec<'_> {} + +impl Drop for LlmExecutionResponseCodec<'_> { + fn drop(&mut self) { + unsafe { (self.host.v3.async_completion_release)(self.completion) }; + } +} + +impl LlmExecutionResponseCodec<'_> { + /// Decode an opaque response into Relay's normalized response model. + pub fn decode(&self, response: &Json) -> Result { + native_codec_call(&self.host.v3.v1, |out| unsafe { + let response = HostString::from_json(&self.host.v3.v1, response) + .ok_or_else(|| "failed to serialize LLM response".to_string())?; + let status = (self.host.async_completion_llm_response_codec_decode)( + self.completion, + response.as_ptr(), + out, + ); + codec_status(&self.host.v3.v1, status) + }) + } +} + +impl LlmExecutionRequestContext<'_> { + /// Resolve the active request codec capability. + #[must_use] + pub fn resolve_codec(&self) -> Option<&LlmExecutionRequestCodec<'_>> { + self.resolved.as_ref() + } +} + +impl LlmExecutionResponseContext<'_> { + /// Resolve the active response codec capability. + #[must_use] + pub fn resolve_codec(&self) -> Option<&LlmExecutionResponseCodec<'_>> { + self.resolved.as_ref() + } +} + impl<'a> LlmSanitizeRequestContext<'a> { /// Resolve the active request codec capability. #[must_use] @@ -545,6 +775,7 @@ pub type NemoRelayNativeLlmExecutionCb = unsafe extern "C" fn( user_data: *mut c_void, name: *const NemoRelayNativeString, request_json: *const NemoRelayNativeString, + context: NemoRelayNativeLlmExecutionContext, next_fn: NemoRelayNativeLlmNextFn, next_ctx: *mut c_void, out_json: *mut *mut NemoRelayNativeString, @@ -555,6 +786,7 @@ pub type NemoRelayNativeLlmStreamExecutionCb = unsafe extern "C" fn( user_data: *mut c_void, name: *const NemoRelayNativeString, request_json: *const NemoRelayNativeString, + context: NemoRelayNativeLlmExecutionContext, next_fn: NemoRelayNativeLlmStreamNextFn, next_ctx: *mut c_void, out_stream: *mut NemoRelayNativeLlmStreamV1, @@ -1015,6 +1247,7 @@ pub type NemoRelayNativeAsyncNextResultCb = unsafe extern "C" fn( pub type NemoRelayNativeAsyncStreamMiddlewareCb = unsafe extern "C" fn( user_data: *mut c_void, invocation_json: *const NemoRelayNativeString, + context: *const NemoRelayNativeLlmExecutionContext, next: *const NemoRelayNativeAsyncNext, stream: *const NemoRelayNativeAsyncStream, ) -> u32; @@ -1047,6 +1280,20 @@ pub type NemoRelayNativeAsyncMiddlewareCb = unsafe extern "C" fn( completion: *const NemoRelayNativeAsyncCompletion, ) -> u32; +/// Completion-based native LLM execution callback. +/// +/// `invocation_json` and `context` are borrowed for the call. The callback may +/// retain the invocation-scoped codec handles by returning `Pending`; they +/// remain valid until the completion settles. The callback owns `next` and +/// must release it exactly once after its final use. +pub type NemoRelayNativeAsyncLlmExecutionCb = unsafe extern "C" fn( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + context: *const NemoRelayNativeLlmExecutionContext, + next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, +) -> u32; + /// ABI-v3 host extension appended to [`NemoRelayNativeHostApiV1`]. /// /// Its first field is the complete v1/v2 table, so legacy plugins can keep @@ -1377,6 +1624,49 @@ pub struct NemoRelayNativeHostApiV6 { ) -> NemoRelayStatus, } +/// ABI-v7 host table for context-aware LLM execution intercept callbacks. +/// +/// The inherited function-pointer table has the same fields as ABI v6, but +/// its LLM execution callback typedefs include +/// [`NemoRelayNativeLlmExecutionContext`]. The distinct table version prevents +/// either side from invoking a callback compiled with the old argument layout. +#[repr(C)] +#[derive(Clone, Copy)] +pub struct NemoRelayNativeHostApiV7 { + /// ABI-v6 table compiled with the ABI-v7 callback typedefs. + pub v6: NemoRelayNativeHostApiV6, + /// Registers a completion-based asynchronous LLM execution intercept. + pub plugin_context_register_async_llm_execution_intercept: + unsafe extern "C" fn( + ctx: *mut NemoRelayNativePluginContext, + name: *const NemoRelayNativeString, + priority: i32, + cb: NemoRelayNativeAsyncLlmExecutionCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, + ) -> NemoRelayStatus, + /// Retains an output stream for an execution codec facade. + pub async_stream_retain: + unsafe extern "C" fn(stream: *const NemoRelayNativeAsyncStream) -> NemoRelayStatus, + /// Decodes an LLM request through the request codec owned by `stream`. + /// + /// The operation fails after the output stream settles or is cancelled. + pub async_stream_llm_request_codec_decode: unsafe extern "C" fn( + stream: *const NemoRelayNativeAsyncStream, + request_json: *const NemoRelayNativeString, + out: *mut *mut NemoRelayNativeString, + ) -> NemoRelayStatus, + /// Encodes an annotated request through the request codec owned by `stream`. + /// + /// The operation fails after the output stream settles or is cancelled. + pub async_stream_llm_request_codec_encode: unsafe extern "C" fn( + stream: *const NemoRelayNativeAsyncStream, + annotated_json: *const NemoRelayNativeString, + original_json: *const NemoRelayNativeString, + out: *mut *mut NemoRelayNativeString, + ) -> NemoRelayStatus, +} + unsafe impl Send for NemoRelayNativeHostApiV3 {} unsafe impl Sync for NemoRelayNativeHostApiV3 {} // SAFETY: the v4 host table is immutable after construction. Its function @@ -1390,6 +1680,9 @@ unsafe impl Sync for NemoRelayNativeHostApiV5 {} // SAFETY: the v6 host table is immutable and its log function is thread-safe. unsafe impl Send for NemoRelayNativeHostApiV6 {} unsafe impl Sync for NemoRelayNativeHostApiV6 {} +// SAFETY: the v7 table is immutable and contains only thread-safe host functions. +unsafe impl Send for NemoRelayNativeHostApiV7 {} +unsafe impl Sync for NemoRelayNativeHostApiV7 {} // The host API table is immutable after construction. Function pointers and // the null-terminated version string pointer are safe to share across threads. @@ -2885,6 +3178,7 @@ impl<'a> PluginContext<'a> { /// `cb`, `user_data`, and `free_fn` must remain valid for every host /// callback invocation until the host deregisters the callback or calls /// `free_fn`. `free_fn` must match the allocation behind `user_data`. + /// Registration consumes that ownership even when ABI v7 is unavailable. pub unsafe fn register_llm_execution_intercept_raw( &mut self, name: &str, @@ -2893,6 +3187,14 @@ impl<'a> PluginContext<'a> { user_data: *mut c_void, free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { + if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT + || self.host.struct_size < std::mem::size_of::() + { + if let Some(free_fn) = free_fn { + unsafe { free_fn(user_data) }; + } + return NemoRelayStatus::InvalidArg; + } self.with_name_and_callback(name, user_data, free_fn, |host, name| unsafe { (host.plugin_context_register_llm_execution_intercept)( self.raw, name, priority, cb, user_data, free_fn, @@ -2906,6 +3208,7 @@ impl<'a> PluginContext<'a> { /// `cb`, `user_data`, and `free_fn` must remain valid for every host /// callback invocation until the host deregisters the callback or calls /// `free_fn`. `free_fn` must match the allocation behind `user_data`. + /// Registration consumes that ownership even when ABI v7 is unavailable. pub unsafe fn register_llm_stream_execution_intercept_raw( &mut self, name: &str, @@ -2914,6 +3217,14 @@ impl<'a> PluginContext<'a> { user_data: *mut c_void, free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { + if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT + || self.host.struct_size < std::mem::size_of::() + { + if let Some(free_fn) = free_fn { + unsafe { free_fn(user_data) }; + } + return NemoRelayStatus::InvalidArg; + } self.with_name_and_callback(name, user_data, free_fn, |host, name| unsafe { (host.plugin_context_register_llm_stream_execution_intercept)( self.raw, name, priority, cb, user_data, free_fn, @@ -2945,6 +3256,16 @@ impl<'a> PluginContext<'a> { user_data: *mut c_void, free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { + if matches!( + kind, + NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept + | NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept + ) { + if let Some(free_fn) = free_fn { + unsafe { free_fn(user_data) }; + } + return NemoRelayStatus::InvalidArg; + } if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE || self.host.struct_size < std::mem::size_of::() { @@ -2968,6 +3289,38 @@ impl<'a> PluginContext<'a> { }) } + /// Registers a completion-based asynchronous LLM execution intercept. + /// + /// # Safety + /// `cb`, `user_data`, and `free_fn` must remain valid until the host + /// deregisters the callback or invokes `free_fn`. This call consumes the + /// `user_data` ownership even when it rejects the host ABI. A callback + /// returning `Pending` must settle and release its completion and `next` + /// references exactly once. + pub unsafe fn register_async_llm_execution_intercept_raw( + &mut self, + name: &str, + priority: i32, + cb: NemoRelayNativeAsyncLlmExecutionCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, + ) -> NemoRelayStatus { + if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT + || self.host.struct_size < std::mem::size_of::() + { + if let Some(free_fn) = free_fn { + unsafe { free_fn(user_data) }; + } + return NemoRelayStatus::InvalidArg; + } + let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV7) }; + self.with_name_and_callback(name, user_data, free_fn, |_, name| unsafe { + (host.plugin_context_register_async_llm_execution_intercept)( + self.raw, name, priority, cb, user_data, free_fn, + ) + }) + } + /// Registers an incremental completion-based LLM stream intercept. /// /// # Safety @@ -2978,7 +3331,9 @@ impl<'a> PluginContext<'a> { /// Retry only [`NemoRelayStatus::Backpressured`] operations. The output /// stream owns the callback lifetime. `next` may be invoked /// repeatedly or concurrently until that stream settles; Relay then - /// rejects or cancels unfinished and later calls. + /// rejects or cancels unfinished and later calls. This execution-specific + /// callback requires native ABI v7 because its callback context and + /// stream-scoped request-codec operations are part of that ABI. pub unsafe fn register_async_stream_middleware_raw( &mut self, name: &str, @@ -2987,17 +3342,22 @@ impl<'a> PluginContext<'a> { 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::() + if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } return NemoRelayStatus::InvalidArg; } - let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV3) }; + let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV7) }; self.with_name_and_callback(name, user_data, free_fn, |_, name| unsafe { - (host.plugin_context_register_async_stream_middleware)( + (host + .v6 + .v5 + .v4 + .v3 + .plugin_context_register_async_stream_middleware)( self.raw, name, priority, cb, user_data, free_fn, ) }) @@ -3238,11 +3598,21 @@ enum OwnedHostApi { V3(NemoRelayNativeHostApiV3), V4(NemoRelayNativeHostApiV4), V5(NemoRelayNativeHostApiV5), + V6(NemoRelayNativeHostApiV6), + V7(NemoRelayNativeHostApiV7), } impl OwnedHostApi { unsafe fn copy_from(host: &NemoRelayNativeHostApiV1) -> Self { - if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT + if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT + && host.struct_size >= std::mem::size_of::() + { + Self::V7(unsafe { *(host as *const _ as *const NemoRelayNativeHostApiV7) }) + } else if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING + && host.struct_size >= std::mem::size_of::() + { + Self::V6(unsafe { *(host as *const _ as *const NemoRelayNativeHostApiV6) }) + } else if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT && host.struct_size >= std::mem::size_of::() { Self::V5(unsafe { *(host as *const _ as *const NemoRelayNativeHostApiV5) }) @@ -3265,6 +3635,8 @@ impl OwnedHostApi { Self::V3(host) => &host.v1, Self::V4(host) => &host.v3.v1, Self::V5(host) => &host.v4.v3.v1, + Self::V6(host) => &host.v5.v4.v3.v1, + Self::V7(host) => &host.v6.v5.v4.v3.v1, } } } @@ -3546,8 +3918,7 @@ where P: NativePlugin, F: FnOnce() -> P, { - let supported_abi = (NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY..=NEMO_RELAY_NATIVE_ABI_VERSION) - .contains(&host_ref.abi_version); + let supported_abi = host_ref.abi_version == NEMO_RELAY_NATIVE_ABI_VERSION; if !supported_abi { return NemoRelayStatus::InvalidArg; } diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index ae6d44a78..bb8dd6e80 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -26,17 +26,20 @@ use nemo_relay_plugin::{ LlmJsonAsyncStream, LlmJsonStream, LlmNext, LlmRequest, LlmRequestInterceptOutcome, LlmStream, LlmStreamNext, LogSeverity, MetricKind, MetricMeasurement, MetricValueType, NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE, - NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, - NativeExecutorConfig, NativePlugin, NemoRelayNativeAsyncCallbackState, - NemoRelayNativeAsyncCompletion, NemoRelayNativeAsyncLlmStreamOpenCb, + NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, + NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, NativeExecutorConfig, NativePlugin, + NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncCompletion, + NemoRelayNativeAsyncLlmExecutionCb, NemoRelayNativeAsyncLlmStreamOpenCb, NemoRelayNativeAsyncLlmStreamPullCb, NemoRelayNativeAsyncMiddlewareCb, NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, NemoRelayNativeAsyncNextResultCb, NemoRelayNativeAsyncNextStreamCb, NemoRelayNativeAsyncStream, NemoRelayNativeAsyncStreamMiddlewareCb, NemoRelayNativeConditionalMiddlewareCb, NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, - NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeLlmAsyncStream, - NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, + NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, + NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, + NemoRelayNativeLlmExecutionRequestContext, NemoRelayNativeLlmExecutionResponseContext, NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, @@ -337,6 +340,22 @@ impl RegisteredAsync { } } +struct RegisteredAsyncLlmExecution { + name: String, + priority: i32, + cb: NemoRelayNativeAsyncLlmExecutionCb, + user_data: usize, + free_fn: NemoRelayNativeFreeFn, +} + +impl RegisteredAsyncLlmExecution { + unsafe fn free(self) { + if let Some(free_fn) = self.free_fn { + unsafe { free_fn(self.user_data as *mut c_void) }; + } + } +} + struct RegisteredAsyncStream { name: String, priority: i32, @@ -391,6 +410,7 @@ impl_captured_registration!( RegisteredLlmStreamExecution, RegisteredLlmRequestIntercept, RegisteredAsync, + RegisteredAsyncLlmExecution, RegisteredAsyncStream, ); @@ -455,6 +475,8 @@ static LLM_STREAM_EXECUTION_REGISTRATION: Mutex> = Mutex::new(None); static ASYNC_REGISTRATIONS: Mutex> = Mutex::new(Vec::new()); +static ASYNC_LLM_EXECUTION_REGISTRATION: Mutex> = + Mutex::new(None); static ASYNC_STREAM_REGISTRATION: Mutex> = Mutex::new(None); static ASYNC_PUSH_BACKPRESSURE: AtomicUsize = AtomicUsize::new(0); static ASYNC_COMPLETION_RETAINS: AtomicUsize = AtomicUsize::new(0); @@ -466,7 +488,7 @@ static UNAVAILABLE_CONTEXT_GATE_CALLS: AtomicUsize = AtomicUsize::new(0); #[test] fn native_abi_struct_sizes_are_self_describing() { - assert_eq!(NEMO_RELAY_NATIVE_ABI_VERSION, 6); + assert_eq!(NEMO_RELAY_NATIVE_ABI_VERSION, 7); assert_eq!( size_of::(), test_host().struct_size @@ -541,6 +563,7 @@ fn assert_native_abi_platform_layout() { assert_type_layout::(8, 616); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, log), 608); + assert_native_abi_v7_layout(8, 648, 616, 624, 632, 640); assert_type_layout::(8, 56); assert_eq!(plugin_offsets(), [0, 8, 16, 24, 32, 40, 48]); assert_type_layout::(8, 40); @@ -601,6 +624,10 @@ fn assert_native_abi_platform_layout() { ), 296 ); + assert_type_layout::(4, 304); + assert_eq!(offset_of!(NemoRelayNativeHostApiV6, v5), 0); + assert_eq!(offset_of!(NemoRelayNativeHostApiV6, log), 300); + assert_native_abi_v7_layout(4, 320, 304, 308, 312, 316); assert_type_layout::(4, 28); assert_eq!(plugin_offsets(), [0, 4, 8, 12, 16, 20, 24]); assert_type_layout::(4, 20); @@ -612,6 +639,43 @@ fn assert_type_layout(expected_alignment: usize, expected_size: usize) { assert_eq!(size_of::(), expected_size); } +fn assert_native_abi_v7_layout( + expected_alignment: usize, + expected_size: usize, + registration_offset: usize, + retain_offset: usize, + decode_offset: usize, + encode_offset: usize, +) { + assert_type_layout::(expected_alignment, expected_size); + assert_eq!(offset_of!(NemoRelayNativeHostApiV7, v6), 0); + assert_eq!( + offset_of!( + NemoRelayNativeHostApiV7, + plugin_context_register_async_llm_execution_intercept + ), + registration_offset + ); + assert_eq!( + offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), + retain_offset + ); + assert_eq!( + offset_of!( + NemoRelayNativeHostApiV7, + async_stream_llm_request_codec_decode + ), + decode_offset + ); + assert_eq!( + offset_of!( + NemoRelayNativeHostApiV7, + async_stream_llm_request_codec_encode + ), + encode_offset + ); +} + #[test] fn native_abi_v4_extension_is_append_only() { #[cfg(target_pointer_width = "64")] @@ -675,6 +739,22 @@ fn native_abi_v6_logging_extension_is_append_only() { ); } +#[test] +fn native_abi_v7_execution_context_extension_preserves_the_host_table() { + assert_eq!(offset_of!(NemoRelayNativeHostApiV7, v6), 0); + assert_eq!( + offset_of!( + NemoRelayNativeHostApiV7, + plugin_context_register_async_llm_execution_intercept + ), + size_of::() + ); + assert_eq!( + offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), + size_of::() + size_of::() + ); +} + fn host_api_v4_offsets() -> [usize; 12] { [ offset_of!(NemoRelayNativeHostApiV4, v3), @@ -1038,6 +1118,7 @@ unsafe extern "C" fn passthrough_llm_execution_cb( _user_data: *mut c_void, _name: *const NemoRelayNativeString, _request_json: *const NemoRelayNativeString, + _context: NemoRelayNativeLlmExecutionContext, _next_fn: nemo_relay_plugin::NemoRelayNativeLlmNextFn, _next_ctx: *mut c_void, _out_json: *mut *mut NemoRelayNativeString, @@ -1049,6 +1130,7 @@ unsafe extern "C" fn passthrough_llm_stream_execution_cb( _user_data: *mut c_void, _name: *const NemoRelayNativeString, _request_json: *const NemoRelayNativeString, + _context: NemoRelayNativeLlmExecutionContext, _next_fn: nemo_relay_plugin::NemoRelayNativeLlmStreamNextFn, _next_ctx: *mut c_void, _out_stream: *mut NemoRelayNativeLlmStreamV1, @@ -1068,6 +1150,7 @@ unsafe extern "C" fn pending_async_middleware_cb( unsafe extern "C" fn pending_async_stream_middleware_cb( _user_data: *mut c_void, _invocation_json: *const NemoRelayNativeString, + _context: *const NemoRelayNativeLlmExecutionContext, _next: *const NemoRelayNativeAsyncNext, _stream: *const NemoRelayNativeAsyncStream, ) -> u32 { @@ -1863,6 +1946,8 @@ struct MockAsyncOutput { events: Mutex>, events_cv: Condvar, cancelled: AtomicBool, + settled: AtomicBool, + retains: AtomicUsize, releases: AtomicUsize, } @@ -1872,6 +1957,8 @@ impl MockAsyncOutput { events: Mutex::new(Vec::new()), events_cv: Condvar::new(), cancelled: AtomicBool::new(false), + settled: AtomicBool::new(false), + retains: AtomicUsize::new(0), releases: AtomicUsize::new(0), } } @@ -2082,6 +2169,38 @@ unsafe extern "C" fn capture_register_async_middleware( NemoRelayStatus::Ok } +unsafe extern "C" fn capture_register_async_llm_execution( + _ctx: *mut NemoRelayNativePluginContext, + name: *const NemoRelayNativeString, + priority: i32, + cb: NemoRelayNativeAsyncLlmExecutionCb, + user_data: *mut c_void, + free_fn: NemoRelayNativeFreeFn, +) -> NemoRelayStatus { + let status = *REGISTRATION_STATUS.lock().unwrap(); + if status != NemoRelayStatus::Ok { + if let Some(free_fn) = free_fn { + unsafe { free_fn(user_data) }; + } + return status; + } + let name = match required_host_string(&test_host(), name) { + Ok(name) => name, + Err(status) => return status, + }; + replace_registration( + &ASYNC_LLM_EXECUTION_REGISTRATION, + RegisteredAsyncLlmExecution { + name, + priority, + cb, + user_data: user_data as usize, + free_fn, + }, + ); + NemoRelayStatus::Ok +} + unsafe extern "C" fn capture_async_stream_push( stream: *const NemoRelayNativeAsyncStream, chunk_json: *const NemoRelayNativeString, @@ -2115,6 +2234,7 @@ unsafe extern "C" fn capture_async_stream_finish( return NemoRelayStatus::NullPointer; } let stream = unsafe { &*stream.cast::() }; + stream.settled.store(true, Ordering::SeqCst); stream .events .lock() @@ -2132,6 +2252,7 @@ unsafe extern "C" fn capture_async_stream_reject( return NemoRelayStatus::NullPointer; } let stream = unsafe { &*stream.cast::() }; + stream.settled.store(true, Ordering::SeqCst); let message = read_host_string(&test_host(), message).unwrap(); stream .events @@ -2165,6 +2286,53 @@ unsafe extern "C" fn capture_async_stream_release(stream: *const NemoRelayNative } } +unsafe extern "C" fn capture_async_stream_retain( + stream: *const NemoRelayNativeAsyncStream, +) -> NemoRelayStatus { + if stream.is_null() { + return NemoRelayStatus::NullPointer; + } + unsafe { &*stream.cast::() } + .retains + .fetch_add(1, Ordering::SeqCst); + NemoRelayStatus::Ok +} + +fn mock_async_stream_is_active(stream: *const NemoRelayNativeAsyncStream) -> bool { + let Some(stream) = (unsafe { stream.cast::().as_ref() }) else { + return false; + }; + !stream.cancelled.load(Ordering::SeqCst) && !stream.settled.load(Ordering::SeqCst) +} + +unsafe extern "C" fn capture_async_stream_request_decode( + stream: *const NemoRelayNativeAsyncStream, + _request_json: *const NemoRelayNativeString, + out: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + if !mock_async_stream_is_active(stream) { + return NemoRelayStatus::InvalidArg; + } + write_json(&test_host(), &json!({}), out) +} + +unsafe extern "C" fn capture_async_stream_request_encode( + stream: *const NemoRelayNativeAsyncStream, + _annotated_json: *const NemoRelayNativeString, + original_json: *const NemoRelayNativeString, + out: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + if !mock_async_stream_is_active(stream) { + return NemoRelayStatus::InvalidArg; + } + if out.is_null() || original_json.is_null() { + return NemoRelayStatus::NullPointer; + } + let original = read_host_string(&test_host(), original_json).unwrap(); + let value = serde_json::from_str(&original).unwrap(); + write_json(&test_host(), &value, out) +} + unsafe extern "C" fn unavailable_async_next_stream( _next: *const NemoRelayNativeAsyncNext, _invocation_json: *const NemoRelayNativeString, @@ -2240,19 +2408,37 @@ unsafe extern "C" fn capture_async_next_result( } unsafe extern "C" fn capture_async_request_decode( - _completion: *const NemoRelayNativeAsyncCompletion, + completion: *const NemoRelayNativeAsyncCompletion, _request_json: *const NemoRelayNativeString, out: *mut *mut NemoRelayNativeString, ) -> NemoRelayStatus { + if completion.is_null() + || unsafe { &*completion.cast::() } + .settled + .lock() + .unwrap() + .is_some() + { + return NemoRelayStatus::InvalidArg; + } write_json(&test_host(), &json!({}), out) } unsafe extern "C" fn capture_async_request_encode( - _completion: *const NemoRelayNativeAsyncCompletion, + completion: *const NemoRelayNativeAsyncCompletion, _annotated_json: *const NemoRelayNativeString, original_json: *const NemoRelayNativeString, out: *mut *mut NemoRelayNativeString, ) -> NemoRelayStatus { + if completion.is_null() + || unsafe { &*completion.cast::() } + .settled + .lock() + .unwrap() + .is_some() + { + return NemoRelayStatus::InvalidArg; + } if out.is_null() || original_json.is_null() { return NemoRelayStatus::NullPointer; } @@ -2262,10 +2448,19 @@ unsafe extern "C" fn capture_async_request_encode( } unsafe extern "C" fn capture_async_response_decode( - _completion: *const NemoRelayNativeAsyncCompletion, + completion: *const NemoRelayNativeAsyncCompletion, _response_json: *const NemoRelayNativeString, out: *mut *mut NemoRelayNativeString, ) -> NemoRelayStatus { + if completion.is_null() + || unsafe { &*completion.cast::() } + .settled + .lock() + .unwrap() + .is_some() + { + return NemoRelayStatus::InvalidArg; + } write_json(&test_host(), &json!({}), out) } @@ -2830,7 +3025,7 @@ unsafe extern "C" fn capture_plugin_log( fn test_host_v6() -> NemoRelayNativeHostApiV6 { let mut v5 = test_host_v5(); - v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION; + v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING; v5.v4.v3.v1.struct_size = size_of::(); NemoRelayNativeHostApiV6 { v5, @@ -2838,6 +3033,19 @@ fn test_host_v6() -> NemoRelayNativeHostApiV6 { } } +fn test_host_v7() -> NemoRelayNativeHostApiV7 { + let mut v6 = test_host_v6(); + v6.v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION; + v6.v5.v4.v3.v1.struct_size = size_of::(); + NemoRelayNativeHostApiV7 { + v6, + plugin_context_register_async_llm_execution_intercept: capture_register_async_llm_execution, + async_stream_retain: capture_async_stream_retain, + async_stream_llm_request_codec_decode: capture_async_stream_request_decode, + async_stream_llm_request_codec_encode: capture_async_stream_request_encode, + } +} + fn test_llm_request() -> LlmRequest { LlmRequest { headers: Map::new(), @@ -2886,6 +3094,66 @@ fn invoke_async_registration( result } +fn invoke_async_llm_execution_registration( + host: &NemoRelayNativeHostApiV4, + registration: &RegisteredAsyncLlmExecution, + invocation: Json, + next: &MockAsyncNext, +) -> std::result::Result { + let completion = MockAsyncCompletion::new(); + let invocation = json_host_string(&host.v3.v1, invocation); + let execution_context = TestExecutionContext::new(true); + let state = unsafe { + (registration.cb)( + registration.user_data as *mut c_void, + invocation, + execution_context.as_ptr(), + next.raw(), + completion.raw(), + ) + }; + unsafe { (host.v3.v1.string_free)(invocation) }; + assert_eq!( + NemoRelayNativeAsyncCallbackState::try_from(state), + Ok(NemoRelayNativeAsyncCallbackState::Pending) + ); + let result = completion.wait(); + completion.wait_for_release(); + assert!(completion.releases.load(Ordering::SeqCst) >= 1); + result +} + +struct TestExecutionContext { + response: Option>, + context: NemoRelayNativeLlmExecutionContext, +} + +impl TestExecutionContext { + fn new(unary: bool) -> Self { + let request_codec = NonNull::::dangling().as_ptr(); + let response = unary.then(|| { + Box::new(NemoRelayNativeLlmExecutionResponseContext { + codec_kind: NemoRelayNativeLlmCodecKind::Opaque, + codec_id: ptr::null(), + codec: NonNull::::dangling().as_ptr(), + }) + }); + let context = NemoRelayNativeLlmExecutionContext { + request_codec: NemoRelayNativeLlmExecutionRequestContext { + codec_kind: NemoRelayNativeLlmCodecKind::Opaque, + codec_id: ptr::null(), + codec: request_codec, + }, + response_codec: response.as_deref().map_or(ptr::null(), ptr::from_ref), + }; + Self { response, context } + } + + fn as_ptr(&self) -> *const NemoRelayNativeLlmExecutionContext { + ptr::from_ref(&self.context) + } +} + fn begin_test() -> MutexGuard<'static, ()> { let guard = TEST_LOCK .lock() @@ -2903,6 +3171,7 @@ fn reset_state() { for registration in ASYNC_REGISTRATIONS.lock().unwrap().drain(..) { unsafe { registration.free() }; } + clear_registration(&ASYNC_LLM_EXECUTION_REGISTRATION); clear_registration(&ASYNC_STREAM_REGISTRATION); clear_registration(&SUBSCRIBER_REGISTRATION); clear_registration(&EVENT_SANITIZE_REGISTRATION); @@ -4139,8 +4408,9 @@ fn typed_subscriber_registration_decodes_events() { #[allow(clippy::cognitive_complexity)] // One table-style test deliberately exercises every surface. fn typed_async_middleware_registers_and_round_trips_every_surface() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let host_v4 = &host.v6.v5.v4; + let mut ctx = test_context(&host.v6.v5.v4.v3.v1); ctx.register_mark_sanitize_guardrail("mark-async", 1, |_event, mut fields| async move { tokio::time::sleep(Duration::from_millis(1)).await; @@ -4244,13 +4514,21 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { ctx.register_llm_execution_intercept( "llm-execution-async", 13, - |_name, request, next| async move { next.call(request).await }, + |_name, request, context, next| async move { + assert_eq!(context.request_codec().codec, LlmCodecIdentity::Opaque); + assert!(context.request_codec().resolve_codec().is_some()); + assert!(context.response_codec().is_some()); + next.call(request).await + }, ) .unwrap(); ctx.register_llm_stream_execution_intercept( "llm-stream-async", 14, - |_name, request, next| async move { + |_name, request, context, next| async move { + assert_eq!(context.request_codec().codec, LlmCodecIdentity::Opaque); + assert!(context.request_codec().resolve_codec().is_some()); + assert!(context.response_codec().is_none()); let stream = next.call(request).await?; let transformed = stream.map(|item| item.map(|chunk| json!({ "wrapped": chunk }))); Ok(Box::pin(transformed) as LlmJsonAsyncStream) @@ -4278,7 +4556,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { ) }) .collect::>(); - assert_eq!(metadata.len(), 14); + assert_eq!(metadata.len(), 13); assert_eq!(metadata[0].1, "mark-async"); assert_eq!(metadata[0].2, 1); assert_eq!( @@ -4291,6 +4569,14 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept ); assert!(metadata[11].3); + { + let registration = ASYNC_LLM_EXECUTION_REGISTRATION.lock().unwrap(); + let registration = registration.as_ref().unwrap(); + assert_eq!( + (registration.name.as_str(), registration.priority), + ("llm-execution-async", 13) + ); + } { let stream_registration = ASYNC_STREAM_REGISTRATION.lock().unwrap(); let stream_registration = stream_registration.as_ref().unwrap(); @@ -4313,7 +4599,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { let metadata_registration = take_async_registration(NemoRelayNativeAsyncMiddlewareKind::EventMetadataInjector); let result = invoke_async_registration( - &host, + host_v4, &metadata_registration, json!({ "event": event }), None, @@ -4329,7 +4615,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { ] { let registration = take_async_registration(kind); let result = invoke_async_registration( - &host, + host_v4, ®istration, json!({ "event": event, "fields": { "data": { "visible": true } } }), None, @@ -4349,7 +4635,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { ] { let registration = take_async_registration(kind); let result = invoke_async_registration( - &host, + host_v4, ®istration, json!({ "name": "calculator", "value": { "x": 1 } }), None, @@ -4363,7 +4649,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { take_async_registration(NemoRelayNativeAsyncMiddlewareKind::ToolConditionalExecution); assert_eq!( invoke_async_registration( - &host, + host_v4, ®istration, json!({ "name": "calculator", "value": {} }), None, @@ -4390,7 +4676,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { let registration = take_async_registration(NemoRelayNativeAsyncMiddlewareKind::ToolExecutionIntercept); let outcome = invoke_async_registration( - &host, + host_v4, ®istration, json!({ "name": "calculator", @@ -4409,7 +4695,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { take_async_registration(NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeRequest); assert_eq!( invoke_async_registration( - &host, + host_v4, ®istration, json!({ "request": request, @@ -4427,7 +4713,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { take_async_registration(NemoRelayNativeAsyncMiddlewareKind::LlmSanitizeResponse); assert_eq!( invoke_async_registration( - &host, + host_v4, ®istration, json!({ "response": { "answer": 42 }, @@ -4445,7 +4731,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { take_async_registration(NemoRelayNativeAsyncMiddlewareKind::LlmConditionalExecution); assert_eq!( invoke_async_registration( - &host, + host_v4, ®istration, json!({ "request": test_llm_request() }), None @@ -4458,7 +4744,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { let registration = take_async_registration(NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept); let result = invoke_async_registration( - &host, + host_v4, ®istration, json!({ "name": "provider", "request": test_llm_request(), "annotated": null }), None, @@ -4467,14 +4753,17 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { assert_eq!(result["request"]["headers"]["x-tested"], true); unsafe { registration.free() }; - let registration = - take_async_registration(NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept); + let registration = ASYNC_LLM_EXECUTION_REGISTRATION + .lock() + .unwrap() + .take() + .unwrap(); assert_eq!( - invoke_async_registration( - &host, + invoke_async_llm_execution_registration( + host_v4, ®istration, json!({ "name": "provider", "request": test_llm_request() }), - Some(&next), + &next, ) .unwrap()["content"], json!({ "prompt": "hello" }) @@ -4485,18 +4774,20 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { let output = MockAsyncOutput::new(); ASYNC_PUSH_BACKPRESSURE.store(2, Ordering::SeqCst); let invocation = json_host_string( - &host.v3.v1, + &host_v4.v3.v1, json!({ "name": "provider", "request": test_llm_request() }), ); + let execution_context = TestExecutionContext::new(false); let state = unsafe { (stream_registration.cb)( stream_registration.user_data as *mut c_void, invocation, + execution_context.as_ptr(), next.raw(), output.raw(), ) }; - unsafe { (host.v3.v1.string_free)(invocation) }; + unsafe { (host_v4.v3.v1.string_free)(invocation) }; assert_eq!( NemoRelayNativeAsyncCallbackState::try_from(state), Ok(NemoRelayNativeAsyncCallbackState::Pending) @@ -4510,7 +4801,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { ] ); output.wait_for_release(); - assert_eq!(output.releases.load(Ordering::SeqCst), 1); + assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(pull_stream.releases.load(Ordering::SeqCst), 1); assert_eq!(next.releases.load(Ordering::SeqCst), 3); unsafe { stream_registration.free() }; @@ -4521,6 +4812,168 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { assert_eq!(live_host_strings(), 0); } +#[test] +fn typed_async_unary_execution_codecs_expire_after_completion_settles() { + let _guard = begin_test(); + let host = test_host_v7(); + let host_v4 = &host.v6.v5.v4; + let mut ctx = test_context(&host_v4.v3.v1); + let (context_tx, context_rx) = std::sync::mpsc::sync_channel(1); + ctx.register_llm_execution_intercept( + "retain-unary-context", + 0, + move |_name, _request, context, _next| { + context_tx.send(context).unwrap(); + async move { Ok(json!({"answer": 42})) } + }, + ) + .unwrap(); + + let registration = ASYNC_LLM_EXECUTION_REGISTRATION + .lock() + .unwrap() + .take() + .unwrap(); + let completion = MockAsyncCompletion::new(); + let next = MockAsyncNext { + calls: AtomicUsize::new(0), + releases: AtomicUsize::new(0), + pull_stream: ptr::null(), + }; + let invocation = json_host_string( + &host_v4.v3.v1, + json!({"name": "provider", "request": test_llm_request()}), + ); + let execution_context = TestExecutionContext::new(true); + let state = unsafe { + (registration.cb)( + registration.user_data as *mut c_void, + invocation, + execution_context.as_ptr(), + next.raw(), + completion.raw(), + ) + }; + unsafe { (host_v4.v3.v1.string_free)(invocation) }; + assert_eq!( + NemoRelayNativeAsyncCallbackState::try_from(state), + Ok(NemoRelayNativeAsyncCallbackState::Pending) + ); + assert_eq!(completion.wait().unwrap(), json!({"answer": 42})); + completion.wait_for_release(); + + let context = context_rx.recv_timeout(Duration::from_secs(5)).unwrap(); + assert_eq!(ASYNC_COMPLETION_RETAINS.load(Ordering::SeqCst), 2); + assert!( + context + .request_codec() + .resolve_codec() + .unwrap() + .decode(&test_llm_request()) + .is_err() + ); + assert!( + context + .response_codec() + .unwrap() + .resolve_codec() + .unwrap() + .decode(&json!({"answer": 42})) + .is_err() + ); + drop(context); + let deadline = Instant::now() + Duration::from_secs(5); + while completion.releases.load(Ordering::SeqCst) < 3 { + assert!( + Instant::now() < deadline, + "retained codec facades were not released" + ); + std::thread::yield_now(); + } + assert_eq!(next.releases.load(Ordering::SeqCst), 1); + unsafe { registration.free() }; +} + +#[test] +fn typed_async_stream_execution_codec_expires_after_stream_finishes() { + let _guard = begin_test(); + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; + let mut ctx = test_context(host_v1); + let (context_tx, context_rx) = std::sync::mpsc::sync_channel(1); + ctx.register_llm_stream_execution_intercept( + "retain-stream-context", + 0, + move |_name, _request, context, _next| { + context_tx.send(context).unwrap(); + async move { + Ok( + Box::pin(futures::stream::iter(vec![Ok(json!({"chunk": 1}))])) + as LlmJsonAsyncStream, + ) + } + }, + ) + .unwrap(); + + let registration = ASYNC_STREAM_REGISTRATION.lock().unwrap().take().unwrap(); + let output = MockAsyncOutput::new(); + let next = MockAsyncNext { + calls: AtomicUsize::new(0), + releases: AtomicUsize::new(0), + pull_stream: ptr::null(), + }; + let invocation = json_host_string( + host_v1, + json!({"name": "provider", "request": test_llm_request()}), + ); + let execution_context = TestExecutionContext::new(false); + let state = unsafe { + (registration.cb)( + registration.user_data as *mut c_void, + invocation, + execution_context.as_ptr(), + next.raw(), + output.raw(), + ) + }; + unsafe { (host_v1.string_free)(invocation) }; + assert_eq!( + NemoRelayNativeAsyncCallbackState::try_from(state), + Ok(NemoRelayNativeAsyncCallbackState::Pending) + ); + assert_eq!( + output.wait_terminal(), + vec![ + MockOutputEvent::Chunk(json!({"chunk": 1})), + MockOutputEvent::Finished, + ] + ); + + let context = context_rx.recv_timeout(Duration::from_secs(5)).unwrap(); + assert_eq!(output.retains.load(Ordering::SeqCst), 1); + assert!(context.response_codec().is_none()); + assert!( + context + .request_codec() + .resolve_codec() + .unwrap() + .decode(&test_llm_request()) + .is_err() + ); + drop(context); + let deadline = Instant::now() + Duration::from_secs(5); + while output.releases.load(Ordering::SeqCst) < 2 { + assert!( + Instant::now() < deadline, + "retained stream codec facade was not released" + ); + std::thread::yield_now(); + } + assert_eq!(next.releases.load(Ordering::SeqCst), 1); + unsafe { registration.free() }; +} + #[test] fn typed_async_llm_sanitize_context_decodes_oci_genai_builtin_identity() { let _guard = begin_test(); @@ -4651,8 +5104,8 @@ fn typed_async_llm_sanitize_context_rejects_unknown_builtin_identity() { #[test] fn typed_async_registration_failure_rolls_back_callback_state() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let mut ctx = test_context(&host.v6.v5.v4.v3.v1); *REGISTRATION_STATUS.lock().unwrap() = NemoRelayStatus::InvalidArg; let unary_drops = Arc::new(AtomicUsize::new(0)); @@ -4670,7 +5123,7 @@ fn typed_async_registration_failure_rolls_back_callback_state() { let result = ctx.register_llm_stream_execution_intercept( "rejected-stream", 0, - move |_name, _request, _next| { + move |_name, _request, _context, _next| { let _ = &stream_probe; async move { Ok(Box::pin(futures::stream::empty()) as LlmJsonAsyncStream) } }, @@ -4748,8 +5201,9 @@ fn typed_async_callbacks_isolate_errors_panics_and_invalid_input() { #[test] fn typed_async_continuations_are_concurrent_and_executor_owned() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let host_v4 = &host.v6.v5.v4; + let mut ctx = test_context(&host_v4.v3.v1); ctx.register_tool_execution_intercept("concurrent", 0, |context, next| async move { assert_eq!(context.tool_name, "tool"); assert_eq!(context.tool_call_id, None); @@ -4779,7 +5233,7 @@ fn typed_async_continuations_are_concurrent_and_executor_owned() { let registration = take_async_registration(NemoRelayNativeAsyncMiddlewareKind::ToolExecutionIntercept); let result = invoke_async_registration( - &host, + host_v4, ®istration, json!({ "name": "tool", "value": {} }), Some(&next), @@ -4804,7 +5258,7 @@ fn typed_async_continuations_are_concurrent_and_executor_owned() { ctx.register_llm_stream_execution_intercept( "stream-open-error", 0, - |_name, request, next| async move { next.call(request).await }, + |_name, request, _context, next| async move { next.call(request).await }, ) .unwrap(); let registration = ASYNC_STREAM_REGISTRATION.lock().unwrap().take().unwrap(); @@ -4815,18 +5269,20 @@ fn typed_async_continuations_are_concurrent_and_executor_owned() { pull_stream: ptr::null(), }; let invocation = json_host_string( - &host.v3.v1, + &host_v4.v3.v1, json!({ "name": "provider", "request": test_llm_request() }), ); + let execution_context = TestExecutionContext::new(false); unsafe { (registration.cb)( registration.user_data as *mut c_void, invocation, + execution_context.as_ptr(), next.raw(), output.raw(), ) }; - unsafe { (host.v3.v1.string_free)(invocation) }; + unsafe { (host_v4.v3.v1.string_free)(invocation) }; assert_eq!( output.wait_terminal(), vec![MockOutputEvent::Rejected( @@ -4834,7 +5290,7 @@ fn typed_async_continuations_are_concurrent_and_executor_owned() { )] ); output.wait_for_release(); - assert_eq!(output.releases.load(Ordering::SeqCst), 1); + assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; } @@ -5064,12 +5520,13 @@ fn typed_async_executor_drop_inside_tokio_runtime_drains_accepted_tasks() { #[test] fn typed_async_stream_cancellation_while_polling_releases_output() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; + let mut ctx = test_context(host_v1); let started = Arc::new(AtomicBool::new(false)); ctx.register_llm_stream_execution_intercept("cancel-poll", 0, { let started = Arc::clone(&started); - move |_name, _request, _next| { + move |_name, _request, _context, _next| { let started = Arc::clone(&started); async move { Ok(Box::pin(futures::stream::poll_fn(move |_| { @@ -5088,18 +5545,20 @@ fn typed_async_stream_cancellation_while_polling_releases_output() { pull_stream: ptr::null(), }; let invocation = json_host_string( - &host.v3.v1, + host_v1, json!({ "name": "provider", "request": test_llm_request() }), ); + let execution_context = TestExecutionContext::new(false); unsafe { (registration.cb)( registration.user_data as *mut c_void, invocation, + execution_context.as_ptr(), next.raw(), output.raw(), ) }; - unsafe { (host.v3.v1.string_free)(invocation) }; + unsafe { (host_v1.string_free)(invocation) }; let deadline = Instant::now() + Duration::from_secs(5); while !started.load(Ordering::SeqCst) { assert!(Instant::now() < deadline, "returned stream was not polled"); @@ -5114,12 +5573,13 @@ fn typed_async_stream_cancellation_while_polling_releases_output() { #[test] fn typed_async_stream_restores_callback_scope_while_polling_returned_stream() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; + let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-scope", 0, - |_name, _request, _next| async move { + |_name, _request, _context, _next| async move { Ok(Box::pin(futures::stream::iter([Ok(json!({ "chunk": 1 }))])) as LlmJsonAsyncStream) }, ) @@ -5132,18 +5592,20 @@ fn typed_async_stream_restores_callback_scope_while_polling_returned_stream() { pull_stream: ptr::null(), }; let invocation = json_host_string( - &host.v3.v1, + host_v1, json!({ "name": "provider", "request": test_llm_request() }), ); + let execution_context = TestExecutionContext::new(false); unsafe { (registration.cb)( registration.user_data as *mut c_void, invocation, + execution_context.as_ptr(), next.raw(), output.raw(), ) }; - unsafe { (host.v3.v1.string_free)(invocation) }; + unsafe { (host_v1.string_free)(invocation) }; assert_eq!( output.wait_terminal(), vec![ @@ -5159,12 +5621,13 @@ fn typed_async_stream_restores_callback_scope_while_polling_returned_stream() { #[test] fn typed_async_stream_rejects_item_errors_and_releases_output() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; + let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-error", 0, - |_name, _request, _next| async move { + |_name, _request, _context, _next| async move { Ok(Box::pin(futures::stream::iter(vec![ Ok(json!({ "chunk": 1 })), Err("stream item failed".into()), @@ -5180,18 +5643,20 @@ fn typed_async_stream_rejects_item_errors_and_releases_output() { pull_stream: ptr::null(), }; let invocation = json_host_string( - &host.v3.v1, + host_v1, json!({ "name": "provider", "request": test_llm_request() }), ); + let execution_context = TestExecutionContext::new(false); unsafe { (registration.cb)( registration.user_data as *mut c_void, invocation, + execution_context.as_ptr(), next.raw(), output.raw(), ) }; - unsafe { (host.v3.v1.string_free)(invocation) }; + unsafe { (host_v1.string_free)(invocation) }; assert_eq!( output.wait_terminal(), vec![ @@ -5200,7 +5665,7 @@ fn typed_async_stream_rejects_item_errors_and_releases_output() { ] ); output.wait_for_release(); - assert_eq!(output.releases.load(Ordering::SeqCst), 1); + assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; } @@ -5208,12 +5673,13 @@ fn typed_async_stream_rejects_item_errors_and_releases_output() { #[test] fn typed_async_stream_rejects_poll_panics_and_releases_output() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; + let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-panic", 0, - |_name, _request, _next| async move { + |_name, _request, _context, _next| async move { let mut polled = false; Ok(Box::pin(futures::stream::poll_fn(move |_| { if polled { @@ -5233,18 +5699,20 @@ fn typed_async_stream_rejects_poll_panics_and_releases_output() { pull_stream: ptr::null(), }; let invocation = json_host_string( - &host.v3.v1, + host_v1, json!({ "name": "provider", "request": test_llm_request() }), ); + let execution_context = TestExecutionContext::new(false); unsafe { (registration.cb)( registration.user_data as *mut c_void, invocation, + execution_context.as_ptr(), next.raw(), output.raw(), ) }; - unsafe { (host.v3.v1.string_free)(invocation) }; + unsafe { (host_v1.string_free)(invocation) }; assert_eq!( output.wait_terminal(), vec![ @@ -5253,7 +5721,7 @@ fn typed_async_stream_rejects_poll_panics_and_releases_output() { ] ); output.wait_for_release(); - assert_eq!(output.releases.load(Ordering::SeqCst), 1); + assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; } @@ -5261,12 +5729,13 @@ fn typed_async_stream_rejects_poll_panics_and_releases_output() { #[test] fn typed_async_stream_propagates_downstream_pull_errors() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; + let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-downstream-error", 0, - |_name, request, next| async move { next.call(request).await }, + |_name, request, _context, next| async move { next.call(request).await }, ) .unwrap(); let registration = ASYNC_STREAM_REGISTRATION.lock().unwrap().take().unwrap(); @@ -5282,24 +5751,26 @@ fn typed_async_stream_propagates_downstream_pull_errors() { pull_stream: pull_stream.raw(), }; let invocation = json_host_string( - &host.v3.v1, + host_v1, json!({ "name": "provider", "request": test_llm_request() }), ); + let execution_context = TestExecutionContext::new(false); unsafe { (registration.cb)( registration.user_data as *mut c_void, invocation, + execution_context.as_ptr(), next.raw(), output.raw(), ) }; - unsafe { (host.v3.v1.string_free)(invocation) }; + unsafe { (host_v1.string_free)(invocation) }; assert_eq!( output.wait_terminal(), vec![MockOutputEvent::Rejected("downstream pull failed".into())] ); output.wait_for_release(); - assert_eq!(output.releases.load(Ordering::SeqCst), 1); + assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(pull_stream.releases.load(Ordering::SeqCst), 1); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; @@ -5308,12 +5779,13 @@ fn typed_async_stream_propagates_downstream_pull_errors() { #[test] fn typed_async_stream_rejects_missing_continuation() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; + let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-null-next", 0, - |_name, _request, _next| async move { + |_name, _request, _context, _next| async move { panic!("stream callback must not run without a continuation") }, ) @@ -5321,18 +5793,20 @@ fn typed_async_stream_rejects_missing_continuation() { let registration = ASYNC_STREAM_REGISTRATION.lock().unwrap().take().unwrap(); let output = MockAsyncOutput::new(); let invocation = json_host_string( - &host.v3.v1, + host_v1, json!({ "name": "provider", "request": test_llm_request() }), ); + let execution_context = TestExecutionContext::new(false); let state = unsafe { (registration.cb)( registration.user_data as *mut c_void, invocation, + execution_context.as_ptr(), ptr::null(), output.raw(), ) }; - unsafe { (host.v3.v1.string_free)(invocation) }; + unsafe { (host_v1.string_free)(invocation) }; assert_eq!( NemoRelayNativeAsyncCallbackState::try_from(state), Ok(NemoRelayNativeAsyncCallbackState::Pending) @@ -5437,8 +5911,8 @@ fn raw_event_sanitize_registrations_cover_every_surface() { #[test] fn raw_callback_registrations_preserve_every_middleware_shape() { let _guard = begin_test(); - let host = test_host_v5(); - let mut ctx = test_context(&host.v4.v3.v1); + let host = test_host_v7(); + let mut ctx = test_context(&host.v6.v5.v4.v3.v1); unsafe { assert_eq!( @@ -5578,10 +6052,10 @@ fn raw_callback_registrations_preserve_every_middleware_shape() { } #[test] -fn raw_async_callback_registrations_use_the_v3_extension_tables() { +fn raw_async_callback_registrations_use_the_versioned_extension_tables() { let _guard = begin_test(); - let host = test_host_v4(); - let mut ctx = test_context(&host.v3.v1); + let host = test_host_v7(); + let mut ctx = test_context(&host.v6.v5.v4.v3.v1); assert_eq!( unsafe { @@ -5778,10 +6252,7 @@ impl NativePlugin for RegisteringPlugin { ctx: &mut PluginContext<'_>, ) -> nemo_relay_plugin::Result<()> { assert_eq!(plugin_config.get("enabled"), Some(&json!(true))); - assert_eq!( - ctx.host_api().abi_version, - NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT - ); + assert_eq!(ctx.host_api().abi_version, NEMO_RELAY_NATIVE_ABI_VERSION); assert!(ctx.runtime().scope_stack_active()); ctx.register_subscriber("registered", |_event: &Event| {})?; let status = unsafe { @@ -6139,8 +6610,8 @@ fn exported_plugin_default_validate_returns_empty_diagnostics() { #[test] fn exported_plugin_register_installs_callbacks_and_propagates_errors() { let _guard = begin_test(); - let host = test_host_v5(); - let host_v1 = &host.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut plugin = NemoRelayNativePluginV1::default(); assert_eq!( @@ -6261,14 +6732,13 @@ fn exported_entry_symbol_validates_args_before_constructor() { } #[test] -fn exported_entry_symbol_accepts_supported_prior_host_versions() { +fn exported_entry_symbol_rejects_prior_host_versions() { let _guard = begin_test(); CONSTRUCTOR_CALLS.store(0, Ordering::SeqCst); for abi_version in [ NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE, - NEMO_RELAY_NATIVE_ABI_VERSION, ] { let mut host = test_host(); host.abi_version = abi_version; @@ -6276,12 +6746,20 @@ fn exported_entry_symbol_accepts_supported_prior_host_versions() { assert_eq!( unsafe { constructor_counting_entry(&host, &mut plugin) }, - NemoRelayStatus::Ok + NemoRelayStatus::InvalidArg ); - unsafe { drop_exported_plugin(&host, plugin) }; } - assert_eq!(CONSTRUCTOR_CALLS.load(Ordering::SeqCst), 3); + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; + let mut plugin = NemoRelayNativePluginV1::default(); + assert_eq!( + unsafe { constructor_counting_entry(host_v1, &mut plugin) }, + NemoRelayStatus::Ok + ); + unsafe { drop_exported_plugin(host_v1, plugin) }; + + assert_eq!(CONSTRUCTOR_CALLS.load(Ordering::SeqCst), 1); } #[test] diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index 5afe4c44a..a954344ce 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1716,8 +1716,8 @@ fn deregister_llm_request_intercept(name: &str) -> PyResult { /// Register an LLM execution intercept that can replace the LLM call. /// -/// ``callable``: ``async (native: Any, next) -> Any`` — middleware intercept function. -/// Call ``await next(native)`` to invoke the next intercept or original +/// ``callable``: ``async (name, request, context, next) -> Any`` — middleware intercept function. +/// Call ``await next(request)`` to invoke the next intercept or original /// implementation; skip calling ``next`` to short-circuit. #[pyfunction] fn register_llm_execution_intercept( @@ -1741,9 +1741,9 @@ fn deregister_llm_execution_intercept(name: &str) -> PyResult { /// Register an LLM stream-execution intercept that can replace the streaming LLM call. /// -/// ``callable``: ``async (native: Any, next) -> AsyncIterator[Any]`` — +/// ``callable``: ``async (name, request, context, next) -> AsyncIterator[Any]`` — /// middleware streaming intercept function. -/// Call ``await next(native)`` to invoke the next intercept or original +/// Call ``await next(request)`` to invoke the next intercept or original /// streaming implementation; skip calling ``next`` to short-circuit. #[pyfunction] fn register_llm_stream_execution_intercept( diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 481c640de..1f84e07a0 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -31,12 +31,12 @@ use nemo_relay::api::runtime::subscriber_dispatcher::{ }; use nemo_relay::api::runtime::{ EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, - LlmExecutionNextFn, LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, - LlmSanitizeRequestFn, LlmSanitizeResponseContext, LlmSanitizeResponseFn, - LlmStreamExecutionNextFn, LlmStreamInner, MiddlewareContinuationContext, PropagationContext, - ScopeStackHandle, ToolConditionalFn, ToolExecutionContext, ToolExecutionNextFn, - ToolInterceptFn, ToolSanitizeFn, capture_propagation_context, capture_traceparent, - current_scope_stack, + LlmExecutionContext, LlmExecutionFn, LlmExecutionNextFn, LlmJsonStream, LlmRequestInterceptFn, + LlmSanitizeRequestContext, LlmSanitizeRequestFn, LlmSanitizeResponseContext, + LlmSanitizeResponseFn, LlmStreamExecutionFn, LlmStreamExecutionNextFn, LlmStreamInner, + MiddlewareContinuationContext, PropagationContext, ScopeStackHandle, ToolConditionalFn, + ToolExecutionContext, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + capture_propagation_context, capture_traceparent, current_scope_stack, }; use nemo_relay::error::{FlowError, Result as FlowResult}; use pyo3::exceptions::PyRuntimeError; @@ -57,7 +57,7 @@ use nemo_relay::codec::traits::{LlmCodec, LlmResponseCodec}; use crate::convert::{json_to_py, py_to_json}; use crate::py_types::{ PyAnnotatedLLMRequest, PyAnnotatedLLMResponse, PyLLMRequest, PyLLMRequestInterceptOutcome, - PyLlmSanitizeRequestContext, PyLlmSanitizeResponseContext, PyScopeStack, + PyLlmExecutionContext, PyLlmSanitizeRequestContext, PyLlmSanitizeResponseContext, PyScopeStack, PyToolExecutionContext, PyToolExecutionInterceptOutcome, PyToolExecutionResult, TOOL_EXECUTION_INTERCEPT_RESULT_ERROR, }; @@ -1309,22 +1309,15 @@ pub fn wrap_py_tool_exec_intercept_fn( ) } -/// Wrap a Python callable `(name, LlmRequest, next) -> dict` for LLM execution intercepts. -pub fn wrap_py_llm_exec_intercept_fn( - py_fn: Py, -) -> Arc< - dyn Fn( - &str, - LlmRequest, - LlmExecutionNextFn, - ) -> Pin> + Send>> - + Send - + Sync, -> { +/// Wrap a Python callable `(name, request, context, next) -> dict` for LLM execution intercepts. +pub fn wrap_py_llm_exec_intercept_fn(py_fn: Py) -> LlmExecutionFn { let py_fn = Arc::new(py_fn); let task_locals = capture_python_task_locals(); Arc::new( - move |name: &str, request: LlmRequest, next: LlmExecutionNextFn| { + move |name: &str, + request: LlmRequest, + context: LlmExecutionContext, + next: LlmExecutionNextFn| { let py_fn = py_fn.clone(); let name = name.to_string(); let task_locals = task_locals_with_running_loop(task_locals.as_ref()); @@ -1334,6 +1327,7 @@ pub fn wrap_py_llm_exec_intercept_fn( copy_middleware_invocation(py, task_locals) .map_err(|error| FlowError::Internal(error.to_string()))?; let py_req = PyLLMRequest { inner: request }; + let py_context = PyLlmExecutionContext { inner: context }; let py_next = PyLlmNextFn { inner: next, context: MiddlewareContinuationContext::capture(), @@ -1350,10 +1344,13 @@ pub fn wrap_py_llm_exec_intercept_fn( loop_affine_callback(py, py_fn.bind(py), task_locals.as_ref(), false) .map_err(|error| FlowError::Internal(error.to_string()))?; let result = match invocation_context.as_ref() { - Some(context) => { - context.call_method1("run", (callback.bind(py), &name, py_req, py_next)) - } - None => callback.bind(py).call1((&name, py_req, py_next)), + Some(context) => context.call_method1( + "run", + (callback.bind(py), &name, py_req, py_context, py_next), + ), + None => callback + .bind(py) + .call1((&name, py_req, py_context, py_next)), } .map_err(|e: PyErr| FlowError::Internal(e.to_string()))?; split_py_object_or_future_with_locals( @@ -1373,28 +1370,22 @@ pub fn wrap_py_llm_exec_intercept_fn( ) } -/// Wrap a Python callable `(LlmRequest, next) -> AsyncIterator[Any]` for LLM +/// Wrap a Python callable `(name, request, context, next) -> AsyncIterator[Any]` for LLM /// stream execution intercepts. /// /// The Python callable may return the async iterator directly or return an /// awaitable that resolves to one. The resulting iterator is drained on the /// Tokio runtime and forwarded into a Rust `Stream>`. -pub fn wrap_py_llm_stream_exec_intercept_fn( - py_fn: Py, -) -> Arc< - dyn Fn( - &str, - LlmRequest, - LlmStreamExecutionNextFn, - ) -> Pin> + Send>> - + Send - + Sync, -> { +pub fn wrap_py_llm_stream_exec_intercept_fn(py_fn: Py) -> LlmStreamExecutionFn { let py_fn = Arc::new(py_fn); let task_locals = capture_python_task_locals(); Arc::new( - move |_name: &str, request: LlmRequest, next: LlmStreamExecutionNextFn| { + move |name: &str, + request: LlmRequest, + context: LlmExecutionContext, + next: LlmStreamExecutionNextFn| { let py_fn = py_fn.clone(); + let name = name.to_string(); let task_locals = task_locals_with_running_loop(task_locals.as_ref()); Box::pin(async move { let (outcome, invocation_task_locals) = Python::attach(|py| { @@ -1402,6 +1393,7 @@ pub fn wrap_py_llm_stream_exec_intercept_fn( copy_middleware_invocation(py, task_locals) .map_err(|error| FlowError::Internal(error.to_string()))?; let py_req = PyLLMRequest { inner: request }; + let py_context = PyLlmExecutionContext { inner: context }; let py_next = PyLlmStreamNextFn { inner: next, context: MiddlewareContinuationContext::capture(), @@ -1418,10 +1410,13 @@ pub fn wrap_py_llm_stream_exec_intercept_fn( loop_affine_callback(py, py_fn.bind(py), task_locals.as_ref(), false) .map_err(|error| FlowError::Internal(error.to_string()))?; let result = match invocation_context.as_ref() { - Some(context) => { - context.call_method1("run", (callback.bind(py), py_req, py_next)) - } - None => callback.bind(py).call1((py_req, py_next)), + Some(context) => context.call_method1( + "run", + (callback.bind(py), &name, py_req, py_context, py_next), + ), + None => callback + .bind(py) + .call1((&name, py_req, py_context, py_next)), } .map_err(python_callback_error)?; let outcome = split_py_object_or_future_with_locals( diff --git a/crates/python/src/py_types/core.rs b/crates/python/src/py_types/core.rs index 616176f2e..d514d506e 100644 --- a/crates/python/src/py_types/core.rs +++ b/crates/python/src/py_types/core.rs @@ -19,7 +19,7 @@ use nemo_relay::api::event::{ use nemo_relay::api::llm::LlmRequestInterceptOutcome; use nemo_relay::api::runtime::subscriber_dispatcher::PublicationBuffer; use nemo_relay::api::runtime::{ - LlmSanitizeRequestContext, LlmSanitizeResponseContext, PropagationContext, + LlmExecutionContext, LlmSanitizeRequestContext, LlmSanitizeResponseContext, PropagationContext, ThreadScopeStackBinding, ToolExecutionContext, }; use nemo_relay::api::tool::{ToolExecutionInterceptOutcome, ToolExecutionResult}; @@ -105,6 +105,35 @@ impl PyLlmSanitizeResponseContext { } } +/// Codec capabilities for one managed LLM execution intercept invocation. +#[pyclass(name = "LlmExecutionContext", frozen)] +pub struct PyLlmExecutionContext { + pub(crate) inner: LlmExecutionContext, +} + +#[pymethods] +impl PyLlmExecutionContext { + /// Request codec identity and optional decode/encode capability. + #[getter] + fn request_codec(&self) -> PyLlmSanitizeRequestContext { + PyLlmSanitizeRequestContext { + inner: self.inner.request_codec().clone(), + } + } + + /// Unary response codec identity and optional decode capability. + /// + /// Streaming execution returns ``None`` because Relay does not have a + /// complete-response codec contract for response chunks. + #[getter] + fn response_codec(&self) -> Option { + self.inner + .response_codec() + .cloned() + .map(|inner| PyLlmSanitizeResponseContext { inner }) + } +} + // --------------------------------------------------------------------------- // LlmStream (async iterator) // --------------------------------------------------------------------------- diff --git a/crates/python/src/py_types/mod.rs b/crates/python/src/py_types/mod.rs index 1168922b9..6e56ec7f2 100644 --- a/crates/python/src/py_types/mod.rs +++ b/crates/python/src/py_types/mod.rs @@ -153,6 +153,7 @@ fn register_runtime_types(m: &Bound<'_, PyModule>) -> PyResult<()> { fn register_llm_types(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; + m.add_class::()?; m.add_class::()?; m.add_class::()?; m.add_class::()?; diff --git a/crates/python/tests/coverage/coverage_tests.rs b/crates/python/tests/coverage/coverage_tests.rs index 66281d94b..759278d47 100644 --- a/crates/python/tests/coverage/coverage_tests.rs +++ b/crates/python/tests/coverage/coverage_tests.rs @@ -23,7 +23,9 @@ use crate::py_callable::{ }; use nemo_relay::api::event::{BaseEvent, Event, EventCategory, ScopeCategory, ScopeEvent}; use nemo_relay::api::llm::LlmRequest; -use nemo_relay::api::runtime::{LlmExecutionNextFn, LlmStreamExecutionNextFn, ToolExecutionNextFn}; +use nemo_relay::api::runtime::{ + LlmExecutionContext, LlmExecutionNextFn, LlmStreamExecutionNextFn, ToolExecutionNextFn, +}; fn load_module<'py>(py: Python<'py>, code: &str) -> Bound<'py, PyModule> { let code = CString::new(code).unwrap(); @@ -451,10 +453,10 @@ def llm_conditional(request): def llm_request_intercept(name, request, annotated): return Outcome(request, annotated) -async def llm_execution_intercept(name, request, next): +async def llm_execution_intercept(name, request, context, next): return await next(request) -async def llm_stream_execution_intercept(request, next): +async def llm_stream_execution_intercept(name, request, context, next): return await next(request) def tool_request_intercept(name, value): @@ -817,7 +819,7 @@ async def tool_intercept(context, next): async def llm_exec(request): return {"model": request.content["model"]} -async def llm_intercept(name, request, next): +async def llm_intercept(name, request, context, next): result = await next(request) result["wrapped"] = True return result @@ -883,9 +885,14 @@ async def llm_intercept(name, request, next): Box::pin(async move { Ok(json!({"model": request.content["model"]})) }) }); assert_eq!( - llm_intercept("llm", make_request(), llm_next) - .await - .unwrap(), + llm_intercept( + "llm", + make_request(), + LlmExecutionContext::new(Default::default(), Some(Default::default())), + llm_next, + ) + .await + .unwrap(), json!({"model": "test-model", "wrapped": true}) ); Ok(()) @@ -906,7 +913,7 @@ async def llm_stream(request): yield {"chunk": 1} yield {"chunk": 2} -async def llm_stream_intercept(request, next): +async def llm_stream_intercept(name, request, context, next): return await next(request) "#, ); @@ -934,9 +941,14 @@ async def llm_stream_intercept(request, next): )) }) }); - let mut stream = stream_intercept("llm", make_request(), stream_next) - .await - .unwrap(); + let mut stream = stream_intercept( + "llm", + make_request(), + LlmExecutionContext::new(Default::default(), None), + stream_next, + ) + .await + .unwrap(); let mut seen = Vec::new(); while let Some(chunk) = stream.next().await { seen.push(chunk.unwrap()); @@ -965,14 +977,14 @@ async def tool_intercept_fail(context, next): async def llm_exec_fail(request): raise RuntimeError("llm exec boom") -async def llm_intercept_fail(name, request, next): +async def llm_intercept_fail(name, request, context, next): raise RuntimeError("llm intercept boom") async def llm_stream_fail(request): raise RuntimeError("stream fail") yield {"never": True} -def llm_stream_intercept_sync(request, next): +def llm_stream_intercept_sync(name, request, context, next): class _Iter: def __init__(self): self.done = False @@ -987,7 +999,7 @@ def llm_stream_intercept_sync(request, next): return {"chunk": "sync"} return _Iter() -async def llm_stream_intercept_fail(request, next): +async def llm_stream_intercept_fail(name, request, context, next): raise RuntimeError("stream intercept boom") "#, ); @@ -1050,11 +1062,16 @@ async def llm_stream_intercept_fail(request, next): Box::pin(async move { Ok(json!({"model": request.content["model"]})) }) }); assert!( - llm_intercept("llm", make_request(), llm_next) - .await - .unwrap_err() - .to_string() - .contains("llm intercept boom") + llm_intercept( + "llm", + make_request(), + LlmExecutionContext::new(Default::default(), Some(Default::default())), + llm_next, + ) + .await + .unwrap_err() + .to_string() + .contains("llm intercept boom") ); let stream_exec = wrap_py_llm_stream_exec_fn(llm_stream_fail_py); @@ -1079,9 +1096,14 @@ async def llm_stream_intercept_fail(request, next): )) }) }); - let mut stream = stream_intercept("llm", make_request(), stream_next) - .await - .unwrap(); + let mut stream = stream_intercept( + "llm", + make_request(), + LlmExecutionContext::new(Default::default(), None), + stream_next, + ) + .await + .unwrap(); assert_eq!( stream.next().await.unwrap().unwrap(), json!({"chunk": "sync"}) @@ -1097,7 +1119,14 @@ async def llm_stream_intercept_fail(request, next): )) }) }); - let err = match failing_stream_intercept("llm", make_request(), stream_next).await { + let err = match failing_stream_intercept( + "llm", + make_request(), + LlmExecutionContext::new(Default::default(), None), + stream_next, + ) + .await + { Ok(_) => panic!("expected stream intercept failure"), Err(err) => err, }; diff --git a/crates/python/tests/coverage/py_api_coverage_tests.rs b/crates/python/tests/coverage/py_api_coverage_tests.rs index dc3020231..507d5eed7 100644 --- a/crates/python/tests/coverage/py_api_coverage_tests.rs +++ b/crates/python/tests/coverage/py_api_coverage_tests.rs @@ -393,7 +393,7 @@ async def llm_exec(request): }] } -async def llm_exec_intercept(name, request, next): +async def llm_exec_intercept(name, request, context, next): response = await next(request) response["from_intercept"] = True return response @@ -404,7 +404,7 @@ def llm_stream_exec(request): yield {"delta": 2} return gen() -async def llm_stream_intercept(request, next): +async def llm_stream_intercept(name, request, context, next): stream = await next(request) async def gen(): diff --git a/crates/python/tests/coverage/py_callable_coverage_tests.rs b/crates/python/tests/coverage/py_callable_coverage_tests.rs index 6f5c1eebb..11d1d1046 100644 --- a/crates/python/tests/coverage/py_callable_coverage_tests.rs +++ b/crates/python/tests/coverage/py_callable_coverage_tests.rs @@ -126,7 +126,7 @@ def sync_tool_intercept(context, next): def sync_llm_exec(request): return {"model": request.content["model"], "mode": "sync"} -def sync_llm_intercept(name, request, next): +def sync_llm_intercept(name, request, context, next): return {"name": name, "model": request.content["model"], "mode": "sync"} def request_echo(name, request, annotated): @@ -219,9 +219,17 @@ class RaisingResponseCodec: Box::pin(async move { Ok(json!({"model": request.content["model"]})) }) }); assert_eq!( - llm_intercept("llm", make_request(), llm_next) - .await - .unwrap(), + llm_intercept( + "llm", + make_request(), + nemo_relay::api::runtime::LlmExecutionContext::new( + Default::default(), + Some(Default::default()), + ), + llm_next, + ) + .await + .unwrap(), json!({"name": "llm", "model": "test-model", "mode": "sync"}) ); }); diff --git a/crates/python/tests/coverage/py_plugin_coverage_tests.rs b/crates/python/tests/coverage/py_plugin_coverage_tests.rs index 9edd8296c..3ee2f03d1 100644 --- a/crates/python/tests/coverage/py_plugin_coverage_tests.rs +++ b/crates/python/tests/coverage/py_plugin_coverage_tests.rs @@ -374,10 +374,10 @@ def llm_conditional(request): def llm_request_intercept(name, request, annotated): return Outcome(request, annotated) -async def llm_execution_intercept(name, request, next): +async def llm_execution_intercept(name, request, context, next): return await next(request) -async def llm_stream_execution_intercept(request, next): +async def llm_stream_execution_intercept(name, request, context, next): return await next(request) def tool_request_intercept(name, value): @@ -571,10 +571,10 @@ def llm_conditional(request): def llm_request_intercept(name, request, annotated): return Outcome(request, annotated) -async def llm_execution_intercept(name, request, next): +async def llm_execution_intercept(name, request, context, next): return await next(request) -async def llm_stream_execution_intercept(request, next): +async def llm_stream_execution_intercept(name, request, context, next): return await next(request) def tool_request_intercept(name, value): @@ -776,10 +776,10 @@ def llm_conditional(request): def llm_request_intercept(name, request, annotated): return Outcome(request, annotated) -async def llm_execution_intercept(name, request, next): +async def llm_execution_intercept(name, request, context, next): return await next(request) -async def llm_stream_execution_intercept(request, next): +async def llm_stream_execution_intercept(name, request, context, next): return await next(request) def tool_request_intercept(name, value): diff --git a/crates/worker-proto/README.md b/crates/worker-proto/README.md index b78356175..2b38e45a0 100644 --- a/crates/worker-proto/README.md +++ b/crates/worker-proto/README.md @@ -30,6 +30,13 @@ for earlier releases must regenerate their bindings, rebuild, and declare `compat.relay` beginning at `0.8.0`. `ToolNext` returns `ToolExecutionResultResponse`, and tool execution intercepts use structural `ToolExecutionInterceptOutcome` messages. +Relay 0.10 adds directional execution codec context to `LlmInvocation` and makes +that context required by the 0.10 worker SDK's unary and streaming LLM execution +callbacks. The protobuf field is additive, but the SDK callback contract is a +release-level source break. Rebuild workers, regenerate custom bindings, and use +a `compat.relay` lower bound of `0.10.0`. The protocol identifier remains +`grpc-v1`. + ## Protocol Surface | Surface | Role | @@ -41,6 +48,7 @@ and tool execution intercepts use structural `ToolExecutionInterceptOutcome` mes | Tool results | `ToolNext` returns `ToolExecutionResultResponse`, and `ToolExecutionInterceptResult` returns `ToolExecutionInterceptOutcome`. Both preserve the application result and optional annotation. Intercept outcomes also include ordered pending marks. These fields use lossless protobuf `JsonValue` wrappers rather than `google.protobuf.Value`. | | Mark options | `EmitMarkRequest.data_schema` carries a `nemo.relay.DataSchema@1` envelope, `severity` carries the log severity, and `category` carries an optional semantic mark category. Omitting these fields preserves legacy behavior. | | Runtime diagnostics | Authenticated `GetRuntimeDiagnostics` returns a bounded active-host `{ code, message, count }` snapshot. Older hosts return gRPC `UNIMPLEMENTED`. | +| LLM execution codec context | `LlmInvocation.execution_codec_context` carries request identity and an invocation-scoped request capability. Unary execution also carries response identity and a response capability; streaming execution omits the response direction. | ## Installation diff --git a/crates/worker-proto/proto/nemo/relay/worker/v1/plugin_worker.proto b/crates/worker-proto/proto/nemo/relay/worker/v1/plugin_worker.proto index b54dd0338..b0a597156 100644 --- a/crates/worker-proto/proto/nemo/relay/worker/v1/plugin_worker.proto +++ b/crates/worker-proto/proto/nemo/relay/worker/v1/plugin_worker.proto @@ -245,9 +245,6 @@ message Registration { int32 priority = 3; bool break_chain = 4; reserved 5; - // Requests invocation-scoped codec identities and supported codec operations - // for this execution interceptor. Older hosts ignore this field. - bool llm_execution_codec_context = 6; } message InvokeRequest { @@ -291,8 +288,9 @@ message LlmInvocation { LlmSanitizeRequestContext request_sanitize_context = 9; LlmSanitizeResponseContext response_sanitize_context = 10; } - // Present only for an execution interceptor that opted into codec context. - // Older workers ignore this additive field. + // Present for LLM execution interceptors. Older workers ignore this + // additive field at the wire boundary, but worker SDK callback signatures + // are release-versioned and must match the Relay release they target. LlmExecutionCodecContext execution_codec_context = 11; } diff --git a/crates/worker-proto/tests/proto_tests.rs b/crates/worker-proto/tests/proto_tests.rs index 8a8189f76..d90d34e44 100644 --- a/crates/worker-proto/tests/proto_tests.rs +++ b/crates/worker-proto/tests/proto_tests.rs @@ -8,8 +8,8 @@ use nemo_relay_worker_proto::v1::{ GetRuntimeDiagnosticsRequest, GetRuntimeDiagnosticsResponse, HandshakeRequest, HealthRequest, InvokeRequest, JsonEnvelope, JsonValue, LlmCodecIdentity, LlmCodecKind, LlmExecutionCodecContext, LlmInvocation, LlmSanitizeRequestContext, LlmSanitizeResponseContext, - RegisterConditionalMiddlewareGuardrailRequest, Registration, RegistrationSurface, - RuntimeDiagnostic, ScopeType, ToolExecutionResult as ProtoToolExecutionResult, invoke_request, + RegisterConditionalMiddlewareGuardrailRequest, RegistrationSurface, RuntimeDiagnostic, + ScopeType, ToolExecutionResult as ProtoToolExecutionResult, invoke_request, }; use nemo_relay_worker_proto::{ WORKER_PROTOCOL_GRPC_V1, decode_json_envelope, decode_json_value, json_envelope, json_value, @@ -17,18 +17,6 @@ use nemo_relay_worker_proto::{ use prost::Message; use serde_json::json; -#[derive(Clone, PartialEq, Message)] -struct LegacyRegistration { - #[prost(string, tag = "1")] - local_name: String, - #[prost(int32, tag = "2")] - surface: i32, - #[prost(int32, tag = "3")] - priority: i32, - #[prost(bool, tag = "4")] - break_chain: bool, -} - #[derive(Clone, PartialEq, Message)] struct LegacyLlmInvocation { #[prost(string, tag = "1")] @@ -168,31 +156,9 @@ fn request_field_numbers_are_stable() { } #[test] -fn execution_codec_context_fields_are_additive_and_stable() { - let legacy_registration = LegacyRegistration { - local_name: "legacy".into(), - surface: RegistrationSurface::LlmExecutionIntercept as i32, - priority: 7, - break_chain: false, - }; - let decoded_by_new_host = Registration::decode(legacy_registration.encode_to_vec().as_slice()) - .expect("new host must decode a legacy registration"); - assert_eq!(decoded_by_new_host.local_name, "legacy"); - assert!(!decoded_by_new_host.llm_execution_codec_context); - - let contextual_registration = Registration { - llm_execution_codec_context: true, - ..Default::default() - }; - assert_eq!(contextual_registration.encode_to_vec(), b"\x30\x01"); - let decoded_by_legacy_host = - LegacyRegistration::decode(contextual_registration.encode_to_vec().as_slice()) - .expect("legacy host must ignore the additive registration field"); - assert_eq!(decoded_by_legacy_host, LegacyRegistration::default()); - +fn execution_codec_context_invocation_field_is_additive_and_stable() { let legacy_invocation = LegacyLlmInvocation { model_name: "legacy-model".into(), - ..Default::default() }; let decoded_by_new_worker = LlmInvocation::decode(legacy_invocation.encode_to_vec().as_slice()) .expect("new worker must decode a legacy invocation"); diff --git a/crates/worker/README.md b/crates/worker/README.md index 5c12ee806..e52368b6d 100644 --- a/crates/worker/README.md +++ b/crates/worker/README.md @@ -26,6 +26,15 @@ for an earlier release must rebuild with this SDK and declare `compat.relay` beg `ToolExecutionResult`, which keeps an optional opaque annotation beside the application result. +Relay 0.10 makes LLM execution codec context part of every unary and streaming +execution callback. Rebuild workers with the 0.10 SDK and raise their +`compat.relay` lower bound to `0.10.0`. The callback receives +`LlmExecutionContext` immediately before `next`. Its request direction reports +the active codec and may resolve request decode/encode operations. Unary +callbacks also receive response decode context; streaming callbacks deliberately +receive no response codec because Relay does not decode incomplete chunks. The +wire protocol remains named `grpc-v1`. + ## Authoring Surface | Surface | Role | @@ -33,6 +42,7 @@ result. | `WorkerPlugin` | Defines plugin identity, validation, registration, and multiple-component behavior in the worker process. | | `PluginContext` | Installs typed handlers for all 16 supported registration surfaces. | | `PluginRuntime` and continuations | Emit marks, manage scopes, and call the remaining tool or LLM execution chain through the authenticated host service. | +| `LlmExecutionContext` | Reports request and unary-response codec identity and exposes invocation-scoped codec operations when Relay resolved a codec. | | Canonical tool results | Preserve application results and opaque annotations across tool callbacks and continuations. | | `serve_plugin` | Starts the Tokio gRPC server from the activation identity, local endpoints, and token supplied by Relay. | diff --git a/crates/worker/src/lib.rs b/crates/worker/src/lib.rs index 2cad4082e..802f21f5e 100644 --- a/crates/worker/src/lib.rs +++ b/crates/worker/src/lib.rs @@ -390,62 +390,42 @@ impl WorkerResponseCodec { /// Invocation-scoped codec identities and operations for an LLM execution interceptor. /// -/// `is_available` is false when this SDK is running against a Relay host that -/// predates execution codec context. A supported host can still report no -/// active codec for either direction. -/// /// This context identifies the codecs selected when Relay created the managed /// invocation. Rewriting a request into another provider's wire format does /// not change these identities; codec operations reject incompatible payloads. #[derive(Clone)] pub struct LlmExecutionContext { - available: bool, - request_codec_identity: LlmCodecIdentity, - response_codec_identity: LlmCodecIdentity, - request_codec: Option, - response_codec: Option, + request_codec: LlmSanitizeRequestContext, + response_codec: Option, } impl std::fmt::Debug for LlmExecutionContext { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter .debug_struct("LlmExecutionContext") - .field("available", &self.available) - .field("request_codec_identity", &self.request_codec_identity) - .field("response_codec_identity", &self.response_codec_identity) + .field("request_codec", &self.request_codec.codec) + .field( + "response_codec", + &self.response_codec.as_ref().map(|context| &context.codec), + ) .finish_non_exhaustive() } } impl LlmExecutionContext { - /// Whether the Relay host supplied execution codec context. - #[must_use] - pub fn is_available(&self) -> bool { - self.available - } - - /// Identity of the active request codec, or `None` when no codec is active. + /// Request codec identity and invocation-scoped operations. #[must_use] - pub fn request_codec_identity(&self) -> &LlmCodecIdentity { - &self.request_codec_identity + pub fn request_codec(&self) -> &LlmSanitizeRequestContext { + &self.request_codec } - /// Identity of the active response codec, or `None` when no codec is active. - #[must_use] - pub fn response_codec_identity(&self) -> &LlmCodecIdentity { - &self.response_codec_identity - } - - /// Invocation-scoped request codec proxy, when the host supplied one. - #[must_use] - pub fn request_codec(&self) -> Option { - self.request_codec.clone() - } - - /// Invocation-scoped response codec proxy, when the host supplied one. + /// Unary response codec identity and invocation-scoped decode operation. + /// + /// Streaming execution returns `None` because Relay response codecs decode + /// completed provider responses, not individual stream chunks. #[must_use] - pub fn response_codec(&self) -> Option { - self.response_codec.clone() + pub fn response_codec(&self) -> Option<&LlmSanitizeResponseContext> { + self.response_codec.as_ref() } } @@ -887,40 +867,15 @@ impl PluginContext { name: &str, priority: i32, callback: F, - ) where - F: Fn(&str, LlmRequest, LlmNext) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - { - self.push_registration( - name, - RegistrationSurface::LlmExecutionIntercept, - priority, - false, - ); - self.handlers.llm_executions.insert( - name.into(), - Arc::new(move |model, request, _context, next| { - Box::pin(callback(model, request, next)) - }), - ); - } - - /// Registers an LLM execution intercept with invocation-scoped codec access. - pub fn register_llm_execution_intercept_with_context( - &mut self, - name: &str, - priority: i32, - callback: F, ) where F: Fn(&str, LlmRequest, LlmExecutionContext, LlmNext) -> Fut + Send + Sync + 'static, Fut: Future> + Send + 'static, { - self.push_registration_record( + self.push_registration( name, RegistrationSurface::LlmExecutionIntercept, priority, false, - true, ); self.handlers.llm_executions.insert( name.into(), @@ -941,43 +896,15 @@ impl PluginContext { name: &str, priority: i32, callback: F, - ) where - F: Fn(&str, LlmRequest, LlmStreamNext) -> Fut + Send + Sync + 'static, - Fut: Future> + Send + 'static, - { - self.push_registration( - name, - RegistrationSurface::LlmStreamExecutionIntercept, - priority, - false, - ); - self.handlers.llm_stream_executions.insert( - name.into(), - Arc::new(move |model, request, _context, next| { - Box::pin(callback(model, request, next)) - }), - ); - } - - /// Registers a streaming LLM execution intercept with request codec access. - /// - /// The context identifies the response codec but does not expose a response - /// decoder because stream chunks are not complete provider responses. - pub fn register_llm_stream_execution_intercept_with_context( - &mut self, - name: &str, - priority: i32, - callback: F, ) where F: Fn(&str, LlmRequest, LlmExecutionContext, LlmStreamNext) -> Fut + Send + Sync + 'static, Fut: Future> + Send + 'static, { - self.push_registration_record( + self.push_registration( name, RegistrationSurface::LlmStreamExecutionIntercept, priority, false, - true, ); self.handlers.llm_stream_executions.insert( name.into(), @@ -993,24 +920,12 @@ impl PluginContext { surface: RegistrationSurface, priority: i32, break_chain: bool, - ) { - self.push_registration_record(name, surface, priority, break_chain, false); - } - - fn push_registration_record( - &mut self, - name: &str, - surface: RegistrationSurface, - priority: i32, - break_chain: bool, - llm_execution_codec_context: bool, ) { self.handlers.registrations.push(Registration { local_name: name.into(), surface: surface as i32, priority, break_chain, - llm_execution_codec_context, }); } } @@ -2995,43 +2910,38 @@ impl LlmPayload { runtime: &PluginRuntime, invocation_id: &str, ) -> Result { - let Some(context) = self.execution_codec_context.as_ref() else { - return Ok(LlmExecutionContext { - available: false, - request_codec_identity: LlmCodecIdentity::None, - response_codec_identity: LlmCodecIdentity::None, - request_codec: None, - response_codec: None, - }); - }; + let context = require_execution_field( + self.execution_codec_context.as_ref(), + "execution context is missing", + )?; let request = require_execution_field(context.request.as_ref(), "request context is missing")?; - let response = - require_execution_field(context.response.as_ref(), "response context is missing")?; let request_identity = require_execution_field(request.codec.as_ref(), "request codec identity is missing")?; - let response_identity = require_execution_field( - response.codec.as_ref(), - "response codec identity is missing", - )?; + let response_codec = context + .response + .as_ref() + .map(|response| -> Result { + let identity = require_execution_field( + response.codec.as_ref(), + "response codec identity is missing", + )?; + Ok(LlmSanitizeResponseContext { + codec: codec_identity_from_proto(Some(identity)), + runtime: Some(runtime.clone()), + codec_capability_id: response.codec_capability_id.clone(), + invocation_id: Some(invocation_id.to_owned()), + }) + }) + .transpose()?; Ok(LlmExecutionContext { - available: true, - request_codec_identity: codec_identity_from_proto(Some(request_identity)), - response_codec_identity: codec_identity_from_proto(Some(response_identity)), - request_codec: request.codec_capability_id.as_ref().map(|capability_id| { - WorkerRequestCodec { - runtime: runtime.clone(), - capability_id: capability_id.clone(), - invocation_id: invocation_id.to_owned(), - } - }), - response_codec: response.codec_capability_id.as_ref().map(|capability_id| { - WorkerResponseCodec { - runtime: runtime.clone(), - capability_id: capability_id.clone(), - invocation_id: invocation_id.to_owned(), - } - }), + request_codec: LlmSanitizeRequestContext { + codec: codec_identity_from_proto(Some(request_identity)), + runtime: Some(runtime.clone()), + codec_capability_id: request.codec_capability_id.clone(), + invocation_id: Some(invocation_id.to_owned()), + }, + response_codec, }) } diff --git a/crates/worker/tests/unit/execution_context_tests.rs b/crates/worker/tests/unit/execution_context_tests.rs index 0f9211f3c..c5d603c07 100644 --- a/crates/worker/tests/unit/execution_context_tests.rs +++ b/crates/worker/tests/unit/execution_context_tests.rs @@ -27,43 +27,32 @@ fn llm_payload( } #[test] -fn context_registration_is_opt_in_without_a_new_surface() { +fn execution_registration_uses_the_existing_surface() { let mut context = PluginContext::new(); - context.register_llm_execution_intercept("legacy", 7, |_, _, _| async { - Ok(serde_json::json!({"legacy": true})) - }); - context.register_llm_execution_intercept_with_context("context", 7, |_, _, _, _| async { + context.register_llm_execution_intercept("context", 7, |_, _, _, _| async { Ok(serde_json::json!({"context": true})) }); - let legacy = &context.handlers.registrations[0]; - let contextual = &context.handlers.registrations[1]; + let registration = &context.handlers.registrations[0]; assert_eq!( - legacy.surface, + registration.surface, RegistrationSurface::LlmExecutionIntercept as i32 ); - assert_eq!(contextual.surface, legacy.surface); - assert_eq!(contextual.priority, legacy.priority); - assert!(!legacy.llm_execution_codec_context); - assert!(contextual.llm_execution_codec_context); + assert_eq!(registration.priority, 7); } #[test] -fn absent_execution_context_identifies_an_older_host() { +fn absent_execution_context_is_a_release_mismatch() { let payload = llm_payload(None); - let context = payload + let error = payload .execution_context(&disconnected_runtime(), "invocation") - .unwrap(); - assert!(!context.is_available()); - assert_eq!(context.request_codec_identity(), &LlmCodecIdentity::None); - assert_eq!(context.response_codec_identity(), &LlmCodecIdentity::None); - assert!(context.request_codec().is_none()); - assert!(context.response_codec().is_none()); + .unwrap_err(); + assert!(error.to_string().contains("execution context is missing")); } #[test] -fn execution_context_from_a_new_host_preserves_identities_and_capabilities() { +fn execution_context_preserves_directional_identities_and_capabilities() { let codec = nemo_relay_worker_proto::v1::LlmCodecIdentity { kind: LlmCodecKind::Builtin as i32, id: Some("openai_chat".into()), @@ -84,15 +73,92 @@ fn execution_context_from_a_new_host_preserves_identities_and_capabilities() { let context = payload .execution_context(&disconnected_runtime(), "invocation") .unwrap(); - assert!(context.is_available()); assert_eq!( - context.request_codec_identity(), + &context.request_codec().codec, &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) ); + assert!(context.request_codec().resolve_codec().is_some()); + let response = context.response_codec().expect("unary response context"); assert_eq!( - context.response_codec_identity(), + &response.codec, &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) ); - assert!(context.request_codec().is_some()); - assert!(context.response_codec().is_some()); + assert!(response.resolve_codec().is_some()); +} + +#[test] +fn execution_context_distinguishes_absent_and_resolved_opaque_codecs() { + let absent = llm_payload(Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: LlmCodecKind::Unspecified as i32, + id: None, + }), + codec_capability_id: None, + }), + response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: LlmCodecKind::Unspecified as i32, + id: None, + }), + codec_capability_id: None, + }), + }, + ))) + .execution_context(&disconnected_runtime(), "absent-invocation") + .unwrap(); + assert_eq!(absent.request_codec().codec, LlmCodecIdentity::None); + assert!(absent.request_codec().resolve_codec().is_none()); + let absent_response = absent.response_codec().expect("unary response direction"); + assert_eq!(absent_response.codec, LlmCodecIdentity::None); + assert!(absent_response.resolve_codec().is_none()); + + let opaque = llm_payload(Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: LlmCodecKind::Opaque as i32, + id: None, + }), + codec_capability_id: Some("opaque-request".into()), + }), + response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: LlmCodecKind::Opaque as i32, + id: None, + }), + codec_capability_id: Some("opaque-response".into()), + }), + }, + ))) + .execution_context(&disconnected_runtime(), "opaque-invocation") + .unwrap(); + assert_eq!(opaque.request_codec().codec, LlmCodecIdentity::Opaque); + assert!(opaque.request_codec().resolve_codec().is_some()); + let opaque_response = opaque.response_codec().expect("unary response direction"); + assert_eq!(opaque_response.codec, LlmCodecIdentity::Opaque); + assert!(opaque_response.resolve_codec().is_some()); +} + +#[test] +fn streaming_execution_context_has_no_response_codec() { + let payload = llm_payload(Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: LlmCodecKind::Builtin as i32, + id: Some("openai_chat".into()), + }), + codec_capability_id: Some("request-capability".into()), + }), + response: None, + }, + ))); + + let context = payload + .execution_context(&disconnected_runtime(), "invocation") + .unwrap(); + assert!(context.request_codec().resolve_codec().is_some()); + assert!(context.response_codec().is_none()); } diff --git a/crates/worker/tests/worker_sdk_tests.rs b/crates/worker/tests/worker_sdk_tests.rs index 9128eff58..ef2bb42ad 100644 --- a/crates/worker/tests/worker_sdk_tests.rs +++ b/crates/worker/tests/worker_sdk_tests.rs @@ -2043,7 +2043,7 @@ impl WorkerPlugin for CancellationPlugin { let stream_started = self.stream_started.clone(); let stream_cancelled = self.stream_cancelled.clone(); let stream_cancelled_notified = self.stream_cancelled_notified.clone(); - ctx.register_llm_stream_execution_intercept("cancel-stream", 0, move |_, _, _| { + ctx.register_llm_stream_execution_intercept("cancel-stream", 0, move |_, _, _, _| { let stream_started = stream_started.clone(); let stream_cancelled = stream_cancelled.clone(); let stream_cancelled_notified = stream_cancelled_notified.clone(); @@ -2059,7 +2059,7 @@ impl WorkerPlugin for CancellationPlugin { let stream_setup_started = self.stream_setup_started.clone(); let stream_setup_cancelled = self.stream_setup_cancelled.clone(); - ctx.register_llm_stream_execution_intercept("cancel-stream-setup", 0, move |_, _, _| { + ctx.register_llm_stream_execution_intercept("cancel-stream-setup", 0, move |_, _, _, _| { let stream_setup_started = stream_setup_started.clone(); let stream_setup_cancelled = stream_setup_cancelled.clone(); async move { @@ -2498,7 +2498,7 @@ impl WorkerPlugin for SurfacePlugin { ctx.register_llm_execution_intercept( "llm-exec", 1, - move |_, request, next: LlmNext| async move { + move |_, request, _context, next: LlmNext| async move { let next_value = next.call(request).await?; Ok(set_json_field(next_value, "phase", "llm_exec")) }, @@ -2508,7 +2508,7 @@ impl WorkerPlugin for SurfacePlugin { ctx.register_llm_stream_execution_intercept( "llm-stream", 1, - move |_, request, next: LlmStreamNext| { + move |_, request, _context, next: LlmStreamNext| { let runtime = stream_runtime.clone(); async move { let next_stream = next.call(request).await?; @@ -2516,15 +2516,17 @@ impl WorkerPlugin for SurfacePlugin { } }, ); - ctx.register_llm_stream_execution_intercept("llm-stream-error", 1, |_, _, _| async { + ctx.register_llm_stream_execution_intercept("llm-stream-error", 1, |_, _, _, _| async { let stream: JsonStream = Box::pin(tokio_stream::iter(vec![Err( WorkerSdkError::Callback("stream boom".into()), )])); Ok(stream) }); - ctx.register_llm_stream_execution_intercept("llm-stream-open-error", 1, |_, _, _| async { - Err(WorkerSdkError::Callback("stream open boom".into())) - }); + ctx.register_llm_stream_execution_intercept( + "llm-stream-open-error", + 1, + |_, _, _, _| async { Err(WorkerSdkError::Callback("stream open boom".into())) }, + ); Ok(()) } } @@ -3269,6 +3271,21 @@ fn llm_invoke( annotated_request: Option, response: Option, ) -> InvokeRequest { + let execution_codec_context = match surface { + RegistrationSurface::LlmExecutionIntercept => Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(empty_request_codec_context()), + response: Some(empty_response_codec_context()), + }, + )), + RegistrationSurface::LlmStreamExecutionIntercept => Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(empty_request_codec_context()), + response: None, + }, + )), + _ => None, + }; InvokeRequest { activation_id: ACTIVATION_ID.into(), invocation_id: "invoke-1".into(), @@ -3286,7 +3303,7 @@ fn llm_invoke( annotated_request: annotated_request.map(json_env), response: response.map(json_env), sanitize_context: None, - execution_codec_context: None, + execution_codec_context, }, )), } @@ -3317,6 +3334,26 @@ fn llm_invoke_without_request( } } +fn empty_request_codec_context() -> nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { + nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: nemo_relay_worker_proto::v1::LlmCodecKind::Unspecified as i32, + id: None, + }), + codec_capability_id: None, + } +} + +fn empty_response_codec_context() -> nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { + nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: nemo_relay_worker_proto::v1::LlmCodecKind::Unspecified as i32, + id: None, + }), + codec_capability_id: None, + } +} + fn llm_invoke_with_bad_annotation(registration_name: &str) -> InvokeRequest { let mut request = llm_invoke( registration_name, diff --git a/docs/about-nemo-relay/release-notes/index.mdx b/docs/about-nemo-relay/release-notes/index.mdx index a319f3888..7518d8e9c 100644 --- a/docs/about-nemo-relay/release-notes/index.mdx +++ b/docs/about-nemo-relay/release-notes/index.mdx @@ -33,6 +33,22 @@ For the full release history, including individual pull requests, refer to NeMo Relay 0.10 is under development. This page will track its user-visible changes, compatibility updates, and fixed known issues. +### LLM Execution Codec Context + +Relay now passes the selected request and unary-response codec context to every +LLM execution interceptor. Policies can use Relay's host-owned codec to inspect +or safely rewrite the final provider request and decode a completed unary +response without importing or copying Relay's codec implementations. Streaming +interceptors receive request codec access only. + +**Breaking change:** The context is a required callback argument across Rust, +Python, Node.js, Go, C, native plugins, and Rust and Python gRPC workers. Native +plugins must rebuild for internal ABI v7, and affected workers must regenerate +their protobuf bindings and rebuild. Authored compatibility labels remain +`native_api = "1"` and `grpc-v1`; plugin manifests must use a Relay range that +begins at 0.10 or otherwise excludes 0.9. Refer to the [Migration +Guides](/reference/migration-guides) for callback shapes and upgrade steps. + ### Fixes and Other Changes - Scope-stack forks, managed LLM and tool execution callbacks, and diff --git a/docs/build-plugins/about.mdx b/docs/build-plugins/about.mdx index 3e6f067c4..4e99a7095 100644 --- a/docs/build-plugins/about.mdx +++ b/docs/build-plugins/about.mdx @@ -56,7 +56,9 @@ artifact solves a concrete operational problem. Native Rust plugins suit reusable middleware whose callback latency or throughput is important enough to justify platform-specific binaries and full in-process trust. A native plugin uses manifest compatibility `compat.native_api = "1"`; the current SDK -negotiates C host-table ABI v5 and retains frozen v4, v3, and v2 host-table compatibility. +uses C host-table ABI v7. Relay 0.10 intentionally rejects older compiled callback +layouts, so native plugins must rebuild for this release even though the authored +manifest label remains `native_api = "1"`. Those are different version axes, as the [Native ABI Reference](/build-plugins/native/native-abi-reference) explains. @@ -73,6 +75,15 @@ same pair plus Relay-owned pending marks. Native and worker plugins built for an release must rebuild for this contract. Workers retain the `grpc-v1` name and protobuf package; their tool-result fields, not their protocol identity, changed. +Relay 0.10 makes codec context part of every LLM execution callback in language +bindings, native plugins, and worker SDKs. This is a source break. Native and worker +plugins must rebuild and constrain `compat.relay` to `>=0.10.0`; their authored +`native_api = "1"` and `grpc-v1` labels do not change. + +The codec capability belongs to that execution callback. Unary capabilities expire when +the callback settles; streaming request capabilities remain valid until the returned +stream completes or closes. Retaining a codec object does not extend either lifetime. + ## Match Behavior to a Plugin | Desired Behavior | Good Starting Model | Why | diff --git a/docs/build-plugins/language-binding/register-behavior.mdx b/docs/build-plugins/language-binding/register-behavior.mdx index 537ef592b..7ef7e8636 100644 --- a/docs/build-plugins/language-binding/register-behavior.mdx +++ b/docs/build-plugins/language-binding/register-behavior.mdx @@ -153,7 +153,13 @@ this common center. ```python from collections.abc import AsyncIterator, Awaitable, Callable -from nemo_relay import AnnotatedLLMRequest, Json, LLMRequest, LLMRequestInterceptOutcome +from nemo_relay import ( + AnnotatedLLMRequest, + Json, + LLMRequest, + LLMRequestInterceptOutcome, + LlmExecutionContext, +) settings = normalized_config(config) tag = settings["tag"] @@ -199,8 +205,12 @@ context.register_llm_request_intercept( ) async def stream_request( - request: LLMRequest, next_call: Callable[[LLMRequest], Awaitable[AsyncIterator[Json]]] + _name: str, + request: LLMRequest, + context: LlmExecutionContext, + next_call: Callable[[LLMRequest], Awaitable[AsyncIterator[Json]]], ) -> AsyncIterator[Json]: + assert context.response_codec is None async for chunk in await next_call(request): yield {**chunk, "plugin_stream": True} @@ -251,7 +261,8 @@ context.registerLlmRequestIntercept( context.registerLlmStreamExecutionIntercept( 'documentation-stream', execution.priority, - async function* (request, next) { + async function* (request, context, next) { + console.assert(context.responseCodec === null); for await (const chunk of await next(request)) { yield { ...chunk, plugin_stream: true }; } @@ -319,8 +330,9 @@ ctx.register_llm_request_intercept( ctx.register_llm_stream_execution_intercept( "documentation-stream", config.execution.priority, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, context, next| { Box::pin(async move { + debug_assert!(context.response_codec().is_none()); let downstream = next(request).await?; Ok(LlmJsonStream::new(downstream.map(|chunk| { chunk.map(|mut value| { @@ -337,6 +349,12 @@ ctx.register_llm_stream_execution_intercept( +Execution codec objects are scoped to the callback that received them. Unary +capabilities expire when that callback settles. A streaming request capability remains +usable while the returned stream is live and expires on completion, error, close, or +drop. Retaining a codec object past that boundary produces an error; it does not extend +the invocation lifetime. + The LLM intercept returns the [complete outcome](/reference/llm-request-intercept-outcomes) rather than relying on mutation. In particular, it preserves `annotated`. The Rust request is mutable inside its diff --git a/docs/build-plugins/native/about.mdx b/docs/build-plugins/native/about.mdx index 386f8967f..83fa3915c 100644 --- a/docs/build-plugins/native/about.mdx +++ b/docs/build-plugins/native/about.mdx @@ -26,11 +26,16 @@ Three version values answer different questions: |---|---|---| | Package manifest | `manifest_version = 1` | Shape of the authored `relay-plugin.toml` file. | | Manifest native API | `compat.native_api = "1"` | Native plugin package contract accepted by discovery and trust validation. | -| C host-table ABI | v5 | Function table negotiated by the current `nemo-relay-plugin` SDK. The host also exposes frozen v4, v3, and v2 tables for compatible older binaries. | +| C host-table ABI | v7 | Function table required by the current `nemo-relay-plugin` SDK. Relay 0.10 rejects older callback layouts. | The checked example registers a tool execution intercept, so it declares -`compat.relay = ">=0.9.0,<1.0"`. A manifest that admits Relay 0.8 cannot register -the context-carrying callback introduced in Relay 0.9. +`compat.relay = ">=0.10.0,<1.0"`. Every native plugin built with this SDK must use +that lower bound because an older Relay cannot load the ABI v7 function table, even +when the component itself does not register an LLM execution intercept. + +Relay 0.10 adds directional codec context to every LLM execution callback and +advances the internal table to v7. Rebuild all native plugins for this release; +the authored `native_api = "1"` label does not change. Relay 0.8 changed the native API 1 tool-result JSON contract without changing the v4 host-table layout. A tool callback and `ToolNext` continuation return diff --git a/docs/build-plugins/native/build-and-package.mdx b/docs/build-plugins/native/build-and-package.mdx index deb93fe16..dbfcf1eab 100644 --- a/docs/build-plugins/native/build-and-package.mdx +++ b/docs/build-plugins/native/build-and-package.mdx @@ -77,7 +77,7 @@ id = "examples.rust_native_policy" kind = "rust_dynamic" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" native_api = "1" [defaults] diff --git a/docs/build-plugins/native/native-abi-reference.mdx b/docs/build-plugins/native/native-abi-reference.mdx index c1cd43d28..57f173c12 100644 --- a/docs/build-plugins/native/native-abi-reference.mdx +++ b/docs/build-plugins/native/native-abi-reference.mdx @@ -23,7 +23,16 @@ extern "C" fn nemo_relay_register_plugin( ) -> NemoRelayStatus ``` -The current host negotiates ABI v6, then the frozen v5, v4, v3, and legacy v2 tables. +The current host requires ABI v7. It does not fall back to v6 or earlier tables because +v7 changes the LLM execution callback layouts; invoking a stale callback through the +new layout would be unsafe. Rebuild every native plugin for Relay 0.10 and raise its +`compat.relay` lower bound to `0.10.0`. The authored manifest label remains +`compat.native_api = "1"`. + +ABI v7 adds directional request and unary-response codec context to raw, typed, +and asynchronous LLM execution callbacks. The generic asynchronous middleware +callback remains unchanged; ABI v7 appends an execution-specific unary registration, +and the existing execution-specific stream callback gains the context parameter. ABI v6 adds a host-routed operational logging function. Native plugins pass a level, optional target, message, and optional JSON-object fields; Relay applies its operational logging policy and preserves those fields in structured JSONL output. @@ -43,7 +52,9 @@ function signatures and field order are defined by the public | Frozen v1/v2 prefix | Version and struct-size negotiation; host version; string allocation, access, and release; thread-local error reporting; callback-scoped LLM request decode and encode plus response decode; subscriber, five tool, six LLM, and three event-sanitizer registrations; current scope, scope push and pop, mark emission, isolated stack creation and release, thread-stack set, capture, and restore, captured-binding release, active-stack inspection, and scoped binding. | | Frozen v3 extension | Completion resolve, reject, cancellation inspection, and release; one-shot completion-coupled continuation invocation; continuation release; generic async middleware registration; bounded output-stream push, finish, reject, cancellation inspection, and release; downstream stream invocation; stream-middleware registration; and repeated or concurrent unary continuation invocation with independent result callbacks. | | Frozen v4 extension | Completion-scoped LLM request decode and encode plus response decode; pull-based downstream LLM stream open, pull, cancel, and release; completion retain for typed codec facades; output-stream backpressure inspection; extended mark emission; runtime diagnostics; activation-owned runtime capability creation, retain, and release; global runtime-registration discovery; owned conditional middleware guardrail registration and deregistration; and activation-owned and runtime-discovered callback gate registration. The callback registration slots are appended after the original constant-reason slots. | -| Current v5 extension | Context-carrying tool execution intercept registration. The callback receives one JSON object containing `tool_name`, `args`, and `tool_call_id`. | +| v5 extension | Context-carrying tool execution intercept registration. The callback receives one JSON object containing `tool_name`, `args`, and `tool_call_id`. | +| v6 extension | Host-routed operational logging with structured fields. | +| Current v7 extension | Directional codec context for raw and typed LLM execution callbacks, an execution-specific asynchronous unary registration, context on the asynchronous stream-execution callback, and stream retain plus stream-scoped request decode and encode operations. | The prefix and descriptor layout are explicit. A plugin fills the descriptor with its stable kind, component multiplicity, opaque state, callbacks, and destructor. The host @@ -151,7 +162,10 @@ or binding, and restoration. ## Async Completions and Continuations -`PluginContext::register_async_middleware_raw` registers non-stream middleware that can +`PluginContext::register_async_middleware_raw` registers non-stream middleware other +than LLM execution. ABI v7 uses the execution-specific +`plugin_context_register_async_llm_execution_intercept` entry for asynchronous unary +LLM execution callbacks. Both paths can settle later. Return `Complete` only after resolving or rejecting the completion inside the callback. Return `Pending` only after retaining it. A retained completion must settle exactly once and then be released. Release every async `next` reference after its last @@ -216,6 +230,19 @@ error also omits them and records the callback failure. Neither outcome changes request or application response. Never retain a raw handle or resolved typed facade after the sanitizer callback ends. +LLM execution context uses the same identity and opaque codec operations. Its +request direction is always present and may resolve request decode/encode. Unary +execution also carries an optional response direction with response decode. +Streaming execution sets the response direction to unavailable because the +response codecs operate on completed provider responses, not individual chunks. +Raw execution-context handles are borrowed for the owning callback. Typed unary +facades retain their completion and typed streaming request facades retain their +output stream, so they remain memory-safe while the callback future or returned +stream owns them. Codec calls fail after the completion or stream settles. Raw +plugins that need request codec access after an asynchronous stream callback +returns must retain the stream and use the v7 stream-scoped decode and encode +operations; release that stream reference after the last call. + ## Unload Ordering Relay keeps the library loaded while any owned registration or callback can reference diff --git a/docs/build-plugins/native/wrap-execution.mdx b/docs/build-plugins/native/wrap-execution.mdx index 22a3082ae..a628b61c6 100644 --- a/docs/build-plugins/native/wrap-execution.mdx +++ b/docs/build-plugins/native/wrap-execution.mdx @@ -66,7 +66,8 @@ or charging policy before they use this pattern. context.register_llm_execution_intercept( "documentation_llm_execution", config.execution.priority, - move |_name, request, next| async move { + move |_name, request, context, next| async move { + debug_assert!(context.response_codec().is_some()); let repeat = request.content .get("repeat_downstream") .and_then(Json::as_bool) @@ -108,7 +109,8 @@ output promptly. context.register_llm_stream_execution_intercept( "documentation_llm_stream_execution", config.execution.priority, - move |_name, request, next| async move { + move |_name, request, context, next| async move { + debug_assert!(context.response_codec().is_none()); let stream = next.call(request).await?; let mapped: LlmJsonAsyncStream = Box::pin(stream.map(|chunk| { chunk.map(|chunk| match chunk { diff --git a/docs/build-plugins/package-discoverable-plugins.mdx b/docs/build-plugins/package-discoverable-plugins.mdx index 3ec20ff0b..f93672014 100644 --- a/docs/build-plugins/package-discoverable-plugins.mdx +++ b/docs/build-plugins/package-discoverable-plugins.mdx @@ -17,7 +17,7 @@ equivalent binding configuration. | Block | What It Controls | |---|---| | `[plugin]` | Stable package identity and plugin kind. | -| `[compat]` | Supported Relay range plus the native manifest API or worker protocol contract. A typed native plugin that uses the 0.8 SDK should declare a Relay range beginning at 0.8.0. | +| `[compat]` | Supported Relay range plus the native manifest API or worker protocol contract. A typed native plugin built with the current SDK must declare a Relay range beginning at 0.10.0. | | `[defaults]` and `[capabilities]` | Initial enabled state and the features the package declares, such as `plugin_native`, `plugin_worker`, and `config_schema`. | | `[config_schema]` | An optional JSON Schema path resolved relative to the manifest. Packages that use it also declare the `config_schema` capability. | | `[source]` | The artifact covered by integrity verification and, for managed Python workers, the package root used to create the environment. | @@ -26,8 +26,8 @@ equivalent binding configuration. For native packages, `compat.native_api = "1"` is the authored manifest contract. It is not the [C host-table ABI number](/build-plugins/native/native-abi-reference). The -current SDK requests ABI v5 and the host retains frozen v4, v3, and v2 compatibility -for older compiled plugins. For workers, declare `compat.worker_protocol = "grpc-v1"`; the +current SDK requires ABI v7, and Relay 0.10 rejects older compiled callback layouts. +For workers, declare `compat.worker_protocol = "grpc-v1"`; the [handshake](/build-plugins/workers/grpc-v1-protocol) still negotiates the exact protocol, surfaces, authentication token, and lifecycle at startup. @@ -48,7 +48,7 @@ id = "examples.python_grpc_worker" kind = "worker" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" worker_protocol = "grpc-v1" [defaults] @@ -81,7 +81,7 @@ id = "examples.rust_native_policy" kind = "rust_dynamic" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" native_api = "1" [defaults] @@ -113,7 +113,7 @@ id = "examples.rust_grpc_worker" kind = "worker" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" worker_protocol = "grpc-v1" [defaults] diff --git a/docs/build-plugins/workers/about.mdx b/docs/build-plugins/workers/about.mdx index e4557d2e6..8e7d516e1 100644 --- a/docs/build-plugins/workers/about.mdx +++ b/docs/build-plugins/workers/about.mdx @@ -18,6 +18,13 @@ SDK worker, regenerate custom protobuf bindings, and declare `compat.relay` begi `0.8.0`. The protocol remains named `grpc-v1`; this is a changed tool-result contract, not a new protocol family. +Relay 0.10 adds `LlmExecutionContext` before the continuation in every unary and +streaming LLM execution callback. Rebuild every SDK worker, regenerate custom +protobuf bindings, and use a `compat.relay` lower bound of `0.10.0`. The request +direction can expose decode/encode operations; unary callbacks can also decode +the complete response. Streaming callbacks do not receive a response codec. +The protocol remains named `grpc-v1`. + Workers can add a data schema, log severity, and semantic category to a mark, emit validated metric measurements, and read the current host-level runtime diagnostics. These diagnostics are ordered by code and are not attributed to a particular plugin. diff --git a/docs/build-plugins/workers/grpc-v1-protocol.mdx b/docs/build-plugins/workers/grpc-v1-protocol.mdx index 0230e7208..50933d0e6 100644 --- a/docs/build-plugins/workers/grpc-v1-protocol.mdx +++ b/docs/build-plugins/workers/grpc-v1-protocol.mdx @@ -16,6 +16,12 @@ changes the tool-result boundary to structural protobuf messages. Rebuild worker regenerate custom bindings, and declare `compat.relay` beginning at `0.8.0`; an earlier worker cannot decode the current `ToolNext` response or tool-execution outcome. +Relay 0.10 retains those protocol identifiers and adds execution codec context to +`LlmInvocation`. The field is protobuf-compatible, but the 0.10 Rust and Python +SDKs intentionally change every LLM execution callback to receive +`LlmExecutionContext` before its continuation. Rebuild workers, regenerate custom +bindings, and declare `compat.relay` beginning at `0.10.0`. + The protocol consists of the worker-facing service implemented by the plugin process and the host-runtime service implemented by Relay. These are the current service definitions, including cancellation, codecs, and streaming continuations: @@ -130,7 +136,8 @@ An invocation names the activation, invocation, registration, surface, optional continuation, captured scope, and token. Its payload is exactly one event, tool invocation, LLM invocation, or conditional middleware invocation. LLM sanitizer invocations additionally carry codec identity and an opaque invocation-scoped codec -capability. +capability. LLM execution invocations carry separate request and unary-response codec +context. Streaming execution carries only the request direction. For `EVENT_METADATA_INJECTOR`, Relay sends the immutable Event snapshot in `InvokeRequest.event`. The worker returns an `InvokeResponse.json` object containing @@ -171,6 +178,12 @@ message LlmInvocation { LlmSanitizeRequestContext request_sanitize_context = 9; LlmSanitizeResponseContext response_sanitize_context = 10; } + LlmExecutionCodecContext execution_codec_context = 11; +} + +message LlmExecutionCodecContext { + LlmSanitizeRequestContext request = 1; + LlmSanitizeResponseContext response = 2; } message ToolInvocation { @@ -252,6 +265,13 @@ identity deliberately withholds an ID. The capability ID is optional and invocation-scoped. A worker must treat it as a secret and must not use it after the owner invocation ends. +For an LLM execution invocation, `execution_codec_context.request` is required. +Its identity is present even when no codec resolves, and its capability ID is +present only when host codec operations are available. Unary execution also +includes `response`; streaming execution omits it entirely because Relay's +response codecs require a complete provider response. The worker SDK hides +both opaque capability IDs behind `resolve_codec()` proxies. + ## Host-Runtime Service | RPC | Contract | diff --git a/docs/build-plugins/workers/middleware-and-continuations.mdx b/docs/build-plugins/workers/middleware-and-continuations.mdx index 419d8b474..2d5446729 100644 --- a/docs/build-plugins/workers/middleware-and-continuations.mdx +++ b/docs/build-plugins/workers/middleware-and-continuations.mdx @@ -313,6 +313,7 @@ from collections.abc import AsyncIterator from nemo_relay_plugin import ( Json, LlmExecutionCallback, + LlmExecutionContext, LlmNext, LlmRequest, LlmStreamExecutionCallback, @@ -344,7 +345,13 @@ async def tool_execution( pending_marks=marks, ) -async def llm_execution(_name: str, request: LlmRequest, next_call: LlmNext) -> Json: +async def llm_execution( + _name: str, + request: LlmRequest, + context: LlmExecutionContext, + next_call: LlmNext, +) -> Json: + assert context.response_codec is not None content = request.get("content") repeat = isinstance(content, dict) and content.get("repeat_downstream") is True if repeat: @@ -359,8 +366,12 @@ async def llm_execution(_name: str, request: LlmRequest, next_call: LlmNext) -> return await next_call.call(request) async def llm_stream_execution( - _name: str, request: LlmRequest, next_call: LlmStreamNext + _name: str, + request: LlmRequest, + context: LlmExecutionContext, + next_call: LlmStreamNext, ) -> AsyncIterator[Json]: + assert context.response_codec is None async for chunk in next_call.call(request): if isinstance(chunk, dict): yield {**chunk, "plugin_stream": True} @@ -405,7 +416,8 @@ context.register_tool_execution_intercept( context.register_llm_execution_intercept( "documentation_llm_execution", config.execution.priority, - move |_model, request, next| async move { + move |_model, request, context, next| async move { + debug_assert!(context.response_codec().is_some()); if request.content .get("repeat_downstream") .and_then(Json::as_bool) @@ -426,7 +438,8 @@ context.register_llm_execution_intercept( context.register_llm_stream_execution_intercept( "documentation_llm_stream_execution", config.execution.priority, - move |_model, request, next| async move { + move |_model, request, context, next| async move { + debug_assert!(context.response_codec().is_none()); let downstream = next.call(request).await?; let mapped: JsonStream = Box::pin(downstream.map(|chunk| { chunk.map(|mut value| { @@ -443,6 +456,15 @@ context.register_llm_stream_execution_intercept( +`LlmExecutionContext` describes the codecs selected by the host for this +managed invocation. `request_codec` is always present as an identity-bearing +subcontext; `resolve_codec()` returns a request decode/encode proxy only when +Relay resolved a codec. Unary execution also provides `response_codec` with an +optional decode proxy. Streaming execution sets `response_codec` to `None` +instead of implying that individual chunks can be decoded as completed +responses. Codec proxies are invocation-scoped and must not be retained after +the callback or returned stream finishes. + The second unary result is awaited even though the first response is selected. That prevents an unobserved continuation from outliving the worker callback. In the stream case, each error remains an error item and no chunks are requested before the consumer diff --git a/docs/build-plugins/workers/python.mdx b/docs/build-plugins/workers/python.mdx index d756f0e0a..950b6be8e 100644 --- a/docs/build-plugins/workers/python.mdx +++ b/docs/build-plugins/workers/python.mdx @@ -129,7 +129,7 @@ id = "examples.python_grpc_worker" kind = "worker" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" worker_protocol = "grpc-v1" [capabilities] @@ -150,8 +150,8 @@ runtime = "python" entrypoint = "nemo_relay_python_grpc_worker_example.worker:main" ``` -The manifest requires Relay 0.9 because this worker uses the context-aware tool execution -callback contract. The protocol identifier remains `grpc-v1`. `ToolNext.call()` returns +The manifest requires Relay 0.10 because this worker uses the context-aware tool execution +contract and the `LlmExecutionContext` callback shape. The protocol identifier remains `grpc-v1`. `ToolNext.call()` returns `ToolExecutionResult`, whose `result` contains the application payload and whose optional `annotation` remains adjacent opaque metadata. See the [protocol reference](/build-plugins/workers/grpc-v1-protocol) for the complete boundary. diff --git a/docs/build-plugins/workers/rust.mdx b/docs/build-plugins/workers/rust.mdx index 9df51f897..629e0fe67 100644 --- a/docs/build-plugins/workers/rust.mdx +++ b/docs/build-plugins/workers/rust.mdx @@ -68,7 +68,7 @@ id = "examples.rust_grpc_worker" kind = "worker" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" worker_protocol = "grpc-v1" [defaults] @@ -91,8 +91,8 @@ runtime = "rust" entrypoint = "target/debug/" ``` -The manifest requires Relay 0.9 because this worker uses the context-aware tool execution -callback contract. The protocol identifier remains `grpc-v1`. `ToolNext::call` returns +The manifest requires Relay 0.10 because this worker uses the context-aware tool execution +contract and the `LlmExecutionContext` callback shape. The protocol identifier remains `grpc-v1`. `ToolNext::call` returns `ToolExecutionResult`, preserving an application payload and optional opaque annotation separately from Relay-owned pending marks. Rebuild the worker when adopting this SDK release. diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index 424737835..9eca5c7b9 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -11,6 +11,40 @@ upgrade actions as they are identified during the 0.10 development cycle. ## Upgrade to NeMo Relay 0.10 +### Update LLM Execution Intercepts + +Relay 0.10 adds `LlmExecutionContext` immediately before the continuation in +every unary and streaming LLM execution-intercept callback. Update callbacks +as follows: + +| Surface | Relay 0.10 callback shape | +|---|---| +| Rust, Python, native plugin, Rust worker, Python worker | `(name, request, context, next)` | +| Node.js, Go | `(request, context, next)` | +| C | `(user_data, name, request, context, next, next_ctx)` | + +The request direction reports the selected codec and exposes decode and encode +operations when Relay resolved one. Unary execution also exposes response +identity and decode. Streaming execution has no response codec because chunks +are not complete provider responses. + +This is a source and binary compatibility break for execution-intercept users: + +- Recompile language-binding consumers that register an LLM execution + intercept and update their callback argument order. +- Rebuild every native plugin against the 0.10 SDK. Native plugins continue to + declare `compat.native_api = "1"`, but must set `compat.relay` to + `">=0.10.0,<1.0"` or another range that excludes Relay 0.9. Relay 0.10 uses + internal native ABI v7 and rejects older compiled layouts. +- Regenerate and rebuild gRPC workers that register an LLM execution intercept. + The protocol remains `grpc-v1`, but those workers must also exclude Relay 0.9 + in `compat.relay`. + +LLM request intercepts and conditional middleware do not change. For the full +context and lifetime contract, refer to [Build Plugins](/build-plugins/about), +[Native ABI Reference](/build-plugins/native/native-abi-reference), and +[gRPC-v1 Worker Protocol](/build-plugins/workers/grpc-v1-protocol). + ### Refresh Codex Provider Routing After upgrading the CLI, refresh personal Relay-managed Codex installations: diff --git a/examples/language-binding-plugin/node/main.mjs b/examples/language-binding-plugin/node/main.mjs index fa7ec9a9c..7c96a0234 100644 --- a/examples/language-binding-plugin/node/main.mjs +++ b/examples/language-binding-plugin/node/main.mjs @@ -303,14 +303,18 @@ export const documentationPlugin = { pendingMarks: execution.emit_pending_marks ? [{ name: 'documentation-plugin.tool-complete' }] : [], }; }); - context.registerLlmExecutionIntercept('llm-execution', execution.priority, async (request, next) => + context.registerLlmExecutionIntercept('llm-execution', execution.priority, async (request, _context, next) => next(request), ); - context.registerLlmStreamExecutionIntercept('llm-stream', execution.priority, async function* (request, next) { - for await (const chunk of await next(request)) { - yield { ...chunk, plugin_stream: true }; - } - }); + context.registerLlmStreamExecutionIntercept( + 'llm-stream', + execution.priority, + async function* (request, _context, next) { + for await (const chunk of await next(request)) { + yield { ...chunk, plugin_stream: true }; + } + }, + ); } }, }; @@ -341,8 +345,7 @@ export async function main() { plugin.register('documentation-plugin', documentationPlugin); console.log('registered:', plugin.listKinds()); const invalid = plugin.validate(config('invalid')).config.diagnostics; - const disabledInvalid = plugin.validate(config('invalid', false)).config - .diagnostics; + const disabledInvalid = plugin.validate(config('invalid', false)).config.diagnostics; if (disabledInvalid[0]?.code !== 'documentation-plugin.unsupported_mode') { throw new Error('disabled invalid configuration must still be validated'); } diff --git a/examples/language-binding-plugin/node/test-plugin.mjs b/examples/language-binding-plugin/node/test-plugin.mjs index 9a729720c..9a64b8e3c 100644 --- a/examples/language-binding-plugin/node/test-plugin.mjs +++ b/examples/language-binding-plugin/node/test-plugin.mjs @@ -205,11 +205,14 @@ test('LLM policy blocks the configured model', () => { test('LLM stream chunks are transformed', async () => { const intercept = registeredCallbacks().get('registerLlmStreamExecutionIntercept'); - const transformed = intercept({ headers: {}, content: { model: 'allowed-model' } }, async () => - (async function* () { - yield { chunk: 1 }; - yield { chunk: 2 }; - })(), + const transformed = intercept( + { headers: {}, content: { model: 'allowed-model' } }, + { requestCodec: {}, responseCodec: null }, + async () => + (async function* () { + yield { chunk: 1 }; + yield { chunk: 2 }; + })(), ); const chunks = []; for await (const chunk of transformed) chunks.push(chunk); diff --git a/examples/language-binding-plugin/python/main.py b/examples/language-binding-plugin/python/main.py index df9b57f1c..7c554b2fe 100644 --- a/examples/language-binding-plugin/python/main.py +++ b/examples/language-binding-plugin/python/main.py @@ -337,7 +337,9 @@ async def runtime_events( context.register_tool_execution_intercept("runtime-events", 0, runtime_events) async def stream_request( + _name: str, _request: nemo_relay.LLMRequest, + _context: nemo_relay.LlmExecutionContext, next_call: Callable[[nemo_relay.LLMRequest], Awaitable[AsyncIterator[nemo_relay.Json]]], ) -> AsyncIterator[nemo_relay.Json]: async for chunk in await next_call(_request): @@ -366,6 +368,7 @@ async def tool_execution( async def llm_execution( _name: str, request: nemo_relay.LLMRequest, + _context: nemo_relay.LlmExecutionContext, next_call: Callable[[nemo_relay.LLMRequest], Awaitable[nemo_relay.Json]], ) -> nemo_relay.Json: return await next_call(request) diff --git a/examples/language-binding-plugin/rust/src/lib.rs b/examples/language-binding-plugin/rust/src/lib.rs index 76945a0ce..68f55a002 100644 --- a/examples/language-binding-plugin/rust/src/lib.rs +++ b/examples/language-binding-plugin/rust/src/lib.rs @@ -293,14 +293,14 @@ impl Plugin for DocumentationPlugin { context.register_llm_execution_intercept( "llm-execution", settings.execution.priority, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { Box::pin(async move { next(request).await }) }), )?; context.register_llm_stream_execution_intercept( "llm-stream", settings.execution.priority, - Arc::new(move |_name, request, next| { + Arc::new(move |_name, request, _context, next| { Box::pin(async move { let stream = next(request).await?; Ok(LlmJsonStream::new(stream.map(|chunk| { diff --git a/examples/python-grpc-worker-plugin/README.md b/examples/python-grpc-worker-plugin/README.md index eb14a4815..5c50c6aca 100644 --- a/examples/python-grpc-worker-plugin/README.md +++ b/examples/python-grpc-worker-plugin/README.md @@ -11,9 +11,10 @@ It validates the shared documentation configuration, registers every safe invocation-scoped codec proxies, transforms streams lazily, and cleans up marks, scopes, isolated stacks, and cancelled tasks. -The worker targets the Relay 0.8 `grpc-v1` result contract. Its tool continuation returns -`ToolExecutionResult`. Its execution intercept preserves the application result, carries -the upstream annotation under worker metadata, and adds Relay-owned pending marks. +The worker targets Relay 0.10 while retaining the `grpc-v1` protocol name. Its tool +continuation returns `ToolExecutionResult`. Its LLM execution callbacks receive +directional codec context before the continuation; unary execution can decode a complete +response, while streaming execution exposes request codec operations only. Run the example's own test project from this directory: diff --git a/examples/python-grpc-worker-plugin/nemo_relay_python_grpc_worker_example/worker.py b/examples/python-grpc-worker-plugin/nemo_relay_python_grpc_worker_example/worker.py index cc9eeffbe..cf9c67d6b 100644 --- a/examples/python-grpc-worker-plugin/nemo_relay_python_grpc_worker_example/worker.py +++ b/examples/python-grpc-worker-plugin/nemo_relay_python_grpc_worker_example/worker.py @@ -290,7 +290,7 @@ async def tool_execution(context: ToolExecutionContext, next_call: Any) -> ToolE pending_marks=marks, ) - async def llm_execution(_name: str, request: dict[str, Any], next_call: Any) -> Json: + async def llm_execution(_name: str, request: dict[str, Any], _context: Any, next_call: Any) -> Json: content = request.get("content") repeat = isinstance(content, dict) and content.get("repeat_downstream") is True if repeat: @@ -305,6 +305,7 @@ async def llm_execution(_name: str, request: dict[str, Any], next_call: Any) -> async def llm_stream_execution( _name: str, request: dict[str, Any], + _context: Any, next_call: Any, ) -> AsyncIterator[Json]: async for chunk in next_call.call(request): diff --git a/examples/python-grpc-worker-plugin/relay-plugin.toml b/examples/python-grpc-worker-plugin/relay-plugin.toml index 354c50554..bb627d138 100644 --- a/examples/python-grpc-worker-plugin/relay-plugin.toml +++ b/examples/python-grpc-worker-plugin/relay-plugin.toml @@ -8,7 +8,7 @@ id = "examples.python_grpc_worker" kind = "worker" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" worker_protocol = "grpc-v1" [defaults] @@ -25,7 +25,7 @@ manifest_root = "." artifact = "nemo_relay_python_grpc_worker_example/worker.py" [integrity] -sha256 = "sha256:3cd034b104d1d6aec431aad732d7ef2e1a66c7c68bf289f4b2fbac8e881effa3" +sha256 = "sha256:2797df6981b698f90342546a2d0bd725c559355f35c07edd92b53cae49e3ed1c" [load] runtime = "python" diff --git a/examples/python-grpc-worker-plugin/tests/test_worker.py b/examples/python-grpc-worker-plugin/tests/test_worker.py index 2bb19edaf..c21996fc2 100644 --- a/examples/python-grpc-worker-plugin/tests/test_worker.py +++ b/examples/python-grpc-worker-plugin/tests/test_worker.py @@ -27,7 +27,16 @@ pytest.importorskip("grpc") -from nemo_relay_plugin import PluginContext, PluginRuntime, ToolExecutionContext, ToolExecutionResult # noqa: E402 +from nemo_relay_plugin import ( # noqa: E402 + LlmCodecIdentity, + LlmExecutionContext, + LlmSanitizeRequestContext, + LlmSanitizeResponseContext, + PluginContext, + PluginRuntime, + ToolExecutionContext, + ToolExecutionResult, +) EXAMPLE_ROOT = Path(__file__).parents[1] MODULE_NAME = "nemo_relay_python_grpc_worker_example.worker" @@ -88,6 +97,12 @@ def callback(context: MagicMock, method: str, name: str | None = None) -> Any: return calls[0].args[1] +def execution_context(*, streaming: bool = False) -> LlmExecutionContext: + request = LlmSanitizeRequestContext(LlmCodecIdentity("none")) + response = None if streaming else LlmSanitizeResponseContext(LlmCodecIdentity("none")) + return LlmExecutionContext(request_codec=request, response_codec=response) + + def test_manifest_digest_matches_worker_source() -> None: manifest = read_manifest() artifact = EXAMPLE_ROOT / manifest["source"]["artifact"] @@ -99,7 +114,7 @@ def test_manifest_digest_matches_worker_source() -> None: def test_manifest_declares_current_worker_protocol() -> None: manifest = read_manifest() - assert manifest["compat"] == {"relay": ">=0.9.0,<1.0", "worker_protocol": "grpc-v1"} + assert manifest["compat"] == {"relay": ">=0.10.0,<1.0", "worker_protocol": "grpc-v1"} def test_schema_declares_only_supported_groups() -> None: @@ -465,6 +480,7 @@ async def test_llm_execution_can_repeat_continuation(example: Any) -> None: result = await intercept( "allowed-model", {"headers": {}, "content": {"repeat_downstream": True}}, + execution_context(), next_call, ) @@ -481,6 +497,7 @@ async def test_repeated_llm_continuation_ignores_the_second_failure(example: Any result = await intercept( "allowed-model", {"headers": {}, "content": {"repeat_downstream": True}}, + execution_context(), next_call, ) @@ -499,7 +516,12 @@ async def downstream() -> Any: next_call = MagicMock() next_call.call.return_value = downstream() - stream = intercept("allowed-model", {"headers": {}, "content": {}}, next_call) + stream = intercept( + "allowed-model", + {"headers": {}, "content": {}}, + execution_context(streaming=True), + next_call, + ) assert [item async for item in stream] == [ {"chunk": 1, "plugin_stream": True}, diff --git a/examples/rust-grpc-worker-plugin/README.md b/examples/rust-grpc-worker-plugin/README.md index ce5951593..bb4aa8883 100644 --- a/examples/rust-grpc-worker-plugin/README.md +++ b/examples/rust-grpc-worker-plugin/README.md @@ -10,6 +10,11 @@ guide. It validates the shared documentation configuration, registers every safe `grpc-v1` surface, exercises continuations and lazy streams, uses invocation-scoped codecs, and demonstrates marks and scope-stack cleanup. +The worker targets Relay 0.10 while retaining the `grpc-v1` protocol name. Unary and +streaming LLM execution callbacks receive directional codec context before their +continuation. Streaming deliberately has no response codec because chunks are not +complete provider responses. + Run `cargo test` and `cargo build` from this directory. The configuration and schema tests are order-independent. The lifecycle test builds a fresh worker, materializes a digest-checked manifest, activates it through `grpc-v1`, runs diff --git a/examples/rust-grpc-worker-plugin/relay-plugin.toml b/examples/rust-grpc-worker-plugin/relay-plugin.toml index a28b2d535..582050db2 100644 --- a/examples/rust-grpc-worker-plugin/relay-plugin.toml +++ b/examples/rust-grpc-worker-plugin/relay-plugin.toml @@ -8,7 +8,7 @@ id = "examples.rust_grpc_worker" kind = "worker" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" worker_protocol = "grpc-v1" [defaults] diff --git a/examples/rust-grpc-worker-plugin/src/lib.rs b/examples/rust-grpc-worker-plugin/src/lib.rs index 081ef2530..8bb0fa75e 100644 --- a/examples/rust-grpc-worker-plugin/src/lib.rs +++ b/examples/rust-grpc-worker-plugin/src/lib.rs @@ -308,7 +308,7 @@ fn register_execution(context: &mut PluginContext, config: &ExampleConfig) { context.register_llm_execution_intercept( "documentation_llm_execution", config.execution.priority, - move |_model, request, next| async move { + move |_model, request, _context, next| async move { if request .content .get("repeat_downstream") @@ -327,7 +327,7 @@ fn register_execution(context: &mut PluginContext, config: &ExampleConfig) { context.register_llm_stream_execution_intercept( "documentation_llm_stream_execution", config.execution.priority, - move |_model, request, next| async move { + move |_model, request, _context, next| async move { let stream = next.call(request).await?; let mapped: JsonStream = Box::pin(stream.map(|chunk| { chunk.map(|chunk| match chunk { diff --git a/examples/rust-grpc-worker-plugin/tests/config.rs b/examples/rust-grpc-worker-plugin/tests/config.rs index fd67302c4..4f6399ef8 100644 --- a/examples/rust-grpc-worker-plugin/tests/config.rs +++ b/examples/rust-grpc-worker-plugin/tests/config.rs @@ -193,7 +193,7 @@ fn assert_schema_defaults(schema: &Json, value: &Json, path: &str) { #[test] fn manifest_uses_the_rust_worker_load_contract() { let manifest = include_str!("../relay-plugin.toml"); - assert!(manifest.contains("relay = \">=0.9.0,<1.0\"")); + assert!(manifest.contains("relay = \">=0.10.0,<1.0\"")); assert!(manifest.contains("worker_protocol = \"grpc-v1\"")); assert!(manifest.contains("runtime = \"rust\"")); assert!(manifest.contains("entrypoint = \"target/debug/\"")); diff --git a/examples/rust-grpc-worker-plugin/tests/lifecycle.rs b/examples/rust-grpc-worker-plugin/tests/lifecycle.rs index 176584857..66e1085da 100644 --- a/examples/rust-grpc-worker-plugin/tests/lifecycle.rs +++ b/examples/rust-grpc-worker-plugin/tests/lifecycle.rs @@ -237,7 +237,7 @@ id = "{PLUGIN_ID}" kind = "worker" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" worker_protocol = "grpc-v1" [defaults] diff --git a/examples/rust-native-plugin/README.md b/examples/rust-native-plugin/README.md index cba94f590..3f3e63326 100644 --- a/examples/rust-native-plugin/README.md +++ b/examples/rust-native-plugin/README.md @@ -11,6 +11,11 @@ helpers live in separate source modules. Together they register the subscriber, all three event sanitizers, five tool surfaces, and six LLM surfaces exposed by the current typed 0.10.0 SDK. +Relay 0.10 uses native ABI v7. Every LLM execution callback receives directional codec +context before its continuation; streaming execution exposes request codec operations +but no response decoder. The manifest continues to declare `native_api = "1"`, and its +Relay lower bound is `0.10.0` because older compiled callback layouts are rejected. + Run the focused tests and build the shared library from this directory. The configuration tests isolate validation and schema contracts. The lifecycle test builds a fresh `cdylib`, materializes a digest-checked manifest, activates the diff --git a/examples/rust-native-plugin/relay-plugin.toml b/examples/rust-native-plugin/relay-plugin.toml index 00a710683..1d15545f1 100644 --- a/examples/rust-native-plugin/relay-plugin.toml +++ b/examples/rust-native-plugin/relay-plugin.toml @@ -8,7 +8,7 @@ id = "examples.rust_native_policy" kind = "rust_dynamic" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" native_api = "1" [defaults] diff --git a/examples/rust-native-plugin/src/execution.rs b/examples/rust-native-plugin/src/execution.rs index db07adaad..f3b789181 100644 --- a/examples/rust-native-plugin/src/execution.rs +++ b/examples/rust-native-plugin/src/execution.rs @@ -64,7 +64,7 @@ pub(crate) fn register( context.register_llm_execution_intercept( "documentation_llm_execution", config.execution.priority, - move |_name, request, next| async move { + move |_name, request, _context, next| async move { if request .content .get("repeat_downstream") @@ -86,7 +86,7 @@ pub(crate) fn register( context.register_llm_stream_execution_intercept( "documentation_llm_stream_execution", config.execution.priority, - move |_name, request, next| async move { + move |_name, request, _context, next| async move { let stream = next.call(request).await?; let stream: LlmJsonAsyncStream = Box::pin(stream.map(|chunk| { chunk.map(|chunk| match chunk { diff --git a/examples/rust-native-plugin/tests/lifecycle.rs b/examples/rust-native-plugin/tests/lifecycle.rs index e62e8a52d..498541900 100644 --- a/examples/rust-native-plugin/tests/lifecycle.rs +++ b/examples/rust-native-plugin/tests/lifecycle.rs @@ -205,7 +205,7 @@ id = "{PLUGIN_ID}" kind = "rust_dynamic" [compat] -relay = ">=0.9.0,<1.0" +relay = ">=0.10.0,<1.0" native_api = "1" [defaults] diff --git a/go/nemo_relay/adaptive_plugin_test.go b/go/nemo_relay/adaptive_plugin_test.go index 0e88eaa97..84175edc0 100644 --- a/go/nemo_relay/adaptive_plugin_test.go +++ b/go/nemo_relay/adaptive_plugin_test.go @@ -164,7 +164,7 @@ func registerLifecycleInterceptors(ctx *PluginContext, pluginKind string) error return ctx.RegisterLlmExecutionIntercept( "llm_exec", 7, - func(requestJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(requestJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { responseJSON, err := next(requestJSON) if err != nil { return nil, err @@ -195,7 +195,7 @@ func registerLifecycleStreamPlugin(streamPluginKind string) error { return ctx.RegisterLlmStreamExecutionIntercept( "llm_stream_exec", 7, - func(requestJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(requestJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { responseJSON, err := next(requestJSON) if err != nil { return nil, err @@ -555,12 +555,12 @@ func TestPluginFuncsAndClosedContextBranches(t *testing.T) { return closed.RegisterToolRequestIntercept("tool_request", 1, false, func(name string, args json.RawMessage) json.RawMessage { return args }) }}, {"llm execution", func() error { - return closed.RegisterLlmExecutionIntercept("llm_exec", 1, func(request json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + return closed.RegisterLlmExecutionIntercept("llm_exec", 1, func(request json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { return next(request) }) }}, {"llm stream", func() error { - return closed.RegisterLlmStreamExecutionIntercept("llm_stream_exec", 1, func(request json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + return closed.RegisterLlmStreamExecutionIntercept("llm_stream_exec", 1, func(request json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { return next(request) }) }}, diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index 7f4094922..d887c379d 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -38,6 +38,10 @@ typedef struct NemoRelayLlmSanitizeResponseContext { const char* codec_id; const FfiLlmSanitizeResponseCodec* codec; } NemoRelayLlmSanitizeResponseContext; +typedef struct NemoRelayLlmExecutionContext { + NemoRelayLlmSanitizeRequestContext request_codec; + const NemoRelayLlmSanitizeResponseContext* response_codec; +} NemoRelayLlmExecutionContext; typedef void (*NemoRelayFreeFn)(void* user_data); typedef char* (*NemoRelayToolSanitizeFn)(void* user_data, const char* name, const char* args_json); @@ -56,7 +60,7 @@ typedef struct FfiPluginContext FfiPluginContext; typedef char* (*NemoRelayToolExecNextFn)(const char* args_json, void* next_ctx); typedef char* (*NemoRelayToolExecInterceptCb)(void* user_data, const char* context_json, NemoRelayToolExecNextFn next_fn, void* next_ctx); 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); +typedef char* (*NemoRelayLlmExecInterceptCb)(void* user_data, const char* name, const char* native_json, NemoRelayLlmExecutionContext context, NemoRelayLlmExecNextFn next_fn, void* next_ctx); // Helper to call the tool exec next function pointer from Go static inline char* callToolExecNext(NemoRelayToolExecNextFn next_fn, const char* args_json, void* next_ctx) { @@ -246,14 +250,26 @@ type LLMSanitizeResponseContext struct { resolved *LLMResponseSanitizeCodec } +// LLMExecutionContext provides invocation-scoped codec access to an LLM +// execution intercept. RequestCodec is always present. ResponseCodec is +// available for unary execution and nil for streaming execution. +type LLMExecutionContext struct { + RequestCodec LLMSanitizeRequestContext + ResponseCodec *LLMSanitizeResponseContext +} + // ResolveCodec returns the active callback-scoped response codec, if any. func (context LLMSanitizeResponseContext) ResolveCodec() *LLMResponseSanitizeCodec { return context.resolved } -// ErrLLMSanitizeCodecExpired is returned when a callback-scoped codec -// capability is used after its sanitizer callback has returned. -var ErrLLMSanitizeCodecExpired = errors.New("LLM sanitizer codec capability is no longer active") +// ErrLLMCodecExpired is returned when a callback-scoped codec capability is +// used after its sanitizer or execution callback has returned. +var ErrLLMCodecExpired = errors.New("LLM codec capability is no longer active") + +// ErrLLMSanitizeCodecExpired is retained as an alias for callers that already +// compare the sanitizer-specific error value. +var ErrLLMSanitizeCodecExpired = ErrLLMCodecExpired type llmSanitizeCodecInvocation struct { mu sync.RWMutex @@ -327,7 +343,7 @@ type LLMExecutionFunc func(requestJSON json.RawMessage) (json.RawMessage, error) // as JSON and a `next` function. Call `next` to invoke the next intercept in // the chain (or the original LLM implementation if this is the innermost // intercept). Skip calling `next` to short-circuit the chain entirely. -type LLMExecutionInterceptFunc func(requestJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) +type LLMExecutionInterceptFunc func(requestJSON json.RawMessage, context LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) // CollectorFunc is a callback invoked with each intercepted chunk during a // streaming LLM response. It is used to accumulate chunks on the Go side for @@ -917,9 +933,20 @@ func goToolExecInterceptTrampoline(userData unsafe.Pointer, contextJSON *C.char, } //export goLlmExecInterceptTrampoline -func goLlmExecInterceptTrampoline(userData unsafe.Pointer, nativeJSON *C.char, nextFn C.NemoRelayLlmExecNextFn, nextCtx unsafe.Pointer) *C.char { +func goLlmExecInterceptTrampoline(userData unsafe.Pointer, name *C.char, nativeJSON *C.char, context C.NemoRelayLlmExecutionContext, nextFn C.NemoRelayLlmExecNextFn, nextCtx unsafe.Pointer) *C.char { + // Go intentionally omits the interceptor name from its public callback shape. + _ = name fn := lookupClosure(userData).(LLMExecutionInterceptFunc) goJSON := json.RawMessage(C.GoString(nativeJSON)) + invocation := newLLMSanitizeCodecInvocation() + defer invocation.invalidate() + goContext := LLMExecutionContext{ + RequestCodec: llmSanitizeRequestContextFromC(context.request_codec, invocation), + } + if context.response_codec != nil { + responseCodec := llmSanitizeResponseContextFromC(*context.response_codec, invocation) + goContext.ResponseCodec = &responseCodec + } goNext := func(reqJSON json.RawMessage) (json.RawMessage, error) { cJSON := C.CString(string(reqJSON)) @@ -933,7 +960,7 @@ func goLlmExecInterceptTrampoline(userData unsafe.Pointer, nativeJSON *C.char, n return json.RawMessage(C.GoString(result)), nil } - result, err := fn(goJSON, goNext) + result, err := fn(goJSON, goContext, goNext) if err != nil { setLastErrorMessage(err.Error()) return nil diff --git a/go/nemo_relay/intercepts/intercepts_test.go b/go/nemo_relay/intercepts/intercepts_test.go index 4d909b494..fb32e98d4 100644 --- a/go/nemo_relay/intercepts/intercepts_test.go +++ b/go/nemo_relay/intercepts/intercepts_test.go @@ -117,7 +117,7 @@ func runGlobalLLMInterceptShorthandChecks(t *testing.T) { } if err := intercepts.RegisterLlmExecution("intercepts_llm_exec", 1, - func(request json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(request json.RawMessage, _ nemo_relay.LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { result, err := next(request) if err != nil { return nil, err @@ -154,7 +154,7 @@ func runGlobalLLMInterceptShorthandChecks(t *testing.T) { } if err := intercepts.RegisterLlmStreamExecution("intercepts_llm_stream", 1, - func(request json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(request json.RawMessage, _ nemo_relay.LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { return next(request) }, ); err != nil { @@ -208,14 +208,14 @@ func runScopeLocalLLMInterceptShorthandChecks(t *testing.T, scopeUUID string) { t.Fatalf("ScopeRegisterLlmRequest failed: %v", err) } if err := intercepts.ScopeRegisterLlmExecution(scopeUUID, "intercepts_scope_llm_exec", 1, - func(request json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(request json.RawMessage, _ nemo_relay.LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { return next(request) }, ); err != nil { t.Fatalf("ScopeRegisterLlmExecution failed: %v", err) } if err := intercepts.ScopeRegisterLlmStreamExecution(scopeUUID, "intercepts_scope_llm_stream", 1, - func(request json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(request json.RawMessage, _ nemo_relay.LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { return next(request) }, ); err != nil { diff --git a/go/nemo_relay/llm_test.go b/go/nemo_relay/llm_test.go index a663808d3..16398ef67 100644 --- a/go/nemo_relay/llm_test.go +++ b/go/nemo_relay/llm_test.go @@ -735,20 +735,133 @@ func TestLlmRequestInterceptRegisterDeregister(t *testing.T) { } func TestLlmExecutionInterceptRegisterDeregister(t *testing.T) { + observed := false err := RegisterLlmExecutionIntercept("go_llm_exec", 1, - func(nativeJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(nativeJSON json.RawMessage, context LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + if context.ResponseCodec == nil { + t.Fatal("unary execution context is missing its response codec") + } + if context.RequestCodec.Codec.CodecKind != LLMCodecNone || context.RequestCodec.ResolveCodec() != nil { + t.Fatalf("unexpected absent request codec context: %#v", context.RequestCodec) + } + if context.ResponseCodec.Codec.CodecKind != LLMCodecNone || context.ResponseCodec.ResolveCodec() != nil { + t.Fatalf("unexpected absent response codec context: %#v", context.ResponseCodec) + } + observed = true return next(nativeJSON) }, ) if err != nil { t.Fatalf(llmRegisterFailed, err) } - DeregisterLlmExecutionIntercept("go_llm_exec") + defer DeregisterLlmExecutionIntercept("go_llm_exec") + + response, err := LlmCallExecute( + "go_llm_execution_context_absent", + makeRequest(), + func(json.RawMessage) (json.RawMessage, error) { + return json.RawMessage(`{"ok":true}`), nil + }, + ) + if err != nil { + t.Fatalf(llmCallExecuteFailed, err) + } + if string(response) != `{"ok":true}` { + t.Fatalf("unexpected response: %s", response) + } + if !observed { + t.Fatal("execution intercept did not observe the absent codec context") + } +} + +func TestLlmExecutionInterceptResolvesDirectionalCodecs(t *testing.T) { + const interceptName = "go_llm_execution_codec_context" + _ = DeregisterLlmExecutionIntercept(interceptName) + defer DeregisterLlmExecutionIntercept(interceptName) + + var ( + retainedRequestCodec *LLMRequestSanitizeCodec + retainedResponseCodec *LLMResponseSanitizeCodec + retainedRequest LLMRequestDTO + ) + err := RegisterLlmExecutionIntercept( + interceptName, + 1, + func(nativeJSON json.RawMessage, context LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + if context.RequestCodec.Codec.CodecKind != LLMCodecOpaque { + return nil, fmt.Errorf("unexpected request codec identity: %#v", context.RequestCodec.Codec) + } + if context.ResponseCodec == nil || + context.ResponseCodec.Codec.CodecKind != LLMCodecBuiltin || + context.ResponseCodec.Codec.CodecID == nil || + *context.ResponseCodec.Codec.CodecID != "openai_chat" { + return nil, fmt.Errorf("unexpected response codec context: %#v", context.ResponseCodec) + } + requestCodec := context.RequestCodec.ResolveCodec() + responseCodec := context.ResponseCodec.ResolveCodec() + if requestCodec == nil || responseCodec == nil { + return nil, errors.New("execution codec capability did not resolve") + } + var request LLMRequestDTO + if err := json.Unmarshal(nativeJSON, &request); err != nil { + return nil, err + } + annotated, err := requestCodec.Decode(request) + if err != nil { + return nil, err + } + encoded, err := requestCodec.Encode(annotated, request) + if err != nil { + return nil, err + } + encodedJSON, err := json.Marshal(encoded) + if err != nil { + return nil, err + } + response, err := next(encodedJSON) + if err != nil { + return nil, err + } + if _, err := responseCodec.Decode(response); err != nil { + return nil, err + } + retainedRequestCodec = requestCodec + retainedResponseCodec = responseCodec + retainedRequest = request + return response, nil + }, + ) + if err != nil { + t.Fatalf(llmRegisterFailed, err) + } + + response, err := LlmCallExecute( + "execution_codec_context", + makeRequest(), + requireEncodedModelExecutor(t), + WithLLMCodec(llmRequestResponseCodec()), + WithLLMResponseCodec(NewOpenAIChatCodec()), + ) + if err != nil { + t.Fatalf(llmCallExecuteFailed, err) + } + if retainedRequestCodec == nil || retainedResponseCodec == nil { + t.Fatal("execution intercept did not retain both codec capabilities") + } + if _, err := retainedRequestCodec.Decode(retainedRequest); !errors.Is(err, ErrLLMCodecExpired) { + t.Fatalf("retained request codec must expire after execution callback, got %v", err) + } + if _, err := retainedResponseCodec.Decode(response); !errors.Is(err, ErrLLMCodecExpired) { + t.Fatalf("retained response codec must expire after execution callback, got %v", err) + } } func TestLlmStreamExecutionInterceptRegisterDeregister(t *testing.T) { err := RegisterLlmStreamExecutionIntercept("go_llm_sexec", 1, - func(nativeJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(nativeJSON json.RawMessage, context LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + if context.ResponseCodec != nil { + t.Fatal("stream execution context unexpectedly exposes a response codec") + } return next(nativeJSON) }, ) @@ -762,7 +875,7 @@ func TestLlmStreamExecutionInterceptCanCallNext(t *testing.T) { request := makeRequest() err := RegisterLlmStreamExecutionIntercept("go_llm_stream_exec_next", 1, - func(nativeJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(nativeJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { nextResult, err := next(nativeJSON) if err != nil { return nil, err @@ -839,7 +952,7 @@ func TestLlmRequestInterceptModifies(t *testing.T) { func TestLlmExecutionInterceptReplaces(t *testing.T) { RegisterLlmExecutionIntercept("go_llm_exec_rep", 1, - func(nativeJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(nativeJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { return json.RawMessage(`{"from_intercept": true}`), nil }, ) @@ -888,7 +1001,7 @@ func TestLlmCallableErrorPropagation(t *testing.T) { func TestLlmFullPipelineInterceptsAndExecute(t *testing.T) { // Register an execution intercept RegisterLlmExecutionIntercept("go_llm_pipe_exec_int", 1, - func(nativeJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(nativeJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { result, err := next(nativeJSON) if err != nil { return nil, err @@ -1021,7 +1134,7 @@ func TestLlmConditionalGuardrailSelectiveReject(t *testing.T) { func TestLlmExecutionInterceptWrapsCallable(t *testing.T) { RegisterLlmExecutionIntercept("go_llm_wrap_exec", 1, - func(nativeJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(nativeJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { result, err := next(nativeJSON) if err != nil { return nil, err @@ -1057,7 +1170,7 @@ func TestLlmExecutionInterceptWrapsCallable(t *testing.T) { func TestLlmExecutionInterceptSeesNextError(t *testing.T) { RegisterLlmExecutionIntercept("go_llm_wrap_exec_err", 1, - func(nativeJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(nativeJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { return next(nativeJSON) }, ) diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index 1f419679b..485c83774 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -42,6 +42,7 @@ typedef struct FfiLlmSanitizeRequestCodec FfiLlmSanitizeRequestCodec; typedef struct FfiLlmSanitizeResponseCodec FfiLlmSanitizeResponseCodec; typedef struct NemoRelayLlmSanitizeRequestContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeRequestCodec* codec; } NemoRelayLlmSanitizeRequestContext; typedef struct NemoRelayLlmSanitizeResponseContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeResponseCodec* codec; } NemoRelayLlmSanitizeResponseContext; +typedef struct NemoRelayLlmExecutionContext { NemoRelayLlmSanitizeRequestContext request_codec; const NemoRelayLlmSanitizeResponseContext* response_codec; } NemoRelayLlmExecutionContext; typedef void (*NemoRelayFreeFn)(void* user_data); @@ -173,7 +174,7 @@ typedef int32_t (*NemoRelayLlmRequestInterceptCb)(void* user_data, const char* n 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_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); +typedef char* (*NemoRelayLlmExecInterceptCb)(void* user_data, const char* name, const char* native_json, NemoRelayLlmExecutionContext context, 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_deregister_llm_execution_intercept(const char* name); @@ -328,7 +329,7 @@ extern char* goLlmConditionalTrampoline(void*, const FfiLLMRequest*); extern char* goLlmExecTrampoline(void*, const char*); extern char* goToolExecInterceptTrampoline(void*, const char*, NemoRelayToolExecNextFn, void*); extern char* goToolExecInterceptContextTrampoline(void*, const char*, NemoRelayToolExecNextFn, void*); -extern char* goLlmExecInterceptTrampoline(void*, const char*, NemoRelayLlmExecNextFn, void*); +extern char* goLlmExecInterceptTrampoline(void*, const char*, const char*, NemoRelayLlmExecutionContext, NemoRelayLlmExecNextFn, void*); // Codec trampolines (used at execute time, not registration) extern char* goCodecDecodeTrampoline(void*, const FfiLLMRequest*); diff --git a/go/nemo_relay/plugin.go b/go/nemo_relay/plugin.go index 641cc7546..ed733ce0d 100644 --- a/go/nemo_relay/plugin.go +++ b/go/nemo_relay/plugin.go @@ -14,6 +14,7 @@ typedef struct FfiLlmSanitizeRequestCodec FfiLlmSanitizeRequestCodec; typedef struct FfiLlmSanitizeResponseCodec FfiLlmSanitizeResponseCodec; typedef struct NemoRelayLlmSanitizeRequestContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeRequestCodec* codec; } NemoRelayLlmSanitizeRequestContext; typedef struct NemoRelayLlmSanitizeResponseContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeResponseCodec* codec; } NemoRelayLlmSanitizeResponseContext; +typedef struct NemoRelayLlmExecutionContext { NemoRelayLlmSanitizeRequestContext request_codec; const NemoRelayLlmSanitizeResponseContext* response_codec; } NemoRelayLlmExecutionContext; typedef void (*NemoRelayFreeFn)(void* user_data); typedef char* (*NemoRelayPluginValidateCb)(void* user_data, const char* plugin_config_json); @@ -28,7 +29,7 @@ typedef char* (*NemoRelayLlmSanitizeResponseCb)(void* user_data, const char* res typedef char* (*NemoRelayLlmConditionalCb)(void* user_data, const void* request); typedef int32_t (*NemoRelayLlmRequestInterceptCb)(void* user_data, const char* name, const void* request, const char* annotated_json, char** out_outcome_json); 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); +typedef char* (*NemoRelayLlmExecInterceptCb)(void* user_data, const char* name, const char* native_json, NemoRelayLlmExecutionContext context, NemoRelayLlmExecNextFn next_fn, void* next_ctx); typedef char* (*NemoRelayToolExecNextFn)(const char* args_json, void* next_ctx); typedef char* (*NemoRelayToolExecInterceptCb)(void* user_data, const char* context_json, NemoRelayToolExecNextFn next_fn, void* next_ctx); @@ -72,7 +73,7 @@ extern char* goToolConditionalTrampoline(void*, const char*, const char*); extern void* goLlmRequestTrampoline(void*, const void*, NemoRelayLlmSanitizeRequestContext); extern char* goLlmResponseTrampoline(void*, const char*, NemoRelayLlmSanitizeResponseContext); extern char* goLlmConditionalTrampoline(void*, const void*); -extern char* goLlmExecInterceptTrampoline(void*, const char*, NemoRelayLlmExecNextFn, void*); +extern char* goLlmExecInterceptTrampoline(void*, const char*, const char*, NemoRelayLlmExecutionContext, NemoRelayLlmExecNextFn, void*); extern int32_t goLlmRequestInterceptTrampoline(void*, const char*, const void*, const char*, char**); extern char* goToolExecInterceptTrampoline(void*, const char*, NemoRelayToolExecNextFn, void*); extern char* goToolExecInterceptContextTrampoline(void*, const char*, NemoRelayToolExecNextFn, void*); diff --git a/go/nemo_relay/scope_local_test.go b/go/nemo_relay/scope_local_test.go index 87c2dcc87..46f10c97d 100644 --- a/go/nemo_relay/scope_local_test.go +++ b/go/nemo_relay/scope_local_test.go @@ -1191,7 +1191,7 @@ func assertScopeLocalLLMWrappersDeregister(t *testing.T, scopeUUID string, reque &executionInterceptCalls, func() error { return ScopeRegisterLlmExecutionIntercept(scopeUUID, "llm_scope_exec_int", 1, - func(requestJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(requestJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { executionInterceptCalls++ return next(requestJSON) }, @@ -1207,7 +1207,7 @@ func assertScopeLocalLLMStreamWrapperDeregisters(t *testing.T, scopeUUID string, t.Helper() err := ScopeRegisterLlmStreamExecutionIntercept(scopeUUID, "llm_scope_stream_int", 1, - func(requestJSON json.RawMessage, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { + func(requestJSON json.RawMessage, _ LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) { nextResult, err := next(requestJSON) if err != nil { return nil, err diff --git a/python/nemo_relay/__init__.py b/python/nemo_relay/__init__.py index dac5b169c..75a55d51f 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -101,6 +101,7 @@ async def main(): AtofStreamSinkConfig, LLMAttributes, LlmCodecIdentity, + LlmExecutionContext, LLMHandle, LLMRequest, LLMRequestInterceptOutcome, @@ -282,17 +283,23 @@ class EventSanitizeFields(TypedDict): 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 +#: callback receives the logical LLM name, request, execution context, and next callable. It may #: await the next callable or return a replacement JSON-compatible response. LlmExecutionIntercept: TypeAlias = Callable[ - [str, LLMRequest, Callable[[LLMRequest], Awaitable[Json]]], + [str, LLMRequest, LlmExecutionContext, Callable[[LLMRequest], Awaitable[Json]]], Json | Awaitable[Json], ] #: Execution intercept callback that wraps streaming LLM execution. The -#: callback receives the current request and a next callable that returns an -#: async iterator of chunks. It may return or await a replacement iterator. +#: callback receives the logical LLM name, current request, execution context, +#: and a next callable that returns an async iterator of chunks. It may return +#: or await a replacement iterator. LlmStreamExecutionIntercept: TypeAlias = Callable[ - [LLMRequest, Callable[[LLMRequest], Awaitable[AsyncIterator[Json]]]], + [ + str, + LLMRequest, + LlmExecutionContext, + Callable[[LLMRequest], Awaitable[AsyncIterator[Json]]], + ], AsyncIterator[Json] | Awaitable[AsyncIterator[Json]], ] @@ -835,6 +842,7 @@ def worker() -> None: "LlmSanitizeRequestGuardrail", "LlmSanitizeResponseGuardrail", "LlmCodecIdentity", + "LlmExecutionContext", "LlmSanitizeRequestContext", "LlmSanitizeResponseContext", "LlmSanitizeRequestCodec", diff --git a/python/nemo_relay/__init__.pyi b/python/nemo_relay/__init__.pyi index 51a43938c..326ce1909 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -68,6 +68,9 @@ from nemo_relay._native import ( from nemo_relay._native import ( LlmCodecIdentity as LlmCodecIdentity, ) +from nemo_relay._native import ( + LlmExecutionContext as LlmExecutionContext, +) from nemo_relay._native import ( LLMHandle as LLMHandle, ) @@ -345,25 +348,30 @@ Return: The complete canonical outcome passed to later middleware. """ LlmExecutionIntercept: TypeAlias = Callable[ - [str, LLMRequest, Callable[[LLMRequest], Awaitable[Json]]], + [str, LLMRequest, LlmExecutionContext, Callable[[LLMRequest], Awaitable[Json]]], Json | Awaitable[Json], ] """Execution intercept callback that wraps non-streaming LLM execution. Arguments: - The logical LLM name, current request, and next callable. + The logical LLM name, current request, execution context, and next callable. Return: A JSON-compatible response, either directly or as an awaitable. """ LlmStreamExecutionIntercept: TypeAlias = Callable[ - [LLMRequest, Callable[[LLMRequest], Awaitable[AsyncIterator[Json]]]], + [ + str, + LLMRequest, + LlmExecutionContext, + Callable[[LLMRequest], Awaitable[AsyncIterator[Json]]], + ], AsyncIterator[Json] | Awaitable[AsyncIterator[Json]], ] """Execution intercept callback that wraps streaming LLM execution. Arguments: - The current request and next callable. + The logical LLM name, current request, execution context, and next callable. Return: An async iterator of JSON chunks, either directly or as an awaitable. diff --git a/python/nemo_relay/_native.pyi b/python/nemo_relay/_native.pyi index 452456ba2..e22297b13 100644 --- a/python/nemo_relay/_native.pyi +++ b/python/nemo_relay/_native.pyi @@ -105,11 +105,16 @@ _LlmRequestIntercept: TypeAlias = Callable[ "LLMRequestInterceptOutcome | Awaitable[LLMRequestInterceptOutcome]", ] _LlmExecutionIntercept: TypeAlias = Callable[ - [str, "LLMRequest", Callable[["LLMRequest"], Awaitable[_Json]]], + [str, "LLMRequest", "LlmExecutionContext", Callable[["LLMRequest"], Awaitable[_Json]]], _Json | Awaitable[_Json], ] _LlmStreamExecutionIntercept: TypeAlias = Callable[ - ["LLMRequest", Callable[["LLMRequest"], Awaitable[AsyncIterator[_Json]]]], + [ + str, + "LLMRequest", + "LlmExecutionContext", + Callable[["LLMRequest"], Awaitable[AsyncIterator[_Json]]], + ], AsyncIterator[_Json] | Awaitable[AsyncIterator[_Json]], ] @@ -121,6 +126,14 @@ class LlmCodecIdentity: @property def id(self) -> str | None: ... +class LlmExecutionContext: + """Codec capabilities for one managed LLM execution intercept invocation.""" + + @property + def request_codec(self) -> LlmSanitizeRequestContext: ... + @property + def response_codec(self) -> LlmSanitizeResponseContext | None: ... + class LlmSanitizeRequestContext: """Per-call context passed to an LLM request sanitizer callback.""" diff --git a/python/nemo_relay/intercepts.py b/python/nemo_relay/intercepts.py index 09b249886..12c8f87c5 100644 --- a/python/nemo_relay/intercepts.py +++ b/python/nemo_relay/intercepts.py @@ -237,9 +237,10 @@ def register_llm_execution(name: str, priority: int, fn: LlmExecutionIntercept) Args: name: Unique intercept name used for later replacement or removal. priority: Execution order for the intercept. Lower values run first. - fn: Callable invoked as ``fn(name, request, next_call)``. The callback - may call ``next_call(request)`` to continue execution, modify the - result, or short-circuit the provider call. + fn: Callable invoked as ``fn(name, request, context, next_call)``. The + context exposes the active request and unary-response codecs. The + callback may call ``next_call(request)`` to continue execution, + modify the result, or short-circuit the provider call. Returns: None: This function returns after the intercept is registered. @@ -281,9 +282,10 @@ def register_llm_stream_execution( Args: name: Unique intercept name used for later replacement or removal. priority: Execution order for the intercept. Lower values run first. - fn: Callable invoked as ``fn(request, next_call)`` that returns an - async iterator of JSON chunks, either by delegating to - ``next_call(request)`` or by replacing the stream entirely. + fn: Callable invoked as ``fn(name, request, context, next_call)`` that + returns an async iterator of JSON chunks. The streaming context + exposes the request codec and no response codec. The callback may + delegate to ``next_call(request)`` or replace the stream entirely. Returns: None: This function returns after the intercept is registered. diff --git a/python/nemo_relay/scope_local.py b/python/nemo_relay/scope_local.py index 007dfa584..107eca109 100644 --- a/python/nemo_relay/scope_local.py +++ b/python/nemo_relay/scope_local.py @@ -676,8 +676,9 @@ def register_llm_execution(scope_handle: ScopeHandle, name: str, priority: int, this scope is popped. name: Unique intercept name within the owning scope. priority: Execution order for the intercept. Lower values run first. - fn: Callable invoked as ``fn(name, request, next_call)`` that may call - ``next_call(request)`` to continue execution or short-circuit it. + fn: Callable invoked as ``fn(name, request, context, next_call)`` that + may inspect the active codecs, call ``next_call(request)`` to + continue execution, or short-circuit it. Returns: None: This function returns after the scope-local intercept is @@ -721,9 +722,10 @@ def register_llm_stream_execution( this scope is popped. name: Unique intercept name within the owning scope. priority: Execution order for the intercept. Lower values run first. - fn: Callable invoked as ``fn(request, next_call)`` that returns an - async iterator of chunks, either by delegating to ``next_call`` or - by replacing the stream entirely. + fn: Callable invoked as ``fn(name, request, context, next_call)`` that + returns an async iterator of chunks. Streaming exposes the request + codec but no response codec. The callback may delegate to + ``next_call`` or replace the stream entirely. Returns: None: This function returns after the scope-local intercept is diff --git a/python/plugin/src/nemo_relay_plugin/__init__.py b/python/plugin/src/nemo_relay_plugin/__init__.py index b71d5ea44..e42fccae7 100644 --- a/python/plugin/src/nemo_relay_plugin/__init__.py +++ b/python/plugin/src/nemo_relay_plugin/__init__.py @@ -72,12 +72,9 @@ LlmSanitizeResponseCallback: LLM response sanitizer callback. LlmConditionalCallback: LLM execution guardrail callback. LlmRequestCallback: LLM request intercept callback. - LlmExecutionCallback: Unary LLM execution intercept callback. - LlmExecutionWithContextCallback: Unary LLM execution intercept callback - with codec context. - LlmStreamExecutionCallback: Streaming LLM execution intercept callback. - LlmStreamExecutionWithContextCallback: Streaming LLM execution intercept - callback with request codec context. + LlmExecutionCallback: Unary LLM execution intercept callback with codec context. + LlmStreamExecutionCallback: Streaming LLM execution intercept callback with + request codec context. Public authoring types: WorkerPlugin: Base validation and registration contract for a plugin. @@ -108,7 +105,6 @@ LlmConditionalCallback, LlmExecutionCallback, LlmExecutionContext, - LlmExecutionWithContextCallback, LlmNext, LlmOptimizationContribution, LlmOptimizationDataSchema, @@ -125,7 +121,6 @@ LlmSanitizeResponseCallback, LlmSanitizeResponseContext, LlmStreamExecutionCallback, - LlmStreamExecutionWithContextCallback, LlmStreamNext, LogSeverity, MetricKind, @@ -174,7 +169,6 @@ "LlmCodecIdentity", "LlmExecutionCallback", "LlmExecutionContext", - "LlmExecutionWithContextCallback", "LogSeverity", "MetricKind", "MetricMeasurement", @@ -196,7 +190,6 @@ "LlmSanitizeResponseCallback", "LlmStreamNext", "LlmStreamExecutionCallback", - "LlmStreamExecutionWithContextCallback", "PluginContext", "PluginRuntime", "RuntimeDiagnostic", diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index ba7f103fe..b5a724470 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -260,18 +260,15 @@ async def decode(self, response: Json) -> Json: @dataclass(frozen=True) class LlmExecutionContext: - """Invocation-scoped codec identities and operations for LLM execution middleware. + """Directional codec context for one LLM execution invocation. The identities describe the codecs selected when Relay created the managed invocation. Rewriting a request into another provider's wire format does not select a new codec; incompatible codec operations fail. """ - available: bool - request_codec_identity: LlmCodecIdentity - response_codec_identity: LlmCodecIdentity - request_codec: WorkerRequestCodec | None = field(default=None, repr=False, compare=False) - response_codec: WorkerResponseCodec | None = field(default=None, repr=False, compare=False) + request_codec: LlmSanitizeRequestContext + response_codec: LlmSanitizeResponseContext | None def _llm_codec_identity(invocation: pb.LlmInvocation) -> LlmCodecIdentity: @@ -304,33 +301,38 @@ def _llm_execution_context( invocation_id: str, ) -> LlmExecutionContext: if not invocation.HasField("execution_codec_context"): - return LlmExecutionContext( - available=False, - request_codec_identity=LlmCodecIdentity("none"), - response_codec_identity=LlmCodecIdentity("none"), - ) + raise WorkerSdkError("malformed LLM execution codec context: execution context is missing") context = invocation.execution_codec_context if not context.HasField("request") or not context.request.HasField("codec"): raise WorkerSdkError("malformed LLM execution codec context: request codec identity is missing") - if not context.HasField("response") or not context.response.HasField("codec"): - raise WorkerSdkError("malformed LLM execution codec context: response codec identity is missing") request_id = context.request.codec_capability_id if context.request.HasField("codec_capability_id") else None - response_id = context.response.codec_capability_id if context.response.HasField("codec_capability_id") else None - request_identity = _codec_identity( - context.request.codec.kind, - context.request.codec.id if context.request.codec.HasField("id") else None, - ) - response_identity = _codec_identity( - context.response.codec.kind, - context.response.codec.id if context.response.codec.HasField("id") else None, + request_context = LlmSanitizeRequestContext( + codec=_codec_identity( + context.request.codec.kind, + context.request.codec.id if context.request.codec.HasField("id") else None, + ), + _runtime=runtime, + _capability_id=request_id, + _invocation_id=invocation_id, ) + response_context: LlmSanitizeResponseContext | None = None + if context.HasField("response"): + if not context.response.HasField("codec"): + raise WorkerSdkError("malformed LLM execution codec context: response codec identity is missing") + response_id = context.response.codec_capability_id if context.response.HasField("codec_capability_id") else None + response_context = LlmSanitizeResponseContext( + codec=_codec_identity( + context.response.codec.kind, + context.response.codec.id if context.response.codec.HasField("id") else None, + ), + _runtime=runtime, + _capability_id=response_id, + _invocation_id=invocation_id, + ) return LlmExecutionContext( - available=True, - request_codec_identity=request_identity, - response_codec_identity=response_identity, - request_codec=(WorkerRequestCodec(runtime, request_id, invocation_id) if request_id else None), - response_codec=(WorkerResponseCodec(runtime, response_id, invocation_id) if response_id else None), + request_codec=request_context, + response_codec=response_context, ) @@ -1107,15 +1109,8 @@ def register(self, ctx: PluginContext, config: Json) -> None | Awaitable[None]: [str, LlmRequest, AnnotatedLlmRequest | None], LlmRequestInterceptOutcome | Awaitable[LlmRequestInterceptOutcome], ] -LlmExecutionCallback: TypeAlias = Callable[[str, LlmRequest, "LlmNext"], Json | Awaitable[Json]] -LlmExecutionWithContextCallback: TypeAlias = Callable[ - [str, LlmRequest, LlmExecutionContext, "LlmNext"], Json | Awaitable[Json] -] +LlmExecutionCallback: TypeAlias = Callable[[str, LlmRequest, LlmExecutionContext, "LlmNext"], Json | Awaitable[Json]] LlmStreamExecutionCallback: TypeAlias = Callable[ - [str, LlmRequest, "LlmStreamNext"], - Iterable[Json] | AsyncIterator[Json] | Awaitable[Iterable[Json] | AsyncIterator[Json]], -] -LlmStreamExecutionWithContextCallback: TypeAlias = Callable[ [str, LlmRequest, LlmExecutionContext, "LlmStreamNext"], Iterable[Json] | AsyncIterator[Json] | Awaitable[Iterable[Json] | AsyncIterator[Json]], ] @@ -1140,8 +1135,8 @@ class _Handlers: llm_sanitize_responses: dict[str, LlmSanitizeResponseCallback] llm_conditionals: dict[str, LlmConditionalCallback] llm_requests: dict[str, LlmRequestCallback] - llm_executions: dict[str, LlmExecutionWithContextCallback] - llm_stream_executions: dict[str, LlmStreamExecutionWithContextCallback] + llm_executions: dict[str, LlmExecutionCallback] + llm_stream_executions: dict[str, LlmStreamExecutionCallback] @classmethod def empty(cls) -> _Handlers: @@ -1545,32 +1540,14 @@ def register_llm_execution_intercept( Args: name: Component-local registration name. - callback: Function receiving ``(model_name, request, next_call)`` - and returning response JSON, directly or through an awaitable. + callback: Function receiving ``(model_name, request, context, + next_call)`` and returning response JSON, directly or through + an awaitable. It can call :meth:`LlmNext.call` zero, one, or multiple times while the invocation is active. priority: Execution order. Lower values run first. """ self._push_registration(name, pb.LLM_EXECUTION_INTERCEPT, priority, False) - self._handlers.llm_executions[name] = lambda model, request, _context, next_call: callback( - model, request, next_call - ) - - def register_llm_execution_intercept_with_context( - self, - name: str, - callback: LlmExecutionWithContextCallback, - *, - priority: int = 0, - ) -> None: - """Register LLM execution middleware with invocation-scoped codecs.""" - self._push_registration( - name, - pb.LLM_EXECUTION_INTERCEPT, - priority, - False, - llm_execution_codec_context=True, - ) self._handlers.llm_executions[name] = callback def register_llm_stream_execution_intercept( @@ -1584,12 +1561,13 @@ def register_llm_stream_execution_intercept( Args: name: Component-local registration name. - callback: Function receiving ``(model_name, request, next_call)``. - Return an iterable, an async iterator, or an awaitable resolving - to either. Every yielded item must be JSON. Strings, byte - sequences, mappings, and scalar values are not valid streams. - The callback can call :meth:`LlmStreamNext.call` zero, one, or - multiple times while the invocation is active. + callback: Function receiving ``(model_name, request, context, + next_call)``. Return an iterable, an async iterator, or an + awaitable resolving to either. Every yielded item must be JSON. + Strings, byte sequences, mappings, and scalar values are not + valid streams. The callback can call + :meth:`LlmStreamNext.call` zero, one, or multiple times while + the invocation is active. priority: Execution order. Lower values run first. Streaming behavior: @@ -1598,29 +1576,6 @@ def register_llm_stream_execution_intercept( error. """ self._push_registration(name, pb.LLM_STREAM_EXECUTION_INTERCEPT, priority, False) - self._handlers.llm_stream_executions[name] = lambda model, request, _context, next_call: callback( - model, request, next_call - ) - - def register_llm_stream_execution_intercept_with_context( - self, - name: str, - callback: LlmStreamExecutionWithContextCallback, - *, - priority: int = 0, - ) -> None: - """Register streaming middleware with request codec access. - - The context identifies the response codec but does not expose a - response decoder because stream chunks are not complete responses. - """ - self._push_registration( - name, - pb.LLM_STREAM_EXECUTION_INTERCEPT, - priority, - False, - llm_execution_codec_context=True, - ) self._handlers.llm_stream_executions[name] = callback def _push_registration( @@ -1629,8 +1584,6 @@ def _push_registration( surface: int, priority: int, break_chain: bool, - *, - llm_execution_codec_context: bool = False, ) -> None: if any( registration.local_name == name and registration.surface == surface @@ -1643,7 +1596,6 @@ def _push_registration( surface=surface, priority=priority, break_chain=break_chain, - llm_execution_codec_context=llm_execution_codec_context, ) ) diff --git a/python/tests/plugin/test_public_api_docstrings.py b/python/tests/plugin/test_public_api_docstrings.py index fb6859568..08b01c87d 100644 --- a/python/tests/plugin/test_public_api_docstrings.py +++ b/python/tests/plugin/test_public_api_docstrings.py @@ -36,9 +36,7 @@ "LlmConditionalCallback", "LlmRequestCallback", "LlmExecutionCallback", - "LlmExecutionWithContextCallback", "LlmStreamExecutionCallback", - "LlmStreamExecutionWithContextCallback", } diff --git a/python/tests/plugin/test_worker_sdk.py b/python/tests/plugin/test_worker_sdk.py index 492739ccc..7f2c9d3e8 100644 --- a/python/tests/plugin/test_worker_sdk.py +++ b/python/tests/plugin/test_worker_sdk.py @@ -637,11 +637,11 @@ def llm_request(name: str, request: Json, annotated: Json | None) -> LlmRequestI pending_marks=[PendingMarkSpec("worker.pending", data={"source": "python"})], ) - async def llm_execution(name: str, request: Json, next_call: Any) -> Json: + async def llm_execution(name: str, request: Json, _context: Any, next_call: Any) -> Json: result = await next_call.call(_tag_llm_request(request, f"llm_execute_{name}")) return _tag(result, "llm_execution") - async def llm_stream_execution(name: str, request: Json, next_call: Any) -> AsyncIterator[Json]: + async def llm_stream_execution(name: str, request: Json, _context: Any, next_call: Any) -> AsyncIterator[Json]: stream = next_call.call(_tag_llm_request(request, f"llm_stream_{name}")) async for chunk in stream: yield _tag(chunk, "llm_stream_execution") @@ -736,7 +736,7 @@ def test_generated_proto_matches_worker_contract() -> None: "Shutdown", } assert pb.InvokeRequest.DESCRIPTOR.fields_by_name["auth_token"].number == 7 - assert pb.Registration.DESCRIPTOR.fields_by_name["llm_execution_codec_context"].number == 6 + assert "llm_execution_codec_context" not in pb.Registration.DESCRIPTOR.fields_by_name execution_context = pb.LlmInvocation.DESCRIPTOR.fields_by_name["execution_codec_context"] assert execution_context.number == 11 assert execution_context.containing_oneof is None @@ -1116,28 +1116,21 @@ def test_plugin_context_registers_llm_sanitizers_under_standard_names() -> None: ] -def test_execution_codec_context_is_opt_in_on_the_existing_surface() -> None: +def test_execution_codec_context_uses_the_existing_surface() -> None: context = PluginContext() - async def legacy(_name: str, request: Json, next_call: Any) -> Json: + async def execution(_name: str, request: Json, _context: Any, next_call: Any) -> Json: return await next_call.call(request) - async def contextual(_name: str, request: Json, _context: Any, next_call: Any) -> Json: - return await next_call.call(request) - - context.register_llm_execution_intercept("legacy", legacy, priority=7) - context.register_llm_execution_intercept_with_context("contextual", contextual, priority=7) + context.register_llm_execution_intercept("execution", execution, priority=7) - legacy_registration, contextual_registration = context._handlers.registrations - assert legacy_registration.surface == pb.LLM_EXECUTION_INTERCEPT - assert contextual_registration.surface == legacy_registration.surface - assert contextual_registration.priority == legacy_registration.priority - assert not legacy_registration.llm_execution_codec_context - assert contextual_registration.llm_execution_codec_context + registration = context._handlers.registrations[0] + assert registration.surface == pb.LLM_EXECUTION_INTERCEPT + assert registration.priority == 7 -async def test_contextual_execution_callback_handles_old_and_new_hosts() -> None: - seen: list[bool] = [] +async def test_execution_callback_receives_directional_codec_context() -> None: + seen: list[plugin_api.LlmExecutionContext] = [] class ContextualExecutionPlugin(WorkerPlugin): plugin_id = "tests.contextual_execution" @@ -1146,25 +1139,26 @@ def register(self, ctx: PluginContext, config: Json) -> None: del config async def execution(name: str, request: Json, context: Any, next_call: Any) -> Json: - del name - seen.append(context.available) - if context.available: - assert context.request_codec_identity == plugin_api.LlmCodecIdentity("builtin", "openai_chat") - assert context.response_codec_identity == plugin_api.LlmCodecIdentity("builtin", "openai_chat") - assert context.request_codec is not None - assert context.response_codec is not None - await context.request_codec.decode(request) + assert name == "model" + assert request == {"content": {"model": "gpt-test"}} + seen.append(context) + assert context.request_codec.codec == plugin_api.LlmCodecIdentity("builtin", "openai_chat") + assert context.response_codec is not None + assert context.response_codec.codec == plugin_api.LlmCodecIdentity("builtin", "openai_chat") + request_codec = context.request_codec.resolve_codec() + assert request_codec is not None + await request_codec.decode(request) result = await next_call.call(request) - if context.available: - await context.response_codec.decode(result) + response_codec = context.response_codec.resolve_codec() + assert response_codec is not None + await response_codec.decode(result) return result - ctx.register_llm_execution_intercept_with_context("execution", execution) + ctx.register_llm_execution_intercept("execution", execution) host = RecordingHostStub() service = _service(ContextualExecutionPlugin(), host) - register = await _register(service) - assert register.registrations[0].llm_execution_codec_context + await _register(service) new_context = pb.LlmExecutionCodecContext( request=pb.LlmSanitizeRequestContext( @@ -1176,19 +1170,17 @@ async def execution(name: str, request: Json, context: Any, next_call: Any) -> J codec_capability_id="response-capability", ), ) - for execution_context in [None, new_context]: - payload = _llm_payload(request={"content": {"model": "gpt-test"}}) - if execution_context is not None: - payload.execution_codec_context.CopyFrom(execution_context) - result = await _invoke_json_async( - service, - "execution", - pb.LLM_EXECUTION_INTERCEPT, - payload=payload, - ) - assert "next_llm" in result + payload = _llm_payload(request={"content": {"model": "gpt-test"}}) + payload.execution_codec_context.CopyFrom(new_context) + result = await _invoke_json_async( + service, + "execution", + pb.LLM_EXECUTION_INTERCEPT, + payload=payload, + ) + assert "next_llm" in result - assert seen == [False, True] + assert len(seen) == 1 codec_capabilities = [ request.codec_capability_id for request in host.requests @@ -1197,6 +1189,46 @@ async def execution(name: str, request: Json, context: Any, next_call: Any) -> J assert codec_capabilities == ["request-capability", "response-capability"] +def test_execution_context_distinguishes_absent_and_resolved_opaque_codecs() -> None: + runtime = PluginRuntime( + activation_id=ACTIVATION_ID, + auth_token=AUTH_TOKEN, + host_stub=RecordingHostStub(), + ) + + absent_invocation = pb.LlmInvocation( + execution_codec_context=pb.LlmExecutionCodecContext( + request=pb.LlmSanitizeRequestContext(codec=pb.LlmCodecIdentity()), + response=pb.LlmSanitizeResponseContext(codec=pb.LlmCodecIdentity()), + ) + ) + absent = plugin_api._llm_execution_context(absent_invocation, runtime, "absent-invocation") + assert absent.request_codec.codec == plugin_api.LlmCodecIdentity("none") + assert absent.request_codec.resolve_codec() is None + assert absent.response_codec is not None + assert absent.response_codec.codec == plugin_api.LlmCodecIdentity("none") + assert absent.response_codec.resolve_codec() is None + + opaque_invocation = pb.LlmInvocation( + execution_codec_context=pb.LlmExecutionCodecContext( + request=pb.LlmSanitizeRequestContext( + codec=pb.LlmCodecIdentity(kind=pb.LLM_CODEC_KIND_OPAQUE), + codec_capability_id="opaque-request", + ), + response=pb.LlmSanitizeResponseContext( + codec=pb.LlmCodecIdentity(kind=pb.LLM_CODEC_KIND_OPAQUE), + codec_capability_id="opaque-response", + ), + ) + ) + opaque = plugin_api._llm_execution_context(opaque_invocation, runtime, "opaque-invocation") + assert opaque.request_codec.codec == plugin_api.LlmCodecIdentity("opaque") + assert opaque.request_codec.resolve_codec() is not None + assert opaque.response_codec is not None + assert opaque.response_codec.codec == plugin_api.LlmCodecIdentity("opaque") + assert opaque.response_codec.resolve_codec() is not None + + async def test_llm_sanitizers_receive_codec_context_and_can_omit_payloads() -> None: seen: list[tuple[str, LlmSanitizeRequestContext | LlmSanitizeResponseContext]] = [] @@ -2321,8 +2353,8 @@ class FailingStreamPlugin(WorkerPlugin): def register(self, ctx: PluginContext, config: Json) -> None: del config - async def fail(name: str, request: Json, next_call: Any) -> AsyncIterator[Json]: - del name, request, next_call + async def fail(name: str, request: Json, context: Any, next_call: Any) -> AsyncIterator[Json]: + del name, request, context, next_call raise RuntimeError("stream boom") yield {} @@ -2351,8 +2383,8 @@ class SyncStreamPlugin(WorkerPlugin): def register(self, ctx: PluginContext, config: Json) -> None: del config - def stream(name: str, request: Json, next_call: Any) -> list[Json]: - del name, request, next_call + def stream(name: str, request: Json, context: Any, next_call: Any) -> list[Json]: + del name, request, context, next_call return [{"sync": True}] ctx.register_llm_stream_execution_intercept("sync_stream", stream) @@ -2381,8 +2413,8 @@ class InvalidStreamPlugin(WorkerPlugin): def register(self, ctx: PluginContext, config: Json) -> None: del config - def stream(name: str, request: Json, next_call: Any) -> Any: - del name, request, next_call + def stream(name: str, request: Json, context: Any, next_call: Any) -> Any: + del name, request, context, next_call return invalid_stream ctx.register_llm_stream_execution_intercept("invalid_stream", stream) @@ -2883,8 +2915,8 @@ class CancelStreamPlugin(WorkerPlugin): def register(self, ctx: PluginContext, config: Json) -> None: del config - async def llm_stream(model_name: str, request: Json, next_call: Any) -> AsyncIterator[Json]: - del model_name, request, next_call + async def llm_stream(model_name: str, request: Json, context: Any, next_call: Any) -> AsyncIterator[Json]: + del model_name, request, context, next_call started.set() try: await asyncio.Event().wait() @@ -2937,8 +2969,8 @@ class BufferedStreamPlugin(WorkerPlugin): def register(self, ctx: PluginContext, config: Json) -> None: del config - async def llm_stream(model_name: str, request: Json, next_call: Any) -> AsyncIterator[Json]: - del model_name, request, next_call + async def llm_stream(model_name: str, request: Json, context: Any, next_call: Any) -> AsyncIterator[Json]: + del model_name, request, context, next_call for index in range(4): yield {"index": index} finished.set() @@ -2989,8 +3021,8 @@ class CancelledStreamPlugin(WorkerPlugin): def register(self, ctx: PluginContext, config: Json) -> None: del config - async def llm_stream(model_name: str, request: Json, next_call: Any) -> AsyncIterator[Json]: - del model_name, request, next_call + async def llm_stream(model_name: str, request: Json, context: Any, next_call: Any) -> AsyncIterator[Json]: + del model_name, request, context, next_call if False: yield {} raise asyncio.CancelledError @@ -3466,6 +3498,13 @@ def _invoke_request( continuation_id: str = "next-1", **kwargs: Any, ) -> Any: + llm = kwargs.get("llm") + if llm is not None and surface in {pb.LLM_EXECUTION_INTERCEPT, pb.LLM_STREAM_EXECUTION_INTERCEPT}: + if not llm.HasField("execution_codec_context"): + context = pb.LlmExecutionCodecContext(request=pb.LlmSanitizeRequestContext(codec=pb.LlmCodecIdentity())) + if surface == pb.LLM_EXECUTION_INTERCEPT: + context.response.CopyFrom(pb.LlmSanitizeResponseContext(codec=pb.LlmCodecIdentity())) + llm.execution_codec_context.CopyFrom(context) return pb.InvokeRequest( activation_id=ACTIVATION_ID, invocation_id=invocation_id, diff --git a/python/tests/test_adaptive.py b/python/tests/test_adaptive.py index e4b030fa5..68f5b5ea7 100644 --- a/python/tests/test_adaptive.py +++ b/python/tests/test_adaptive.py @@ -17,6 +17,7 @@ AnnotatedLLMRequest, Json, JsonObject, + LlmExecutionContext, LLMRequest, LLMRequestInterceptOutcome, ScopeType, @@ -310,7 +311,10 @@ def intercept( return LLMRequestInterceptOutcome(LLMRequest(headers, request.content), annotated) async def llm_exec_intercept( - _name: str, request: LLMRequest, next_call: Callable[[LLMRequest], Awaitable[Json]] + _name: str, + request: LLMRequest, + _context: LlmExecutionContext, + next_call: Callable[[LLMRequest], Awaitable[Json]], ) -> Json: response = await next_call(request) assert isinstance(response, dict) @@ -318,7 +322,9 @@ async def llm_exec_intercept( return response async def llm_stream_exec_intercept( + _name: str, request: LLMRequest, + _context: LlmExecutionContext, next_call: Callable[[LLMRequest], Awaitable[AsyncIterator[Json]]], ) -> AsyncIterator[Json]: stream = await next_call(request) diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index 900316e8a..0dcd63bfe 100644 --- a/python/tests/test_llm.py +++ b/python/tests/test_llm.py @@ -325,6 +325,180 @@ def sanitize_response(response, context): assert request_codec_used is True assert response_codec_used is True + async def test_execution_intercept_receives_directional_codecs(self) -> None: + observed = False + + async def execution_intercept(name, request, context, next_call): + nonlocal observed + assert name == "py_llm_execution_context" + assert context.request_codec.codec.kind == "builtin" + assert context.request_codec.codec.id == "openai_chat" + request_codec = context.request_codec.resolve_codec() + assert request_codec is not None + assert request_codec.decode(request).model == "test-model" + + assert context.response_codec is not None + assert context.response_codec.codec.kind == "builtin" + assert context.response_codec.codec.id == "openai_chat" + response = await next_call(request) + response_codec = context.response_codec.resolve_codec() + assert response_codec is not None + assert response_codec.decode_response(response).model == "test-model" + observed = True + return response + + codec = OpenAIChatCodec() + response = { + "id": "chatcmpl-execution-context", + "model": "test-model", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}}], + } + intercepts.register_llm_execution("py_llm_execution_context", 1, execution_intercept) + try: + result = await llm.execute( + "py_llm_execution_context", + make_request(), + lambda _request: response, + codec=codec, + response_codec=codec, + ) + finally: + intercepts.deregister_llm_execution("py_llm_execution_context") + + assert result == response + assert observed + + async def test_execution_intercept_distinguishes_absent_and_opaque_codecs(self) -> None: + class OpaqueCodec: + def __init__(self) -> None: + self.inner = OpenAIChatCodec() + + def decode(self, request): + return self.inner.decode(request) + + def encode(self, annotated, original): + return self.inner.encode(annotated, original) + + def decode_response(self, response): + return self.inner.decode_response(response) + + seen = [] + + async def execution_intercept(name, request, context, next_call): + seen.append(name) + if name == "py_llm_execution_context_absent": + assert context.request_codec.codec.kind == "none" + assert context.request_codec.resolve_codec() is None + assert context.response_codec is not None + assert context.response_codec.codec.kind == "none" + assert context.response_codec.resolve_codec() is None + return await next_call(request) + + assert name == "py_llm_execution_context_opaque" + assert context.request_codec.codec.kind == "opaque" + request_codec = context.request_codec.resolve_codec() + assert request_codec is not None + encoded = request_codec.encode(request_codec.decode(request), request) + + assert context.response_codec is not None + assert context.response_codec.codec.kind == "opaque" + response_codec = context.response_codec.resolve_codec() + assert response_codec is not None + response = await next_call(encoded) + assert response_codec.decode_response(response).model == "test-model" + return response + + response = { + "id": "chatcmpl-execution-context-matrix", + "model": "test-model", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}}], + } + codec = OpaqueCodec() + intercepts.register_llm_execution("py_llm_execution_context_matrix", 1, execution_intercept) + try: + absent = await llm.execute( + "py_llm_execution_context_absent", + make_request(), + lambda _request: response, + ) + opaque = await llm.execute( + "py_llm_execution_context_opaque", + make_request(), + lambda _request: response, + codec=codec, + response_codec=codec, + ) + finally: + intercepts.deregister_llm_execution("py_llm_execution_context_matrix") + + assert absent == response + assert opaque == response + assert seen == ["py_llm_execution_context_absent", "py_llm_execution_context_opaque"] + + async def test_execution_codec_capability_expires_when_callback_settles(self) -> None: + retained_codec = None + + async def execution_intercept(_name, request, context, next_call): + nonlocal retained_codec + retained_codec = context.request_codec.resolve_codec() + assert retained_codec is not None + return await next_call(request) + + codec = OpenAIChatCodec() + intercepts.register_llm_execution("py_llm_execution_codec_expiry", 1, execution_intercept) + try: + await llm.execute( + "py_llm_execution_codec_expiry", + make_request(), + lambda _request: {"ok": True}, + codec=codec, + ) + finally: + intercepts.deregister_llm_execution("py_llm_execution_codec_expiry") + + assert retained_codec is not None + with pytest.raises(RuntimeError, match="LLM execution codec capability is no longer active"): + retained_codec.decode(make_request()) + + async def test_stream_execution_context_has_no_response_codec(self) -> None: + observed = False + retained_codec = None + + async def execution_intercept(name, request, context, next_call): + nonlocal observed, retained_codec + assert name == "py_llm_stream_execution_context" + assert context.request_codec.codec.kind == "builtin" + retained_codec = context.request_codec.resolve_codec() + assert retained_codec is not None + assert context.response_codec is None + observed = True + return await next_call(request) + + async def provider(_request): + yield {"token": "ok"} + + codec = OpenAIChatCodec() + intercepts.register_llm_stream_execution("py_llm_stream_execution_context", 1, execution_intercept) + try: + stream = await llm.stream_execute( + "py_llm_stream_execution_context", + make_request(), + provider, + lambda _chunk: None, + lambda: {}, + codec=codec, + response_codec=codec, + ) + assert retained_codec is not None + assert retained_codec.decode(make_request()).model == "test-model" + assert [chunk async for chunk in stream] == [{"token": "ok"}] + with pytest.raises(RuntimeError, match="LLM execution codec capability is no longer active"): + retained_codec.decode(make_request()) + finally: + intercepts.deregister_llm_stream_execution("py_llm_stream_execution_context") + + assert observed + def test_none_omits_payload_and_short_circuits_later_sanitizers(self) -> None: events = [] later_called = False @@ -604,12 +778,12 @@ def test_execution_intercept(self) -> None: intercepts.register_llm_execution( "py_llm_exec", 1, - lambda name, request, next: {"intercepted": True}, + lambda name, request, context, next: {"intercepted": True}, ) assert intercepts.deregister_llm_execution("py_llm_exec") def test_stream_execution_intercept(self) -> None: - def stream_fn(request, next): + def stream_fn(name, request, context, next): async def gen(): yield {"token": "test"} @@ -636,7 +810,7 @@ async def test_execution_callback_capture_traceparent_matches_llm_scope(self) -> observed = [] subscribers.register("py_llm_capture_traceparent", events.append) - async def execution_intercept(_name, request, next_handler): + async def execution_intercept(_name, request, _context, next_handler): context = capture_propagation_context() explicit_root = "018f13f0-7c1a-7a80-8000-000000000799" rooted = capture_propagation_context_with_root(explicit_root) @@ -677,7 +851,7 @@ async def test_execution_callback_capture_traceparent_preserves_imported_root(se observed = [] subscribers.register("py_llm_capture_propagated_trace_root", events.append) - async def execution_intercept(_name, request, next_handler): + async def execution_intercept(_name, request, _context, next_handler): observed.append(capture_traceparent()) return await next_handler(request) @@ -836,7 +1010,7 @@ async def test_cancelling_execute_cancels_pending_execution_intercept(self) -> N provider_calls: list[LLMRequest] = [] events: list[Event] = [] - async def middleware(_name, request, next): + async def middleware(_name, request, _context, next): started.set() try: await release.wait() @@ -877,7 +1051,7 @@ async def test_cancelling_stream_execute_cancels_pending_stream_intercept(self) provider_calls: list[LLMRequest] = [] events: list[Event] = [] - async def middleware(request, next): + async def middleware(_name, request, _context, next): started.set() try: await release.wait() @@ -937,7 +1111,7 @@ def request_intercept(_name, request, annotated): observed.append(("request", request_id.get())) return LLMRequestInterceptOutcome(request, annotated) - def execution_intercept(_name, _request, _next): + def execution_intercept(_name, _request, _context, _next): observed.append(("execution", request_id.get())) return {"ok": True} @@ -968,7 +1142,7 @@ async def test_sync_stream_intercept_preserves_async_caller_context(self) -> Non request_id = contextvars.ContextVar("llm_stream_middleware_request_id", default="registration") observed: list[tuple[str, str]] = [] - def middleware(request, next): + def middleware(_name, request, _context, next): observed.append(("callback", request_id.get())) async def generate(): @@ -1037,7 +1211,7 @@ async def close() -> None: return close() - def middleware(_request, _next): + def middleware(_name, _request, _context, _next): observed.append(("callback", request_id.get())) return CustomIterator() @@ -1177,7 +1351,7 @@ async def test_execution_intercept_replaces(self) -> None: intercepts.register_llm_execution( "py_llm_exec_rep", 1, - lambda name, request, next: {"from_intercept": True}, + lambda name, request, context, next: {"from_intercept": True}, ) def original_func(request): @@ -1191,7 +1365,7 @@ def original_func(request): intercepts.deregister_llm_execution("py_llm_exec_rep") async def test_execution_intercept_can_await_next(self) -> None: - async def middleware(name, request, next): + async def middleware(name, request, context, next): updated = LLMRequest(request.headers, {**request.content, "model": "via-next"}) result = await next(updated) result["from_intercept"] = True @@ -1212,7 +1386,7 @@ async def test_execution_intercept_rejects_next_after_settlement(self) -> None: captured_next = None provider_calls = 0 - async def middleware(_name, _request, next): + async def middleware(_name, _request, _context, next): nonlocal captured_next captured_next = next return {"source": "intercept"} @@ -1235,7 +1409,7 @@ def provider(_request): assert provider_calls == 0 async def test_stream_execution_intercept_can_await_next(self) -> None: - def middleware(request, next): + def middleware(_name, request, _context, next): async def gen(): updated = LLMRequest(request.headers, {**request.content, "prefix": "wrapped"}) stream = await next(updated) @@ -1268,7 +1442,7 @@ async def test_stream_execution_intercept_rejects_next_after_settlement(self) -> captured_next = None provider_calls = 0 - async def middleware(_request, next): + async def middleware(_name, _request, _context, next): nonlocal captured_next captured_next = next @@ -1301,7 +1475,7 @@ async def stream(): assert provider_calls == 0 async def test_stream_execution_intercept_async_function_is_supported(self) -> None: - def middleware(request, next): + def middleware(_name, request, _context, next): updated = LLMRequest(request.headers, {**request.content, "prefix": "async"}) async def gen(): @@ -1510,7 +1684,7 @@ async def test_stream_execution_intercept_rejects_invalid_iterator(self) -> None intercepts.register_llm_stream_execution( "py_llm_stream_bad_iter", 1, - cast(intercepts.LlmStreamExecutionIntercept, lambda request, next: object()), + cast(intercepts.LlmStreamExecutionIntercept, lambda name, request, context, next: object()), ) try: stream = await llm.stream_execute( @@ -1529,7 +1703,7 @@ async def test_stream_execution_intercept_handles_iterator_that_stops_in___anext intercepts.register_llm_stream_execution( "py_llm_stream_direct_stop", 1, - lambda request, next: _ImmediateStopAsyncIter(), + lambda name, request, context, next: _ImmediateStopAsyncIter(), ) try: stream = await llm.stream_execute( @@ -1550,7 +1724,7 @@ async def test_stream_execution_intercept_propagates_direct___anext__error(self) intercepts.register_llm_stream_execution( "py_llm_stream_direct_error", 1, - lambda request, next: _BrokenAsyncIter(), + lambda name, request, context, next: _BrokenAsyncIter(), ) try: stream = await llm.stream_execute( @@ -1569,7 +1743,7 @@ async def test_stream_execution_intercept_failure_emits_exception_type(self) -> events = [] subscribers.register("py_llm_stream_intercept_failure_sub", events.append) - def failing_middleware(request, next) -> Never: + def failing_middleware(name, request, context, next) -> Never: raise ValueError("stream intercept boom") intercepts.register_llm_stream_execution( diff --git a/python/tests/test_scope_local.py b/python/tests/test_scope_local.py index 712715579..fb92a7868 100644 --- a/python/tests/test_scope_local.py +++ b/python/tests/test_scope_local.py @@ -19,6 +19,7 @@ Event, Json, JsonObject, + LlmExecutionContext, LLMRequest, LLMRequestInterceptOutcome, LlmSanitizeRequestContext, @@ -647,7 +648,9 @@ def test_register_and_deregister_scope_local_wrappers(self) -> None: request = LLMRequest({}, {"messages": [], "model": "scope-local"}) async def stream_intercept( + _name: str, request_inner: LLMRequest, + _context: LlmExecutionContext, next_fn: Callable[[LLMRequest], Awaitable[AsyncIterator[Json]]], ) -> AsyncIterator[Json]: if request_inner.content.get("emit_test_chunk"): @@ -703,7 +706,7 @@ async def stream_intercept( handle, "sl_llm_exec_cov", 1, - lambda name, req, next_fn: {"intercepted": True}, + lambda name, req, context, next_fn: {"intercepted": True}, ) assert scope_local.deregister_llm_execution(handle, "sl_llm_exec_cov") is True @@ -757,7 +760,12 @@ def intercept(_name: str, req: LLMRequest, annotated: AnnotatedLLMRequest | None async def test_scope_local_llm_execution_intercept_can_await_next(self) -> None: request = LLMRequest({}, {"messages": [], "model": "scope-local"}) - async def middleware(_name: str, req: LLMRequest, next_fn: Callable[[LLMRequest], Awaitable[Json]]) -> Json: + async def middleware( + _name: str, + req: LLMRequest, + _context: LlmExecutionContext, + next_fn: Callable[[LLMRequest], Awaitable[Json]], + ) -> Json: updated = LLMRequest(req.headers, {**req.content, "model": "via-scope-local"}) result = await next_fn(updated) assert isinstance(result, dict) From 63b90deb94395a0cf5cd9c87e946b23c786412d1 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 22 Sep 2026 11:16:24 -0400 Subject: [PATCH 03/22] docs: clarify execution context migration Signed-off-by: Alex Fournier --- crates/ffi/nemo_relay.h | 30 ++++++++++++------- crates/ffi/src/api/llm_registry.rs | 14 +++++---- crates/ffi/src/api/scope_registry.rs | 13 +++++--- crates/ffi/src/callable.rs | 3 +- docs/about-nemo-relay/release-notes/index.mdx | 16 +++++----- docs/reference/migration-guides.mdx | 21 ++++++++----- 6 files changed, 62 insertions(+), 35 deletions(-) diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 5b5368bd9..dedbf26fc 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -386,7 +386,8 @@ typedef NemoRelayStatus (*NemoRelayLlmRequestInterceptCb)(void *user_data, * * `request_codec` is always present. `response_codec` is non-null for unary * execution and null for streaming execution, where Relay has no completed - * response to decode. + * response to decode. Pointers reachable from this value are borrowed and + * valid only until the intercept callback returns. */ typedef struct NemoRelayLlmExecutionContext { /** @@ -1501,14 +1502,15 @@ NemoRelayStatus nemo_relay_deregister_llm_request_intercept(const char *name); /** * Register an LLM execution intercept following the middleware chain pattern. - * The callback receives `(request, next_fn, next_ctx)` — call + * The callback receives `(name, request, context, next_fn, next_ctx)` — call * `next_fn(request, next_ctx)` to invoke the next intercept or the original * LLM call, or skip calling it to short-circuit. * * # Parameters * - `name`: Unique intercept name. * - `priority`: Execution priority (lower runs first). - * - `exec_cb`: Middleware callback receiving request and a next function. + * - `exec_cb`: Middleware callback receiving the LLM name, request, codec + * context, and a next function. * - `exec_user_data`: Opaque pointer for the execution callback. * - `exec_free`: Optional destructor for `exec_user_data`. * @@ -1531,14 +1533,17 @@ NemoRelayStatus nemo_relay_deregister_llm_execution_intercept(const char *name); /** * Register an LLM streaming execution intercept following the middleware chain - * pattern. The callback receives `(request, next_fn, next_ctx)` — call + * pattern. The callback receives + * `(name, request, context, next_fn, next_ctx)` — call * `next_fn(request, next_ctx)` to invoke the next intercept or the original - * streaming LLM call, or skip calling it to short-circuit. + * streaming LLM call, or skip calling it to short-circuit. The response codec + * in `context` is null because chunks are not complete provider responses. * * # Parameters * - `name`: Unique intercept name. * - `priority`: Execution priority (lower runs first). - * - `exec_cb`: Middleware callback receiving request and a next function. + * - `exec_cb`: Middleware callback receiving the LLM name, request, request + * codec context, and a next function. * - `exec_user_data`: Opaque pointer for the execution callback. * - `exec_free`: Optional destructor for `exec_user_data`. * @@ -2952,13 +2957,15 @@ NemoRelayStatus nemo_relay_scope_deregister_llm_request_intercept(const char *sc /** * Register a scope-local LLM execution intercept following the middleware - * chain pattern. + * chain pattern. The callback receives + * `(name, request, context, next_fn, next_ctx)`. * * # Parameters * - `scope_uuid`: UUID of the target scope (null-terminated C string). * - `name`: Unique intercept name. * - `priority`: Execution priority (lower runs first). - * - `exec_cb`: Middleware callback receiving request and a next function. + * - `exec_cb`: Middleware callback receiving the LLM name, request, codec + * context, and a next function. * - `exec_user_data`: Opaque pointer for the execution callback. * - `exec_free`: Optional destructor for `exec_user_data`. * @@ -2983,13 +2990,16 @@ NemoRelayStatus nemo_relay_scope_deregister_llm_execution_intercept(const char * /** * Register a scope-local LLM streaming execution intercept following the - * middleware chain pattern. + * middleware chain pattern. The callback receives + * `(name, request, context, next_fn, next_ctx)`. The response codec in + * `context` is null. * * # Parameters * - `scope_uuid`: UUID of the target scope (null-terminated C string). * - `name`: Unique intercept name. * - `priority`: Execution priority (lower runs first). - * - `exec_cb`: Middleware callback receiving request and a next function. + * - `exec_cb`: Middleware callback receiving the LLM name, request, request + * codec context, and a next function. * - `exec_user_data`: Opaque pointer for the execution callback. * - `exec_free`: Optional destructor for `exec_user_data`. * diff --git a/crates/ffi/src/api/llm_registry.rs b/crates/ffi/src/api/llm_registry.rs index 5c9ba7420..21cfa4ce8 100644 --- a/crates/ffi/src/api/llm_registry.rs +++ b/crates/ffi/src/api/llm_registry.rs @@ -233,14 +233,15 @@ pub unsafe extern "C" fn nemo_relay_deregister_llm_request_intercept( } /// Register an LLM execution intercept following the middleware chain pattern. -/// The callback receives `(request, next_fn, next_ctx)` — call +/// The callback receives `(name, request, context, next_fn, next_ctx)` — call /// `next_fn(request, next_ctx)` to invoke the next intercept or the original /// LLM call, or skip calling it to short-circuit. /// /// # Parameters /// - `name`: Unique intercept name. /// - `priority`: Execution priority (lower runs first). -/// - `exec_cb`: Middleware callback receiving request and a next function. +/// - `exec_cb`: Middleware callback receiving the LLM name, request, codec +/// context, and a next function. /// - `exec_user_data`: Opaque pointer for the execution callback. /// - `exec_free`: Optional destructor for `exec_user_data`. /// @@ -286,14 +287,17 @@ pub unsafe extern "C" fn nemo_relay_deregister_llm_execution_intercept( } /// Register an LLM streaming execution intercept following the middleware chain -/// pattern. The callback receives `(request, next_fn, next_ctx)` — call +/// pattern. The callback receives +/// `(name, request, context, next_fn, next_ctx)` — call /// `next_fn(request, next_ctx)` to invoke the next intercept or the original -/// streaming LLM call, or skip calling it to short-circuit. +/// streaming LLM call, or skip calling it to short-circuit. The response codec +/// in `context` is null because chunks are not complete provider responses. /// /// # Parameters /// - `name`: Unique intercept name. /// - `priority`: Execution priority (lower runs first). -/// - `exec_cb`: Middleware callback receiving request and a next function. +/// - `exec_cb`: Middleware callback receiving the LLM name, request, request +/// codec context, and a next function. /// - `exec_user_data`: Opaque pointer for the execution callback. /// - `exec_free`: Optional destructor for `exec_user_data`. /// diff --git a/crates/ffi/src/api/scope_registry.rs b/crates/ffi/src/api/scope_registry.rs index 574d471ec..d2d543ecc 100644 --- a/crates/ffi/src/api/scope_registry.rs +++ b/crates/ffi/src/api/scope_registry.rs @@ -616,13 +616,15 @@ pub unsafe extern "C" fn nemo_relay_scope_deregister_llm_request_intercept( } /// Register a scope-local LLM execution intercept following the middleware -/// chain pattern. +/// chain pattern. The callback receives +/// `(name, request, context, next_fn, next_ctx)`. /// /// # Parameters /// - `scope_uuid`: UUID of the target scope (null-terminated C string). /// - `name`: Unique intercept name. /// - `priority`: Execution priority (lower runs first). -/// - `exec_cb`: Middleware callback receiving request and a next function. +/// - `exec_cb`: Middleware callback receiving the LLM name, request, codec +/// context, and a next function. /// - `exec_user_data`: Opaque pointer for the execution callback. /// - `exec_free`: Optional destructor for `exec_user_data`. /// @@ -678,13 +680,16 @@ pub unsafe extern "C" fn nemo_relay_scope_deregister_llm_execution_intercept( } /// Register a scope-local LLM streaming execution intercept following the -/// middleware chain pattern. +/// middleware chain pattern. The callback receives +/// `(name, request, context, next_fn, next_ctx)`. The response codec in +/// `context` is null. /// /// # Parameters /// - `scope_uuid`: UUID of the target scope (null-terminated C string). /// - `name`: Unique intercept name. /// - `priority`: Execution priority (lower runs first). -/// - `exec_cb`: Middleware callback receiving request and a next function. +/// - `exec_cb`: Middleware callback receiving the LLM name, request, request +/// codec context, and a next function. /// - `exec_user_data`: Opaque pointer for the execution callback. /// - `exec_free`: Optional destructor for `exec_user_data`. /// diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 4af0e5a0e..20daa8135 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -206,7 +206,8 @@ pub struct NemoRelayLlmSanitizeResponseContext { /// /// `request_codec` is always present. `response_codec` is non-null for unary /// execution and null for streaming execution, where Relay has no completed -/// response to decode. +/// response to decode. Pointers reachable from this value are borrowed and +/// valid only until the intercept callback returns. #[repr(C)] #[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmExecutionContext { diff --git a/docs/about-nemo-relay/release-notes/index.mdx b/docs/about-nemo-relay/release-notes/index.mdx index 7518d8e9c..cb2c92203 100644 --- a/docs/about-nemo-relay/release-notes/index.mdx +++ b/docs/about-nemo-relay/release-notes/index.mdx @@ -41,13 +41,15 @@ or safely rewrite the final provider request and decode a completed unary response without importing or copying Relay's codec implementations. Streaming interceptors receive request codec access only. -**Breaking change:** The context is a required callback argument across Rust, -Python, Node.js, Go, C, native plugins, and Rust and Python gRPC workers. Native -plugins must rebuild for internal ABI v7, and affected workers must regenerate -their protobuf bindings and rebuild. Authored compatibility labels remain -`native_api = "1"` and `grpc-v1`; plugin manifests must use a Relay range that -begins at 0.10 or otherwise excludes 0.9. Refer to the [Migration -Guides](/reference/migration-guides) for callback shapes and upgrade steps. +**Breaking change:** Callback signatures change across Rust, Python, Node.js, +Go, C, native plugins, and Rust and Python gRPC workers. Python +language-binding streaming and public C callbacks also gain the logical LLM +name. Native plugins must rebuild for internal ABI v7, and affected workers +must regenerate their protobuf bindings and rebuild. Authored compatibility +labels remain `native_api = "1"` and `grpc-v1`; plugin manifests must use a +Relay range that begins at 0.10 or otherwise excludes 0.9. Refer to the +[Migration Guides](/reference/migration-guides) for callback shapes and upgrade +steps. ### Fixes and Other Changes diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index 9eca5c7b9..62c0840d6 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -13,21 +13,26 @@ upgrade actions as they are identified during the 0.10 development cycle. ### Update LLM Execution Intercepts -Relay 0.10 adds `LlmExecutionContext` immediately before the continuation in -every unary and streaming LLM execution-intercept callback. Update callbacks -as follows: +Relay 0.10 changes every unary and streaming LLM execution-intercept callback. +Update callbacks as follows: -| Surface | Relay 0.10 callback shape | -|---|---| -| Rust, Python, native plugin, Rust worker, Python worker | `(name, request, context, next)` | -| Node.js, Go | `(request, context, next)` | -| C | `(user_data, name, request, context, next, next_ctx)` | +| Surface | Relay 0.9 | Relay 0.10 | +|---|---|---| +| Rust, Python unary, typed native plugin, Rust worker, Python worker | `(name, request, next)` | `(name, request, context, next)` | +| Python language-binding streaming | `(request, next)` | `(name, request, context, next)` | +| Node.js, Go | `(request, next)` | `(request, context, next)` | +| C | `(user_data, request, next, next_ctx)` | `(user_data, name, request, context, next, next_ctx)` | +| Raw native callback | `(..., name, request, next, ...)` | `(..., name, request, context, next, ...)` | The request direction reports the selected codec and exposes decode and encode operations when Relay resolved one. Unary execution also exposes response identity and decode. Streaming execution has no response codec because chunks are not complete provider responses. +Codec access is invocation-scoped. Do not cache the context or a resolved codec: +unary access expires when the callback settles, and streaming request access +expires when the returned stream closes. + This is a source and binary compatibility break for execution-intercept users: - Recompile language-binding consumers that register an LLM execution From cd8c5bec3e33012534bcdfb76b48c842e43caf49 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 22 Sep 2026 15:11:42 -0400 Subject: [PATCH 04/22] fix: harden execution codec context validation Signed-off-by: Alex Fournier --- crates/core/src/plugin/dynamic/worker.rs | 31 +-- .../core/tests/unit/dynamic_worker_tests.rs | 243 ------------------ crates/worker/src/lib.rs | 11 +- .../tests/unit/execution_context_tests.rs | 62 ++++- crates/worker/tests/worker_sdk_tests.rs | 38 +-- python/plugin/src/nemo_relay_plugin/_api.py | 22 +- python/tests/plugin/test_worker_sdk.py | 68 ++++- 7 files changed, 177 insertions(+), 298 deletions(-) diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index aeaf86879..74fcca2eb 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -2104,11 +2104,7 @@ impl WorkerPluginCallback { let _completion = WorkerStreamCompletionSignal(completion_tx); let _codec_capabilities = codec_capabilities; let result = tokio::select! { - // Stream setup can include awaiting the downstream provider through - // `next`, so the control-plane timeout must not cap it. Dropping the - // caller closes `rx`; the sibling branch then cancels the worker and - // releases the continuation and codec capabilities. - result = client.invoke_stream(worker_rpc_request(invoke)) => result, + result = worker_rpc(client.invoke_stream(worker_rpc_request(invoke))) => result, _ = tx.closed() => { guard.cancel("host stopped consuming the worker stream"); guard.finish(); @@ -2261,15 +2257,9 @@ impl WorkerPluginCallback { async fn invoke_async(&self, request: InvokeRequest) -> FlowResult { let callback_name = request.registration_name.clone(); let surface = request.surface; - // A continuation-bearing callback can legitimately include downstream - // provider latency. The caller still owns cancellation through the - // invocation guard, but the control-plane timeout must not cap `next`. - let result = if request.continuation_id.is_empty() { - self.invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) - .await - } else { - self.invoke_async_without_timeout(request).await - }; + let result = self + .invoke_async_with_timeout(request, WORKER_RPC_TIMEOUT) + .await; if let Err(error) = &result { let surface_name = RegistrationSurface::try_from(surface) .map(|surface| surface.as_str_name()) @@ -2286,19 +2276,6 @@ impl WorkerPluginCallback { result } - async fn invoke_async_without_timeout( - &self, - request: InvokeRequest, - ) -> FlowResult { - let mut guard = WorkerInvocationGuard::new(self, &request); - let mut client = self.client.clone(); - let result = client.invoke(worker_rpc_request(request)).await; - guard.finish(); - result - .map(|response| response.into_inner()) - .map_err(|err| worker_status_to_flow("worker invoke failed", err)) - } - async fn invoke_async_with_timeout( &self, request: InvokeRequest, diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index b661b5991..e863ae34a 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -1614,212 +1614,6 @@ async fn callback_timeout_sends_explicit_worker_cancellation() { assert!(cancellation.reason.contains("timed out")); } -#[tokio::test(start_paused = true)] -async fn continuation_bearing_worker_tool_callback_allows_slow_next() { - enable_operational_logs(); - let host_state = shared_worker_host_state(); - let (started_tx, started_rx) = oneshot::channel(); - let started_tx = Arc::new(Mutex::new(Some(started_tx))); - let (callback, _shutdown, mut cancel_rx) = fake_callback_service_with_handlers( - { - let host_state = Arc::clone(&host_state); - let started_tx = Arc::clone(&started_tx); - move |request| { - let host_state = Arc::clone(&host_state); - let started_tx = Arc::clone(&started_tx); - Box::pin(async move { - let continuation = continuation_for(&host_state, &request.continuation_id); - let Continuation::Tool { next, .. } = continuation else { - panic!("expected a tool continuation"); - }; - if let Some(started) = started_tx.lock().unwrap().take() { - let _ = started.send(()); - } - let response = next(json!({"input": "slow"})) - .await - .expect("slow downstream tool must complete"); - InvokeResponse { - result: Some(InvokeResult::ToolExecution(ToolExecutionInterceptResult { - outcome: Some(ProtoToolExecutionInterceptOutcome { - result: Some(JsonValue { - json: serde_json::to_vec(&response.result).unwrap(), - }), - annotation: None, - pending_marks: None, - }), - })), - } - }) - } - }, - |_| Box::pin(tokio_stream::empty()), - ) - .await; - *host_state.lock().unwrap() = Some(callback.host_state.clone()); - - let callback_task = callback.clone(); - let result = tokio::spawn(async move { - callback_task - .invoke_tool_execution( - "slow-tool-next", - "lookup", - json!({"input": "original"}), - Some("call-1"), - Arc::new(|value| { - Box::pin(async move { - tokio::time::sleep(WORKER_RPC_TIMEOUT + Duration::from_secs(1)).await; - Ok(ToolExecutionResult::new(json!({"slow": value}))) - }) - }), - ) - .await - }); - started_rx.await.expect("worker tool callback must start"); - tokio::time::advance(WORKER_RPC_TIMEOUT + Duration::from_secs(2)).await; - - assert_eq!( - result.await.unwrap().unwrap().result, - json!({"slow": {"input": "slow"}}) - ); - assert_no_worker_cancellation(&mut cancel_rx, "slow downstream tool execution"); -} - -#[tokio::test(start_paused = true)] -async fn continuation_bearing_worker_callback_allows_slow_next() { - enable_operational_logs(); - let host_state = shared_worker_host_state(); - let (started_tx, started_rx) = oneshot::channel(); - let started_tx = Arc::new(Mutex::new(Some(started_tx))); - let (callback, _shutdown, mut cancel_rx) = fake_callback_service_with_handlers( - { - let host_state = Arc::clone(&host_state); - let started_tx = Arc::clone(&started_tx); - move |request| { - let host_state = Arc::clone(&host_state); - let started_tx = Arc::clone(&started_tx); - Box::pin(async move { - let continuation = continuation_for(&host_state, &request.continuation_id); - let Continuation::Llm { next, .. } = continuation else { - panic!("expected an LLM continuation"); - }; - if let Some(started) = started_tx.lock().unwrap().take() { - let _ = started.send(()); - } - tokio::time::sleep(WORKER_RPC_TIMEOUT + Duration::from_secs(1)).await; - let response = next(valid_llm_request()) - .await - .expect("slow downstream provider must complete"); - InvokeResponse { - result: Some(InvokeResult::Json(JsonResult { - value: Some(json_envelope(JSON_SCHEMA, &response).unwrap()), - error: None, - })), - } - }) - } - }, - |_| Box::pin(tokio_stream::empty()), - ) - .await; - *host_state.lock().unwrap() = Some(callback.host_state.clone()); - - let callback_task = callback.clone(); - let result = tokio::spawn(async move { - callback_task - .invoke_llm_execution( - "slow-next", - "model", - valid_llm_request(), - openai_execution_codec_context(), - Arc::new(|_| Box::pin(async { Ok(json!({"slow": "completed"})) })), - ) - .await - }); - started_rx.await.expect("worker callback must start"); - tokio::time::advance(WORKER_RPC_TIMEOUT + Duration::from_secs(2)).await; - - assert_eq!(result.await.unwrap().unwrap(), json!({"slow": "completed"})); - assert_no_worker_cancellation(&mut cancel_rx, "slow downstream execution"); -} - -#[tokio::test(start_paused = true)] -async fn continuation_bearing_worker_stream_allows_slow_next() { - enable_operational_logs(); - let host_state = shared_worker_host_state(); - let (started_tx, started_rx) = oneshot::channel(); - let started_tx = Arc::new(Mutex::new(Some(started_tx))); - let (callback, _shutdown, mut cancel_rx) = fake_callback_service_with_async_stream_handler( - |_| { - Box::pin(async { - InvokeResponse { - result: Some(InvokeResult::Empty(EmptyResult {})), - } - }) - }, - { - let host_state = Arc::clone(&host_state); - let started_tx = Arc::clone(&started_tx); - move |request| { - let host_state = Arc::clone(&host_state); - let started_tx = Arc::clone(&started_tx); - Box::pin(async move { - let continuation = continuation_for(&host_state, &request.continuation_id); - let Continuation::LlmStream { next, .. } = continuation else { - panic!("expected an LLM stream continuation"); - }; - if let Some(started) = started_tx.lock().unwrap().take() { - let _ = started.send(()); - } - tokio::time::sleep(WORKER_RPC_TIMEOUT + Duration::from_secs(1)).await; - let stream = next(valid_llm_request()) - .await - .expect("slow downstream stream must open"); - Box::pin(stream.map(|item| { - item.map(|value| StreamChunk { - item: Some(StreamItem::Value( - json_envelope(JSON_SCHEMA, &value) - .expect("stream value must encode"), - )), - }) - .map_err(|error| Status::internal(error.to_string())) - })) as FakeInvokeStream - }) - } - }, - ) - .await; - *host_state.lock().unwrap() = Some(callback.host_state.clone()); - - let callback_task = callback.clone(); - let result = tokio::spawn(async move { - callback_task - .invoke_llm_stream_execution( - "slow-stream-next", - "model", - valid_llm_request(), - openai_stream_execution_codec_context(), - Arc::new(|_| { - Box::pin(async { - Ok(LlmJsonStream::new(tokio_stream::iter(vec![Ok( - json!({"slow_stream": "completed"}), - )]))) - }) - }), - ) - .await - }); - started_rx.await.expect("worker stream callback must start"); - tokio::time::advance(WORKER_RPC_TIMEOUT + Duration::from_secs(2)).await; - - let mut stream = result.await.unwrap().expect("slow worker stream must open"); - assert_eq!( - stream.next().await.unwrap().unwrap(), - json!({"slow_stream": "completed"}) - ); - assert!(stream.next().await.is_none()); - assert_no_worker_cancellation(&mut cancel_rx, "slow downstream stream setup"); -} - #[tokio::test(flavor = "multi_thread")] async fn dropping_callback_future_cancels_worker_and_cleans_host_state() { enable_operational_logs(); @@ -3475,29 +3269,6 @@ fn assert_response_codec_expired( ); } -fn continuation_for(state: &SharedWorkerHostState, continuation_id: &str) -> Continuation { - state - .lock() - .unwrap() - .as_ref() - .expect("host state") - .continuation(continuation_id) - .expect("continuation must remain active") -} - -fn assert_no_worker_cancellation( - cancel_rx: &mut mpsc::UnboundedReceiver, - operation: &str, -) { - assert!( - matches!( - cancel_rx.try_recv(), - Err(tokio::sync::mpsc::error::TryRecvError::Empty) - ), - "{operation} must not trigger worker cancellation" - ); -} - async fn fake_callback_service( invoke: impl Fn(InvokeRequest) -> InvokeResponse + Send + Sync + 'static, ) -> (WorkerPluginCallback, oneshot::Sender<()>) { @@ -3527,20 +3298,6 @@ async fn fake_callback_service_with_handlers( (callback, shutdown_tx, cancel_rx) } -async fn fake_callback_service_with_async_stream_handler( - invoke: impl Fn(InvokeRequest) -> FakeInvokeFuture + Send + Sync + 'static, - invoke_stream: impl Fn(InvokeRequest) -> FakeInvokeStreamFuture + Send + Sync + 'static, -) -> ( - WorkerPluginCallback, - oneshot::Sender<()>, - mpsc::UnboundedReceiver, -) { - let (client, shutdown_tx, cancel_rx, _register_calls) = - fake_worker_client_with_async_handlers(invoke, invoke_stream).await; - let (callback, shutdown_tx) = callback_for_client(client, shutdown_tx); - (callback, shutdown_tx, cancel_rx) -} - fn callback_for_client( client: PluginWorkerClient, shutdown_tx: oneshot::Sender<()>, diff --git a/crates/worker/src/lib.rs b/crates/worker/src/lib.rs index 802f21f5e..2661157ae 100644 --- a/crates/worker/src/lib.rs +++ b/crates/worker/src/lib.rs @@ -2066,7 +2066,7 @@ impl PluginWorker for WorkerService { .ok_or_else(|| Status::not_found("stream execution handler not registered"))?; let payload = llm_payload(request.payload).map_err(status_from_sdk)?; let execution_context = payload - .execution_context(&self.runtime, &invocation_id) + .execution_context(&self.runtime, &invocation_id, false) .map_err(status_from_sdk)?; let request_value = required_json::(payload.request, "llm request").map_err(status_from_sdk)?; @@ -2705,7 +2705,8 @@ impl WorkerService { scope: &Option, ) -> Result { let payload = llm_payload(request.payload)?; - let execution_context = payload.execution_context(&self.runtime, &request.invocation_id)?; + let execution_context = + payload.execution_context(&self.runtime, &request.invocation_id, true)?; let request_value = required_json::(payload.request, "llm request")?; let handler = self.llm_execution(&request.registration_name)?; let next = LlmNext { @@ -2909,6 +2910,7 @@ impl LlmPayload { &self, runtime: &PluginRuntime, invocation_id: &str, + response_required: bool, ) -> Result { let context = require_execution_field( self.execution_codec_context.as_ref(), @@ -2934,6 +2936,11 @@ impl LlmPayload { }) }) .transpose()?; + if response_required && response_codec.is_none() { + return Err(WorkerSdkError::InvalidInput( + "malformed LLM execution codec context: response context is missing".into(), + )); + } Ok(LlmExecutionContext { request_codec: LlmSanitizeRequestContext { codec: codec_identity_from_proto(Some(request_identity)), diff --git a/crates/worker/tests/unit/execution_context_tests.rs b/crates/worker/tests/unit/execution_context_tests.rs index c5d603c07..bdf75dc4f 100644 --- a/crates/worker/tests/unit/execution_context_tests.rs +++ b/crates/worker/tests/unit/execution_context_tests.rs @@ -46,11 +46,63 @@ fn absent_execution_context_is_a_release_mismatch() { let payload = llm_payload(None); let error = payload - .execution_context(&disconnected_runtime(), "invocation") + .execution_context(&disconnected_runtime(), "invocation", true) .unwrap_err(); assert!(error.to_string().contains("execution context is missing")); } +#[test] +fn malformed_execution_context_fields_are_rejected() { + let request = || nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity::default()), + codec_capability_id: None, + }; + let response = || nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { + codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity::default()), + codec_capability_id: None, + }; + let cases = [ + ( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: None, + response: Some(response()), + }, + "request context is missing", + ), + ( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext::default()), + response: Some(response()), + }, + "request codec identity is missing", + ), + ( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(request()), + response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext::default()), + }, + "response codec identity is missing", + ), + ( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(request()), + response: None, + }, + "response context is missing", + ), + ]; + + for (context, expected) in cases { + let error = llm_payload(Some(Box::new(context))) + .execution_context(&disconnected_runtime(), "invocation", true) + .unwrap_err(); + assert!( + error.to_string().contains(expected), + "expected '{expected}' in '{error}'" + ); + } +} + #[test] fn execution_context_preserves_directional_identities_and_capabilities() { let codec = nemo_relay_worker_proto::v1::LlmCodecIdentity { @@ -71,7 +123,7 @@ fn execution_context_preserves_directional_identities_and_capabilities() { ))); let context = payload - .execution_context(&disconnected_runtime(), "invocation") + .execution_context(&disconnected_runtime(), "invocation", true) .unwrap(); assert_eq!( &context.request_codec().codec, @@ -106,7 +158,7 @@ fn execution_context_distinguishes_absent_and_resolved_opaque_codecs() { }), }, ))) - .execution_context(&disconnected_runtime(), "absent-invocation") + .execution_context(&disconnected_runtime(), "absent-invocation", true) .unwrap(); assert_eq!(absent.request_codec().codec, LlmCodecIdentity::None); assert!(absent.request_codec().resolve_codec().is_none()); @@ -132,7 +184,7 @@ fn execution_context_distinguishes_absent_and_resolved_opaque_codecs() { }), }, ))) - .execution_context(&disconnected_runtime(), "opaque-invocation") + .execution_context(&disconnected_runtime(), "opaque-invocation", true) .unwrap(); assert_eq!(opaque.request_codec().codec, LlmCodecIdentity::Opaque); assert!(opaque.request_codec().resolve_codec().is_some()); @@ -157,7 +209,7 @@ fn streaming_execution_context_has_no_response_codec() { ))); let context = payload - .execution_context(&disconnected_runtime(), "invocation") + .execution_context(&disconnected_runtime(), "invocation", false) .unwrap(); assert!(context.request_codec().resolve_codec().is_some()); assert!(context.response_codec().is_none()); diff --git a/crates/worker/tests/worker_sdk_tests.rs b/crates/worker/tests/worker_sdk_tests.rs index ef2bb42ad..b74f1158b 100644 --- a/crates/worker/tests/worker_sdk_tests.rs +++ b/crates/worker/tests/worker_sdk_tests.rs @@ -3271,21 +3271,7 @@ fn llm_invoke( annotated_request: Option, response: Option, ) -> InvokeRequest { - let execution_codec_context = match surface { - RegistrationSurface::LlmExecutionIntercept => Some(Box::new( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { - request: Some(empty_request_codec_context()), - response: Some(empty_response_codec_context()), - }, - )), - RegistrationSurface::LlmStreamExecutionIntercept => Some(Box::new( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { - request: Some(empty_request_codec_context()), - response: None, - }, - )), - _ => None, - }; + let execution_codec_context = execution_codec_context_for(surface); InvokeRequest { activation_id: ACTIVATION_ID.into(), invocation_id: "invoke-1".into(), @@ -3309,6 +3295,26 @@ fn llm_invoke( } } +fn execution_codec_context_for( + surface: RegistrationSurface, +) -> Option> { + match surface { + RegistrationSurface::LlmExecutionIntercept => Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(empty_request_codec_context()), + response: Some(empty_response_codec_context()), + }, + )), + RegistrationSurface::LlmStreamExecutionIntercept => Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(empty_request_codec_context()), + response: None, + }, + )), + _ => None, + } +} + fn llm_invoke_without_request( registration_name: &str, surface: RegistrationSurface, @@ -3328,7 +3334,7 @@ fn llm_invoke_without_request( annotated_request: None, response: None, sanitize_context: None, - execution_codec_context: None, + execution_codec_context: execution_codec_context_for(surface), }, )), } diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index b5a724470..a9e1aaa6a 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -299,11 +299,15 @@ def _llm_execution_context( invocation: pb.LlmInvocation, runtime: "PluginRuntime", invocation_id: str, + *, + response_required: bool, ) -> LlmExecutionContext: if not invocation.HasField("execution_codec_context"): raise WorkerSdkError("malformed LLM execution codec context: execution context is missing") context = invocation.execution_codec_context - if not context.HasField("request") or not context.request.HasField("codec"): + if not context.HasField("request"): + raise WorkerSdkError("malformed LLM execution codec context: request context is missing") + if not context.request.HasField("codec"): raise WorkerSdkError("malformed LLM execution codec context: request codec identity is missing") request_id = context.request.codec_capability_id if context.request.HasField("codec_capability_id") else None @@ -330,6 +334,8 @@ def _llm_execution_context( _capability_id=response_id, _invocation_id=invocation_id, ) + elif response_required: + raise WorkerSdkError("malformed LLM execution codec context: response context is missing") return LlmExecutionContext( request_codec=request_context, response_codec=response_context, @@ -2540,7 +2546,12 @@ async def _produce_stream(self, request: Any, queue: asyncio.Queue[Any]) -> None handler = self._handler(self._handlers.llm_stream_executions, request.registration_name) payload = _require_payload(request, "llm") llm_request = _decode_required_envelope(payload.request, "llm request", LLM_REQUEST_SCHEMA) - execution_context = _llm_execution_context(payload, self._runtime, request.invocation_id) + execution_context = _llm_execution_context( + payload, + self._runtime, + request.invocation_id, + response_required=False, + ) next_call = LlmStreamNext(self._runtime, request.continuation_id) with _bind_invocation_scope(request): stream = await _maybe_await(handler(payload.model_name, llm_request, execution_context, next_call)) @@ -2716,7 +2727,12 @@ async def _invoke_llm_result(self, request: Any) -> Any: self._handler(self._handlers.llm_executions, request.registration_name)( payload.model_name, _decode_required_envelope(payload.request, "llm request", LLM_REQUEST_SCHEMA), - _llm_execution_context(payload, self._runtime, request.invocation_id), + _llm_execution_context( + payload, + self._runtime, + request.invocation_id, + response_required=True, + ), LlmNext(self._runtime, request.continuation_id), ) ) diff --git a/python/tests/plugin/test_worker_sdk.py b/python/tests/plugin/test_worker_sdk.py index 7f2c9d3e8..adb20f1db 100644 --- a/python/tests/plugin/test_worker_sdk.py +++ b/python/tests/plugin/test_worker_sdk.py @@ -1202,7 +1202,12 @@ def test_execution_context_distinguishes_absent_and_resolved_opaque_codecs() -> response=pb.LlmSanitizeResponseContext(codec=pb.LlmCodecIdentity()), ) ) - absent = plugin_api._llm_execution_context(absent_invocation, runtime, "absent-invocation") + absent = plugin_api._llm_execution_context( + absent_invocation, + runtime, + "absent-invocation", + response_required=True, + ) assert absent.request_codec.codec == plugin_api.LlmCodecIdentity("none") assert absent.request_codec.resolve_codec() is None assert absent.response_codec is not None @@ -1221,7 +1226,12 @@ def test_execution_context_distinguishes_absent_and_resolved_opaque_codecs() -> ), ) ) - opaque = plugin_api._llm_execution_context(opaque_invocation, runtime, "opaque-invocation") + opaque = plugin_api._llm_execution_context( + opaque_invocation, + runtime, + "opaque-invocation", + response_required=True, + ) assert opaque.request_codec.codec == plugin_api.LlmCodecIdentity("opaque") assert opaque.request_codec.resolve_codec() is not None assert opaque.response_codec is not None @@ -1229,6 +1239,60 @@ def test_execution_context_distinguishes_absent_and_resolved_opaque_codecs() -> assert opaque.response_codec.resolve_codec() is not None +@pytest.mark.parametrize( + ("invocation", "expected"), + [ + (pb.LlmInvocation(), "execution context is missing"), + ( + pb.LlmInvocation(execution_codec_context=pb.LlmExecutionCodecContext()), + "request context is missing", + ), + ( + pb.LlmInvocation( + execution_codec_context=pb.LlmExecutionCodecContext( + request=pb.LlmSanitizeRequestContext(), + ) + ), + "request codec identity is missing", + ), + ( + pb.LlmInvocation( + execution_codec_context=pb.LlmExecutionCodecContext( + request=pb.LlmSanitizeRequestContext(codec=pb.LlmCodecIdentity()), + response=pb.LlmSanitizeResponseContext(), + ) + ), + "response codec identity is missing", + ), + ( + pb.LlmInvocation( + execution_codec_context=pb.LlmExecutionCodecContext( + request=pb.LlmSanitizeRequestContext(codec=pb.LlmCodecIdentity()), + ) + ), + "response context is missing", + ), + ], +) +def test_execution_context_rejects_malformed_unary_contexts( + invocation: pb.LlmInvocation, + expected: str, +) -> None: + runtime = PluginRuntime( + activation_id=ACTIVATION_ID, + auth_token=AUTH_TOKEN, + host_stub=RecordingHostStub(), + ) + + with pytest.raises(WorkerSdkError, match=expected): + plugin_api._llm_execution_context( + invocation, + runtime, + "invocation", + response_required=True, + ) + + async def test_llm_sanitizers_receive_codec_context_and_can_omit_payloads() -> None: seen: list[tuple[str, LlmSanitizeRequestContext | LlmSanitizeResponseContext]] = [] From bff02a7f9bf5b366fd08d133174b0efffa409846 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 22 Sep 2026 15:29:44 -0400 Subject: [PATCH 05/22] test: update execution context callbacks Signed-off-by: Alex Fournier --- crates/node/tests/llm_tests.mjs | 2 +- go/nemo_relay/callbacks.go | 7 ++++--- go/nemo_relay/intercepts/intercepts.go | 14 ++++++++------ go/nemo_relay/nemo_relay.go | 13 +++++++------ python/tests/test_llm.py | 2 +- 5 files changed, 21 insertions(+), 17 deletions(-) diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index 6ac75ede2..49d8152b6 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -1312,7 +1312,7 @@ describe('LLM intercepts', () => { const events = []; const observed = []; registerSubscriber('node_llm_exec_propagated_w3c', (event) => events.push(event)); - registerLlmExecutionIntercept('node_llm_exec_propagated_w3c', 10, async (request, next) => { + registerLlmExecutionIntercept('node_llm_exec_propagated_w3c', 10, async (request, _context, next) => { const context = lib.capturePropagationContext(); const rootless = lib.captureRootlessPropagationContext(); const explicitRoot = '018f13f0-7c1a-7a80-8000-000000000799'; diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index d887c379d..626500d84 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -340,9 +340,10 @@ type LLMExecutionFunc func(requestJSON json.RawMessage) (json.RawMessage, error) // LLMExecutionInterceptFunc is a callback for LLM execution intercepts // following the middleware chain pattern. It receives the serialized LLMRequest -// as JSON and a `next` function. Call `next` to invoke the next intercept in -// the chain (or the original LLM implementation if this is the innermost -// intercept). Skip calling `next` to short-circuit the chain entirely. +// as JSON, the invocation's codec context, and a `next` function. Call `next` +// to invoke the next intercept in the chain (or the original LLM implementation +// if this is the innermost intercept). Skip calling `next` to short-circuit the +// chain entirely. type LLMExecutionInterceptFunc func(requestJSON json.RawMessage, context LLMExecutionContext, next func(json.RawMessage) (json.RawMessage, error)) (json.RawMessage, error) // CollectorFunc is a callback invoked with each intercepted chunk during a diff --git a/go/nemo_relay/intercepts/intercepts.go b/go/nemo_relay/intercepts/intercepts.go index 7ffbdd2bf..6f647a23b 100644 --- a/go/nemo_relay/intercepts/intercepts.go +++ b/go/nemo_relay/intercepts/intercepts.go @@ -87,9 +87,10 @@ func DeregisterLlmRequest(name string) error { // --- LLM Execution --- // RegisterLlmExecution registers an LLM execution intercept following the -// middleware chain pattern. execFn is called with the request and a next -// function. Call next to continue the chain or skip it to short-circuit. This -// is a shorthand for [nemo_relay.RegisterLlmExecutionIntercept]. +// middleware chain pattern. execFn is called with the request, codec context, +// and a next function. Call next to continue the chain or skip it to +// short-circuit. This is a shorthand for +// [nemo_relay.RegisterLlmExecutionIntercept]. func RegisterLlmExecution(name string, priority int32, execFn nemo_relay.LLMExecutionInterceptFunc) error { return nemo_relay.RegisterLlmExecutionIntercept(name, priority, execFn) } @@ -103,9 +104,10 @@ func DeregisterLlmExecution(name string) error { // --- LLM Stream Execution --- // RegisterLlmStreamExecution registers a streaming LLM execution intercept -// following the middleware chain pattern. execFn is called with the request and -// a next function. Call next to continue the chain or skip it to short-circuit. -// This is a shorthand for [nemo_relay.RegisterLlmStreamExecutionIntercept]. +// following the middleware chain pattern. execFn is called with the request, +// codec context, and a next function. Call next to continue the chain or skip +// it to short-circuit. This is a shorthand for +// [nemo_relay.RegisterLlmStreamExecutionIntercept]. func RegisterLlmStreamExecution(name string, priority int32, execFn nemo_relay.LLMExecutionInterceptFunc) error { return nemo_relay.RegisterLlmStreamExecutionIntercept(name, priority, execFn) } diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index 485c83774..553278239 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -1899,9 +1899,10 @@ func DeregisterLlmRequestIntercept(name string) error { } // RegisterLlmExecutionIntercept registers an execution intercept following -// the middleware chain pattern. execFn is called with the request parameters -// and a `next` function. Call `next` to invoke the next intercept or original -// implementation; skip calling `next` to short-circuit the chain. +// the middleware chain pattern. execFn is called with the request parameters, +// codec context, and a `next` function. Call `next` to invoke the next +// intercept or original implementation; skip calling `next` to short-circuit +// the chain. func RegisterLlmExecutionIntercept(name string, priority int32, execFn LLMExecutionInterceptFunc) error { execID := registerClosure(execFn) cName := C.CString(name) @@ -1924,9 +1925,9 @@ func DeregisterLlmExecutionIntercept(name string) error { // RegisterLlmStreamExecutionIntercept registers an execution intercept for // streaming LLM calls following the middleware chain pattern. execFn is called -// with the request parameters and a `next` function. Call `next` to invoke the -// next intercept or original implementation; skip calling `next` to -// short-circuit. +// with the request parameters, codec context, and a `next` function. Call +// `next` to invoke the next intercept or original implementation; skip calling +// `next` to short-circuit. func RegisterLlmStreamExecutionIntercept(name string, priority int32, execFn LLMExecutionInterceptFunc) error { execID := registerClosure(execFn) cName := C.CString(name) diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index 0dcd63bfe..4f573cf80 100644 --- a/python/tests/test_llm.py +++ b/python/tests/test_llm.py @@ -936,7 +936,7 @@ async def test_execution_callback_preserves_imported_w3c_trace_context(self) -> observed = [] subscribers.register("py_llm_capture_propagated_w3c", events.append) - async def execution_intercept(_name, request, next_handler): + async def execution_intercept(_name, request, _context, next_handler): context = capture_propagation_context() rootless = capture_rootless_propagation_context() explicit_root = capture_propagation_context_with_root(root_uuid) From 017273ae0966352f7f75da12c9831d77f5500c9f Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 22 Sep 2026 16:14:25 -0400 Subject: [PATCH 06/22] test: consolidate execution context coverage Signed-off-by: Alex Fournier --- .../tests/integration/worker_plugin_tests.rs | 10 - .../core/tests/unit/dynamic_worker_tests.rs | 22 -- crates/core/tests/unit/llm_api_tests.rs | 229 +----------------- crates/ffi/tests/unit/callable_tests.rs | 54 ++--- crates/node/tests/llm_tests.mjs | 76 ++---- crates/plugin/tests/typed_callbacks.rs | 16 -- .../tests/unit/execution_context_tests.rs | 175 +++++-------- python/tests/plugin/test_worker_sdk.py | 84 +++---- python/tests/test_llm.py | 138 ++++------- 9 files changed, 171 insertions(+), 633 deletions(-) diff --git a/crates/core/tests/integration/worker_plugin_tests.rs b/crates/core/tests/integration/worker_plugin_tests.rs index f29b2d62d..b2345dc4f 100644 --- a/crates/core/tests/integration/worker_plugin_tests.rs +++ b/crates/core/tests/integration/worker_plugin_tests.rs @@ -1745,16 +1745,6 @@ async fn python_worker_execution_codec_context_round_trips_host_codecs() { request_codec: Arc::new(RuntimeOpenAiChatCodec), response_codec: Arc::new(RuntimeOpenAiChatCodec), }, - // A third invocation proves the spawned worker remains responsive after - // exercising both directional codec capabilities. - Case { - name: "post-round-trip-health", - identity_kind: "builtin", - identity_id: "openai_chat", - answer: "healthy", - request_codec: Arc::new(OpenAIChatCodec), - response_codec: Arc::new(OpenAIChatCodec), - }, ] { let provider_only = json!({"lane": case.name, "preserve": true}); let expected_provider_only = provider_only.clone(); diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index e863ae34a..decf7316c 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -2747,28 +2747,6 @@ async fn host_runtime_service_reports_poisoned_internal_locks() { } }); let codec = Arc::new(OpenAIChatCodec); - let sanitizer_error = callback - .invoke_llm_sanitize_request( - "poisoned-sanitizer-codec-context", - valid_llm_request(), - LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), - ) - .await - .expect_err("poisoned sanitizer codec setup must fail before invoking the worker"); - assert!(matches!( - sanitizer_error, - FlowError::Internal(message) if message.contains("codec lock poisoned") - )); - assert!( - callback - .host_state - .scope_stacks - .lock() - .expect("scope stack lock") - .is_empty(), - "failed sanitizer codec setup must remove its invocation scope stack" - ); - let context = LlmExecutionContext::new( LlmSanitizeRequestContext::for_request_codec(Some(codec.clone())), Some(LlmSanitizeResponseContext::for_response_codec(Some(codec))), diff --git a/crates/core/tests/unit/llm_api_tests.rs b/crates/core/tests/unit/llm_api_tests.rs index a68ceb179..9919cd17b 100644 --- a/crates/core/tests/unit/llm_api_tests.rs +++ b/crates/core/tests/unit/llm_api_tests.rs @@ -5,7 +5,6 @@ #![allow(clippy::await_holding_lock)] -use std::collections::BTreeSet; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Barrier, Mutex}; @@ -22,10 +21,8 @@ use super::{ use crate::api::event::{Event, ScopeCategory}; use crate::api::optimization::finalize_optimization_summary; use crate::api::registry::{ - RuntimeRegistrationKind, deregister_conditional_middleware_guardrail, deregister_llm_execution_intercept, deregister_llm_stream_execution_intercept, - register_conditional_middleware_guardrail, register_llm_execution_intercept, - register_llm_stream_execution_intercept, scope_register_llm_execution_intercept, + register_llm_execution_intercept, register_llm_stream_execution_intercept, }; use crate::api::registry::{ deregister_llm_sanitize_request_guardrail, deregister_llm_sanitize_response_guardrail, @@ -312,60 +309,6 @@ fn response_sanitizer_context_preserves_all_codec_identity_states() { ); } -#[test] -fn managed_execution_passes_codec_context_to_downstream_interceptor() { - let _guard = lock_global_runtime(); - reset_global(); - set_thread_scope_stack(create_scope_stack()); - - let observations = Arc::new(Mutex::new(Vec::new())); - let captured = Arc::clone(&observations); - register_llm_execution_intercept( - "execution-codec-context-outer", - 1, - Arc::new(move |_name, request, context, next| { - let captured = Arc::clone(&captured); - Box::pin(async move { - assert_openai_execution_context(&context); - captured.lock().unwrap().push("before-next"); - let result = next(request).await; - assert_openai_execution_context(&context); - captured.lock().unwrap().push("after-next"); - result - }) - }), - ) - .unwrap(); - let captured = Arc::clone(&observations); - register_llm_execution_intercept( - "execution-codec-context-inner", - 2, - Arc::new(move |_name, request, context, next| { - let captured = Arc::clone(&captured); - Box::pin(async move { - assert_openai_execution_context(&context); - captured.lock().unwrap().push("inner"); - next(request).await - }) - }), - ) - .unwrap(); - - tokio::runtime::Runtime::new().unwrap().block_on(async { - execute_openai_call( - "execution-codec-context", - Arc::new(|_| Box::pin(async { Ok(json!({"ok": true})) })), - ) - .await - .unwrap(); - }); - - assert!(deregister_llm_execution_intercept("execution-codec-context-outer").unwrap()); - assert!(deregister_llm_execution_intercept("execution-codec-context-inner").unwrap()); - let observations = observations.lock().unwrap(); - assert_eq!(*observations, ["before-next", "inner", "after-next"]); -} - #[test] fn managed_execution_codec_context_decodes_encodes_and_decodes_response() { let _guard = lock_global_runtime(); @@ -690,127 +633,6 @@ fn dropping_unconsumed_stream_expires_execution_codec_facade() { assert!(deregister_llm_stream_execution_intercept("execution-codec-stream-drop").unwrap()); } -#[test] -fn codec_context_preserves_scope_ordering_and_conditional_gating() { - let _guard = lock_global_runtime(); - reset_global(); - set_thread_scope_stack(create_scope_stack()); - - let scope = push_scope( - PushScopeParams::builder() - .name("execution-codec-context-scope") - .scope_type(ScopeType::Custom) - .build(), - ) - .unwrap(); - let calls = Arc::new(Mutex::new(Vec::new())); - - let captured = Arc::clone(&calls); - register_llm_execution_intercept( - "execution-codec-global-first", - 10, - Arc::new(move |_name, request, context, next| { - let captured = Arc::clone(&captured); - Box::pin(async move { - assert_openai_execution_context(&context); - captured.lock().unwrap().push("global-first-enter"); - let result = next(request).await; - captured.lock().unwrap().push("global-first-exit"); - result - }) - }), - ) - .unwrap(); - - register_llm_execution_intercept( - "execution-codec-gated", - 15, - Arc::new(move |_name, _request, _context, _next| { - Box::pin(async move { panic!("conditionally disabled interceptor must not execute") }) - }), - ) - .unwrap(); - let gate_kinds = BTreeSet::from([RuntimeRegistrationKind::LlmExecutionIntercept]); - register_conditional_middleware_guardrail( - "execution-codec-context-gate", - gate_kinds, - "execution-codec-gated", - Arc::new(|_, _| Some("disabled for regression test".into())), - ) - .unwrap(); - - let captured = Arc::clone(&calls); - register_llm_execution_intercept( - "execution-codec-global-second", - 20, - Arc::new(move |_name, request, context, next| { - let captured = Arc::clone(&captured); - Box::pin(async move { - assert_openai_execution_context(&context); - captured.lock().unwrap().push("global-second-enter"); - let result = next(request).await; - captured.lock().unwrap().push("global-second-exit"); - result - }) - }), - ) - .unwrap(); - - let captured = Arc::clone(&calls); - scope_register_llm_execution_intercept( - &scope.uuid, - "execution-codec-scope", - 30, - Arc::new(move |_name, request, context, next| { - let captured = Arc::clone(&captured); - Box::pin(async move { - assert_openai_execution_context(&context); - captured.lock().unwrap().push("scope-enter"); - let result = next(request).await; - captured.lock().unwrap().push("scope-exit"); - result - }) - }), - ) - .unwrap(); - - let captured = Arc::clone(&calls); - let response = tokio::runtime::Runtime::new().unwrap().block_on(async { - execute_openai_call( - "execution-codec-context-mixed-chain", - Arc::new(move |_| { - let captured = Arc::clone(&captured); - Box::pin(async move { - captured.lock().unwrap().push("provider"); - Ok(json!({"ok": true})) - }) - }), - ) - .await - .unwrap() - }); - - assert_eq!(response, json!({"ok": true})); - assert!(deregister_conditional_middleware_guardrail("execution-codec-context-gate").unwrap()); - assert!(deregister_llm_execution_intercept("execution-codec-global-first").unwrap()); - assert!(deregister_llm_execution_intercept("execution-codec-gated").unwrap()); - assert!(deregister_llm_execution_intercept("execution-codec-global-second").unwrap()); - pop_scope(PopScopeParams::builder().handle_uuid(&scope.uuid).build()).unwrap(); - - assert_eq!( - calls.lock().unwrap().as_slice(), - [ - "global-first-enter", - "global-second-enter", - "scope-enter", - "provider", - "scope-exit", - "global-second-exit", - "global-first-exit", - ] - ); -} - #[test] fn managed_execution_distinguishes_absent_codecs_from_streaming_response_unavailability() { let _guard = lock_global_runtime(); @@ -990,55 +812,6 @@ impl LlmResponseCodec for ProjectionFailingCodec { } } -#[test] -fn execution_context_preserves_runtime_and_opaque_codec_identities() { - let runtime_request: Arc = Arc::new(RuntimeIdentityCodec); - let runtime_response: Arc = Arc::new(RuntimeIdentityCodec); - let runtime_context = - LlmExecutionContext::for_unary_codecs(Some(runtime_request), &Some(runtime_response)); - assert_eq!( - runtime_context.request_codec().codec(), - &LlmCodecIdentity::Runtime("com.example.chat.v1".into()) - ); - assert_eq!( - runtime_context.response_codec().unwrap().codec(), - &LlmCodecIdentity::Runtime("com.example.chat.v1".into()) - ); - assert!(runtime_context.request_codec().resolve_codec().is_some()); - assert!( - runtime_context - .response_codec() - .unwrap() - .resolve_codec() - .is_some() - ); - - let opaque_request: Arc = Arc::new(ProjectionFailingCodec { - projection_attempts: Arc::new(AtomicUsize::new(0)), - }); - let opaque_response: Arc = Arc::new(ProjectionFailingCodec { - projection_attempts: Arc::new(AtomicUsize::new(0)), - }); - let opaque_context = - LlmExecutionContext::for_unary_codecs(Some(opaque_request), &Some(opaque_response)); - assert_eq!( - opaque_context.request_codec().codec(), - &LlmCodecIdentity::Opaque - ); - assert_eq!( - opaque_context.response_codec().unwrap().codec(), - &LlmCodecIdentity::Opaque - ); - assert!(opaque_context.request_codec().resolve_codec().is_some()); - assert!( - opaque_context - .response_codec() - .unwrap() - .resolve_codec() - .is_some() - ); -} - fn emit_compaction() { event( EmitMarkEventParams::builder() diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 6dd25b6b3..9e02de0f7 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -332,7 +332,7 @@ impl nemo_relay::codec::traits::LlmResponseCodec for OpaqueExecutionCodec { } } -unsafe extern "C" fn llm_exec_absent_context_cb( +unsafe extern "C" fn llm_exec_codec_context_cb( _user_data: *mut libc::c_void, name: *const c_char, native_json: *const c_char, @@ -340,42 +340,33 @@ unsafe extern "C" fn llm_exec_absent_context_cb( next_fn: NemoRelayLlmExecNextFn, next_ctx: *mut libc::c_void, ) -> *mut c_char { + let kind = context.request_codec.codec_kind; + let expected_name = match kind { + NemoRelayLlmSanitizeCodecKind::None => "ffi-none", + NemoRelayLlmSanitizeCodecKind::Opaque => "ffi-opaque", + other => panic!("unexpected execution codec kind: {other:?}"), + }; assert_eq!( unsafe { CStr::from_ptr(name) }.to_str().unwrap(), - "ffi-none" + expected_name ); + assert!(context.request_codec.codec_id.is_null()); assert_eq!( - context.request_codec.codec_kind, - NemoRelayLlmSanitizeCodecKind::None + context.request_codec.codec.is_null(), + kind == NemoRelayLlmSanitizeCodecKind::None ); - assert!(context.request_codec.codec_id.is_null()); - assert!(context.request_codec.codec.is_null()); assert!(!context.response_codec.is_null()); let response = unsafe { &*context.response_codec }; - assert_eq!(response.codec_kind, NemoRelayLlmSanitizeCodecKind::None); + assert_eq!(response.codec_kind, kind); assert!(response.codec_id.is_null()); - assert!(response.codec.is_null()); - unsafe { next_fn(native_json, next_ctx) } -} - -unsafe extern "C" fn llm_exec_opaque_context_cb( - _user_data: *mut libc::c_void, - name: *const c_char, - native_json: *const c_char, - context: NemoRelayLlmExecutionContext, - next_fn: NemoRelayLlmExecNextFn, - next_ctx: *mut libc::c_void, -) -> *mut c_char { assert_eq!( - unsafe { CStr::from_ptr(name) }.to_str().unwrap(), - "ffi-opaque" + response.codec.is_null(), + kind == NemoRelayLlmSanitizeCodecKind::None ); - assert_eq!( - context.request_codec.codec_kind, - NemoRelayLlmSanitizeCodecKind::Opaque - ); - assert!(context.request_codec.codec_id.is_null()); - assert!(!context.request_codec.codec.is_null()); + + if kind == NemoRelayLlmSanitizeCodecKind::None { + return unsafe { next_fn(native_json, next_ctx) }; + } let request: LlmRequest = serde_json::from_str(unsafe { CStr::from_ptr(native_json) }.to_str().unwrap()).unwrap(); @@ -402,11 +393,6 @@ unsafe extern "C" fn llm_exec_opaque_context_cb( unsafe { drop(Box::from_raw(encoded)) }; assert!(!result.is_null()); - assert!(!context.response_codec.is_null()); - let response = unsafe { &*context.response_codec }; - assert_eq!(response.codec_kind, NemoRelayLlmSanitizeCodecKind::Opaque); - assert!(response.codec_id.is_null()); - assert!(!response.codec.is_null()); let decoded_ptr = unsafe { crate::api::nemo_relay_llm_sanitize_response_codec_decode(response.codec, result) }; @@ -787,7 +773,7 @@ fn assert_llm_exec_callbacks(runtime: &tokio::runtime::Runtime) { assert_eq!(intercepted["intercepted"], json!(true)); let absent_intercept = - wrap_llm_exec_intercept_fn(llm_exec_absent_context_cb, std::ptr::null_mut(), None); + wrap_llm_exec_intercept_fn(llm_exec_codec_context_cb, std::ptr::null_mut(), None); let absent_next: LlmExecutionNextFn = Arc::new(|request| Box::pin(async move { Ok(json!({"model": request.content["model"]})) })); let absent = runtime @@ -804,7 +790,7 @@ fn assert_llm_exec_callbacks(runtime: &tokio::runtime::Runtime) { assert_eq!(absent["model"], json!("test-model")); let opaque_intercept = - wrap_llm_exec_intercept_fn(llm_exec_opaque_context_cb, std::ptr::null_mut(), None); + wrap_llm_exec_intercept_fn(llm_exec_codec_context_cb, std::ptr::null_mut(), None); let opaque_next: LlmExecutionNextFn = Arc::new(|request| { Box::pin(async move { assert_eq!(request.content["encoded"], json!(true)); diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index 49d8152b6..ef555d12b 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -1633,31 +1633,41 @@ describe('LLM intercepts', () => { deregisterLlmExecutionIntercept('node_llm_exec_repl'); }); - it('execution intercept receives directional codec context', async () => { + it('execution context exposes codec states, operations, and callback lifetime', async () => { const codec = new lib.OpenAIChatCodec(); const response = { id: 'chatcmpl-execution-context', model: 'test-model', choices: [{ index: 0, message: { role: 'assistant', content: 'ok' }, finish_reason: 'stop' }], }; - let observed = false; + const observed = []; + let retainedCodec; registerLlmExecutionIntercept('node_llm_execution_context', 10, async (request, context, next) => { + observed.push(context.requestCodec.codec.kind); + assert.notEqual(context.responseCodec, null); + + if (context.requestCodec.codec.kind === 'none') { + assert.equal(context.requestCodec.resolveCodec(), null); + assert.deepEqual(context.responseCodec.codec, { kind: 'none' }); + assert.equal(context.responseCodec.resolveCodec(), null); + return next(request); + } + assert.deepEqual(context.requestCodec.codec, { kind: 'opaque' }); const requestCodec = context.requestCodec.resolveCodec(); assert.notEqual(requestCodec, null); assert.equal(requestCodec.decode(request).model, 'test-model'); + retainedCodec = requestCodec; - assert.notEqual(context.responseCodec, null); assert.deepEqual(context.responseCodec.codec, { kind: 'opaque' }); const result = await next(request); const responseCodec = context.responseCodec.resolveCodec(); assert.notEqual(responseCodec, null); assert.equal(responseCodec.decodeResponse(result).model, 'test-model'); - observed = true; return result; }); try { - const result = await llmCallExecute( + const opaque = await llmCallExecute( 'node_llm_execution_context', makeNative(), () => response, @@ -1670,67 +1680,23 @@ describe('LLM intercepts', () => { ({ annotated, original }) => codec.encode(annotated, original), codec.decodeResponse.bind(codec), ); - assert.deepEqual(result, response); - assert.equal(observed, true); - } finally { - deregisterLlmExecutionIntercept('node_llm_execution_context'); - } - }); - - it('execution intercept reports absent codecs without capabilities', async () => { - let observed = false; - registerLlmExecutionIntercept('node_llm_execution_context_absent', 10, async (request, context, next) => { - assert.deepEqual(context.requestCodec.codec, { kind: 'none' }); - assert.equal(context.requestCodec.resolveCodec(), null); - assert.notEqual(context.responseCodec, null); - assert.deepEqual(context.responseCodec.codec, { kind: 'none' }); - assert.equal(context.responseCodec.resolveCodec(), null); - observed = true; - return next(request); - }); - try { - const response = await llmCallExecute( + const absent = await llmCallExecute( 'node_llm_execution_context_absent', makeNative(), - () => ({ ok: true }), - null, - null, - null, - null, - null, - ); - assert.deepEqual(response, { ok: true }); - assert.equal(observed, true); - } finally { - deregisterLlmExecutionIntercept('node_llm_execution_context_absent'); - } - }); - - it('execution codec capability expires when the callback settles', async () => { - const codec = new lib.OpenAIChatCodec(); - let retainedCodec; - registerLlmExecutionIntercept('node_llm_execution_codec_expiry', 10, async (request, context, next) => { - retainedCodec = context.requestCodec.resolveCodec(); - assert.notEqual(retainedCodec, null); - return next(request); - }); - try { - await llmCallExecute( - 'node_llm_execution_codec_expiry', - makeNative(), - () => ({ ok: true }), + () => response, null, null, null, null, null, - codec.decode.bind(codec), - ({ annotated, original }) => codec.encode(annotated, original), ); + assert.deepEqual(opaque, response); + assert.deepEqual(absent, response); } finally { - deregisterLlmExecutionIntercept('node_llm_execution_codec_expiry'); + deregisterLlmExecutionIntercept('node_llm_execution_context'); } + assert.deepEqual(observed, ['opaque', 'none']); assert.notEqual(retainedCodec, undefined); assert.throws(() => retainedCodec.decode(makeNative()), /LLM execution codec capability is no longer active/i); }); diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index bb8dd6e80..22677c040 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -739,22 +739,6 @@ fn native_abi_v6_logging_extension_is_append_only() { ); } -#[test] -fn native_abi_v7_execution_context_extension_preserves_the_host_table() { - assert_eq!(offset_of!(NemoRelayNativeHostApiV7, v6), 0); - assert_eq!( - offset_of!( - NemoRelayNativeHostApiV7, - plugin_context_register_async_llm_execution_intercept - ), - size_of::() - ); - assert_eq!( - offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), - size_of::() + size_of::() - ); -} - fn host_api_v4_offsets() -> [usize; 12] { [ offset_of!(NemoRelayNativeHostApiV4, v3), diff --git a/crates/worker/tests/unit/execution_context_tests.rs b/crates/worker/tests/unit/execution_context_tests.rs index bdf75dc4f..80769fbee 100644 --- a/crates/worker/tests/unit/execution_context_tests.rs +++ b/crates/worker/tests/unit/execution_context_tests.rs @@ -26,31 +26,6 @@ fn llm_payload( } } -#[test] -fn execution_registration_uses_the_existing_surface() { - let mut context = PluginContext::new(); - context.register_llm_execution_intercept("context", 7, |_, _, _, _| async { - Ok(serde_json::json!({"context": true})) - }); - - let registration = &context.handlers.registrations[0]; - assert_eq!( - registration.surface, - RegistrationSurface::LlmExecutionIntercept as i32 - ); - assert_eq!(registration.priority, 7); -} - -#[test] -fn absent_execution_context_is_a_release_mismatch() { - let payload = llm_payload(None); - - let error = payload - .execution_context(&disconnected_runtime(), "invocation", true) - .unwrap_err(); - assert!(error.to_string().contains("execution context is missing")); -} - #[test] fn malformed_execution_context_fields_are_rejected() { let request = || nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { @@ -62,38 +37,39 @@ fn malformed_execution_context_fields_are_rejected() { codec_capability_id: None, }; let cases = [ + (None, "execution context is missing"), ( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + Some(nemo_relay_worker_proto::v1::LlmExecutionCodecContext { request: None, response: Some(response()), - }, + }), "request context is missing", ), ( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + Some(nemo_relay_worker_proto::v1::LlmExecutionCodecContext { request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext::default()), response: Some(response()), - }, + }), "request codec identity is missing", ), ( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + Some(nemo_relay_worker_proto::v1::LlmExecutionCodecContext { request: Some(request()), response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext::default()), - }, + }), "response codec identity is missing", ), ( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + Some(nemo_relay_worker_proto::v1::LlmExecutionCodecContext { request: Some(request()), response: None, - }, + }), "response context is missing", ), ]; for (context, expected) in cases { - let error = llm_payload(Some(Box::new(context))) + let error = llm_payload(context.map(Box::new)) .execution_context(&disconnected_runtime(), "invocation", true) .unwrap_err(); assert!( @@ -104,93 +80,58 @@ fn malformed_execution_context_fields_are_rejected() { } #[test] -fn execution_context_preserves_directional_identities_and_capabilities() { - let codec = nemo_relay_worker_proto::v1::LlmCodecIdentity { - kind: LlmCodecKind::Builtin as i32, - id: Some("openai_chat".into()), - }; - let payload = llm_payload(Some(Box::new( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { - request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { - codec: Some(codec.clone()), - codec_capability_id: Some("request-capability".into()), - }), - response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { - codec: Some(codec), - codec_capability_id: Some("response-capability".into()), - }), - }, - ))); - - let context = payload - .execution_context(&disconnected_runtime(), "invocation", true) - .unwrap(); - assert_eq!( - &context.request_codec().codec, - &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) - ); - assert!(context.request_codec().resolve_codec().is_some()); - let response = context.response_codec().expect("unary response context"); - assert_eq!( - &response.codec, - &LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) - ); - assert!(response.resolve_codec().is_some()); -} +fn execution_context_preserves_codec_identities_and_capability_availability() { + let cases = [ + ( + LlmCodecKind::Unspecified, + None, + None, + LlmCodecIdentity::None, + false, + ), + ( + LlmCodecKind::Opaque, + None, + Some("opaque"), + LlmCodecIdentity::Opaque, + true, + ), + ( + LlmCodecKind::Builtin, + Some("openai_chat"), + Some("builtin"), + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat), + true, + ), + ]; -#[test] -fn execution_context_distinguishes_absent_and_resolved_opaque_codecs() { - let absent = llm_payload(Some(Box::new( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { - request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { - codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { - kind: LlmCodecKind::Unspecified as i32, - id: None, + for (kind, id, capability, expected, resolves) in cases { + let codec = || nemo_relay_worker_proto::v1::LlmCodecIdentity { + kind: kind as i32, + id: id.map(str::to_owned), + }; + let capability = |direction: &str| capability.map(|value| format!("{value}-{direction}")); + let context = llm_payload(Some(Box::new( + nemo_relay_worker_proto::v1::LlmExecutionCodecContext { + request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { + codec: Some(codec()), + codec_capability_id: capability("request"), }), - codec_capability_id: None, - }), - response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { - codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { - kind: LlmCodecKind::Unspecified as i32, - id: None, + response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { + codec: Some(codec()), + codec_capability_id: capability("response"), }), - codec_capability_id: None, - }), - }, - ))) - .execution_context(&disconnected_runtime(), "absent-invocation", true) - .unwrap(); - assert_eq!(absent.request_codec().codec, LlmCodecIdentity::None); - assert!(absent.request_codec().resolve_codec().is_none()); - let absent_response = absent.response_codec().expect("unary response direction"); - assert_eq!(absent_response.codec, LlmCodecIdentity::None); - assert!(absent_response.resolve_codec().is_none()); + }, + ))) + .execution_context(&disconnected_runtime(), "invocation", true) + .unwrap(); - let opaque = llm_payload(Some(Box::new( - nemo_relay_worker_proto::v1::LlmExecutionCodecContext { - request: Some(nemo_relay_worker_proto::v1::LlmSanitizeRequestContext { - codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { - kind: LlmCodecKind::Opaque as i32, - id: None, - }), - codec_capability_id: Some("opaque-request".into()), - }), - response: Some(nemo_relay_worker_proto::v1::LlmSanitizeResponseContext { - codec: Some(nemo_relay_worker_proto::v1::LlmCodecIdentity { - kind: LlmCodecKind::Opaque as i32, - id: None, - }), - codec_capability_id: Some("opaque-response".into()), - }), - }, - ))) - .execution_context(&disconnected_runtime(), "opaque-invocation", true) - .unwrap(); - assert_eq!(opaque.request_codec().codec, LlmCodecIdentity::Opaque); - assert!(opaque.request_codec().resolve_codec().is_some()); - let opaque_response = opaque.response_codec().expect("unary response direction"); - assert_eq!(opaque_response.codec, LlmCodecIdentity::Opaque); - assert!(opaque_response.resolve_codec().is_some()); + assert_eq!(context.request_codec().codec, expected); + assert_eq!(context.request_codec().resolve_codec().is_some(), resolves); + let response = context.response_codec().expect("unary response context"); + assert_eq!(response.codec, expected); + assert_eq!(response.resolve_codec().is_some(), resolves); + } } #[test] diff --git a/python/tests/plugin/test_worker_sdk.py b/python/tests/plugin/test_worker_sdk.py index adb20f1db..63a695dfa 100644 --- a/python/tests/plugin/test_worker_sdk.py +++ b/python/tests/plugin/test_worker_sdk.py @@ -736,14 +736,9 @@ def test_generated_proto_matches_worker_contract() -> None: "Shutdown", } assert pb.InvokeRequest.DESCRIPTOR.fields_by_name["auth_token"].number == 7 - assert "llm_execution_codec_context" not in pb.Registration.DESCRIPTOR.fields_by_name execution_context = pb.LlmInvocation.DESCRIPTOR.fields_by_name["execution_codec_context"] assert execution_context.number == 11 assert execution_context.containing_oneof is None - assert {field.name for field in pb.LlmInvocation.DESCRIPTOR.oneofs_by_name["sanitize_context"].fields} == { - "request_sanitize_context", - "response_sanitize_context", - } assert pb.HealthRequest.DESCRIPTOR.fields_by_name["activation_id"].number == 1 assert pb.HealthRequest.DESCRIPTOR.fields_by_name["auth_token"].number == 2 assert pb.SUBSCRIBER == 1 @@ -1116,19 +1111,6 @@ def test_plugin_context_registers_llm_sanitizers_under_standard_names() -> None: ] -def test_execution_codec_context_uses_the_existing_surface() -> None: - context = PluginContext() - - async def execution(_name: str, request: Json, _context: Any, next_call: Any) -> Json: - return await next_call.call(request) - - context.register_llm_execution_intercept("execution", execution, priority=7) - - registration = context._handlers.registrations[0] - assert registration.surface == pb.LLM_EXECUTION_INTERCEPT - assert registration.priority == 7 - - async def test_execution_callback_receives_directional_codec_context() -> None: seen: list[plugin_api.LlmExecutionContext] = [] @@ -1189,54 +1171,46 @@ async def execution(name: str, request: Json, context: Any, next_call: Any) -> J assert codec_capabilities == ["request-capability", "response-capability"] -def test_execution_context_distinguishes_absent_and_resolved_opaque_codecs() -> None: +@pytest.mark.parametrize( + ("proto_kind", "capability_prefix", "expected_kind", "resolves"), + [ + (pb.LLM_CODEC_KIND_UNSPECIFIED, None, "none", False), + (pb.LLM_CODEC_KIND_OPAQUE, "opaque", "opaque", True), + ], +) +def test_execution_context_distinguishes_absent_and_resolved_opaque_codecs( + proto_kind: int, + capability_prefix: str | None, + expected_kind: str, + resolves: bool, +) -> None: runtime = PluginRuntime( activation_id=ACTIVATION_ID, auth_token=AUTH_TOKEN, host_stub=RecordingHostStub(), ) - - absent_invocation = pb.LlmInvocation( - execution_codec_context=pb.LlmExecutionCodecContext( - request=pb.LlmSanitizeRequestContext(codec=pb.LlmCodecIdentity()), - response=pb.LlmSanitizeResponseContext(codec=pb.LlmCodecIdentity()), - ) - ) - absent = plugin_api._llm_execution_context( - absent_invocation, - runtime, - "absent-invocation", - response_required=True, - ) - assert absent.request_codec.codec == plugin_api.LlmCodecIdentity("none") - assert absent.request_codec.resolve_codec() is None - assert absent.response_codec is not None - assert absent.response_codec.codec == plugin_api.LlmCodecIdentity("none") - assert absent.response_codec.resolve_codec() is None - - opaque_invocation = pb.LlmInvocation( + request = pb.LlmSanitizeRequestContext(codec=pb.LlmCodecIdentity(kind=proto_kind)) + response = pb.LlmSanitizeResponseContext(codec=pb.LlmCodecIdentity(kind=proto_kind)) + if capability_prefix is not None: + request.codec_capability_id = f"{capability_prefix}-request" + response.codec_capability_id = f"{capability_prefix}-response" + invocation = pb.LlmInvocation( execution_codec_context=pb.LlmExecutionCodecContext( - request=pb.LlmSanitizeRequestContext( - codec=pb.LlmCodecIdentity(kind=pb.LLM_CODEC_KIND_OPAQUE), - codec_capability_id="opaque-request", - ), - response=pb.LlmSanitizeResponseContext( - codec=pb.LlmCodecIdentity(kind=pb.LLM_CODEC_KIND_OPAQUE), - codec_capability_id="opaque-response", - ), + request=request, + response=response, ) ) - opaque = plugin_api._llm_execution_context( - opaque_invocation, + context = plugin_api._llm_execution_context( + invocation, runtime, - "opaque-invocation", + "invocation", response_required=True, ) - assert opaque.request_codec.codec == plugin_api.LlmCodecIdentity("opaque") - assert opaque.request_codec.resolve_codec() is not None - assert opaque.response_codec is not None - assert opaque.response_codec.codec == plugin_api.LlmCodecIdentity("opaque") - assert opaque.response_codec.resolve_codec() is not None + assert context.request_codec.codec == plugin_api.LlmCodecIdentity(expected_kind) + assert (context.request_codec.resolve_codec() is not None) is resolves + assert context.response_codec is not None + assert context.response_codec.codec == plugin_api.LlmCodecIdentity(expected_kind) + assert (context.response_codec.resolve_codec() is not None) is resolves @pytest.mark.parametrize( diff --git a/python/tests/test_llm.py b/python/tests/test_llm.py index 4f573cf80..b8459b207 100644 --- a/python/tests/test_llm.py +++ b/python/tests/test_llm.py @@ -325,50 +325,7 @@ def sanitize_response(response, context): assert request_codec_used is True assert response_codec_used is True - async def test_execution_intercept_receives_directional_codecs(self) -> None: - observed = False - - async def execution_intercept(name, request, context, next_call): - nonlocal observed - assert name == "py_llm_execution_context" - assert context.request_codec.codec.kind == "builtin" - assert context.request_codec.codec.id == "openai_chat" - request_codec = context.request_codec.resolve_codec() - assert request_codec is not None - assert request_codec.decode(request).model == "test-model" - - assert context.response_codec is not None - assert context.response_codec.codec.kind == "builtin" - assert context.response_codec.codec.id == "openai_chat" - response = await next_call(request) - response_codec = context.response_codec.resolve_codec() - assert response_codec is not None - assert response_codec.decode_response(response).model == "test-model" - observed = True - return response - - codec = OpenAIChatCodec() - response = { - "id": "chatcmpl-execution-context", - "model": "test-model", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}}], - } - intercepts.register_llm_execution("py_llm_execution_context", 1, execution_intercept) - try: - result = await llm.execute( - "py_llm_execution_context", - make_request(), - lambda _request: response, - codec=codec, - response_codec=codec, - ) - finally: - intercepts.deregister_llm_execution("py_llm_execution_context") - - assert result == response - assert observed - - async def test_execution_intercept_distinguishes_absent_and_opaque_codecs(self) -> None: + async def test_execution_context_exposes_codec_states_operations_and_lifetime(self) -> None: class OpaqueCodec: def __init__(self) -> None: self.inner = OpenAIChatCodec() @@ -382,30 +339,40 @@ def encode(self, annotated, original): def decode_response(self, response): return self.inner.decode_response(response) - seen = [] + expected = { + "py_llm_execution_context_builtin": ("builtin", "openai_chat"), + "py_llm_execution_context_opaque": ("opaque", None), + "py_llm_execution_context_absent": ("none", None), + } + seen: list[str] = [] + retained_codec = None async def execution_intercept(name, request, context, next_call): + nonlocal retained_codec seen.append(name) - if name == "py_llm_execution_context_absent": - assert context.request_codec.codec.kind == "none" + expected_kind, expected_id = expected[name] + assert context.request_codec.codec.kind == expected_kind + assert context.response_codec is not None + assert context.response_codec.codec.kind == expected_kind + + if expected_id is not None: + assert context.request_codec.codec.id == expected_id + assert context.response_codec.codec.id == expected_id + + if expected_kind == "none": assert context.request_codec.resolve_codec() is None - assert context.response_codec is not None - assert context.response_codec.codec.kind == "none" assert context.response_codec.resolve_codec() is None return await next_call(request) - assert name == "py_llm_execution_context_opaque" - assert context.request_codec.codec.kind == "opaque" request_codec = context.request_codec.resolve_codec() assert request_codec is not None - encoded = request_codec.encode(request_codec.decode(request), request) + request = request_codec.encode(request_codec.decode(request), request) + retained_codec = retained_codec or request_codec - assert context.response_codec is not None - assert context.response_codec.codec.kind == "opaque" - response_codec = context.response_codec.resolve_codec() - assert response_codec is not None - response = await next_call(encoded) - assert response_codec.decode_response(response).model == "test-model" + response = await next_call(request) + resolved_response_codec = context.response_codec.resolve_codec() + assert resolved_response_codec is not None + assert resolved_response_codec.decode_response(response).model == "test-model" return response response = { @@ -413,49 +380,28 @@ async def execution_intercept(name, request, context, next_call): "model": "test-model", "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}}], } - codec = OpaqueCodec() intercepts.register_llm_execution("py_llm_execution_context_matrix", 1, execution_intercept) try: - absent = await llm.execute( - "py_llm_execution_context_absent", - make_request(), - lambda _request: response, - ) - opaque = await llm.execute( - "py_llm_execution_context_opaque", - make_request(), - lambda _request: response, - codec=codec, - response_codec=codec, - ) + results = [] + for name, codec in ( + ("py_llm_execution_context_builtin", OpenAIChatCodec()), + ("py_llm_execution_context_opaque", OpaqueCodec()), + ("py_llm_execution_context_absent", None), + ): + codec_options = {} if codec is None else {"codec": codec, "response_codec": codec} + results.append( + await llm.execute( + name, + make_request(), + lambda _request: response, + **codec_options, + ) + ) finally: intercepts.deregister_llm_execution("py_llm_execution_context_matrix") - assert absent == response - assert opaque == response - assert seen == ["py_llm_execution_context_absent", "py_llm_execution_context_opaque"] - - async def test_execution_codec_capability_expires_when_callback_settles(self) -> None: - retained_codec = None - - async def execution_intercept(_name, request, context, next_call): - nonlocal retained_codec - retained_codec = context.request_codec.resolve_codec() - assert retained_codec is not None - return await next_call(request) - - codec = OpenAIChatCodec() - intercepts.register_llm_execution("py_llm_execution_codec_expiry", 1, execution_intercept) - try: - await llm.execute( - "py_llm_execution_codec_expiry", - make_request(), - lambda _request: {"ok": True}, - codec=codec, - ) - finally: - intercepts.deregister_llm_execution("py_llm_execution_codec_expiry") - + assert results == [response, response, response] + assert seen == list(expected) assert retained_codec is not None with pytest.raises(RuntimeError, match="LLM execution codec capability is no longer active"): retained_codec.decode(make_request()) From 6a88331c54610a086c1fa4f4a69af9049276f05f Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 23 Sep 2026 15:41:09 -0400 Subject: [PATCH 07/22] fix: align execution codec context with ABI v6 Signed-off-by: Alex Fournier --- crates/core/src/api/runtime.rs | 10 +- crates/core/src/api/runtime/callbacks.rs | 30 +++-- .../src/api/runtime/llm_execution_context.rs | 34 ++--- crates/core/src/plugin/dynamic.rs | 2 +- crates/core/src/plugin/dynamic/native.rs | 55 +++----- .../tests/fixtures/native_plugin/src/lib.rs | 31 +---- .../tests/integration/native_plugin_tests.rs | 19 ++- crates/core/tests/unit/native_plugin_tests.rs | 50 ++----- crates/ffi/nemo_relay.h | 14 +- crates/ffi/src/callable.rs | 10 +- crates/node/plugin.d.ts | 2 + crates/node/root-types.d.ts | 10 +- crates/plugin/README.md | 8 +- crates/plugin/src/async_sdk.rs | 74 +++++----- crates/plugin/src/lib.rs | 110 +++++++-------- crates/plugin/tests/typed_callbacks.rs | 126 ++++++++---------- crates/python/src/py_types/mod.rs | 8 ++ crates/worker/src/lib.rs | 32 +++-- docs/about-nemo-relay/release-notes/index.mdx | 2 +- docs/build-plugins/about.mdx | 4 +- docs/build-plugins/native/about.mdx | 10 +- .../native/native-abi-reference.mdx | 27 ++-- .../package-discoverable-plugins.mdx | 2 +- docs/reference/migration-guides.mdx | 2 +- examples/rust-native-plugin/README.md | 4 +- go/nemo_relay/callbacks.go | 12 +- python/nemo_relay/__init__.py | 4 + python/nemo_relay/__init__.pyi | 6 + python/nemo_relay/_native.pyi | 7 +- .../plugin/src/nemo_relay_plugin/__init__.py | 6 + python/plugin/src/nemo_relay_plugin/_api.py | 17 ++- 31 files changed, 362 insertions(+), 366 deletions(-) diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index a86713a57..9f29d4315 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -14,11 +14,11 @@ pub mod subscriber_dispatcher; pub use callbacks::{ BuiltinLlmCodec, ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmCollectorFn, LlmConditionalFn, LlmExecutionFn, - LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmRequestInterceptFn, - LlmSanitizeRequestContext, LlmSanitizeRequestFn, LlmSanitizeResponseContext, - LlmSanitizeResponseFn, LlmStreamExecutionFn, LlmStreamExecutionNextFn, LlmStreamInner, - ToolConditionalFn, ToolExecutionContext, ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, - ToolSanitizeFn, + LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmRequestCodecContext, + LlmRequestInterceptFn, LlmResponseCodecContext, LlmSanitizeRequestContext, + LlmSanitizeRequestFn, LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, + LlmStreamExecutionNextFn, LlmStreamInner, ToolConditionalFn, ToolExecutionContext, + ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, }; #[doc(hidden)] pub use continuation_context::MiddlewareContinuationContext; diff --git a/crates/core/src/api/runtime/callbacks.rs b/crates/core/src/api/runtime/callbacks.rs index 8248fdf3e..1044e7410 100644 --- a/crates/core/src/api/runtime/callbacks.rs +++ b/crates/core/src/api/runtime/callbacks.rs @@ -247,27 +247,27 @@ pub(crate) type ToolExecutionOutcomeNextFn = Arc< + Sync, >; -/// Per-call codec context for LLM request sanitize guardrails. +/// Per-call codec identity and capability for an LLM request. /// /// The context distinguishes no codec, Relay built-ins, runtime-registered /// codecs, and active codecs with no stable identity. #[derive(Clone, Default)] -pub struct LlmSanitizeRequestContext { +pub struct LlmRequestCodecContext { /// Identity of the codec active for this payload direction. codec: LlmCodecIdentity, request_codec: Option>, } -impl std::fmt::Debug for LlmSanitizeRequestContext { +impl std::fmt::Debug for LlmRequestCodecContext { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter - .debug_struct("LlmSanitizeRequestContext") + .debug_struct("LlmRequestCodecContext") .field("codec", &self.codec) .finish_non_exhaustive() } } -impl LlmSanitizeRequestContext { +impl LlmRequestCodecContext { /// Construct a context that carries only a codec identity. /// /// Identity-only contexts do not carry a codec handle, so @@ -281,7 +281,7 @@ impl LlmSanitizeRequestContext { } } - /// Construct request-sanitizer context from the active request codec. + /// Construct request codec context from the active request codec. #[must_use] pub fn for_request_codec(codec: Option>) -> Self { let identity = codec @@ -308,27 +308,27 @@ impl LlmSanitizeRequestContext { } } -/// Per-call codec context for LLM response sanitize guardrails. +/// Per-call codec identity and capability for an LLM response. /// /// The context distinguishes no codec, Relay built-ins, runtime-registered /// codecs, and active codecs with no stable identity. #[derive(Clone, Default)] -pub struct LlmSanitizeResponseContext { +pub struct LlmResponseCodecContext { /// Identity of the codec active for this payload direction. codec: LlmCodecIdentity, response_codec: Option>, } -impl std::fmt::Debug for LlmSanitizeResponseContext { +impl std::fmt::Debug for LlmResponseCodecContext { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter - .debug_struct("LlmSanitizeResponseContext") + .debug_struct("LlmResponseCodecContext") .field("codec", &self.codec) .finish_non_exhaustive() } } -impl LlmSanitizeResponseContext { +impl LlmResponseCodecContext { /// Construct a context that carries only a codec identity. /// /// Identity-only contexts do not carry a codec handle, so @@ -342,7 +342,7 @@ impl LlmSanitizeResponseContext { } } - /// Construct response-sanitizer context from the active response codec. + /// Construct response codec context from the active response codec. #[must_use] pub fn for_response_codec(codec: Option>) -> Self { let identity = codec @@ -369,6 +369,12 @@ impl LlmSanitizeResponseContext { } } +/// Backward-compatible name for request codec context supplied to sanitizers. +pub type LlmSanitizeRequestContext = LlmRequestCodecContext; + +/// Backward-compatible name for response codec context supplied to sanitizers. +pub type LlmSanitizeResponseContext = LlmResponseCodecContext; + /// Sanitize an LLM request before the runtime records it. /// /// LLM request sanitizers affect the serialized request payload emitted on diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index e7c52235a..06a4b6b22 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -6,7 +6,7 @@ use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; -use super::callbacks::{LlmSanitizeRequestContext, LlmSanitizeResponseContext}; +use super::callbacks::{LlmRequestCodecContext, LlmResponseCodecContext}; use crate::api::llm::LlmRequest; use crate::codec::request::AnnotatedLlmRequest; use crate::codec::response::AnnotatedLlmResponse; @@ -110,16 +110,16 @@ impl LlmResponseCodec for RevocableResponseCodec { /// after expiry. #[derive(Clone, Debug, Default)] pub struct LlmExecutionContext { - request_codec: LlmSanitizeRequestContext, - response_codec: Option, + request_codec: LlmRequestCodecContext, + response_codec: Option, } impl LlmExecutionContext { /// Construct an execution context from its directional codec contexts. #[must_use] pub fn new( - request_codec: LlmSanitizeRequestContext, - response_codec: Option, + request_codec: LlmRequestCodecContext, + response_codec: Option, ) -> Self { Self { request_codec, @@ -133,8 +133,8 @@ impl LlmExecutionContext { response_codec: &Option>, ) -> Self { Self::new( - LlmSanitizeRequestContext::for_request_codec(request_codec), - Some(LlmSanitizeResponseContext::for_response_codec( + LlmRequestCodecContext::for_request_codec(request_codec), + Some(LlmResponseCodecContext::for_response_codec( response_codec.clone(), )), ) @@ -143,7 +143,7 @@ impl LlmExecutionContext { /// Construct the context for a streaming managed execution. pub(crate) fn for_streaming_codec(request_codec: Option>) -> Self { Self::new( - LlmSanitizeRequestContext::for_request_codec(request_codec), + LlmRequestCodecContext::for_request_codec(request_codec), None, ) } @@ -156,25 +156,25 @@ impl LlmExecutionContext { pub(crate) fn lease(&self) -> (Self, LlmExecutionCodecLeaseGuard) { let gate = Arc::new(ExecutionCodecGate::new()); let request_codec = match self.request_codec.resolve_codec() { - Some(codec) => LlmSanitizeRequestContext::for_request_codec(Some(Arc::new( - RevocableRequestCodec { + Some(codec) => { + LlmRequestCodecContext::for_request_codec(Some(Arc::new(RevocableRequestCodec { codec, gate: Arc::clone(&gate), - }, - ))), - None => LlmSanitizeRequestContext::with_identity(self.request_codec.codec().clone()), + }))) + } + None => LlmRequestCodecContext::with_identity(self.request_codec.codec().clone()), }; let response_codec = self.response_codec .as_ref() .map(|context| match context.resolve_codec() { - Some(codec) => LlmSanitizeResponseContext::for_response_codec(Some(Arc::new( + Some(codec) => LlmResponseCodecContext::for_response_codec(Some(Arc::new( RevocableResponseCodec { codec, gate: Arc::clone(&gate), }, ))), - None => LlmSanitizeResponseContext::with_identity(context.codec().clone()), + None => LlmResponseCodecContext::with_identity(context.codec().clone()), }); ( @@ -185,7 +185,7 @@ impl LlmExecutionContext { /// Return the request-direction codec identity and revocable capability. #[must_use] - pub fn request_codec(&self) -> &LlmSanitizeRequestContext { + pub fn request_codec(&self) -> &LlmRequestCodecContext { &self.request_codec } @@ -194,7 +194,7 @@ impl LlmExecutionContext { /// Streaming execution returns `None` because Relay does not expose a /// completed-response codec for individual stream chunks. #[must_use] - pub fn response_codec(&self) -> Option<&LlmSanitizeResponseContext> { + pub fn response_codec(&self) -> Option<&LlmResponseCodecContext> { self.response_codec.as_ref() } } diff --git a/crates/core/src/plugin/dynamic.rs b/crates/core/src/plugin/dynamic.rs index a6b496073..dadc4a9e8 100644 --- a/crates/core/src/plugin/dynamic.rs +++ b/crates/core/src/plugin/dynamic.rs @@ -169,7 +169,7 @@ pub(super) fn validate_native_abi_compatibility( })?; if version_requirement_matches_minor(&requirement, 0, 9) { return Err(PluginError::InvalidConfig(format!( - "dynamic native plugin '{plugin_kind}' uses native ABI v7 and must declare compat.relay = \">=0.10,<1.0\" or another range that excludes Relay 0.9" + "dynamic native plugin '{plugin_kind}' uses native ABI v6 and must declare compat.relay = \">=0.10,<1.0\" or another range that excludes Relay 0.9" ))); } Ok(()) diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 1fe534178..92663d2f2 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -57,7 +57,7 @@ use chrono::{DateTime, Utc}; use libloading::{Library, Symbol}; use nemo_relay_plugin::{ NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, - NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT, NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, + NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT, NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncCompletion, NemoRelayNativeAsyncLlmExecutionCb, @@ -67,20 +67,20 @@ use nemo_relay_plugin::{ NemoRelayNativeAsyncStreamMiddlewareCb, NemoRelayNativeConditionalMiddlewareCb, NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, - NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, - NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, - NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, - NemoRelayNativeLlmExecutionRequestContext, NemoRelayNativeLlmExecutionResponseContext, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, - NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, - NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, - NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, - NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, NemoRelayNativePluginContext, - NemoRelayNativePluginEntry, NemoRelayNativePluginRuntime, NemoRelayNativePluginV1, - NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, - NemoRelayNativeScopeType, NemoRelayNativeString, NemoRelayNativeToolConditionalCb, - NemoRelayNativeToolExecutionCb, NemoRelayNativeToolExecutionContextCb, - NemoRelayNativeToolJsonCb, NemoRelayNativeWithScopeStackCb, NemoRelayStatus, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeLlmAsyncStream, + NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, + NemoRelayNativeLlmExecutionContext, NemoRelayNativeLlmExecutionRequestContext, + NemoRelayNativeLlmExecutionResponseContext, NemoRelayNativeLlmRequestCodec, + NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, + NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, + NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, + NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, + NemoRelayNativePluginContext, NemoRelayNativePluginEntry, NemoRelayNativePluginRuntime, + NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, + NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, NemoRelayNativeString, + NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, + NemoRelayNativeToolExecutionContextCb, NemoRelayNativeToolJsonCb, + NemoRelayNativeWithScopeStackCb, NemoRelayStatus, }; use serde_json::{Map, Value as Json}; use sha2::{Digest, Sha256}; @@ -947,18 +947,6 @@ 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_v7) - .v6 - .v5 - .v4 - .v3 - .v1 as *const NemoRelayNativeHostApiV1 -} - -#[cfg(test)] -fn native_host_api_v6() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); &HOST_API.get_or_init(build_native_host_api_v6).v5.v4.v3.v1 as *const NemoRelayNativeHostApiV1 } @@ -1120,20 +1108,11 @@ fn build_native_host_api_v5() -> NemoRelayNativeHostApiV5 { fn build_native_host_api_v6() -> NemoRelayNativeHostApiV6 { let mut v5 = build_native_host_api_v5(); - v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING; + v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT; v5.v4.v3.v1.struct_size = std::mem::size_of::(); NemoRelayNativeHostApiV6 { v5, log: native_log, - } -} - -fn build_native_host_api_v7() -> NemoRelayNativeHostApiV7 { - let mut v6 = build_native_host_api_v6(); - v6.v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT; - v6.v5.v4.v3.v1.struct_size = std::mem::size_of::(); - NemoRelayNativeHostApiV7 { - v6, plugin_context_register_async_llm_execution_intercept: native_plugin_context_register_async_llm_execution_intercept, async_stream_retain: native_async_stream_retain, @@ -4208,7 +4187,7 @@ unsafe extern "C" fn native_plugin_context_register_async_middleware( | NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept ) { set_native_last_error( - "LLM execution middleware requires its dedicated ABI-v7 registration function", + "LLM execution middleware requires its dedicated ABI-v6 registration function", ); return NemoRelayStatus::InvalidArg; } diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index fd1385191..f8687746a 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -12,12 +12,12 @@ use nemo_relay_plugin::{ Json, LlmJsonAsyncStream, LlmRequest, LlmRequestInterceptOutcome, MetricKind, MetricMeasurement, MetricValueType, NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, - NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, + NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, NativeExecutorConfig, NativePlugin, NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncMiddlewareCb, NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, NemoRelayNativeAsyncStream, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, - NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeLlmExecutionContext, NemoRelayNativePluginContext, NemoRelayNativePluginV1, NemoRelayNativeString, NemoRelayNativeToolNextFn, NemoRelayStatus, PendingMarkSpec, PluginContext, PluginRuntime, RuntimeRegistrationKind, ScopeCategory, ScopeType, @@ -547,23 +547,6 @@ pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v5( } } -/// Raw ABI-v6 entry used to verify that the immediately stale callback layout is rejected. -#[unsafe(no_mangle)] -pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v6( - host: *const NemoRelayNativeHostApiV1, - out: *mut NemoRelayNativePluginV1, -) -> NemoRelayStatus { - unsafe { - fixture_compat_entry( - host, - out, - NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, - std::mem::size_of::(), - b"fixture_native_v6", - ) - } -} - /// Raw ABI-v4 entry used to verify fallback for plugins built with the previous SDK. #[unsafe(no_mangle)] pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v4( @@ -962,7 +945,7 @@ unsafe extern "C" fn raw_register_event_sanitize_errors( } struct FixtureAsyncPlugin { - host: Option>, + host: Option>, } impl NativePlugin for FixtureAsyncPlugin { @@ -977,17 +960,17 @@ impl NativePlugin for FixtureAsyncPlugin { ) -> nemo_relay_plugin::Result<()> { let host = ctx.host_api(); if host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION - || host.struct_size < std::mem::size_of::() + || host.struct_size < std::mem::size_of::() { - return Err("fixture async plugin requires ABI v7".into()); + return Err("fixture async plugin requires ABI v6".into()); } self.host = Some(Box::new(unsafe { - *(host as *const _ as *const NemoRelayNativeHostApiV7) + *(host as *const _ as *const NemoRelayNativeHostApiV6) })); let user_data = self .host .as_deref() - .map(|host| (host as *const NemoRelayNativeHostApiV7).cast_mut().cast()) + .map(|host| (host as *const NemoRelayNativeHostApiV6).cast_mut().cast()) .expect("fixture async host was initialized"); let registrations: [( diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index 84c6ffffc..9c5511c04 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -974,7 +974,7 @@ fn native_api_one_does_not_admit_a_stale_abi_v2_binary() { [load_spec("fixture_native", &manifest_ref)], "a native_api=1 manifest must not make an ABI-v2 binary compatible", ); - assert!(error.contains("rejected native ABI 7"), "{error}"); + assert!(error.contains("rejected native ABI 6"), "{error}"); assert!(error.contains("rebuild the plugin"), "{error}"); } @@ -1099,7 +1099,7 @@ async fn native_loader_rejects_manifest_that_admits_pre_zero_eight_relay() { } #[tokio::test] -async fn native_abi_v7_rejects_manifest_that_admits_relay_zero_nine() { +async fn native_abi_v6_rejects_manifest_that_admits_relay_zero_nine() { let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; let fixture = build_fixture_plugin(); let manifest_ref = write_manifest_text(ManifestOptions { @@ -1112,10 +1112,10 @@ async fn native_abi_v7_rejects_manifest_that_admits_relay_zero_nine() { }); let error = expect_native_load_error_from_specs( [load_spec("fixture_native", &manifest_ref)], - "an ABI-v7 native plugin must exclude Relay 0.9", + "an ABI-v6 native plugin must exclude Relay 0.9", ); assert!( - error.contains("uses native ABI v7") && error.contains("excludes Relay 0.9"), + error.contains("uses native ABI v6") && error.contains("excludes Relay 0.9"), "{error}" ); } @@ -1164,19 +1164,18 @@ fn native_loader_rejects_abi_v3_plugins() { let error = expect_native_load_error_from_specs( [load_spec("fixture_native_v3", &manifest_ref)], - "ABI-v3 plugins must be rebuilt for ABI v7", + "ABI-v3 plugins must be rebuilt for ABI v6", ); - assert!(error.contains("rejected native ABI 7"), "{error}"); + assert!(error.contains("rejected native ABI 6"), "{error}"); assert!(error.contains("rebuild the plugin"), "{error}"); } #[test] -fn native_loader_rejects_v2_v4_v5_and_v6_plugins() { +fn native_loader_rejects_v2_v4_and_v5_plugins() { let _guard = NATIVE_PLUGIN_TEST_LOCK.blocking_lock(); let fixture = build_fixture_plugin(); for (plugin_id, symbol) in [ - ("fixture_native_v6", "nemo_relay_fixture_native_plugin_v6"), ("fixture_native_v5", "nemo_relay_fixture_native_plugin_v5"), ("fixture_native_v4", "nemo_relay_fixture_native_plugin_v4"), ("fixture_native_v2", "nemo_relay_fixture_native_plugin_v2"), @@ -1192,9 +1191,9 @@ fn native_loader_rejects_v2_v4_v5_and_v6_plugins() { let error = expect_native_load_error_from_specs( [load_spec(plugin_id, &manifest_ref)], - "stale native plugins must be rebuilt for ABI v7", + "stale native plugins must be rebuilt for ABI v6", ); - assert!(error.contains("rejected native ABI 7"), "{error}"); + assert!(error.contains("rejected native ABI 6"), "{error}"); assert!(error.contains("rebuild the plugin"), "{error}"); } } diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index 3aac06e8d..c032e449f 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -634,7 +634,6 @@ fn assert_native_digest_edges() { fn assert_native_host_api_versions() { let current = native_host_api(); - let frozen_v6 = native_host_api_v6(); let frozen_v5 = native_host_api_v5(); let frozen_v4 = native_host_api_v4(); let frozen_v3 = native_host_api_v3(); @@ -642,11 +641,6 @@ fn assert_native_host_api_versions() { assert_native_host_api_descriptor( current, NEMO_RELAY_NATIVE_ABI_VERSION, - std::mem::size_of::(), - ); - assert_native_host_api_descriptor( - frozen_v6, - NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, std::mem::size_of::(), ); assert_native_host_api_descriptor( @@ -669,7 +663,6 @@ fn assert_native_host_api_versions() { NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, std::mem::size_of::(), ); - assert_native_host_api_v7_layout(); assert_native_host_api_v6_layout(); assert_native_host_api_v5_layout(); assert_native_host_api_v4_layout(); @@ -791,46 +784,30 @@ fn assert_native_host_api_v6_layout() { #[cfg(target_pointer_width = "64")] { assert_eq!(std::mem::align_of::(), 8); - assert_eq!(std::mem::size_of::(), 616); + assert_eq!(std::mem::size_of::(), 648); assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, log), 608); - } - #[cfg(target_pointer_width = "32")] - { - assert_eq!(std::mem::align_of::(), 4); - assert_eq!(std::mem::size_of::(), 304); - assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, v5), 0); - assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, log), 300); - } -} - -fn assert_native_host_api_v7_layout() { - #[cfg(target_pointer_width = "64")] - { - assert_eq!(std::mem::align_of::(), 8); - assert_eq!(std::mem::size_of::(), 648); - assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV7, v6), 0); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, plugin_context_register_async_llm_execution_intercept ), 616 ); assert_eq!( - std::mem::offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), + std::mem::offset_of!(NemoRelayNativeHostApiV6, async_stream_retain), 624 ); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, async_stream_llm_request_codec_decode ), 632 ); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, async_stream_llm_request_codec_encode ), 640 @@ -838,30 +815,31 @@ fn assert_native_host_api_v7_layout() { } #[cfg(target_pointer_width = "32")] { - assert_eq!(std::mem::align_of::(), 4); - assert_eq!(std::mem::size_of::(), 320); - assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV7, v6), 0); + assert_eq!(std::mem::align_of::(), 4); + assert_eq!(std::mem::size_of::(), 320); + assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, v5), 0); + assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, log), 300); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, plugin_context_register_async_llm_execution_intercept ), 304 ); assert_eq!( - std::mem::offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), + std::mem::offset_of!(NemoRelayNativeHostApiV6, async_stream_retain), 308 ); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, async_stream_llm_request_codec_decode ), 312 ); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, async_stream_llm_request_codec_encode ), 316 @@ -1886,7 +1864,7 @@ fn assert_native_json_output_and_host_api() { 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::() ); } diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index dedbf26fc..6cf20a83a 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -381,6 +381,16 @@ typedef NemoRelayStatus (*NemoRelayLlmRequestInterceptCb)(void *user_data, const char *annotated_json, char **out_outcome_json); +/** + * General name for request codec context used by execution intercepts. + */ +typedef struct NemoRelayLlmSanitizeRequestContext NemoRelayLlmRequestCodecContext; + +/** + * General name for response codec context used by execution intercepts. + */ +typedef struct NemoRelayLlmSanitizeResponseContext NemoRelayLlmResponseCodecContext; + /** * Directional codec context supplied to an LLM execution intercept. * @@ -393,11 +403,11 @@ typedef struct NemoRelayLlmExecutionContext { /** * Active request codec identity and capability. */ - struct NemoRelayLlmSanitizeRequestContext request_codec; + NemoRelayLlmRequestCodecContext request_codec; /** * Active unary-response codec context, or null for streaming execution. */ - const struct NemoRelayLlmSanitizeResponseContext *response_codec; + const NemoRelayLlmResponseCodecContext *response_codec; } NemoRelayLlmExecutionContext; /** diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 20daa8135..7c32c129c 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -202,6 +202,12 @@ pub struct NemoRelayLlmSanitizeResponseContext { pub codec: *const crate::types::FfiLlmSanitizeResponseCodec, } +/// General name for request codec context used by execution intercepts. +pub type NemoRelayLlmRequestCodecContext = NemoRelayLlmSanitizeRequestContext; + +/// General name for response codec context used by execution intercepts. +pub type NemoRelayLlmResponseCodecContext = NemoRelayLlmSanitizeResponseContext; + /// Directional codec context supplied to an LLM execution intercept. /// /// `request_codec` is always present. `response_codec` is non-null for unary @@ -212,9 +218,9 @@ pub struct NemoRelayLlmSanitizeResponseContext { #[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmExecutionContext { /// Active request codec identity and capability. - pub request_codec: NemoRelayLlmSanitizeRequestContext, + pub request_codec: NemoRelayLlmRequestCodecContext, /// Active unary-response codec context, or null for streaming execution. - pub response_codec: *const NemoRelayLlmSanitizeResponseContext, + pub response_codec: *const NemoRelayLlmResponseCodecContext, } /// LLM request sanitizer. It receives the request first and its codec context diff --git a/crates/node/plugin.d.ts b/crates/node/plugin.d.ts index 86d0bb26f..f4a52648b 100644 --- a/crates/node/plugin.d.ts +++ b/crates/node/plugin.d.ts @@ -29,7 +29,9 @@ export type { LlmOptimizationModelTransition, LlmOptimizationTokenImpact, LlmOptimizationTokens, + LlmRequestCodecContext, LlmRequestInterceptOutcome, + LlmResponseCodecContext, LlmSanitizeRequestContext, LlmSanitizeResponseContext, } from './index'; diff --git a/crates/node/root-types.d.ts b/crates/node/root-types.d.ts index ef4f88edc..e635d0d04 100644 --- a/crates/node/root-types.d.ts +++ b/crates/node/root-types.d.ts @@ -25,12 +25,18 @@ export interface LlmSanitizeResponseContext { resolveCodec(): import('./typed').LlmResponseCodec | null; } +/** General name for request codec context used outside sanitizer callbacks. */ +export type LlmRequestCodecContext = LlmSanitizeRequestContext; + +/** General name for response codec context used outside sanitizer callbacks. */ +export type LlmResponseCodecContext = LlmSanitizeResponseContext; + /** Codec capabilities for one managed LLM execution intercept invocation. */ export interface LlmExecutionContext { /** Request codec identity plus optional decode and encode capability. */ - requestCodec: LlmSanitizeRequestContext; + requestCodec: LlmRequestCodecContext; /** Unary response codec identity plus optional decode capability; `null` for streaming execution. */ - responseCodec: LlmSanitizeResponseContext | null; + responseCodec: LlmResponseCodecContext | null; } /** Schema tag attached to an opaque optimization contribution payload. */ diff --git a/crates/plugin/README.md b/crates/plugin/README.md index 81da76a94..db6ca1d49 100644 --- a/crates/plugin/README.md +++ b/crates/plugin/README.md @@ -32,7 +32,7 @@ the dynamic-library boundary on the stable C-compatible ABI. | `PluginContext` | Installs component-owned subscribers, guardrails, intercepts, continuations, and streams. | | `PluginRuntime` | Emits marks and manages Relay-owned scopes and scope stacks through typed host helpers. | | `nemo_relay_plugin!` | Exports the one versioned native entry point used by the loader. | -| Native ABI v7 | Keeps C-compatible host and plugin tables behind the safe Rust interface. ABI v7 adds directional codec context to LLM execution callbacks and intentionally rejects plugins compiled with an older callback layout. | +| Native ABI v6 | Keeps C-compatible host and plugin tables behind the safe Rust interface. Relay 0.10 finalizes ABI v6 with directional codec context for LLM execution callbacks and requires native plugins to rebuild against that layout. | | Typed async middleware | Drives guardrails, sanitizers, and intercepts on a per-component SDK-owned Tokio executor. Subscribers and raw ABI registrations remain synchronous. | | Async continuations and streams | `ToolNext`, `LlmNext`, and `LlmStreamNext` support repeated or concurrent downstream calls. Streaming LLM continuations use a pull-based host handle. | | Tool results | `ToolNext` returns `ToolExecutionResult`, which keeps an application result and optional annotation together. | @@ -104,10 +104,10 @@ context-aware tool execution intercept must rebuild and set table. Under Relay 0.9, typed async plugins that did not use this registration could retain `compat.relay = ">=0.8.0,<1.0"`. -Relay 0.10 advances the internal table to ABI v7 and makes +Relay 0.10 finalizes the internal ABI v6 table and makes `LlmExecutionContext` part of every unary and streaming LLM execution callback. -Because this changes callback layouts, the 0.10 host rejects every native plugin -compiled against an older table. Rebuild the plugin with the 0.10 SDK and set +Because this changes callback layouts, the 0.10 host rejects v2-v5 tables. +Rebuild every plugin with the finalized 0.10 v6 SDK and set `compat.relay = ">=0.10.0,<1.0"`. The authored manifest contract remains `compat.native_api = "1"`. diff --git a/crates/plugin/src/async_sdk.rs b/crates/plugin/src/async_sdk.rs index ec416c35f..13a415ef2 100644 --- a/crates/plugin/src/async_sdk.rs +++ b/crates/plugin/src/async_sdk.rs @@ -154,10 +154,10 @@ unsafe impl Send for HostV4 {} unsafe impl Sync for HostV4 {} #[derive(Clone, Copy)] -struct HostV7(NemoRelayNativeHostApiV7); +struct HostV6(NemoRelayNativeHostApiV6); -unsafe impl Send for HostV7 {} -unsafe impl Sync for HostV7 {} +unsafe impl Send for HostV6 {} +unsafe impl Sync for HostV6 {} struct Completion { host: HostV4, @@ -246,7 +246,7 @@ impl CompletionRef { self, codec: LlmCodecIdentity, resolved: bool, - ) -> Result> { + ) -> Result> { let resolved = if resolved { let status = unsafe { (self.host.0.async_completion_retain)(self.raw) }; status_result(status, "retain native async completion capability")?; @@ -260,14 +260,14 @@ impl CompletionRef { } else { None }; - Ok(LlmExecutionRequestContext { codec, resolved }) + Ok(LlmRequestCodecContext { codec, resolved }) } fn execution_response_context( self, codec: LlmCodecIdentity, resolved: bool, - ) -> Result> { + ) -> Result> { let resolved = if resolved { let status = unsafe { (self.host.0.async_completion_retain)(self.raw) }; status_result(status, "retain native async completion capability")?; @@ -279,13 +279,13 @@ impl CompletionRef { } else { None }; - Ok(LlmExecutionResponseContext { codec, resolved }) + Ok(LlmResponseCodecContext { codec, resolved }) } } #[derive(Clone, Copy)] struct StreamRef { - host: HostV7, + host: HostV6, raw: *const NemoRelayNativeAsyncStream, } @@ -296,7 +296,7 @@ impl StreamRef { self, codec: LlmCodecIdentity, resolved: bool, - ) -> Result> { + ) -> Result> { let resolved = if resolved { let status = unsafe { (self.host.0.async_stream_retain)(self.raw) }; status_result(status, "retain native async stream capability")?; @@ -310,7 +310,7 @@ impl StreamRef { } else { None }; - Ok(LlmExecutionRequestContext { codec, resolved }) + Ok(LlmRequestCodecContext { codec, resolved }) } } @@ -960,7 +960,7 @@ type StreamAdapter = dyn Fn(Json, LlmExecutionContext<'static>, LlmStreamNext) -> StreamFuture + Send + Sync; struct StreamCallbackState { - host: HostV7, + host: HostV6, executor: Arc, adapter: Box, } @@ -972,7 +972,7 @@ unsafe extern "C" fn drop_stream_callback(user_data: *mut c_void) { } struct OutputStream { - host: HostV7, + host: HostV6, raw: *const NemoRelayNativeAsyncStream, } @@ -981,19 +981,18 @@ unsafe impl Sync for OutputStream {} impl OutputStream { fn cancelled(&self) -> bool { - unsafe { (self.host.0.v6.v5.v4.v3.async_stream_is_cancelled)(self.raw) } + unsafe { (self.host.0.v5.v4.v3.async_stream_is_cancelled)(self.raw) } } async fn push(&self, value: &Json) -> Result<()> { - let value = HostString::from_json(&self.host.0.v6.v5.v4.v3.v1, value) + let value = HostString::from_json(&self.host.0.v5.v4.v3.v1, value) .ok_or_else(|| "failed to serialize native stream chunk".to_string())?; loop { if self.cancelled() { return Err("native stream consumer cancelled".into()); } - let status = unsafe { - (self.host.0.v6.v5.v4.v3.async_stream_push_json)(self.raw, value.as_ptr()) - }; + let status = + unsafe { (self.host.0.v5.v4.v3.async_stream_push_json)(self.raw, value.as_ptr()) }; match status { NemoRelayStatus::Ok => return Ok(()), NemoRelayStatus::Backpressured => { @@ -1006,20 +1005,19 @@ impl OutputStream { fn finish(&self) -> Result<()> { status_result( - unsafe { (self.host.0.v6.v5.v4.v3.async_stream_finish)(self.raw) }, + unsafe { (self.host.0.v5.v4.v3.async_stream_finish)(self.raw) }, "finish native stream", ) } async fn reject(&self, error: &str) { - if let Some(error) = HostString::new(&self.host.0.v6.v5.v4.v3.v1, error) { + if let Some(error) = HostString::new(&self.host.0.v5.v4.v3.v1, error) { loop { if self.cancelled() { break; } - let status = unsafe { - (self.host.0.v6.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()) - }; + let status = + unsafe { (self.host.0.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()) }; match status { NemoRelayStatus::Backpressured => { tokio::time::sleep(CANCELLATION_POLL_INTERVAL).await; @@ -1031,9 +1029,9 @@ impl OutputStream { } fn reject_once(&self, error: &str) { - if let Some(error) = HostString::new(&self.host.0.v6.v5.v4.v3.v1, error) { + if let Some(error) = HostString::new(&self.host.0.v5.v4.v3.v1, error) { unsafe { - (self.host.0.v6.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()); + (self.host.0.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()); } } } @@ -1041,7 +1039,7 @@ impl OutputStream { impl Drop for OutputStream { fn drop(&mut self) { - unsafe { (self.host.0.v6.v5.v4.v3.async_stream_release)(self.raw) }; + unsafe { (self.host.0.v5.v4.v3.async_stream_release)(self.raw) }; } } @@ -1062,11 +1060,11 @@ unsafe extern "C" fn stream_trampoline( return NemoRelayNativeAsyncCallbackState::Pending as u32; } let next = LlmStreamNext(Arc::new(NextInner { - host: HostV4(state.host.0.v6.v5.v4), + host: HostV4(state.host.0.v5.v4), raw: next, })); let invocation = read_json_value( - &state.host.0.v6.v5.v4.v3.v1, + &state.host.0.v5.v4.v3.v1, invocation_json, "stream invocation", ) @@ -1082,8 +1080,8 @@ unsafe extern "C" fn stream_trampoline( context, ) }); - let bindings = ScopePollBinding::capture(state.host.0.v6.v5.v4.v3.v1).and_then(|future| { - ScopePollBinding::capture(state.host.0.v6.v5.v4.v3.v1).map(|stream| (future, stream)) + let bindings = ScopePollBinding::capture(state.host.0.v5.v4.v3.v1).and_then(|future| { + ScopePollBinding::capture(state.host.0.v5.v4.v3.v1).map(|stream| (future, stream)) }); let future = catch_unwind(AssertUnwindSafe(|| match (invocation, context) { (Ok(invocation), Ok(context)) => (state.adapter)(invocation, context, next), @@ -1094,7 +1092,7 @@ unsafe extern "C" fn stream_trampoline( }); if let Err(error) = state.executor.ensure_started() { output.reject_once(&error); - set_last_error(&state.host.0.v6.v5.v4.v3.v1, &error); + set_last_error(&state.host.0.v5.v4.v3.v1, &error); return NemoRelayNativeAsyncCallbackState::Pending as u32; } let task = async move { @@ -1156,7 +1154,7 @@ unsafe extern "C" fn stream_trampoline( } }; if let Err(error) = state.executor.spawn(task) { - set_last_error(&state.host.0.v6.v5.v4.v3.v1, &error); + set_last_error(&state.host.0.v5.v4.v3.v1, &error); } NemoRelayNativeAsyncCallbackState::Pending as u32 } @@ -1250,7 +1248,7 @@ fn llm_stream_execution_context_from_native( return Err("native LLM stream execution context exposed a response codec".into()); } let request = context.request_codec; - let host = &stream.host.0.v6.v5.v4.v3.v1; + let host = &stream.host.0.v5.v4.v3.v1; let request_codec = stream.execution_request_context( execution_codec_identity(host, request.codec_kind, request.codec_id)?, !request.codec.is_null(), @@ -1273,14 +1271,14 @@ impl PluginContext<'_> { })) } - fn host_v7(&self) -> Result { + fn host_v6(&self) -> Result { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { - return Err("typed LLM execution middleware requires Relay ABI v7".into()); + return Err("typed LLM execution middleware requires Relay ABI v6".into()); } - Ok(HostV7(unsafe { - *(self.host as *const _ as *const NemoRelayNativeHostApiV7) + Ok(HostV6(unsafe { + *(self.host as *const _ as *const NemoRelayNativeHostApiV6) })) } @@ -1802,7 +1800,7 @@ impl PluginContext<'_> { { let callback = Arc::new(callback); let state = Box::into_raw(Box::new(StreamCallbackState { - host: self.host_v7()?, + host: self.host_v6()?, executor: Arc::clone(&self.executor), adapter: Box::new(move |value, context, next| { let callback = Arc::clone(&callback); diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index c10f74983..f98d4f8cc 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -50,14 +50,15 @@ use serde_json::Map; /// Native plugin ABI version supported by this crate. /// -/// Version 7 makes LLM execution intercept callbacks context-aware. +/// Version 6 makes LLM execution intercept callbacks context-aware and adds +/// host-routed operational logging. /// /// This is an intentional callback-layout break. Native plugins must rebuild /// against this Relay release even though their authored `native_api` /// compatibility label remains `1`. -pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 7; +pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 6; /// ABI version that introduced uniform LLM execution codec context. -pub const NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT: u32 = 7; +pub const NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT: u32 = 6; /// ABI version that introduced host-routed operational logging. pub const NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING: u32 = 6; /// ABI version that introduced context-aware raw tool execution intercepts. @@ -96,19 +97,19 @@ unsafe impl Send for LlmSanitizeResponseContext<'_> {} /// response codec context, while streaming execution leaves it unavailable /// until Relay has a completed-response streaming codec contract. pub struct LlmExecutionContext<'a> { - request_codec: LlmExecutionRequestContext<'a>, - response_codec: Option>, + request_codec: LlmRequestCodecContext<'a>, + response_codec: Option>, } /// Request codec context for one LLM execution intercept invocation. -pub struct LlmExecutionRequestContext<'a> { +pub struct LlmRequestCodecContext<'a> { /// Identity of the active request codec. pub codec: LlmCodecIdentity, resolved: Option>, } /// Response codec context for one unary LLM execution intercept invocation. -pub struct LlmExecutionResponseContext<'a> { +pub struct LlmResponseCodecContext<'a> { /// Identity of the active response codec. pub codec: LlmCodecIdentity, resolved: Option>, @@ -117,17 +118,23 @@ pub struct LlmExecutionResponseContext<'a> { impl<'a> LlmExecutionContext<'a> { /// Return the active request codec context. #[must_use] - pub fn request_codec(&self) -> &LlmExecutionRequestContext<'a> { + pub fn request_codec(&self) -> &LlmRequestCodecContext<'a> { &self.request_codec } /// Return the unary response codec context, or `None` for streaming execution. #[must_use] - pub fn response_codec(&self) -> Option<&LlmExecutionResponseContext<'a>> { + pub fn response_codec(&self) -> Option<&LlmResponseCodecContext<'a>> { self.response_codec.as_ref() } } +/// Previous name for request codec context on typed execution intercepts. +pub type LlmExecutionRequestContext<'a> = LlmRequestCodecContext<'a>; + +/// Previous name for response codec context on typed execution intercepts. +pub type LlmExecutionResponseContext<'a> = LlmResponseCodecContext<'a>; + /// Status codes returned by stable native ABI functions. #[repr(i32)] #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -225,7 +232,7 @@ pub struct NemoRelayNativeLlmSanitizeResponseContext { /// Request codec context passed to a native LLM execution intercept. #[repr(C)] #[derive(Debug, Clone, Copy)] -pub struct NemoRelayNativeLlmExecutionRequestContext { +pub struct NemoRelayNativeLlmRequestCodecContext { /// Discriminator for the active request codec. pub codec_kind: NemoRelayNativeLlmCodecKind, /// Optional borrowed built-in or runtime codec identifier. @@ -237,7 +244,7 @@ pub struct NemoRelayNativeLlmExecutionRequestContext { /// Response codec context passed to a native unary LLM execution intercept. #[repr(C)] #[derive(Debug, Clone, Copy)] -pub struct NemoRelayNativeLlmExecutionResponseContext { +pub struct NemoRelayNativeLlmResponseCodecContext { /// Discriminator for the active response codec. pub codec_kind: NemoRelayNativeLlmCodecKind, /// Optional borrowed built-in or runtime codec identifier. @@ -246,14 +253,20 @@ pub struct NemoRelayNativeLlmExecutionResponseContext { pub codec: *const NemoRelayNativeLlmResponseCodec, } +/// Previous native name for request codec context on execution intercepts. +pub type NemoRelayNativeLlmExecutionRequestContext = NemoRelayNativeLlmRequestCodecContext; + +/// Previous native name for response codec context on execution intercepts. +pub type NemoRelayNativeLlmExecutionResponseContext = NemoRelayNativeLlmResponseCodecContext; + /// Codec context passed to native LLM execution intercept callbacks. #[repr(C)] #[derive(Debug, Clone, Copy)] pub struct NemoRelayNativeLlmExecutionContext { /// Request codec context, always present. - pub request_codec: NemoRelayNativeLlmExecutionRequestContext, + pub request_codec: NemoRelayNativeLlmRequestCodecContext, /// Unary response codec context, or null for streaming execution. - pub response_codec: *const NemoRelayNativeLlmExecutionResponseContext, + pub response_codec: *const NemoRelayNativeLlmResponseCodecContext, } /// Safe completion-backed request codec facade for typed native plugins. @@ -355,7 +368,7 @@ enum LlmExecutionRequestCodecOwner { completion: *const NemoRelayNativeAsyncCompletion, }, Stream { - host: NemoRelayNativeHostApiV7, + host: NemoRelayNativeHostApiV6, stream: *const NemoRelayNativeAsyncStream, }, } @@ -378,7 +391,7 @@ impl Drop for LlmExecutionRequestCodec<'_> { (host.v3.async_completion_release)(completion) }, LlmExecutionRequestCodecOwner::Stream { host, stream } => unsafe { - (host.v6.v5.v4.v3.async_stream_release)(stream) + (host.v5.v4.v3.async_stream_release)(stream) }, } } @@ -401,12 +414,12 @@ impl LlmExecutionRequestCodec<'_> { }) } LlmExecutionRequestCodecOwner::Stream { host, stream } => { - native_codec_call(&host.v6.v5.v4.v3.v1, |out| unsafe { - let request = HostString::from_json(&host.v6.v5.v4.v3.v1, request) + native_codec_call(&host.v5.v4.v3.v1, |out| unsafe { + let request = HostString::from_json(&host.v5.v4.v3.v1, request) .ok_or_else(|| "failed to serialize LLM request".to_string())?; let status = (host.async_stream_llm_request_codec_decode)(stream, request.as_ptr(), out); - codec_status(&host.v6.v5.v4.v3.v1, status) + codec_status(&host.v5.v4.v3.v1, status) }) } } @@ -435,10 +448,10 @@ impl LlmExecutionRequestCodec<'_> { }) } LlmExecutionRequestCodecOwner::Stream { host, stream } => { - native_codec_call(&host.v6.v5.v4.v3.v1, |out| unsafe { - let annotated = HostString::from_json(&host.v6.v5.v4.v3.v1, annotated) + native_codec_call(&host.v5.v4.v3.v1, |out| unsafe { + let annotated = HostString::from_json(&host.v5.v4.v3.v1, annotated) .ok_or_else(|| "failed to serialize annotated request".to_string())?; - let original = HostString::from_json(&host.v6.v5.v4.v3.v1, original) + let original = HostString::from_json(&host.v5.v4.v3.v1, original) .ok_or_else(|| "failed to serialize original request".to_string())?; let status = (host.async_stream_llm_request_codec_encode)( stream, @@ -446,7 +459,7 @@ impl LlmExecutionRequestCodec<'_> { original.as_ptr(), out, ); - codec_status(&host.v6.v5.v4.v3.v1, status) + codec_status(&host.v5.v4.v3.v1, status) }) } } @@ -487,7 +500,7 @@ impl LlmExecutionResponseCodec<'_> { } } -impl LlmExecutionRequestContext<'_> { +impl LlmRequestCodecContext<'_> { /// Resolve the active request codec capability. #[must_use] pub fn resolve_codec(&self) -> Option<&LlmExecutionRequestCodec<'_>> { @@ -495,7 +508,7 @@ impl LlmExecutionRequestContext<'_> { } } -impl LlmExecutionResponseContext<'_> { +impl LlmResponseCodecContext<'_> { /// Resolve the active response codec capability. #[must_use] pub fn resolve_codec(&self) -> Option<&LlmExecutionResponseCodec<'_>> { @@ -1607,7 +1620,8 @@ pub struct NemoRelayNativeHostApiV5 { -> NemoRelayStatus, } -/// ABI-v6 host extension for operational logging. +/// ABI-v6 host extension for operational logging and context-aware LLM +/// execution intercept callbacks. /// /// The complete ABI-v5 table is the prefix, preserving layout compatibility. #[repr(C)] @@ -1622,19 +1636,6 @@ pub struct NemoRelayNativeHostApiV6 { message: *const NemoRelayNativeString, fields_json: *const NemoRelayNativeString, ) -> NemoRelayStatus, -} - -/// ABI-v7 host table for context-aware LLM execution intercept callbacks. -/// -/// The inherited function-pointer table has the same fields as ABI v6, but -/// its LLM execution callback typedefs include -/// [`NemoRelayNativeLlmExecutionContext`]. The distinct table version prevents -/// either side from invoking a callback compiled with the old argument layout. -#[repr(C)] -#[derive(Clone, Copy)] -pub struct NemoRelayNativeHostApiV7 { - /// ABI-v6 table compiled with the ABI-v7 callback typedefs. - pub v6: NemoRelayNativeHostApiV6, /// Registers a completion-based asynchronous LLM execution intercept. pub plugin_context_register_async_llm_execution_intercept: unsafe extern "C" fn( @@ -1677,12 +1678,10 @@ unsafe impl Sync for NemoRelayNativeHostApiV4 {} // same thread-safe host function table. unsafe impl Send for NemoRelayNativeHostApiV5 {} unsafe impl Sync for NemoRelayNativeHostApiV5 {} -// SAFETY: the v6 host table is immutable and its log function is thread-safe. +// SAFETY: the v6 host table is immutable and contains only thread-safe host +// functions. unsafe impl Send for NemoRelayNativeHostApiV6 {} unsafe impl Sync for NemoRelayNativeHostApiV6 {} -// SAFETY: the v7 table is immutable and contains only thread-safe host functions. -unsafe impl Send for NemoRelayNativeHostApiV7 {} -unsafe impl Sync for NemoRelayNativeHostApiV7 {} // The host API table is immutable after construction. Function pointers and // the null-terminated version string pointer are safe to share across threads. @@ -3178,7 +3177,7 @@ impl<'a> PluginContext<'a> { /// `cb`, `user_data`, and `free_fn` must remain valid for every host /// callback invocation until the host deregisters the callback or calls /// `free_fn`. `free_fn` must match the allocation behind `user_data`. - /// Registration consumes that ownership even when ABI v7 is unavailable. + /// Registration consumes that ownership even when ABI v6 is unavailable. pub unsafe fn register_llm_execution_intercept_raw( &mut self, name: &str, @@ -3188,7 +3187,7 @@ impl<'a> PluginContext<'a> { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; @@ -3208,7 +3207,7 @@ impl<'a> PluginContext<'a> { /// `cb`, `user_data`, and `free_fn` must remain valid for every host /// callback invocation until the host deregisters the callback or calls /// `free_fn`. `free_fn` must match the allocation behind `user_data`. - /// Registration consumes that ownership even when ABI v7 is unavailable. + /// Registration consumes that ownership even when ABI v6 is unavailable. pub unsafe fn register_llm_stream_execution_intercept_raw( &mut self, name: &str, @@ -3218,7 +3217,7 @@ impl<'a> PluginContext<'a> { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; @@ -3306,14 +3305,14 @@ impl<'a> PluginContext<'a> { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } return NemoRelayStatus::InvalidArg; } - let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV7) }; + let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV6) }; self.with_name_and_callback(name, user_data, free_fn, |_, name| unsafe { (host.plugin_context_register_async_llm_execution_intercept)( self.raw, name, priority, cb, user_data, free_fn, @@ -3332,7 +3331,7 @@ impl<'a> PluginContext<'a> { /// stream owns the callback lifetime. `next` may be invoked /// repeatedly or concurrently until that stream settles; Relay then /// rejects or cancels unfinished and later calls. This execution-specific - /// callback requires native ABI v7 because its callback context and + /// callback requires native ABI v6 because its callback context and /// stream-scoped request-codec operations are part of that ABI. pub unsafe fn register_async_stream_middleware_raw( &mut self, @@ -3343,17 +3342,16 @@ impl<'a> PluginContext<'a> { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } return NemoRelayStatus::InvalidArg; } - let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV7) }; + let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV6) }; self.with_name_and_callback(name, user_data, free_fn, |_, name| unsafe { (host - .v6 .v5 .v4 .v3 @@ -3599,16 +3597,11 @@ enum OwnedHostApi { V4(NemoRelayNativeHostApiV4), V5(NemoRelayNativeHostApiV5), V6(NemoRelayNativeHostApiV6), - V7(NemoRelayNativeHostApiV7), } impl OwnedHostApi { unsafe fn copy_from(host: &NemoRelayNativeHostApiV1) -> Self { - if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - && host.struct_size >= std::mem::size_of::() - { - Self::V7(unsafe { *(host as *const _ as *const NemoRelayNativeHostApiV7) }) - } else if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING + if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING && host.struct_size >= std::mem::size_of::() { Self::V6(unsafe { *(host as *const _ as *const NemoRelayNativeHostApiV6) }) @@ -3636,7 +3629,6 @@ impl OwnedHostApi { Self::V4(host) => &host.v3.v1, Self::V5(host) => &host.v4.v3.v1, Self::V6(host) => &host.v5.v4.v3.v1, - Self::V7(host) => &host.v6.v5.v4.v3.v1, } } } diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index 22677c040..c60fb7ede 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -36,21 +36,21 @@ use nemo_relay_plugin::{ NemoRelayNativeAsyncStreamMiddlewareCb, NemoRelayNativeConditionalMiddlewareCb, NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, - NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, - NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, - NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, - NemoRelayNativeLlmExecutionRequestContext, NemoRelayNativeLlmExecutionResponseContext, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, - NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, - NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, - NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, - NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, NemoRelayNativePluginContext, - NemoRelayNativePluginRuntime, NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, - NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, - NemoRelayNativeString, NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, - NemoRelayNativeToolExecutionContextCb, NemoRelayNativeToolJsonCb, - NemoRelayNativeWithScopeStackCb, NemoRelayStatus, PendingMarkSpec, PluginContext, - PluginRuntime, ScopeType, ToolExecutionInterceptOutcome, ToolExecutionResult, ToolNext, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeLlmAsyncStream, + NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, + NemoRelayNativeLlmExecutionContext, NemoRelayNativeLlmExecutionRequestContext, + NemoRelayNativeLlmExecutionResponseContext, NemoRelayNativeLlmRequestCodec, + NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, + NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, + NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, + NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, + NemoRelayNativePluginContext, NemoRelayNativePluginRuntime, NemoRelayNativePluginV1, + NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, + NemoRelayNativeScopeType, NemoRelayNativeString, NemoRelayNativeToolConditionalCb, + NemoRelayNativeToolExecutionCb, NemoRelayNativeToolExecutionContextCb, + NemoRelayNativeToolJsonCb, NemoRelayNativeWithScopeStackCb, NemoRelayStatus, PendingMarkSpec, + PluginContext, PluginRuntime, ScopeType, ToolExecutionInterceptOutcome, ToolExecutionResult, + ToolNext, }; use serde_json::{Map, json}; @@ -488,7 +488,7 @@ static UNAVAILABLE_CONTEXT_GATE_CALLS: AtomicUsize = AtomicUsize::new(0); #[test] fn native_abi_struct_sizes_are_self_describing() { - assert_eq!(NEMO_RELAY_NATIVE_ABI_VERSION, 7); + assert_eq!(NEMO_RELAY_NATIVE_ABI_VERSION, 6); assert_eq!( size_of::(), test_host().struct_size @@ -560,10 +560,10 @@ fn assert_native_abi_platform_layout() { ), 600 ); - assert_type_layout::(8, 616); + assert_type_layout::(8, 648); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, log), 608); - assert_native_abi_v7_layout(8, 648, 616, 624, 632, 640); + assert_native_abi_v6_execution_layout(8, 648, 616, 624, 632, 640); assert_type_layout::(8, 56); assert_eq!(plugin_offsets(), [0, 8, 16, 24, 32, 40, 48]); assert_type_layout::(8, 40); @@ -624,10 +624,10 @@ fn assert_native_abi_platform_layout() { ), 296 ); - assert_type_layout::(4, 304); + assert_type_layout::(4, 320); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, log), 300); - assert_native_abi_v7_layout(4, 320, 304, 308, 312, 316); + assert_native_abi_v6_execution_layout(4, 320, 304, 308, 312, 316); assert_type_layout::(4, 28); assert_eq!(plugin_offsets(), [0, 4, 8, 12, 16, 20, 24]); assert_type_layout::(4, 20); @@ -639,7 +639,7 @@ fn assert_type_layout(expected_alignment: usize, expected_size: usize) { assert_eq!(size_of::(), expected_size); } -fn assert_native_abi_v7_layout( +fn assert_native_abi_v6_execution_layout( expected_alignment: usize, expected_size: usize, registration_offset: usize, @@ -647,29 +647,28 @@ fn assert_native_abi_v7_layout( decode_offset: usize, encode_offset: usize, ) { - assert_type_layout::(expected_alignment, expected_size); - assert_eq!(offset_of!(NemoRelayNativeHostApiV7, v6), 0); + assert_type_layout::(expected_alignment, expected_size); assert_eq!( offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, plugin_context_register_async_llm_execution_intercept ), registration_offset ); assert_eq!( - offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), + offset_of!(NemoRelayNativeHostApiV6, async_stream_retain), retain_offset ); assert_eq!( offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, async_stream_llm_request_codec_decode ), decode_offset ); assert_eq!( offset_of!( - NemoRelayNativeHostApiV7, + NemoRelayNativeHostApiV6, async_stream_llm_request_codec_encode ), encode_offset @@ -731,7 +730,7 @@ fn native_abi_v5_extension_is_append_only() { } #[test] -fn native_abi_v6_logging_extension_is_append_only() { +fn native_abi_v6_extension_is_append_only() { assert_eq!(offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!( offset_of!(NemoRelayNativeHostApiV6, log), @@ -3014,15 +3013,6 @@ fn test_host_v6() -> NemoRelayNativeHostApiV6 { NemoRelayNativeHostApiV6 { v5, log: capture_plugin_log, - } -} - -fn test_host_v7() -> NemoRelayNativeHostApiV7 { - let mut v6 = test_host_v6(); - v6.v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION; - v6.v5.v4.v3.v1.struct_size = size_of::(); - NemoRelayNativeHostApiV7 { - v6, plugin_context_register_async_llm_execution_intercept: capture_register_async_llm_execution, async_stream_retain: capture_async_stream_retain, async_stream_llm_request_codec_decode: capture_async_stream_request_decode, @@ -4392,9 +4382,9 @@ fn typed_subscriber_registration_decodes_events() { #[allow(clippy::cognitive_complexity)] // One table-style test deliberately exercises every surface. fn typed_async_middleware_registers_and_round_trips_every_surface() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v4 = &host.v6.v5.v4; - let mut ctx = test_context(&host.v6.v5.v4.v3.v1); + let host = test_host_v6(); + let host_v4 = &host.v5.v4; + let mut ctx = test_context(&host.v5.v4.v3.v1); ctx.register_mark_sanitize_guardrail("mark-async", 1, |_event, mut fields| async move { tokio::time::sleep(Duration::from_millis(1)).await; @@ -4799,8 +4789,8 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { #[test] fn typed_async_unary_execution_codecs_expire_after_completion_settles() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v4 = &host.v6.v5.v4; + let host = test_host_v6(); + let host_v4 = &host.v5.v4; let mut ctx = test_context(&host_v4.v3.v1); let (context_tx, context_rx) = std::sync::mpsc::sync_channel(1); ctx.register_llm_execution_intercept( @@ -4881,8 +4871,8 @@ fn typed_async_unary_execution_codecs_expire_after_completion_settles() { #[test] fn typed_async_stream_execution_codec_expires_after_stream_finishes() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut ctx = test_context(host_v1); let (context_tx, context_rx) = std::sync::mpsc::sync_channel(1); ctx.register_llm_stream_execution_intercept( @@ -5088,8 +5078,8 @@ fn typed_async_llm_sanitize_context_rejects_unknown_builtin_identity() { #[test] fn typed_async_registration_failure_rolls_back_callback_state() { let _guard = begin_test(); - let host = test_host_v7(); - let mut ctx = test_context(&host.v6.v5.v4.v3.v1); + let host = test_host_v6(); + let mut ctx = test_context(&host.v5.v4.v3.v1); *REGISTRATION_STATUS.lock().unwrap() = NemoRelayStatus::InvalidArg; let unary_drops = Arc::new(AtomicUsize::new(0)); @@ -5185,8 +5175,8 @@ fn typed_async_callbacks_isolate_errors_panics_and_invalid_input() { #[test] fn typed_async_continuations_are_concurrent_and_executor_owned() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v4 = &host.v6.v5.v4; + let host = test_host_v6(); + let host_v4 = &host.v5.v4; let mut ctx = test_context(&host_v4.v3.v1); ctx.register_tool_execution_intercept("concurrent", 0, |context, next| async move { assert_eq!(context.tool_name, "tool"); @@ -5504,8 +5494,8 @@ fn typed_async_executor_drop_inside_tokio_runtime_drains_accepted_tasks() { #[test] fn typed_async_stream_cancellation_while_polling_releases_output() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut ctx = test_context(host_v1); let started = Arc::new(AtomicBool::new(false)); ctx.register_llm_stream_execution_intercept("cancel-poll", 0, { @@ -5557,8 +5547,8 @@ fn typed_async_stream_cancellation_while_polling_releases_output() { #[test] fn typed_async_stream_restores_callback_scope_while_polling_returned_stream() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-scope", @@ -5605,8 +5595,8 @@ fn typed_async_stream_restores_callback_scope_while_polling_returned_stream() { #[test] fn typed_async_stream_rejects_item_errors_and_releases_output() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-error", @@ -5657,8 +5647,8 @@ fn typed_async_stream_rejects_item_errors_and_releases_output() { #[test] fn typed_async_stream_rejects_poll_panics_and_releases_output() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-panic", @@ -5713,8 +5703,8 @@ fn typed_async_stream_rejects_poll_panics_and_releases_output() { #[test] fn typed_async_stream_propagates_downstream_pull_errors() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-downstream-error", @@ -5763,8 +5753,8 @@ fn typed_async_stream_propagates_downstream_pull_errors() { #[test] fn typed_async_stream_rejects_missing_continuation() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-null-next", @@ -5895,8 +5885,8 @@ fn raw_event_sanitize_registrations_cover_every_surface() { #[test] fn raw_callback_registrations_preserve_every_middleware_shape() { let _guard = begin_test(); - let host = test_host_v7(); - let mut ctx = test_context(&host.v6.v5.v4.v3.v1); + let host = test_host_v6(); + let mut ctx = test_context(&host.v5.v4.v3.v1); unsafe { assert_eq!( @@ -6038,8 +6028,8 @@ fn raw_callback_registrations_preserve_every_middleware_shape() { #[test] fn raw_async_callback_registrations_use_the_versioned_extension_tables() { let _guard = begin_test(); - let host = test_host_v7(); - let mut ctx = test_context(&host.v6.v5.v4.v3.v1); + let host = test_host_v6(); + let mut ctx = test_context(&host.v5.v4.v3.v1); assert_eq!( unsafe { @@ -6594,8 +6584,8 @@ fn exported_plugin_default_validate_returns_empty_diagnostics() { #[test] fn exported_plugin_register_installs_callbacks_and_propagates_errors() { let _guard = begin_test(); - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut plugin = NemoRelayNativePluginV1::default(); assert_eq!( @@ -6734,8 +6724,8 @@ fn exported_entry_symbol_rejects_prior_host_versions() { ); } - let host = test_host_v7(); - let host_v1 = &host.v6.v5.v4.v3.v1; + let host = test_host_v6(); + let host_v1 = &host.v5.v4.v3.v1; let mut plugin = NemoRelayNativePluginV1::default(); assert_eq!( unsafe { constructor_counting_entry(host_v1, &mut plugin) }, diff --git a/crates/python/src/py_types/mod.rs b/crates/python/src/py_types/mod.rs index 6e56ec7f2..c6dc34be1 100644 --- a/crates/python/src/py_types/mod.rs +++ b/crates/python/src/py_types/mod.rs @@ -156,6 +156,14 @@ fn register_llm_types(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + m.add( + "LlmRequestCodecContext", + m.getattr("LlmSanitizeRequestContext")?, + )?; + m.add( + "LlmResponseCodecContext", + m.getattr("LlmSanitizeResponseContext")?, + )?; m.add_class::()?; m.add_class::()?; m.add_class::()?; diff --git a/crates/worker/src/lib.rs b/crates/worker/src/lib.rs index 2661157ae..866cec09d 100644 --- a/crates/worker/src/lib.rs +++ b/crates/worker/src/lib.rs @@ -295,9 +295,9 @@ impl ToolExecutionContext { } } -/// Active codec context supplied to an LLM request sanitizer. +/// Active codec identity and capability for an LLM request. #[derive(Clone)] -pub struct LlmSanitizeRequestContext { +pub struct LlmRequestCodecContext { /// Identity of the active codec. pub codec: LlmCodecIdentity, runtime: Option, @@ -305,9 +305,9 @@ pub struct LlmSanitizeRequestContext { invocation_id: Option, } -/// Active codec context supplied to an LLM response sanitizer. +/// Active codec identity and capability for an LLM response. #[derive(Clone)] -pub struct LlmSanitizeResponseContext { +pub struct LlmResponseCodecContext { /// Identity of the active codec. pub codec: LlmCodecIdentity, runtime: Option, @@ -315,7 +315,7 @@ pub struct LlmSanitizeResponseContext { invocation_id: Option, } -impl LlmSanitizeRequestContext { +impl LlmRequestCodecContext { /// Resolves the active request codec for this callback. #[must_use] pub fn resolve_codec(&self) -> Option { @@ -327,7 +327,7 @@ impl LlmSanitizeRequestContext { } } -impl LlmSanitizeResponseContext { +impl LlmResponseCodecContext { /// Resolves the active response codec for this callback. #[must_use] pub fn resolve_codec(&self) -> Option { @@ -339,6 +339,12 @@ impl LlmSanitizeResponseContext { } } +/// Backward-compatible name for request codec context supplied to sanitizers. +pub type LlmSanitizeRequestContext = LlmRequestCodecContext; + +/// Backward-compatible name for response codec context supplied to sanitizers. +pub type LlmSanitizeResponseContext = LlmResponseCodecContext; + /// Invocation-scoped proxy for the active LLM request codec. #[derive(Clone)] pub struct WorkerRequestCodec { @@ -395,8 +401,8 @@ impl WorkerResponseCodec { /// not change these identities; codec operations reject incompatible payloads. #[derive(Clone)] pub struct LlmExecutionContext { - request_codec: LlmSanitizeRequestContext, - response_codec: Option, + request_codec: LlmRequestCodecContext, + response_codec: Option, } impl std::fmt::Debug for LlmExecutionContext { @@ -415,7 +421,7 @@ impl std::fmt::Debug for LlmExecutionContext { impl LlmExecutionContext { /// Request codec identity and invocation-scoped operations. #[must_use] - pub fn request_codec(&self) -> &LlmSanitizeRequestContext { + pub fn request_codec(&self) -> &LlmRequestCodecContext { &self.request_codec } @@ -424,7 +430,7 @@ impl LlmExecutionContext { /// Streaming execution returns `None` because Relay response codecs decode /// completed provider responses, not individual stream chunks. #[must_use] - pub fn response_codec(&self) -> Option<&LlmSanitizeResponseContext> { + pub fn response_codec(&self) -> Option<&LlmResponseCodecContext> { self.response_codec.as_ref() } } @@ -2923,12 +2929,12 @@ impl LlmPayload { let response_codec = context .response .as_ref() - .map(|response| -> Result { + .map(|response| -> Result { let identity = require_execution_field( response.codec.as_ref(), "response codec identity is missing", )?; - Ok(LlmSanitizeResponseContext { + Ok(LlmResponseCodecContext { codec: codec_identity_from_proto(Some(identity)), runtime: Some(runtime.clone()), codec_capability_id: response.codec_capability_id.clone(), @@ -2942,7 +2948,7 @@ impl LlmPayload { )); } Ok(LlmExecutionContext { - request_codec: LlmSanitizeRequestContext { + request_codec: LlmRequestCodecContext { codec: codec_identity_from_proto(Some(request_identity)), runtime: Some(runtime.clone()), codec_capability_id: request.codec_capability_id.clone(), diff --git a/docs/about-nemo-relay/release-notes/index.mdx b/docs/about-nemo-relay/release-notes/index.mdx index cb2c92203..66ddbd60c 100644 --- a/docs/about-nemo-relay/release-notes/index.mdx +++ b/docs/about-nemo-relay/release-notes/index.mdx @@ -44,7 +44,7 @@ interceptors receive request codec access only. **Breaking change:** Callback signatures change across Rust, Python, Node.js, Go, C, native plugins, and Rust and Python gRPC workers. Python language-binding streaming and public C callbacks also gain the logical LLM -name. Native plugins must rebuild for internal ABI v7, and affected workers +name. Native plugins must rebuild for the finalized internal ABI v6 layout, and affected workers must regenerate their protobuf bindings and rebuild. Authored compatibility labels remain `native_api = "1"` and `grpc-v1`; plugin manifests must use a Relay range that begins at 0.10 or otherwise excludes 0.9. Refer to the diff --git a/docs/build-plugins/about.mdx b/docs/build-plugins/about.mdx index 4e99a7095..aaba58373 100644 --- a/docs/build-plugins/about.mdx +++ b/docs/build-plugins/about.mdx @@ -56,8 +56,8 @@ artifact solves a concrete operational problem. Native Rust plugins suit reusable middleware whose callback latency or throughput is important enough to justify platform-specific binaries and full in-process trust. A native plugin uses manifest compatibility `compat.native_api = "1"`; the current SDK -uses C host-table ABI v7. Relay 0.10 intentionally rejects older compiled callback -layouts, so native plugins must rebuild for this release even though the authored +uses C host-table ABI v6. Relay 0.10 rejects v2-v5 compiled callback layouts, so native +plugins must rebuild for this release even though the authored manifest label remains `native_api = "1"`. Those are different version axes, as the [Native ABI Reference](/build-plugins/native/native-abi-reference) explains. diff --git a/docs/build-plugins/native/about.mdx b/docs/build-plugins/native/about.mdx index 83fa3915c..614407b9b 100644 --- a/docs/build-plugins/native/about.mdx +++ b/docs/build-plugins/native/about.mdx @@ -26,16 +26,16 @@ Three version values answer different questions: |---|---|---| | Package manifest | `manifest_version = 1` | Shape of the authored `relay-plugin.toml` file. | | Manifest native API | `compat.native_api = "1"` | Native plugin package contract accepted by discovery and trust validation. | -| C host-table ABI | v7 | Function table required by the current `nemo-relay-plugin` SDK. Relay 0.10 rejects older callback layouts. | +| C host-table ABI | v6 | Function table required by the current `nemo-relay-plugin` SDK. Relay 0.10 rejects v2-v5 callback layouts. | The checked example registers a tool execution intercept, so it declares `compat.relay = ">=0.10.0,<1.0"`. Every native plugin built with this SDK must use -that lower bound because an older Relay cannot load the ABI v7 function table, even +that lower bound because an older Relay cannot load the finalized ABI v6 function table, even when the component itself does not register an LLM execution intercept. -Relay 0.10 adds directional codec context to every LLM execution callback and -advances the internal table to v7. Rebuild all native plugins for this release; -the authored `native_api = "1"` label does not change. +Relay 0.10 finalizes ABI v6 with directional codec context on every LLM +execution callback. Rebuild all native plugins for this release; the authored +`native_api = "1"` label does not change. Relay 0.8 changed the native API 1 tool-result JSON contract without changing the v4 host-table layout. A tool callback and `ToolNext` continuation return diff --git a/docs/build-plugins/native/native-abi-reference.mdx b/docs/build-plugins/native/native-abi-reference.mdx index 57f173c12..b8639036a 100644 --- a/docs/build-plugins/native/native-abi-reference.mdx +++ b/docs/build-plugins/native/native-abi-reference.mdx @@ -23,19 +23,19 @@ extern "C" fn nemo_relay_register_plugin( ) -> NemoRelayStatus ``` -The current host requires ABI v7. It does not fall back to v6 or earlier tables because -v7 changes the LLM execution callback layouts; invoking a stale callback through the -new layout would be unsafe. Rebuild every native plugin for Relay 0.10 and raise its +The current host requires ABI v6. It does not fall back to v2-v5 tables because the +LLM execution callback layouts changed; invoking a stale callback through the new +layout would be unsafe. Rebuild every native plugin for Relay 0.10 and raise its `compat.relay` lower bound to `0.10.0`. The authored manifest label remains `compat.native_api = "1"`. -ABI v7 adds directional request and unary-response codec context to raw, typed, -and asynchronous LLM execution callbacks. The generic asynchronous middleware -callback remains unchanged; ABI v7 appends an execution-specific unary registration, -and the existing execution-specific stream callback gains the context parameter. -ABI v6 adds a host-routed operational logging function. Native plugins pass a level, optional -target, message, and optional JSON-object fields; Relay applies its operational logging policy -and preserves those fields in structured JSONL output. +ABI v6 adds host-routed operational logging and directional request and unary-response +codec context to raw, typed, and asynchronous LLM execution callbacks. The generic +asynchronous middleware callback remains unchanged; ABI v6 appends an execution-specific +unary registration, and the existing execution-specific stream callback gains the context +parameter. Native plugins pass a log level, optional target, message, and optional JSON-object +fields; Relay applies its operational logging policy and preserves those fields in structured +JSONL output. ABI v4 extends the complete v3 prefix with completion-scoped [codecs](/about-nemo-relay/concepts/codecs) and pull-based downstream LLM streams, plus an activation-owned runtime capability for @@ -53,8 +53,7 @@ function signatures and field order are defined by the public | Frozen v3 extension | Completion resolve, reject, cancellation inspection, and release; one-shot completion-coupled continuation invocation; continuation release; generic async middleware registration; bounded output-stream push, finish, reject, cancellation inspection, and release; downstream stream invocation; stream-middleware registration; and repeated or concurrent unary continuation invocation with independent result callbacks. | | Frozen v4 extension | Completion-scoped LLM request decode and encode plus response decode; pull-based downstream LLM stream open, pull, cancel, and release; completion retain for typed codec facades; output-stream backpressure inspection; extended mark emission; runtime diagnostics; activation-owned runtime capability creation, retain, and release; global runtime-registration discovery; owned conditional middleware guardrail registration and deregistration; and activation-owned and runtime-discovered callback gate registration. The callback registration slots are appended after the original constant-reason slots. | | v5 extension | Context-carrying tool execution intercept registration. The callback receives one JSON object containing `tool_name`, `args`, and `tool_call_id`. | -| v6 extension | Host-routed operational logging with structured fields. | -| Current v7 extension | Directional codec context for raw and typed LLM execution callbacks, an execution-specific asynchronous unary registration, context on the asynchronous stream-execution callback, and stream retain plus stream-scoped request decode and encode operations. | +| Current v6 extension | Host-routed operational logging with structured fields; directional codec context for raw and typed LLM execution callbacks; an execution-specific asynchronous unary registration; context on the asynchronous stream-execution callback; and stream retain plus stream-scoped request decode and encode operations. | The prefix and descriptor layout are explicit. A plugin fills the descriptor with its stable kind, component multiplicity, opaque state, callbacks, and destructor. The host @@ -163,7 +162,7 @@ or binding, and restoration. ## Async Completions and Continuations `PluginContext::register_async_middleware_raw` registers non-stream middleware other -than LLM execution. ABI v7 uses the execution-specific +than LLM execution. ABI v6 uses the execution-specific `plugin_context_register_async_llm_execution_intercept` entry for asynchronous unary LLM execution callbacks. Both paths can settle later. Return `Complete` only after resolving or rejecting the completion inside @@ -240,7 +239,7 @@ facades retain their completion and typed streaming request facades retain their output stream, so they remain memory-safe while the callback future or returned stream owns them. Codec calls fail after the completion or stream settles. Raw plugins that need request codec access after an asynchronous stream callback -returns must retain the stream and use the v7 stream-scoped decode and encode +returns must retain the stream and use the v6 stream-scoped decode and encode operations; release that stream reference after the last call. ## Unload Ordering diff --git a/docs/build-plugins/package-discoverable-plugins.mdx b/docs/build-plugins/package-discoverable-plugins.mdx index f93672014..9034bb552 100644 --- a/docs/build-plugins/package-discoverable-plugins.mdx +++ b/docs/build-plugins/package-discoverable-plugins.mdx @@ -26,7 +26,7 @@ equivalent binding configuration. For native packages, `compat.native_api = "1"` is the authored manifest contract. It is not the [C host-table ABI number](/build-plugins/native/native-abi-reference). The -current SDK requires ABI v7, and Relay 0.10 rejects older compiled callback layouts. +current SDK requires ABI v6, and Relay 0.10 rejects v2-v5 compiled callback layouts. For workers, declare `compat.worker_protocol = "grpc-v1"`; the [handshake](/build-plugins/workers/grpc-v1-protocol) still negotiates the exact protocol, surfaces, authentication token, and lifecycle at startup. diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index 62c0840d6..fa3d0416c 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -40,7 +40,7 @@ This is a source and binary compatibility break for execution-intercept users: - Rebuild every native plugin against the 0.10 SDK. Native plugins continue to declare `compat.native_api = "1"`, but must set `compat.relay` to `">=0.10.0,<1.0"` or another range that excludes Relay 0.9. Relay 0.10 uses - internal native ABI v7 and rejects older compiled layouts. + internal native ABI v6 and rejects v2-v5 compiled layouts. - Regenerate and rebuild gRPC workers that register an LLM execution intercept. The protocol remains `grpc-v1`, but those workers must also exclude Relay 0.9 in `compat.relay`. diff --git a/examples/rust-native-plugin/README.md b/examples/rust-native-plugin/README.md index 3f3e63326..372344304 100644 --- a/examples/rust-native-plugin/README.md +++ b/examples/rust-native-plugin/README.md @@ -11,10 +11,10 @@ helpers live in separate source modules. Together they register the subscriber, all three event sanitizers, five tool surfaces, and six LLM surfaces exposed by the current typed 0.10.0 SDK. -Relay 0.10 uses native ABI v7. Every LLM execution callback receives directional codec +Relay 0.10 uses native ABI v6. Every LLM execution callback receives directional codec context before its continuation; streaming execution exposes request codec operations but no response decoder. The manifest continues to declare `native_api = "1"`, and its -Relay lower bound is `0.10.0` because older compiled callback layouts are rejected. +Relay lower bound is `0.10.0` because Relay 0.9 uses a pre-v6 callback layout. Run the focused tests and build the shared library from this directory. The configuration tests isolate validation and schema contracts. The lifecycle test diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index 626500d84..0a70fea7d 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -250,12 +250,20 @@ type LLMSanitizeResponseContext struct { resolved *LLMResponseSanitizeCodec } +// LLMRequestCodecContext is the general name for request codec context used by +// execution intercepts. The sanitizer-specific name remains source-compatible. +type LLMRequestCodecContext = LLMSanitizeRequestContext + +// LLMResponseCodecContext is the general name for response codec context used +// by execution intercepts. The sanitizer-specific name remains source-compatible. +type LLMResponseCodecContext = LLMSanitizeResponseContext + // LLMExecutionContext provides invocation-scoped codec access to an LLM // execution intercept. RequestCodec is always present. ResponseCodec is // available for unary execution and nil for streaming execution. type LLMExecutionContext struct { - RequestCodec LLMSanitizeRequestContext - ResponseCodec *LLMSanitizeResponseContext + RequestCodec LLMRequestCodecContext + ResponseCodec *LLMResponseCodecContext } // ResolveCodec returns the active callback-scoped response codec, if any. diff --git a/python/nemo_relay/__init__.py b/python/nemo_relay/__init__.py index 75a55d51f..a49ba65f9 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -104,7 +104,9 @@ async def main(): LlmExecutionContext, LLMHandle, LLMRequest, + LlmRequestCodecContext, LLMRequestInterceptOutcome, + LlmResponseCodecContext, LlmSanitizeRequestCodec, LlmSanitizeRequestContext, LlmSanitizeResponseCodec, @@ -843,6 +845,8 @@ def worker() -> None: "LlmSanitizeResponseGuardrail", "LlmCodecIdentity", "LlmExecutionContext", + "LlmRequestCodecContext", + "LlmResponseCodecContext", "LlmSanitizeRequestContext", "LlmSanitizeResponseContext", "LlmSanitizeRequestCodec", diff --git a/python/nemo_relay/__init__.pyi b/python/nemo_relay/__init__.pyi index 326ce1909..f203f8fe6 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -71,6 +71,9 @@ from nemo_relay._native import ( from nemo_relay._native import ( LlmExecutionContext as LlmExecutionContext, ) +from nemo_relay._native import ( + LlmRequestCodecContext as LlmRequestCodecContext, +) from nemo_relay._native import ( LLMHandle as LLMHandle, ) @@ -92,6 +95,9 @@ from nemo_relay._native import ( from nemo_relay._native import ( LlmSanitizeResponseContext as LlmSanitizeResponseContext, ) +from nemo_relay._native import ( + LlmResponseCodecContext as LlmResponseCodecContext, +) from nemo_relay._native import ( LogSeverity as LogSeverity, ) diff --git a/python/nemo_relay/_native.pyi b/python/nemo_relay/_native.pyi index e22297b13..c5434f5d9 100644 --- a/python/nemo_relay/_native.pyi +++ b/python/nemo_relay/_native.pyi @@ -130,9 +130,9 @@ class LlmExecutionContext: """Codec capabilities for one managed LLM execution intercept invocation.""" @property - def request_codec(self) -> LlmSanitizeRequestContext: ... + def request_codec(self) -> LlmRequestCodecContext: ... @property - def response_codec(self) -> LlmSanitizeResponseContext | None: ... + def response_codec(self) -> LlmResponseCodecContext | None: ... class LlmSanitizeRequestContext: """Per-call context passed to an LLM request sanitizer callback.""" @@ -148,6 +148,9 @@ class LlmSanitizeResponseContext: def codec(self) -> LlmCodecIdentity: ... def resolve_codec(self) -> LlmSanitizeResponseCodec | None: ... +LlmRequestCodecContext = LlmSanitizeRequestContext +LlmResponseCodecContext = LlmSanitizeResponseContext + class LlmSanitizeRequestCodec: def decode(self, request: LLMRequest) -> AnnotatedLLMRequest: ... def encode(self, annotated: AnnotatedLLMRequest, original: LLMRequest) -> LLMRequest: ... diff --git a/python/plugin/src/nemo_relay_plugin/__init__.py b/python/plugin/src/nemo_relay_plugin/__init__.py index e42fccae7..32e8bf9e2 100644 --- a/python/plugin/src/nemo_relay_plugin/__init__.py +++ b/python/plugin/src/nemo_relay_plugin/__init__.py @@ -27,6 +27,8 @@ EventSanitizeFields: Mutable event observability fields. LlmRequest: A Relay LLM request represented as a JSON object. LlmCodecIdentity: Typed discriminator for the active LLM codec. + LlmRequestCodecContext: Request-direction codec identity and operations. + LlmResponseCodecContext: Response-direction codec identity and operations. LlmSanitizeRequestContext: Per-call context supplied to an LLM request sanitizer. LlmSanitizeResponseContext: Per-call context supplied to an LLM response sanitizer. LlmExecutionContext: Invocation-scoped codec context supplied to an LLM @@ -115,7 +117,9 @@ LlmOptimizationTokens, LlmRequest, LlmRequestCallback, + LlmRequestCodecContext, LlmRequestInterceptOutcome, + LlmResponseCodecContext, LlmSanitizeRequestCallback, LlmSanitizeRequestContext, LlmSanitizeResponseCallback, @@ -169,6 +173,8 @@ "LlmCodecIdentity", "LlmExecutionCallback", "LlmExecutionContext", + "LlmRequestCodecContext", + "LlmResponseCodecContext", "LogSeverity", "MetricKind", "MetricMeasurement", diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index a9e1aaa6a..0ae5cfc19 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -226,6 +226,13 @@ def resolve_codec(self) -> "WorkerResponseCodec | None": return WorkerResponseCodec(self._runtime, self._capability_id, self._invocation_id) +# General names for contexts that are also supplied to execution interceptors. +# The original class objects remain canonical at runtime for compatibility with +# repr, pickling, and code that inspects ``__name__``. +LlmRequestCodecContext = LlmSanitizeRequestContext +LlmResponseCodecContext = LlmSanitizeResponseContext + + @dataclass(frozen=True) class WorkerRequestCodec: """Invocation-scoped async proxy for an active request codec.""" @@ -267,8 +274,8 @@ class LlmExecutionContext: not select a new codec; incompatible codec operations fail. """ - request_codec: LlmSanitizeRequestContext - response_codec: LlmSanitizeResponseContext | None + request_codec: LlmRequestCodecContext + response_codec: LlmResponseCodecContext | None def _llm_codec_identity(invocation: pb.LlmInvocation) -> LlmCodecIdentity: @@ -311,7 +318,7 @@ def _llm_execution_context( raise WorkerSdkError("malformed LLM execution codec context: request codec identity is missing") request_id = context.request.codec_capability_id if context.request.HasField("codec_capability_id") else None - request_context = LlmSanitizeRequestContext( + request_context = LlmRequestCodecContext( codec=_codec_identity( context.request.codec.kind, context.request.codec.id if context.request.codec.HasField("id") else None, @@ -320,12 +327,12 @@ def _llm_execution_context( _capability_id=request_id, _invocation_id=invocation_id, ) - response_context: LlmSanitizeResponseContext | None = None + response_context: LlmResponseCodecContext | None = None if context.HasField("response"): if not context.response.HasField("codec"): raise WorkerSdkError("malformed LLM execution codec context: response codec identity is missing") response_id = context.response.codec_capability_id if context.response.HasField("codec_capability_id") else None - response_context = LlmSanitizeResponseContext( + response_context = LlmResponseCodecContext( codec=_codec_identity( context.response.codec.kind, context.response.codec.id if context.response.codec.HasField("id") else None, From 49a379d33474d539b010d5b6c00d4eb2d9abf385 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 23 Sep 2026 18:20:40 -0400 Subject: [PATCH 08/22] fix!: advance execution context native ABI to v7 Signed-off-by: Alex Fournier --- crates/core/src/plugin/dynamic.rs | 2 +- crates/core/src/plugin/dynamic/native.rs | 55 +++++--- .../tests/fixtures/native_plugin/src/lib.rs | 31 ++++- .../tests/integration/native_plugin_tests.rs | 19 +-- crates/core/tests/unit/native_plugin_tests.rs | 50 +++++-- crates/plugin/README.md | 8 +- crates/plugin/src/async_sdk.rs | 62 ++++----- crates/plugin/src/lib.rs | 74 ++++++---- crates/plugin/tests/typed_callbacks.rs | 126 ++++++++++-------- docs/about-nemo-relay/release-notes/index.mdx | 2 +- docs/build-plugins/about.mdx | 2 +- docs/build-plugins/native/about.mdx | 6 +- .../native/native-abi-reference.mdx | 23 ++-- .../package-discoverable-plugins.mdx | 2 +- docs/reference/migration-guides.mdx | 2 +- examples/rust-native-plugin/README.md | 4 +- 16 files changed, 281 insertions(+), 187 deletions(-) diff --git a/crates/core/src/plugin/dynamic.rs b/crates/core/src/plugin/dynamic.rs index dadc4a9e8..a6b496073 100644 --- a/crates/core/src/plugin/dynamic.rs +++ b/crates/core/src/plugin/dynamic.rs @@ -169,7 +169,7 @@ pub(super) fn validate_native_abi_compatibility( })?; if version_requirement_matches_minor(&requirement, 0, 9) { return Err(PluginError::InvalidConfig(format!( - "dynamic native plugin '{plugin_kind}' uses native ABI v6 and must declare compat.relay = \">=0.10,<1.0\" or another range that excludes Relay 0.9" + "dynamic native plugin '{plugin_kind}' uses native ABI v7 and must declare compat.relay = \">=0.10,<1.0\" or another range that excludes Relay 0.9" ))); } Ok(()) diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 92663d2f2..1fe534178 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -57,7 +57,7 @@ use chrono::{DateTime, Utc}; use libloading::{Library, Symbol}; use nemo_relay_plugin::{ NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, - NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT, + NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT, NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncCompletion, NemoRelayNativeAsyncLlmExecutionCb, @@ -67,20 +67,20 @@ use nemo_relay_plugin::{ NemoRelayNativeAsyncStreamMiddlewareCb, NemoRelayNativeConditionalMiddlewareCb, NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, - NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeLlmAsyncStream, - NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, - NemoRelayNativeLlmExecutionContext, NemoRelayNativeLlmExecutionRequestContext, - NemoRelayNativeLlmExecutionResponseContext, NemoRelayNativeLlmRequestCodec, - NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, - NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, - NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, - NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, - NemoRelayNativePluginContext, NemoRelayNativePluginEntry, NemoRelayNativePluginRuntime, - NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, - NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, NemoRelayNativeString, - NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, - NemoRelayNativeToolExecutionContextCb, NemoRelayNativeToolJsonCb, - NemoRelayNativeWithScopeStackCb, NemoRelayStatus, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, + NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, + NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, + NemoRelayNativeLlmExecutionRequestContext, NemoRelayNativeLlmExecutionResponseContext, + NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, + NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, + NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, + NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, + NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, NemoRelayNativePluginContext, + NemoRelayNativePluginEntry, NemoRelayNativePluginRuntime, NemoRelayNativePluginV1, + NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, + NemoRelayNativeScopeType, NemoRelayNativeString, NemoRelayNativeToolConditionalCb, + NemoRelayNativeToolExecutionCb, NemoRelayNativeToolExecutionContextCb, + NemoRelayNativeToolJsonCb, NemoRelayNativeWithScopeStackCb, NemoRelayStatus, }; use serde_json::{Map, Value as Json}; use sha2::{Digest, Sha256}; @@ -947,6 +947,18 @@ 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_v7) + .v6 + .v5 + .v4 + .v3 + .v1 as *const NemoRelayNativeHostApiV1 +} + +#[cfg(test)] +fn native_host_api_v6() -> *const NemoRelayNativeHostApiV1 { static HOST_API: OnceLock = OnceLock::new(); &HOST_API.get_or_init(build_native_host_api_v6).v5.v4.v3.v1 as *const NemoRelayNativeHostApiV1 } @@ -1108,11 +1120,20 @@ fn build_native_host_api_v5() -> NemoRelayNativeHostApiV5 { fn build_native_host_api_v6() -> NemoRelayNativeHostApiV6 { let mut v5 = build_native_host_api_v5(); - v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT; + v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING; v5.v4.v3.v1.struct_size = std::mem::size_of::(); NemoRelayNativeHostApiV6 { v5, log: native_log, + } +} + +fn build_native_host_api_v7() -> NemoRelayNativeHostApiV7 { + let mut v6 = build_native_host_api_v6(); + v6.v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT; + v6.v5.v4.v3.v1.struct_size = std::mem::size_of::(); + NemoRelayNativeHostApiV7 { + v6, plugin_context_register_async_llm_execution_intercept: native_plugin_context_register_async_llm_execution_intercept, async_stream_retain: native_async_stream_retain, @@ -4187,7 +4208,7 @@ unsafe extern "C" fn native_plugin_context_register_async_middleware( | NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept ) { set_native_last_error( - "LLM execution middleware requires its dedicated ABI-v6 registration function", + "LLM execution middleware requires its dedicated ABI-v7 registration function", ); return NemoRelayStatus::InvalidArg; } diff --git a/crates/core/tests/fixtures/native_plugin/src/lib.rs b/crates/core/tests/fixtures/native_plugin/src/lib.rs index f8687746a..fd1385191 100644 --- a/crates/core/tests/fixtures/native_plugin/src/lib.rs +++ b/crates/core/tests/fixtures/native_plugin/src/lib.rs @@ -12,12 +12,12 @@ use nemo_relay_plugin::{ Json, LlmJsonAsyncStream, LlmRequest, LlmRequestInterceptOutcome, MetricKind, MetricMeasurement, MetricValueType, NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, - NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, + NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, NativeExecutorConfig, NativePlugin, NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncMiddlewareCb, NemoRelayNativeAsyncMiddlewareKind, NemoRelayNativeAsyncNext, NemoRelayNativeAsyncStream, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, - NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, NemoRelayNativeLlmExecutionContext, NemoRelayNativePluginContext, NemoRelayNativePluginV1, NemoRelayNativeString, NemoRelayNativeToolNextFn, NemoRelayStatus, PendingMarkSpec, PluginContext, PluginRuntime, RuntimeRegistrationKind, ScopeCategory, ScopeType, @@ -547,6 +547,23 @@ pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v5( } } +/// Raw ABI-v6 entry used to verify that the immediately stale callback layout is rejected. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v6( + host: *const NemoRelayNativeHostApiV1, + out: *mut NemoRelayNativePluginV1, +) -> NemoRelayStatus { + unsafe { + fixture_compat_entry( + host, + out, + NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, + std::mem::size_of::(), + b"fixture_native_v6", + ) + } +} + /// Raw ABI-v4 entry used to verify fallback for plugins built with the previous SDK. #[unsafe(no_mangle)] pub unsafe extern "C" fn nemo_relay_fixture_native_plugin_v4( @@ -945,7 +962,7 @@ unsafe extern "C" fn raw_register_event_sanitize_errors( } struct FixtureAsyncPlugin { - host: Option>, + host: Option>, } impl NativePlugin for FixtureAsyncPlugin { @@ -960,17 +977,17 @@ impl NativePlugin for FixtureAsyncPlugin { ) -> nemo_relay_plugin::Result<()> { let host = ctx.host_api(); if host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION - || host.struct_size < std::mem::size_of::() + || host.struct_size < std::mem::size_of::() { - return Err("fixture async plugin requires ABI v6".into()); + return Err("fixture async plugin requires ABI v7".into()); } self.host = Some(Box::new(unsafe { - *(host as *const _ as *const NemoRelayNativeHostApiV6) + *(host as *const _ as *const NemoRelayNativeHostApiV7) })); let user_data = self .host .as_deref() - .map(|host| (host as *const NemoRelayNativeHostApiV6).cast_mut().cast()) + .map(|host| (host as *const NemoRelayNativeHostApiV7).cast_mut().cast()) .expect("fixture async host was initialized"); let registrations: [( diff --git a/crates/core/tests/integration/native_plugin_tests.rs b/crates/core/tests/integration/native_plugin_tests.rs index 9c5511c04..84c6ffffc 100644 --- a/crates/core/tests/integration/native_plugin_tests.rs +++ b/crates/core/tests/integration/native_plugin_tests.rs @@ -974,7 +974,7 @@ fn native_api_one_does_not_admit_a_stale_abi_v2_binary() { [load_spec("fixture_native", &manifest_ref)], "a native_api=1 manifest must not make an ABI-v2 binary compatible", ); - assert!(error.contains("rejected native ABI 6"), "{error}"); + assert!(error.contains("rejected native ABI 7"), "{error}"); assert!(error.contains("rebuild the plugin"), "{error}"); } @@ -1099,7 +1099,7 @@ async fn native_loader_rejects_manifest_that_admits_pre_zero_eight_relay() { } #[tokio::test] -async fn native_abi_v6_rejects_manifest_that_admits_relay_zero_nine() { +async fn native_abi_v7_rejects_manifest_that_admits_relay_zero_nine() { let _guard = NATIVE_PLUGIN_TEST_LOCK.lock().await; let fixture = build_fixture_plugin(); let manifest_ref = write_manifest_text(ManifestOptions { @@ -1112,10 +1112,10 @@ async fn native_abi_v6_rejects_manifest_that_admits_relay_zero_nine() { }); let error = expect_native_load_error_from_specs( [load_spec("fixture_native", &manifest_ref)], - "an ABI-v6 native plugin must exclude Relay 0.9", + "an ABI-v7 native plugin must exclude Relay 0.9", ); assert!( - error.contains("uses native ABI v6") && error.contains("excludes Relay 0.9"), + error.contains("uses native ABI v7") && error.contains("excludes Relay 0.9"), "{error}" ); } @@ -1164,18 +1164,19 @@ fn native_loader_rejects_abi_v3_plugins() { let error = expect_native_load_error_from_specs( [load_spec("fixture_native_v3", &manifest_ref)], - "ABI-v3 plugins must be rebuilt for ABI v6", + "ABI-v3 plugins must be rebuilt for ABI v7", ); - assert!(error.contains("rejected native ABI 6"), "{error}"); + assert!(error.contains("rejected native ABI 7"), "{error}"); assert!(error.contains("rebuild the plugin"), "{error}"); } #[test] -fn native_loader_rejects_v2_v4_and_v5_plugins() { +fn native_loader_rejects_v2_v4_v5_and_v6_plugins() { let _guard = NATIVE_PLUGIN_TEST_LOCK.blocking_lock(); let fixture = build_fixture_plugin(); for (plugin_id, symbol) in [ + ("fixture_native_v6", "nemo_relay_fixture_native_plugin_v6"), ("fixture_native_v5", "nemo_relay_fixture_native_plugin_v5"), ("fixture_native_v4", "nemo_relay_fixture_native_plugin_v4"), ("fixture_native_v2", "nemo_relay_fixture_native_plugin_v2"), @@ -1191,9 +1192,9 @@ fn native_loader_rejects_v2_v4_and_v5_plugins() { let error = expect_native_load_error_from_specs( [load_spec(plugin_id, &manifest_ref)], - "stale native plugins must be rebuilt for ABI v6", + "stale native plugins must be rebuilt for ABI v7", ); - assert!(error.contains("rejected native ABI 6"), "{error}"); + assert!(error.contains("rejected native ABI 7"), "{error}"); assert!(error.contains("rebuild the plugin"), "{error}"); } } diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index c032e449f..3aac06e8d 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -634,6 +634,7 @@ fn assert_native_digest_edges() { fn assert_native_host_api_versions() { let current = native_host_api(); + let frozen_v6 = native_host_api_v6(); let frozen_v5 = native_host_api_v5(); let frozen_v4 = native_host_api_v4(); let frozen_v3 = native_host_api_v3(); @@ -641,6 +642,11 @@ fn assert_native_host_api_versions() { assert_native_host_api_descriptor( current, NEMO_RELAY_NATIVE_ABI_VERSION, + std::mem::size_of::(), + ); + assert_native_host_api_descriptor( + frozen_v6, + NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, std::mem::size_of::(), ); assert_native_host_api_descriptor( @@ -663,6 +669,7 @@ fn assert_native_host_api_versions() { NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, std::mem::size_of::(), ); + assert_native_host_api_v7_layout(); assert_native_host_api_v6_layout(); assert_native_host_api_v5_layout(); assert_native_host_api_v4_layout(); @@ -784,30 +791,46 @@ fn assert_native_host_api_v6_layout() { #[cfg(target_pointer_width = "64")] { assert_eq!(std::mem::align_of::(), 8); - assert_eq!(std::mem::size_of::(), 648); + assert_eq!(std::mem::size_of::(), 616); assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, log), 608); + } + #[cfg(target_pointer_width = "32")] + { + assert_eq!(std::mem::align_of::(), 4); + assert_eq!(std::mem::size_of::(), 304); + assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, v5), 0); + assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, log), 300); + } +} + +fn assert_native_host_api_v7_layout() { + #[cfg(target_pointer_width = "64")] + { + assert_eq!(std::mem::align_of::(), 8); + assert_eq!(std::mem::size_of::(), 648); + assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV7, v6), 0); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, plugin_context_register_async_llm_execution_intercept ), 616 ); assert_eq!( - std::mem::offset_of!(NemoRelayNativeHostApiV6, async_stream_retain), + std::mem::offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), 624 ); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, async_stream_llm_request_codec_decode ), 632 ); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, async_stream_llm_request_codec_encode ), 640 @@ -815,31 +838,30 @@ fn assert_native_host_api_v6_layout() { } #[cfg(target_pointer_width = "32")] { - assert_eq!(std::mem::align_of::(), 4); - assert_eq!(std::mem::size_of::(), 320); - assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, v5), 0); - assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV6, log), 300); + assert_eq!(std::mem::align_of::(), 4); + assert_eq!(std::mem::size_of::(), 320); + assert_eq!(std::mem::offset_of!(NemoRelayNativeHostApiV7, v6), 0); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, plugin_context_register_async_llm_execution_intercept ), 304 ); assert_eq!( - std::mem::offset_of!(NemoRelayNativeHostApiV6, async_stream_retain), + std::mem::offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), 308 ); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, async_stream_llm_request_codec_decode ), 312 ); assert_eq!( std::mem::offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, async_stream_llm_request_codec_encode ), 316 @@ -1864,7 +1886,7 @@ fn assert_native_json_output_and_host_api() { 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::() ); } diff --git a/crates/plugin/README.md b/crates/plugin/README.md index db6ca1d49..dcda060cc 100644 --- a/crates/plugin/README.md +++ b/crates/plugin/README.md @@ -32,7 +32,7 @@ the dynamic-library boundary on the stable C-compatible ABI. | `PluginContext` | Installs component-owned subscribers, guardrails, intercepts, continuations, and streams. | | `PluginRuntime` | Emits marks and manages Relay-owned scopes and scope stacks through typed host helpers. | | `nemo_relay_plugin!` | Exports the one versioned native entry point used by the loader. | -| Native ABI v6 | Keeps C-compatible host and plugin tables behind the safe Rust interface. Relay 0.10 finalizes ABI v6 with directional codec context for LLM execution callbacks and requires native plugins to rebuild against that layout. | +| Native ABI v7 | Keeps C-compatible host and plugin tables behind the safe Rust interface. Relay 0.10 adds directional codec context for LLM execution callbacks and requires native plugins to rebuild against that layout. | | Typed async middleware | Drives guardrails, sanitizers, and intercepts on a per-component SDK-owned Tokio executor. Subscribers and raw ABI registrations remain synchronous. | | Async continuations and streams | `ToolNext`, `LlmNext`, and `LlmStreamNext` support repeated or concurrent downstream calls. Streaming LLM continuations use a pull-based host handle. | | Tool results | `ToolNext` returns `ToolExecutionResult`, which keeps an application result and optional annotation together. | @@ -104,10 +104,10 @@ context-aware tool execution intercept must rebuild and set table. Under Relay 0.9, typed async plugins that did not use this registration could retain `compat.relay = ">=0.8.0,<1.0"`. -Relay 0.10 finalizes the internal ABI v6 table and makes +Relay 0.10 advances the internal table to ABI v7 and makes `LlmExecutionContext` part of every unary and streaming LLM execution callback. -Because this changes callback layouts, the 0.10 host rejects v2-v5 tables. -Rebuild every plugin with the finalized 0.10 v6 SDK and set +Because this changes callback layouts, the 0.10 host rejects v2-v6 tables. +Rebuild every plugin with the 0.10 v7 SDK and set `compat.relay = ">=0.10.0,<1.0"`. The authored manifest contract remains `compat.native_api = "1"`. diff --git a/crates/plugin/src/async_sdk.rs b/crates/plugin/src/async_sdk.rs index 13a415ef2..968921cc2 100644 --- a/crates/plugin/src/async_sdk.rs +++ b/crates/plugin/src/async_sdk.rs @@ -154,10 +154,10 @@ unsafe impl Send for HostV4 {} unsafe impl Sync for HostV4 {} #[derive(Clone, Copy)] -struct HostV6(NemoRelayNativeHostApiV6); +struct HostV7(NemoRelayNativeHostApiV7); -unsafe impl Send for HostV6 {} -unsafe impl Sync for HostV6 {} +unsafe impl Send for HostV7 {} +unsafe impl Sync for HostV7 {} struct Completion { host: HostV4, @@ -285,7 +285,7 @@ impl CompletionRef { #[derive(Clone, Copy)] struct StreamRef { - host: HostV6, + host: HostV7, raw: *const NemoRelayNativeAsyncStream, } @@ -960,7 +960,7 @@ type StreamAdapter = dyn Fn(Json, LlmExecutionContext<'static>, LlmStreamNext) -> StreamFuture + Send + Sync; struct StreamCallbackState { - host: HostV6, + host: HostV7, executor: Arc, adapter: Box, } @@ -972,7 +972,7 @@ unsafe extern "C" fn drop_stream_callback(user_data: *mut c_void) { } struct OutputStream { - host: HostV6, + host: HostV7, raw: *const NemoRelayNativeAsyncStream, } @@ -981,18 +981,19 @@ unsafe impl Sync for OutputStream {} impl OutputStream { fn cancelled(&self) -> bool { - unsafe { (self.host.0.v5.v4.v3.async_stream_is_cancelled)(self.raw) } + unsafe { (self.host.0.v6.v5.v4.v3.async_stream_is_cancelled)(self.raw) } } async fn push(&self, value: &Json) -> Result<()> { - let value = HostString::from_json(&self.host.0.v5.v4.v3.v1, value) + let value = HostString::from_json(&self.host.0.v6.v5.v4.v3.v1, value) .ok_or_else(|| "failed to serialize native stream chunk".to_string())?; loop { if self.cancelled() { return Err("native stream consumer cancelled".into()); } - let status = - unsafe { (self.host.0.v5.v4.v3.async_stream_push_json)(self.raw, value.as_ptr()) }; + let status = unsafe { + (self.host.0.v6.v5.v4.v3.async_stream_push_json)(self.raw, value.as_ptr()) + }; match status { NemoRelayStatus::Ok => return Ok(()), NemoRelayStatus::Backpressured => { @@ -1005,19 +1006,20 @@ impl OutputStream { fn finish(&self) -> Result<()> { status_result( - unsafe { (self.host.0.v5.v4.v3.async_stream_finish)(self.raw) }, + unsafe { (self.host.0.v6.v5.v4.v3.async_stream_finish)(self.raw) }, "finish native stream", ) } async fn reject(&self, error: &str) { - if let Some(error) = HostString::new(&self.host.0.v5.v4.v3.v1, error) { + if let Some(error) = HostString::new(&self.host.0.v6.v5.v4.v3.v1, error) { loop { if self.cancelled() { break; } - let status = - unsafe { (self.host.0.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()) }; + let status = unsafe { + (self.host.0.v6.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()) + }; match status { NemoRelayStatus::Backpressured => { tokio::time::sleep(CANCELLATION_POLL_INTERVAL).await; @@ -1029,9 +1031,9 @@ impl OutputStream { } fn reject_once(&self, error: &str) { - if let Some(error) = HostString::new(&self.host.0.v5.v4.v3.v1, error) { + if let Some(error) = HostString::new(&self.host.0.v6.v5.v4.v3.v1, error) { unsafe { - (self.host.0.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()); + (self.host.0.v6.v5.v4.v3.async_stream_reject)(self.raw, error.as_ptr()); } } } @@ -1039,7 +1041,7 @@ impl OutputStream { impl Drop for OutputStream { fn drop(&mut self) { - unsafe { (self.host.0.v5.v4.v3.async_stream_release)(self.raw) }; + unsafe { (self.host.0.v6.v5.v4.v3.async_stream_release)(self.raw) }; } } @@ -1060,11 +1062,11 @@ unsafe extern "C" fn stream_trampoline( return NemoRelayNativeAsyncCallbackState::Pending as u32; } let next = LlmStreamNext(Arc::new(NextInner { - host: HostV4(state.host.0.v5.v4), + host: HostV4(state.host.0.v6.v5.v4), raw: next, })); let invocation = read_json_value( - &state.host.0.v5.v4.v3.v1, + &state.host.0.v6.v5.v4.v3.v1, invocation_json, "stream invocation", ) @@ -1080,8 +1082,8 @@ unsafe extern "C" fn stream_trampoline( context, ) }); - let bindings = ScopePollBinding::capture(state.host.0.v5.v4.v3.v1).and_then(|future| { - ScopePollBinding::capture(state.host.0.v5.v4.v3.v1).map(|stream| (future, stream)) + let bindings = ScopePollBinding::capture(state.host.0.v6.v5.v4.v3.v1).and_then(|future| { + ScopePollBinding::capture(state.host.0.v6.v5.v4.v3.v1).map(|stream| (future, stream)) }); let future = catch_unwind(AssertUnwindSafe(|| match (invocation, context) { (Ok(invocation), Ok(context)) => (state.adapter)(invocation, context, next), @@ -1092,7 +1094,7 @@ unsafe extern "C" fn stream_trampoline( }); if let Err(error) = state.executor.ensure_started() { output.reject_once(&error); - set_last_error(&state.host.0.v5.v4.v3.v1, &error); + set_last_error(&state.host.0.v6.v5.v4.v3.v1, &error); return NemoRelayNativeAsyncCallbackState::Pending as u32; } let task = async move { @@ -1154,7 +1156,7 @@ unsafe extern "C" fn stream_trampoline( } }; if let Err(error) = state.executor.spawn(task) { - set_last_error(&state.host.0.v5.v4.v3.v1, &error); + set_last_error(&state.host.0.v6.v5.v4.v3.v1, &error); } NemoRelayNativeAsyncCallbackState::Pending as u32 } @@ -1248,7 +1250,7 @@ fn llm_stream_execution_context_from_native( return Err("native LLM stream execution context exposed a response codec".into()); } let request = context.request_codec; - let host = &stream.host.0.v5.v4.v3.v1; + let host = &stream.host.0.v6.v5.v4.v3.v1; let request_codec = stream.execution_request_context( execution_codec_identity(host, request.codec_kind, request.codec_id)?, !request.codec.is_null(), @@ -1271,14 +1273,14 @@ impl PluginContext<'_> { })) } - fn host_v6(&self) -> Result { + fn host_v7(&self) -> Result { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { - return Err("typed LLM execution middleware requires Relay ABI v6".into()); + return Err("typed LLM execution middleware requires Relay ABI v7".into()); } - Ok(HostV6(unsafe { - *(self.host as *const _ as *const NemoRelayNativeHostApiV6) + Ok(HostV7(unsafe { + *(self.host as *const _ as *const NemoRelayNativeHostApiV7) })) } @@ -1800,7 +1802,7 @@ impl PluginContext<'_> { { let callback = Arc::new(callback); let state = Box::into_raw(Box::new(StreamCallbackState { - host: self.host_v6()?, + host: self.host_v7()?, executor: Arc::clone(&self.executor), adapter: Box::new(move |value, context, next| { let callback = Arc::clone(&callback); diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index f98d4f8cc..ff51526d7 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -50,15 +50,14 @@ use serde_json::Map; /// Native plugin ABI version supported by this crate. /// -/// Version 6 makes LLM execution intercept callbacks context-aware and adds -/// host-routed operational logging. +/// Version 7 makes LLM execution intercept callbacks context-aware. /// /// This is an intentional callback-layout break. Native plugins must rebuild /// against this Relay release even though their authored `native_api` /// compatibility label remains `1`. -pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 6; +pub const NEMO_RELAY_NATIVE_ABI_VERSION: u32 = 7; /// ABI version that introduced uniform LLM execution codec context. -pub const NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT: u32 = 6; +pub const NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT: u32 = 7; /// ABI version that introduced host-routed operational logging. pub const NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING: u32 = 6; /// ABI version that introduced context-aware raw tool execution intercepts. @@ -368,7 +367,7 @@ enum LlmExecutionRequestCodecOwner { completion: *const NemoRelayNativeAsyncCompletion, }, Stream { - host: NemoRelayNativeHostApiV6, + host: NemoRelayNativeHostApiV7, stream: *const NemoRelayNativeAsyncStream, }, } @@ -391,7 +390,7 @@ impl Drop for LlmExecutionRequestCodec<'_> { (host.v3.async_completion_release)(completion) }, LlmExecutionRequestCodecOwner::Stream { host, stream } => unsafe { - (host.v5.v4.v3.async_stream_release)(stream) + (host.v6.v5.v4.v3.async_stream_release)(stream) }, } } @@ -414,12 +413,12 @@ impl LlmExecutionRequestCodec<'_> { }) } LlmExecutionRequestCodecOwner::Stream { host, stream } => { - native_codec_call(&host.v5.v4.v3.v1, |out| unsafe { - let request = HostString::from_json(&host.v5.v4.v3.v1, request) + native_codec_call(&host.v6.v5.v4.v3.v1, |out| unsafe { + let request = HostString::from_json(&host.v6.v5.v4.v3.v1, request) .ok_or_else(|| "failed to serialize LLM request".to_string())?; let status = (host.async_stream_llm_request_codec_decode)(stream, request.as_ptr(), out); - codec_status(&host.v5.v4.v3.v1, status) + codec_status(&host.v6.v5.v4.v3.v1, status) }) } } @@ -448,10 +447,10 @@ impl LlmExecutionRequestCodec<'_> { }) } LlmExecutionRequestCodecOwner::Stream { host, stream } => { - native_codec_call(&host.v5.v4.v3.v1, |out| unsafe { - let annotated = HostString::from_json(&host.v5.v4.v3.v1, annotated) + native_codec_call(&host.v6.v5.v4.v3.v1, |out| unsafe { + let annotated = HostString::from_json(&host.v6.v5.v4.v3.v1, annotated) .ok_or_else(|| "failed to serialize annotated request".to_string())?; - let original = HostString::from_json(&host.v5.v4.v3.v1, original) + let original = HostString::from_json(&host.v6.v5.v4.v3.v1, original) .ok_or_else(|| "failed to serialize original request".to_string())?; let status = (host.async_stream_llm_request_codec_encode)( stream, @@ -459,7 +458,7 @@ impl LlmExecutionRequestCodec<'_> { original.as_ptr(), out, ); - codec_status(&host.v5.v4.v3.v1, status) + codec_status(&host.v6.v5.v4.v3.v1, status) }) } } @@ -1620,8 +1619,7 @@ pub struct NemoRelayNativeHostApiV5 { -> NemoRelayStatus, } -/// ABI-v6 host extension for operational logging and context-aware LLM -/// execution intercept callbacks. +/// ABI-v6 host extension for operational logging. /// /// The complete ABI-v5 table is the prefix, preserving layout compatibility. #[repr(C)] @@ -1636,6 +1634,19 @@ pub struct NemoRelayNativeHostApiV6 { message: *const NemoRelayNativeString, fields_json: *const NemoRelayNativeString, ) -> NemoRelayStatus, +} + +/// ABI-v7 host table for context-aware LLM execution intercept callbacks. +/// +/// The inherited function-pointer table has the same fields as ABI v6, but +/// its LLM execution callback typedefs include +/// [`NemoRelayNativeLlmExecutionContext`]. The distinct table version prevents +/// either side from invoking a callback compiled with the old argument layout. +#[repr(C)] +#[derive(Clone, Copy)] +pub struct NemoRelayNativeHostApiV7 { + /// ABI-v6 table compiled with the ABI-v7 callback typedefs. + pub v6: NemoRelayNativeHostApiV6, /// Registers a completion-based asynchronous LLM execution intercept. pub plugin_context_register_async_llm_execution_intercept: unsafe extern "C" fn( @@ -1678,10 +1689,12 @@ unsafe impl Sync for NemoRelayNativeHostApiV4 {} // same thread-safe host function table. unsafe impl Send for NemoRelayNativeHostApiV5 {} unsafe impl Sync for NemoRelayNativeHostApiV5 {} -// SAFETY: the v6 host table is immutable and contains only thread-safe host -// functions. +// SAFETY: the v6 host table is immutable and its log function is thread-safe. unsafe impl Send for NemoRelayNativeHostApiV6 {} unsafe impl Sync for NemoRelayNativeHostApiV6 {} +// SAFETY: the v7 table is immutable and contains only thread-safe host functions. +unsafe impl Send for NemoRelayNativeHostApiV7 {} +unsafe impl Sync for NemoRelayNativeHostApiV7 {} // The host API table is immutable after construction. Function pointers and // the null-terminated version string pointer are safe to share across threads. @@ -3177,7 +3190,7 @@ impl<'a> PluginContext<'a> { /// `cb`, `user_data`, and `free_fn` must remain valid for every host /// callback invocation until the host deregisters the callback or calls /// `free_fn`. `free_fn` must match the allocation behind `user_data`. - /// Registration consumes that ownership even when ABI v6 is unavailable. + /// Registration consumes that ownership even when ABI v7 is unavailable. pub unsafe fn register_llm_execution_intercept_raw( &mut self, name: &str, @@ -3187,7 +3200,7 @@ impl<'a> PluginContext<'a> { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; @@ -3207,7 +3220,7 @@ impl<'a> PluginContext<'a> { /// `cb`, `user_data`, and `free_fn` must remain valid for every host /// callback invocation until the host deregisters the callback or calls /// `free_fn`. `free_fn` must match the allocation behind `user_data`. - /// Registration consumes that ownership even when ABI v6 is unavailable. + /// Registration consumes that ownership even when ABI v7 is unavailable. pub unsafe fn register_llm_stream_execution_intercept_raw( &mut self, name: &str, @@ -3217,7 +3230,7 @@ impl<'a> PluginContext<'a> { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; @@ -3305,14 +3318,14 @@ impl<'a> PluginContext<'a> { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } return NemoRelayStatus::InvalidArg; } - let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV6) }; + let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV7) }; self.with_name_and_callback(name, user_data, free_fn, |_, name| unsafe { (host.plugin_context_register_async_llm_execution_intercept)( self.raw, name, priority, cb, user_data, free_fn, @@ -3331,7 +3344,7 @@ impl<'a> PluginContext<'a> { /// stream owns the callback lifetime. `next` may be invoked /// repeatedly or concurrently until that stream settles; Relay then /// rejects or cancels unfinished and later calls. This execution-specific - /// callback requires native ABI v6 because its callback context and + /// callback requires native ABI v7 because its callback context and /// stream-scoped request-codec operations are part of that ABI. pub unsafe fn register_async_stream_middleware_raw( &mut self, @@ -3342,16 +3355,17 @@ impl<'a> PluginContext<'a> { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT - || self.host.struct_size < std::mem::size_of::() + || self.host.struct_size < std::mem::size_of::() { if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } return NemoRelayStatus::InvalidArg; } - let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV6) }; + let host = unsafe { &*(self.host as *const _ as *const NemoRelayNativeHostApiV7) }; self.with_name_and_callback(name, user_data, free_fn, |_, name| unsafe { (host + .v6 .v5 .v4 .v3 @@ -3597,11 +3611,16 @@ enum OwnedHostApi { V4(NemoRelayNativeHostApiV4), V5(NemoRelayNativeHostApiV5), V6(NemoRelayNativeHostApiV6), + V7(NemoRelayNativeHostApiV7), } impl OwnedHostApi { unsafe fn copy_from(host: &NemoRelayNativeHostApiV1) -> Self { - if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING + if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT + && host.struct_size >= std::mem::size_of::() + { + Self::V7(unsafe { *(host as *const _ as *const NemoRelayNativeHostApiV7) }) + } else if host.abi_version >= NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING && host.struct_size >= std::mem::size_of::() { Self::V6(unsafe { *(host as *const _ as *const NemoRelayNativeHostApiV6) }) @@ -3629,6 +3648,7 @@ impl OwnedHostApi { Self::V4(host) => &host.v3.v1, Self::V5(host) => &host.v4.v3.v1, Self::V6(host) => &host.v5.v4.v3.v1, + Self::V7(host) => &host.v6.v5.v4.v3.v1, } } } diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index c60fb7ede..22677c040 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -36,21 +36,21 @@ use nemo_relay_plugin::{ NemoRelayNativeAsyncStreamMiddlewareCb, NemoRelayNativeConditionalMiddlewareCb, NemoRelayNativeEventSanitizeCb, NemoRelayNativeEventSubscriberCb, NemoRelayNativeFreeFn, NemoRelayNativeHostApiV1, NemoRelayNativeHostApiV3, NemoRelayNativeHostApiV4, - NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeLlmAsyncStream, - NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, - NemoRelayNativeLlmExecutionContext, NemoRelayNativeLlmExecutionRequestContext, - NemoRelayNativeLlmExecutionResponseContext, NemoRelayNativeLlmRequestCodec, - NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, - NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, - NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, - NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, - NemoRelayNativePluginContext, NemoRelayNativePluginRuntime, NemoRelayNativePluginV1, - NemoRelayNativeScopeHandle, NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, - NemoRelayNativeScopeType, NemoRelayNativeString, NemoRelayNativeToolConditionalCb, - NemoRelayNativeToolExecutionCb, NemoRelayNativeToolExecutionContextCb, - NemoRelayNativeToolJsonCb, NemoRelayNativeWithScopeStackCb, NemoRelayStatus, PendingMarkSpec, - PluginContext, PluginRuntime, ScopeType, ToolExecutionInterceptOutcome, ToolExecutionResult, - ToolNext, + NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, + NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, + NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, + NemoRelayNativeLlmExecutionRequestContext, NemoRelayNativeLlmExecutionResponseContext, + NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, + NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, + NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, + NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, + NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, NemoRelayNativePluginContext, + NemoRelayNativePluginRuntime, NemoRelayNativePluginV1, NemoRelayNativeScopeHandle, + NemoRelayNativeScopeStack, NemoRelayNativeScopeStackBinding, NemoRelayNativeScopeType, + NemoRelayNativeString, NemoRelayNativeToolConditionalCb, NemoRelayNativeToolExecutionCb, + NemoRelayNativeToolExecutionContextCb, NemoRelayNativeToolJsonCb, + NemoRelayNativeWithScopeStackCb, NemoRelayStatus, PendingMarkSpec, PluginContext, + PluginRuntime, ScopeType, ToolExecutionInterceptOutcome, ToolExecutionResult, ToolNext, }; use serde_json::{Map, json}; @@ -488,7 +488,7 @@ static UNAVAILABLE_CONTEXT_GATE_CALLS: AtomicUsize = AtomicUsize::new(0); #[test] fn native_abi_struct_sizes_are_self_describing() { - assert_eq!(NEMO_RELAY_NATIVE_ABI_VERSION, 6); + assert_eq!(NEMO_RELAY_NATIVE_ABI_VERSION, 7); assert_eq!( size_of::(), test_host().struct_size @@ -560,10 +560,10 @@ fn assert_native_abi_platform_layout() { ), 600 ); - assert_type_layout::(8, 648); + assert_type_layout::(8, 616); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, log), 608); - assert_native_abi_v6_execution_layout(8, 648, 616, 624, 632, 640); + assert_native_abi_v7_layout(8, 648, 616, 624, 632, 640); assert_type_layout::(8, 56); assert_eq!(plugin_offsets(), [0, 8, 16, 24, 32, 40, 48]); assert_type_layout::(8, 40); @@ -624,10 +624,10 @@ fn assert_native_abi_platform_layout() { ), 296 ); - assert_type_layout::(4, 320); + assert_type_layout::(4, 304); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!(offset_of!(NemoRelayNativeHostApiV6, log), 300); - assert_native_abi_v6_execution_layout(4, 320, 304, 308, 312, 316); + assert_native_abi_v7_layout(4, 320, 304, 308, 312, 316); assert_type_layout::(4, 28); assert_eq!(plugin_offsets(), [0, 4, 8, 12, 16, 20, 24]); assert_type_layout::(4, 20); @@ -639,7 +639,7 @@ fn assert_type_layout(expected_alignment: usize, expected_size: usize) { assert_eq!(size_of::(), expected_size); } -fn assert_native_abi_v6_execution_layout( +fn assert_native_abi_v7_layout( expected_alignment: usize, expected_size: usize, registration_offset: usize, @@ -647,28 +647,29 @@ fn assert_native_abi_v6_execution_layout( decode_offset: usize, encode_offset: usize, ) { - assert_type_layout::(expected_alignment, expected_size); + assert_type_layout::(expected_alignment, expected_size); + assert_eq!(offset_of!(NemoRelayNativeHostApiV7, v6), 0); assert_eq!( offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, plugin_context_register_async_llm_execution_intercept ), registration_offset ); assert_eq!( - offset_of!(NemoRelayNativeHostApiV6, async_stream_retain), + offset_of!(NemoRelayNativeHostApiV7, async_stream_retain), retain_offset ); assert_eq!( offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, async_stream_llm_request_codec_decode ), decode_offset ); assert_eq!( offset_of!( - NemoRelayNativeHostApiV6, + NemoRelayNativeHostApiV7, async_stream_llm_request_codec_encode ), encode_offset @@ -730,7 +731,7 @@ fn native_abi_v5_extension_is_append_only() { } #[test] -fn native_abi_v6_extension_is_append_only() { +fn native_abi_v6_logging_extension_is_append_only() { assert_eq!(offset_of!(NemoRelayNativeHostApiV6, v5), 0); assert_eq!( offset_of!(NemoRelayNativeHostApiV6, log), @@ -3013,6 +3014,15 @@ fn test_host_v6() -> NemoRelayNativeHostApiV6 { NemoRelayNativeHostApiV6 { v5, log: capture_plugin_log, + } +} + +fn test_host_v7() -> NemoRelayNativeHostApiV7 { + let mut v6 = test_host_v6(); + v6.v5.v4.v3.v1.abi_version = NEMO_RELAY_NATIVE_ABI_VERSION; + v6.v5.v4.v3.v1.struct_size = size_of::(); + NemoRelayNativeHostApiV7 { + v6, plugin_context_register_async_llm_execution_intercept: capture_register_async_llm_execution, async_stream_retain: capture_async_stream_retain, async_stream_llm_request_codec_decode: capture_async_stream_request_decode, @@ -4382,9 +4392,9 @@ fn typed_subscriber_registration_decodes_events() { #[allow(clippy::cognitive_complexity)] // One table-style test deliberately exercises every surface. fn typed_async_middleware_registers_and_round_trips_every_surface() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v4 = &host.v5.v4; - let mut ctx = test_context(&host.v5.v4.v3.v1); + let host = test_host_v7(); + let host_v4 = &host.v6.v5.v4; + let mut ctx = test_context(&host.v6.v5.v4.v3.v1); ctx.register_mark_sanitize_guardrail("mark-async", 1, |_event, mut fields| async move { tokio::time::sleep(Duration::from_millis(1)).await; @@ -4789,8 +4799,8 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { #[test] fn typed_async_unary_execution_codecs_expire_after_completion_settles() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v4 = &host.v5.v4; + let host = test_host_v7(); + let host_v4 = &host.v6.v5.v4; let mut ctx = test_context(&host_v4.v3.v1); let (context_tx, context_rx) = std::sync::mpsc::sync_channel(1); ctx.register_llm_execution_intercept( @@ -4871,8 +4881,8 @@ fn typed_async_unary_execution_codecs_expire_after_completion_settles() { #[test] fn typed_async_stream_execution_codec_expires_after_stream_finishes() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut ctx = test_context(host_v1); let (context_tx, context_rx) = std::sync::mpsc::sync_channel(1); ctx.register_llm_stream_execution_intercept( @@ -5078,8 +5088,8 @@ fn typed_async_llm_sanitize_context_rejects_unknown_builtin_identity() { #[test] fn typed_async_registration_failure_rolls_back_callback_state() { let _guard = begin_test(); - let host = test_host_v6(); - let mut ctx = test_context(&host.v5.v4.v3.v1); + let host = test_host_v7(); + let mut ctx = test_context(&host.v6.v5.v4.v3.v1); *REGISTRATION_STATUS.lock().unwrap() = NemoRelayStatus::InvalidArg; let unary_drops = Arc::new(AtomicUsize::new(0)); @@ -5175,8 +5185,8 @@ fn typed_async_callbacks_isolate_errors_panics_and_invalid_input() { #[test] fn typed_async_continuations_are_concurrent_and_executor_owned() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v4 = &host.v5.v4; + let host = test_host_v7(); + let host_v4 = &host.v6.v5.v4; let mut ctx = test_context(&host_v4.v3.v1); ctx.register_tool_execution_intercept("concurrent", 0, |context, next| async move { assert_eq!(context.tool_name, "tool"); @@ -5494,8 +5504,8 @@ fn typed_async_executor_drop_inside_tokio_runtime_drains_accepted_tasks() { #[test] fn typed_async_stream_cancellation_while_polling_releases_output() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut ctx = test_context(host_v1); let started = Arc::new(AtomicBool::new(false)); ctx.register_llm_stream_execution_intercept("cancel-poll", 0, { @@ -5547,8 +5557,8 @@ fn typed_async_stream_cancellation_while_polling_releases_output() { #[test] fn typed_async_stream_restores_callback_scope_while_polling_returned_stream() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-scope", @@ -5595,8 +5605,8 @@ fn typed_async_stream_restores_callback_scope_while_polling_returned_stream() { #[test] fn typed_async_stream_rejects_item_errors_and_releases_output() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-error", @@ -5647,8 +5657,8 @@ fn typed_async_stream_rejects_item_errors_and_releases_output() { #[test] fn typed_async_stream_rejects_poll_panics_and_releases_output() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-panic", @@ -5703,8 +5713,8 @@ fn typed_async_stream_rejects_poll_panics_and_releases_output() { #[test] fn typed_async_stream_propagates_downstream_pull_errors() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-downstream-error", @@ -5753,8 +5763,8 @@ fn typed_async_stream_propagates_downstream_pull_errors() { #[test] fn typed_async_stream_rejects_missing_continuation() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut ctx = test_context(host_v1); ctx.register_llm_stream_execution_intercept( "stream-null-next", @@ -5885,8 +5895,8 @@ fn raw_event_sanitize_registrations_cover_every_surface() { #[test] fn raw_callback_registrations_preserve_every_middleware_shape() { let _guard = begin_test(); - let host = test_host_v6(); - let mut ctx = test_context(&host.v5.v4.v3.v1); + let host = test_host_v7(); + let mut ctx = test_context(&host.v6.v5.v4.v3.v1); unsafe { assert_eq!( @@ -6028,8 +6038,8 @@ fn raw_callback_registrations_preserve_every_middleware_shape() { #[test] fn raw_async_callback_registrations_use_the_versioned_extension_tables() { let _guard = begin_test(); - let host = test_host_v6(); - let mut ctx = test_context(&host.v5.v4.v3.v1); + let host = test_host_v7(); + let mut ctx = test_context(&host.v6.v5.v4.v3.v1); assert_eq!( unsafe { @@ -6584,8 +6594,8 @@ fn exported_plugin_default_validate_returns_empty_diagnostics() { #[test] fn exported_plugin_register_installs_callbacks_and_propagates_errors() { let _guard = begin_test(); - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut plugin = NemoRelayNativePluginV1::default(); assert_eq!( @@ -6724,8 +6734,8 @@ fn exported_entry_symbol_rejects_prior_host_versions() { ); } - let host = test_host_v6(); - let host_v1 = &host.v5.v4.v3.v1; + let host = test_host_v7(); + let host_v1 = &host.v6.v5.v4.v3.v1; let mut plugin = NemoRelayNativePluginV1::default(); assert_eq!( unsafe { constructor_counting_entry(host_v1, &mut plugin) }, diff --git a/docs/about-nemo-relay/release-notes/index.mdx b/docs/about-nemo-relay/release-notes/index.mdx index 66ddbd60c..64b940aed 100644 --- a/docs/about-nemo-relay/release-notes/index.mdx +++ b/docs/about-nemo-relay/release-notes/index.mdx @@ -44,7 +44,7 @@ interceptors receive request codec access only. **Breaking change:** Callback signatures change across Rust, Python, Node.js, Go, C, native plugins, and Rust and Python gRPC workers. Python language-binding streaming and public C callbacks also gain the logical LLM -name. Native plugins must rebuild for the finalized internal ABI v6 layout, and affected workers +name. Native plugins must rebuild for the internal ABI v7 layout, and affected workers must regenerate their protobuf bindings and rebuild. Authored compatibility labels remain `native_api = "1"` and `grpc-v1`; plugin manifests must use a Relay range that begins at 0.10 or otherwise excludes 0.9. Refer to the diff --git a/docs/build-plugins/about.mdx b/docs/build-plugins/about.mdx index aaba58373..352434c3d 100644 --- a/docs/build-plugins/about.mdx +++ b/docs/build-plugins/about.mdx @@ -56,7 +56,7 @@ artifact solves a concrete operational problem. Native Rust plugins suit reusable middleware whose callback latency or throughput is important enough to justify platform-specific binaries and full in-process trust. A native plugin uses manifest compatibility `compat.native_api = "1"`; the current SDK -uses C host-table ABI v6. Relay 0.10 rejects v2-v5 compiled callback layouts, so native +uses C host-table ABI v7. Relay 0.10 rejects v2-v6 compiled callback layouts, so native plugins must rebuild for this release even though the authored manifest label remains `native_api = "1"`. Those are different version axes, as the [Native ABI Reference](/build-plugins/native/native-abi-reference) diff --git a/docs/build-plugins/native/about.mdx b/docs/build-plugins/native/about.mdx index 614407b9b..5a016d63f 100644 --- a/docs/build-plugins/native/about.mdx +++ b/docs/build-plugins/native/about.mdx @@ -26,14 +26,14 @@ Three version values answer different questions: |---|---|---| | Package manifest | `manifest_version = 1` | Shape of the authored `relay-plugin.toml` file. | | Manifest native API | `compat.native_api = "1"` | Native plugin package contract accepted by discovery and trust validation. | -| C host-table ABI | v6 | Function table required by the current `nemo-relay-plugin` SDK. Relay 0.10 rejects v2-v5 callback layouts. | +| C host-table ABI | v7 | Function table required by the current `nemo-relay-plugin` SDK. Relay 0.10 rejects v2-v6 callback layouts. | The checked example registers a tool execution intercept, so it declares `compat.relay = ">=0.10.0,<1.0"`. Every native plugin built with this SDK must use -that lower bound because an older Relay cannot load the finalized ABI v6 function table, even +that lower bound because an older Relay cannot load the ABI v7 function table, even when the component itself does not register an LLM execution intercept. -Relay 0.10 finalizes ABI v6 with directional codec context on every LLM +Relay 0.10 adds ABI v7 with directional codec context on every LLM execution callback. Rebuild all native plugins for this release; the authored `native_api = "1"` label does not change. diff --git a/docs/build-plugins/native/native-abi-reference.mdx b/docs/build-plugins/native/native-abi-reference.mdx index b8639036a..aa84eac0b 100644 --- a/docs/build-plugins/native/native-abi-reference.mdx +++ b/docs/build-plugins/native/native-abi-reference.mdx @@ -23,19 +23,19 @@ extern "C" fn nemo_relay_register_plugin( ) -> NemoRelayStatus ``` -The current host requires ABI v6. It does not fall back to v2-v5 tables because the +The current host requires ABI v7. It does not fall back to v6 or earlier tables because the LLM execution callback layouts changed; invoking a stale callback through the new layout would be unsafe. Rebuild every native plugin for Relay 0.10 and raise its `compat.relay` lower bound to `0.10.0`. The authored manifest label remains `compat.native_api = "1"`. -ABI v6 adds host-routed operational logging and directional request and unary-response -codec context to raw, typed, and asynchronous LLM execution callbacks. The generic -asynchronous middleware callback remains unchanged; ABI v6 appends an execution-specific -unary registration, and the existing execution-specific stream callback gains the context -parameter. Native plugins pass a log level, optional target, message, and optional JSON-object -fields; Relay applies its operational logging policy and preserves those fields in structured -JSONL output. +ABI v7 adds directional request and unary-response codec context to raw, typed, +and asynchronous LLM execution callbacks. The generic asynchronous middleware +callback remains unchanged; ABI v7 appends an execution-specific unary registration, +and the existing execution-specific stream callback gains the context parameter. +ABI v6 adds a host-routed operational logging function. Native plugins pass a level, +optional target, message, and optional JSON-object fields; Relay applies its operational +logging policy and preserves those fields in structured JSONL output. ABI v4 extends the complete v3 prefix with completion-scoped [codecs](/about-nemo-relay/concepts/codecs) and pull-based downstream LLM streams, plus an activation-owned runtime capability for @@ -53,7 +53,8 @@ function signatures and field order are defined by the public | Frozen v3 extension | Completion resolve, reject, cancellation inspection, and release; one-shot completion-coupled continuation invocation; continuation release; generic async middleware registration; bounded output-stream push, finish, reject, cancellation inspection, and release; downstream stream invocation; stream-middleware registration; and repeated or concurrent unary continuation invocation with independent result callbacks. | | Frozen v4 extension | Completion-scoped LLM request decode and encode plus response decode; pull-based downstream LLM stream open, pull, cancel, and release; completion retain for typed codec facades; output-stream backpressure inspection; extended mark emission; runtime diagnostics; activation-owned runtime capability creation, retain, and release; global runtime-registration discovery; owned conditional middleware guardrail registration and deregistration; and activation-owned and runtime-discovered callback gate registration. The callback registration slots are appended after the original constant-reason slots. | | v5 extension | Context-carrying tool execution intercept registration. The callback receives one JSON object containing `tool_name`, `args`, and `tool_call_id`. | -| Current v6 extension | Host-routed operational logging with structured fields; directional codec context for raw and typed LLM execution callbacks; an execution-specific asynchronous unary registration; context on the asynchronous stream-execution callback; and stream retain plus stream-scoped request decode and encode operations. | +| v6 extension | Host-routed operational logging with structured fields. | +| Current v7 extension | Directional codec context for raw and typed LLM execution callbacks, an execution-specific asynchronous unary registration, context on the asynchronous stream-execution callback, and stream retain plus stream-scoped request decode and encode operations. | The prefix and descriptor layout are explicit. A plugin fills the descriptor with its stable kind, component multiplicity, opaque state, callbacks, and destructor. The host @@ -162,7 +163,7 @@ or binding, and restoration. ## Async Completions and Continuations `PluginContext::register_async_middleware_raw` registers non-stream middleware other -than LLM execution. ABI v6 uses the execution-specific +than LLM execution. ABI v7 uses the execution-specific `plugin_context_register_async_llm_execution_intercept` entry for asynchronous unary LLM execution callbacks. Both paths can settle later. Return `Complete` only after resolving or rejecting the completion inside @@ -239,7 +240,7 @@ facades retain their completion and typed streaming request facades retain their output stream, so they remain memory-safe while the callback future or returned stream owns them. Codec calls fail after the completion or stream settles. Raw plugins that need request codec access after an asynchronous stream callback -returns must retain the stream and use the v6 stream-scoped decode and encode +returns must retain the stream and use the v7 stream-scoped decode and encode operations; release that stream reference after the last call. ## Unload Ordering diff --git a/docs/build-plugins/package-discoverable-plugins.mdx b/docs/build-plugins/package-discoverable-plugins.mdx index 9034bb552..4343b0b76 100644 --- a/docs/build-plugins/package-discoverable-plugins.mdx +++ b/docs/build-plugins/package-discoverable-plugins.mdx @@ -26,7 +26,7 @@ equivalent binding configuration. For native packages, `compat.native_api = "1"` is the authored manifest contract. It is not the [C host-table ABI number](/build-plugins/native/native-abi-reference). The -current SDK requires ABI v6, and Relay 0.10 rejects v2-v5 compiled callback layouts. +current SDK requires ABI v7, and Relay 0.10 rejects v2-v6 compiled callback layouts. For workers, declare `compat.worker_protocol = "grpc-v1"`; the [handshake](/build-plugins/workers/grpc-v1-protocol) still negotiates the exact protocol, surfaces, authentication token, and lifecycle at startup. diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index fa3d0416c..5a71d70b1 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -40,7 +40,7 @@ This is a source and binary compatibility break for execution-intercept users: - Rebuild every native plugin against the 0.10 SDK. Native plugins continue to declare `compat.native_api = "1"`, but must set `compat.relay` to `">=0.10.0,<1.0"` or another range that excludes Relay 0.9. Relay 0.10 uses - internal native ABI v6 and rejects v2-v5 compiled layouts. + internal native ABI v7 and rejects v2-v6 compiled layouts. - Regenerate and rebuild gRPC workers that register an LLM execution intercept. The protocol remains `grpc-v1`, but those workers must also exclude Relay 0.9 in `compat.relay`. diff --git a/examples/rust-native-plugin/README.md b/examples/rust-native-plugin/README.md index 372344304..83d00ec7a 100644 --- a/examples/rust-native-plugin/README.md +++ b/examples/rust-native-plugin/README.md @@ -11,10 +11,10 @@ helpers live in separate source modules. Together they register the subscriber, all three event sanitizers, five tool surfaces, and six LLM surfaces exposed by the current typed 0.10.0 SDK. -Relay 0.10 uses native ABI v6. Every LLM execution callback receives directional codec +Relay 0.10 uses native ABI v7. Every LLM execution callback receives directional codec context before its continuation; streaming execution exposes request codec operations but no response decoder. The manifest continues to declare `native_api = "1"`, and its -Relay lower bound is `0.10.0` because Relay 0.9 uses a pre-v6 callback layout. +Relay lower bound is `0.10.0` because Relay 0.9 uses the v5 callback layout. Run the focused tests and build the shared library from this directory. The configuration tests isolate validation and schema contracts. The lifecycle test From 746aa214c1999e53ee560bf5cff41bea2d16c374 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 23 Sep 2026 18:20:46 -0400 Subject: [PATCH 09/22] fix: release expired execution codec leases Signed-off-by: Alex Fournier --- .../src/api/runtime/llm_execution_context.rs | 148 +++++++++++++++--- 1 file changed, 128 insertions(+), 20 deletions(-) diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index 06a4b6b22..a7fd955b0 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -3,8 +3,8 @@ //! Invocation-scoped codec context for LLM execution intercepts. -use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Weak}; use super::callbacks::{LlmRequestCodecContext, LlmResponseCodecContext}; use crate::api::llm::LlmRequest; @@ -16,6 +16,18 @@ use crate::json::Json; const INACTIVE_EXECUTION_CODEC_ERROR: &str = "LLM execution codec capability is no longer active"; +fn inactive_execution_codec_error() -> FlowError { + FlowError::InvalidArgument(INACTIVE_EXECUTION_CODEC_ERROR.into()) +} + +fn upgrade_active_codec(codec: &Weak, gate: &ExecutionCodecGate) -> Result> { + // Upgrade first so a call admitted by the gate keeps the codec alive until + // its synchronous operation completes, even if the lease then expires. + let codec = codec.upgrade().ok_or_else(inactive_execution_codec_error)?; + gate.ensure_active()?; + Ok(codec) +} + #[derive(Debug)] struct ExecutionCodecGate { active: AtomicBool, @@ -32,9 +44,7 @@ impl ExecutionCodecGate { if self.active.load(Ordering::Acquire) { Ok(()) } else { - Err(FlowError::InvalidArgument( - INACTIVE_EXECUTION_CODEC_ERROR.into(), - )) + Err(inactive_execution_codec_error()) } } @@ -46,48 +56,51 @@ impl ExecutionCodecGate { /// Revokes codec capabilities issued to one execution-intercept invocation. pub(crate) struct LlmExecutionCodecLeaseGuard { gate: Arc, + request_codec: Option>, + response_codec: Option>, } impl Drop for LlmExecutionCodecLeaseGuard { fn drop(&mut self) { self.gate.revoke(); + drop(self.request_codec.take()); + drop(self.response_codec.take()); } } struct RevocableRequestCodec { - codec: Arc, + codec: Weak, + identity: super::LlmCodecIdentity, gate: Arc, } impl LlmCodec for RevocableRequestCodec { fn codec_identity(&self) -> super::LlmCodecIdentity { - self.codec.codec_identity() + self.identity.clone() } fn decode(&self, request: &LlmRequest) -> Result { - self.gate.ensure_active()?; - self.codec.decode(request) + upgrade_active_codec(&self.codec, &self.gate)?.decode(request) } fn encode(&self, annotated: &AnnotatedLlmRequest, original: &LlmRequest) -> Result { - self.gate.ensure_active()?; - self.codec.encode(annotated, original) + upgrade_active_codec(&self.codec, &self.gate)?.encode(annotated, original) } } struct RevocableResponseCodec { - codec: Arc, + codec: Weak, + identity: super::LlmCodecIdentity, gate: Arc, } impl LlmResponseCodec for RevocableResponseCodec { fn codec_identity(&self) -> super::LlmCodecIdentity { - self.codec.codec_identity() + self.identity.clone() } fn decode_response(&self, response: &Json) -> Result { - self.gate.ensure_active()?; - self.codec.decode_response(response) + upgrade_active_codec(&self.codec, &self.gate)?.decode_response(response) } } @@ -152,25 +165,33 @@ impl LlmExecutionContext { /// /// The source context retains Relay's selected codecs, but callbacks only /// receive the facades created here. Dropping the returned guard makes all - /// retained facade clones fail without exposing the underlying codec. + /// retained facade clones fail and releases the lease's strong codec + /// references. pub(crate) fn lease(&self) -> (Self, LlmExecutionCodecLeaseGuard) { let gate = Arc::new(ExecutionCodecGate::new()); - let request_codec = match self.request_codec.resolve_codec() { + let leased_request_codec = self.request_codec.resolve_codec(); + let request_codec = match leased_request_codec.as_ref() { Some(codec) => { LlmRequestCodecContext::for_request_codec(Some(Arc::new(RevocableRequestCodec { - codec, + codec: Arc::downgrade(codec), + identity: self.request_codec.codec().clone(), gate: Arc::clone(&gate), }))) } None => LlmRequestCodecContext::with_identity(self.request_codec.codec().clone()), }; + let leased_response_codec = self + .response_codec + .as_ref() + .and_then(LlmResponseCodecContext::resolve_codec); let response_codec = self.response_codec .as_ref() - .map(|context| match context.resolve_codec() { + .map(|context| match leased_response_codec.as_ref() { Some(codec) => LlmResponseCodecContext::for_response_codec(Some(Arc::new( RevocableResponseCodec { - codec, + codec: Arc::downgrade(codec), + identity: context.codec().clone(), gate: Arc::clone(&gate), }, ))), @@ -179,7 +200,11 @@ impl LlmExecutionContext { ( Self::new(request_codec, response_codec), - LlmExecutionCodecLeaseGuard { gate }, + LlmExecutionCodecLeaseGuard { + gate, + request_codec: leased_request_codec, + response_codec: leased_response_codec, + }, ) } @@ -198,3 +223,86 @@ impl LlmExecutionContext { self.response_codec.as_ref() } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::api::runtime::{BuiltinLlmCodec, LlmCodecIdentity}; + + struct DropProbeCodec; + + impl LlmCodec for DropProbeCodec { + fn codec_identity(&self) -> LlmCodecIdentity { + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + } + + fn decode(&self, _request: &LlmRequest) -> Result { + unreachable!("the facade must reject access after lease expiry") + } + + fn encode( + &self, + _annotated: &AnnotatedLlmRequest, + _original: &LlmRequest, + ) -> Result { + unreachable!("the facade must reject access after lease expiry") + } + } + + impl LlmResponseCodec for DropProbeCodec { + fn codec_identity(&self) -> LlmCodecIdentity { + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + } + + fn decode_response(&self, _response: &Json) -> Result { + unreachable!("the facade must reject access after lease expiry") + } + } + + #[test] + fn retained_facades_do_not_keep_backing_codec_alive_after_lease_expiry() { + let backing = Arc::new(DropProbeCodec); + let backing_probe = Arc::downgrade(&backing); + let request_codec: Arc = backing.clone(); + let response_codec: Arc = backing.clone(); + let context = + LlmExecutionContext::for_unary_codecs(Some(request_codec), &Some(response_codec)); + drop(backing); + + let (leased_context, guard) = context.lease(); + let retained_request = leased_context.request_codec().resolve_codec().unwrap(); + let retained_response = leased_context + .response_codec() + .and_then(LlmResponseCodecContext::resolve_codec) + .unwrap(); + drop(leased_context); + drop(context); + + assert!(backing_probe.upgrade().is_some()); + + drop(guard); + + assert!(backing_probe.upgrade().is_none()); + assert_eq!( + retained_request.codec_identity(), + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert_eq!( + retained_response.codec_identity(), + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert!(matches!( + retained_request.decode(&LlmRequest { + headers: serde_json::Map::new(), + content: Json::Null, + }), + Err(FlowError::InvalidArgument(message)) + if message == INACTIVE_EXECUTION_CODEC_ERROR + )); + assert!(matches!( + retained_response.decode_response(&Json::Null), + Err(FlowError::InvalidArgument(message)) + if message == INACTIVE_EXECUTION_CODEC_ERROR + )); + } +} From d9d7c93ad81e9d4efdd3a7435ae749751d684dc3 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 23 Sep 2026 18:20:53 -0400 Subject: [PATCH 10/22] fix: terminate worker streams after errors Signed-off-by: Alex Fournier --- crates/core/src/plugin/dynamic/worker.rs | 8 ++ .../core/tests/unit/dynamic_worker_tests.rs | 118 ++++++++++++++++-- crates/worker/src/lib.rs | 46 ++++--- crates/worker/tests/worker_sdk_tests.rs | 20 ++- 4 files changed, 157 insertions(+), 35 deletions(-) diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index 74fcca2eb..8520fa6e9 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -1721,6 +1721,10 @@ impl Stream for WorkerForwardedLlmStream { } impl LlmStreamInner for WorkerForwardedLlmStream { + fn terminalize(self: Pin<&mut Self>) { + self.get_mut().receiver.take(); + } + fn close(self: Pin<&mut Self>) -> Pin> + Send + '_>> { let this = self.get_mut(); this.receiver.take(); @@ -2132,10 +2136,14 @@ impl WorkerPluginCallback { "worker stream transport failed: {err}" ))), }; + let terminal = result.is_err(); if tx.send(result).await.is_err() { guard.cancel("host stopped consuming the worker stream"); break; } + if terminal { + break; + } } } Err(err) => { diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index decf7316c..4e4ce3f8f 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -1524,7 +1524,12 @@ async fn callback_stream_stops_when_host_receiver_is_dropped() { Box::pin(SignalChunkThenPendingStream { yield_rx, dropped: Some(dropped), - yielded: false, + first_chunk: Some(Ok(StreamChunk { + item: Some(StreamItem::Value( + json_envelope(JSON_SCHEMA, &json!({ "after_receiver_drop": true })) + .expect("test stream chunk should encode"), + )), + })), }) as FakeInvokeStream } }, @@ -1901,6 +1906,65 @@ async fn closing_worker_stream_waits_for_cancellation_and_codec_cleanup() { assert_worker_stream_cancelled_and_cleaned(&mut fixture).await; } +#[tokio::test(flavor = "multi_thread")] +async fn terminal_worker_stream_error_releases_resources_while_host_stream_is_retained() { + enable_operational_logs(); + let mut fixture = pending_worker_stream_with_first_chunk( + "terminal-stream-error", + Ok(StreamChunk { + item: Some(StreamItem::Error(test_worker_error())), + }), + ) + .await; + fixture + .yield_tx + .take() + .expect("yield signal sent once") + .send(()) + .expect("worker stream yield signal should be delivered"); + let error = tokio::time::timeout( + std::time::Duration::from_secs(1), + fixture.stream.as_mut().expect("host stream").next(), + ) + .await + .expect("worker error should arrive promptly") + .expect("worker stream should yield its terminal error") + .expect_err("worker error should remain an error"); + assert!(error.to_string().contains("worker.failed")); + + assert_worker_stream_dropped_and_cleaned(&mut fixture).await; + assert!( + tokio::time::timeout( + std::time::Duration::from_secs(1), + fixture + .stream + .as_mut() + .expect("retained host stream") + .next(), + ) + .await + .expect("host stream should close after its terminal error") + .is_none() + ); + + // Keep the consumer object alive through every cleanup assertion. + drop(fixture.stream.take()); +} + +#[test] +fn terminalizing_forwarded_worker_stream_disconnects_its_producer() { + let (tx, rx) = mpsc::channel(1); + let (_completion_tx, completion_rx) = watch::channel(false); + let mut stream = WorkerForwardedLlmStream { + receiver: Some(tokio_stream::wrappers::ReceiverStream::new(rx)), + completion: completion_rx, + }; + + Pin::new(&mut stream).terminalize(); + + assert!(tx.is_closed()); +} + #[tokio::test] async fn retrying_worker_stream_close_still_waits_after_the_first_future_is_cancelled() { let (_stream_tx, stream_rx) = mpsc::channel(1); @@ -3458,6 +3522,22 @@ struct WorkerStreamLifecycleFixture { } async fn pending_worker_stream_with_codec_context(name: &str) -> WorkerStreamLifecycleFixture { + pending_worker_stream_with_first_chunk( + name, + Ok(StreamChunk { + item: Some(StreamItem::Value( + json_envelope(JSON_SCHEMA, &json!({ "after_receiver_drop": true })) + .expect("test stream chunk should encode"), + )), + }), + ) + .await +} + +async fn pending_worker_stream_with_first_chunk( + name: &str, + first_chunk: std::result::Result, +) -> WorkerStreamLifecycleFixture { let (context_tx, context_rx) = oneshot::channel(); let context_tx = Arc::new(Mutex::new(Some(context_tx))); let (yield_tx, yield_rx) = oneshot::channel(); @@ -3497,7 +3577,7 @@ async fn pending_worker_stream_with_codec_context(name: &str) -> WorkerStreamLif .take() .expect("stream created once"), dropped: worker_stream_dropped_tx.lock().unwrap().take(), - yielded: false, + first_chunk: Some(first_chunk.clone()), }) as FakeInvokeStream } }, @@ -3538,6 +3618,10 @@ async fn assert_worker_stream_cancelled_and_cleaned(fixture: &mut WorkerStreamLi .expect("cancellation channel remains open"); assert_eq!(cancellation.invocation_id, fixture.invocation_id); assert!(cancellation.reason.contains("stopped consuming")); + assert_worker_stream_dropped_and_cleaned(fixture).await; +} + +async fn assert_worker_stream_dropped_and_cleaned(fixture: &mut WorkerStreamLifecycleFixture) { tokio::time::timeout( std::time::Duration::from_secs(1), fixture @@ -3553,12 +3637,30 @@ async fn assert_worker_stream_cancelled_and_cleaned(fixture: &mut WorkerStreamLi &fixture.request_id, &fixture.invocation_id, ); + assert!( + fixture + .callback + .host_state + .continuations + .lock() + .unwrap() + .is_empty() + ); + assert!( + fixture + .callback + .host_state + .scope_stacks + .lock() + .unwrap() + .is_empty() + ); } struct SignalChunkThenPendingStream { yield_rx: oneshot::Receiver<()>, dropped: Option>, - yielded: bool, + first_chunk: Option>, } impl tokio_stream::Stream for SignalChunkThenPendingStream { @@ -3568,18 +3670,12 @@ impl tokio_stream::Stream for SignalChunkThenPendingStream { mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>, ) -> std::task::Poll> { - if self.yielded { + if self.first_chunk.is_none() { return std::task::Poll::Pending; } match Pin::new(&mut self.yield_rx).poll(cx) { std::task::Poll::Ready(_) => { - self.yielded = true; - std::task::Poll::Ready(Some(Ok(StreamChunk { - item: Some(StreamItem::Value( - json_envelope(JSON_SCHEMA, &json!({ "after_receiver_drop": true })) - .expect("test stream chunk should encode"), - )), - }))) + std::task::Poll::Ready(Some(self.first_chunk.take().expect("first chunk"))) } std::task::Poll::Pending => std::task::Poll::Pending, } diff --git a/crates/worker/src/lib.rs b/crates/worker/src/lib.rs index 866cec09d..f6d438417 100644 --- a/crates/worker/src/lib.rs +++ b/crates/worker/src/lib.rs @@ -2157,28 +2157,38 @@ impl PluginWorker for WorkerService { let Some(item) = item else { return; }; - let chunk = match item { - Ok(value) => StreamChunk { - item: Some(nemo_relay_worker_proto::v1::stream_chunk::Item::Value( - match json_envelope(JSON_SCHEMA, &value) { - Ok(value) => value, - Err(err) => { - let _ = - task_tx.send(Err(Status::internal(err.to_string()))).await; - return; - } - }, - )), - }, - Err(err) => StreamChunk { - item: Some(nemo_relay_worker_proto::v1::stream_chunk::Item::Error( - sdk_error_to_worker(err), - )), - }, + let (chunk, terminal) = match item { + Ok(value) => ( + StreamChunk { + item: Some(nemo_relay_worker_proto::v1::stream_chunk::Item::Value( + match json_envelope(JSON_SCHEMA, &value) { + Ok(value) => value, + Err(err) => { + let _ = task_tx + .send(Err(Status::internal(err.to_string()))) + .await; + return; + } + }, + )), + }, + false, + ), + Err(err) => ( + StreamChunk { + item: Some(nemo_relay_worker_proto::v1::stream_chunk::Item::Error( + sdk_error_to_worker(err), + )), + }, + true, + ), }; if task_tx.send(Ok(chunk)).await.is_err() { return; } + if terminal { + return; + } } }); let replaced = self.replace_invocation( diff --git a/crates/worker/tests/worker_sdk_tests.rs b/crates/worker/tests/worker_sdk_tests.rs index b74f1158b..cbf9f9f1b 100644 --- a/crates/worker/tests/worker_sdk_tests.rs +++ b/crates/worker/tests/worker_sdk_tests.rs @@ -1221,7 +1221,7 @@ async fn worker_service_reports_structured_callback_and_payload_errors() { "boom", ); - let stream_err = client + let mut stream_err = client .invoke_stream(Request::new(llm_invoke( "llm-stream-error", RegistrationSurface::LlmStreamExecutionIntercept, @@ -1231,17 +1231,24 @@ async fn worker_service_reports_structured_callback_and_payload_errors() { ))) .await .expect("stream error invoke") - .into_inner() + .into_inner(); + let stream_error_chunk = stream_err .next() .await .expect("stream item") .expect("stream chunk"); - match stream_err.item.expect("stream item") { + match stream_error_chunk.item.expect("stream item") { nemo_relay_worker_proto::v1::stream_chunk::Item::Error(error) => { assert!(error.message.contains("stream boom")); } other => panic!("unexpected stream item: {other:?}"), } + assert!( + tokio::time::timeout(WORKER_TEST_TIMEOUT, stream_err.next()) + .await + .expect("worker stream should terminate after its error") + .is_none() + ); let stream_surface_err = client .invoke_stream(Request::new(tool_invoke( @@ -2517,9 +2524,10 @@ impl WorkerPlugin for SurfacePlugin { }, ); ctx.register_llm_stream_execution_intercept("llm-stream-error", 1, |_, _, _, _| async { - let stream: JsonStream = Box::pin(tokio_stream::iter(vec![Err( - WorkerSdkError::Callback("stream boom".into()), - )])); + let stream: JsonStream = Box::pin( + tokio_stream::once(Err(WorkerSdkError::Callback("stream boom".into()))) + .chain(tokio_stream::pending()), + ); Ok(stream) }); ctx.register_llm_stream_execution_intercept( From c20d11c7aca9984a3f827454c072ec45229b1a71 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 29 Sep 2026 09:32:40 -0700 Subject: [PATCH 11/22] test: cover execution codec facade operations Signed-off-by: Alex Fournier --- crates/plugin/tests/typed_callbacks.rs | 22 ++++++++++++++++--- .../tests/unit/execution_context_tests.rs | 5 +++++ 2 files changed, 24 insertions(+), 3 deletions(-) diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index 22677c040..ecf72bf62 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -4500,9 +4500,20 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { 13, |_name, request, context, next| async move { assert_eq!(context.request_codec().codec, LlmCodecIdentity::Opaque); - assert!(context.request_codec().resolve_codec().is_some()); + let request_codec = context + .request_codec() + .resolve_codec() + .expect("request codec capability"); + let annotated = request_codec.decode(&request)?; + let request = request_codec.encode(&annotated, &request)?; assert!(context.response_codec().is_some()); - next.call(request).await + let response = next.call(request).await?; + let response_codec = context + .response_codec() + .and_then(|codec| codec.resolve_codec()) + .expect("response codec capability"); + let _ = response_codec.decode(&response)?; + Ok(response) }, ) .unwrap(); @@ -4511,7 +4522,12 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { 14, |_name, request, context, next| async move { assert_eq!(context.request_codec().codec, LlmCodecIdentity::Opaque); - assert!(context.request_codec().resolve_codec().is_some()); + let request_codec = context + .request_codec() + .resolve_codec() + .expect("request codec capability"); + let annotated = request_codec.decode(&request)?; + let request = request_codec.encode(&annotated, &request)?; assert!(context.response_codec().is_none()); let stream = next.call(request).await?; let transformed = stream.map(|item| item.map(|chunk| json!({ "wrapped": chunk }))); diff --git a/crates/worker/tests/unit/execution_context_tests.rs b/crates/worker/tests/unit/execution_context_tests.rs index 80769fbee..cf3d81db9 100644 --- a/crates/worker/tests/unit/execution_context_tests.rs +++ b/crates/worker/tests/unit/execution_context_tests.rs @@ -126,6 +126,11 @@ fn execution_context_preserves_codec_identities_and_capability_availability() { .execution_context(&disconnected_runtime(), "invocation", true) .unwrap(); + let debug = format!("{context:?}"); + assert!(debug.contains("request_codec")); + assert!(debug.contains("response_codec")); + assert!(!debug.contains("-request")); + assert!(!debug.contains("-response")); assert_eq!(context.request_codec().codec, expected); assert_eq!(context.request_codec().resolve_codec().is_some(), resolves); let response = context.response_codec().expect("unary response context"); From a007dc6c2d1b7975dccaa28eaa5d4a4e68eb6131 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 29 Sep 2026 10:02:59 -0700 Subject: [PATCH 12/22] test: exercise execution codec lease boundaries Signed-off-by: Alex Fournier --- .../src/api/runtime/llm_execution_context.rs | 36 ++++++++++++------- crates/core/tests/unit/native_plugin_tests.rs | 18 ++++++++++ 2 files changed, 41 insertions(+), 13 deletions(-) diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index a7fd955b0..bd03de6f9 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -229,39 +229,39 @@ mod tests { use super::*; use crate::api::runtime::{BuiltinLlmCodec, LlmCodecIdentity}; - struct DropProbeCodec; + struct LeaseProbeCodec; - impl LlmCodec for DropProbeCodec { + impl LlmCodec for LeaseProbeCodec { fn codec_identity(&self) -> LlmCodecIdentity { LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) } fn decode(&self, _request: &LlmRequest) -> Result { - unreachable!("the facade must reject access after lease expiry") + Ok(AnnotatedLlmRequest::default()) } fn encode( &self, _annotated: &AnnotatedLlmRequest, - _original: &LlmRequest, + original: &LlmRequest, ) -> Result { - unreachable!("the facade must reject access after lease expiry") + Ok(original.clone()) } } - impl LlmResponseCodec for DropProbeCodec { + impl LlmResponseCodec for LeaseProbeCodec { fn codec_identity(&self) -> LlmCodecIdentity { LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) } fn decode_response(&self, _response: &Json) -> Result { - unreachable!("the facade must reject access after lease expiry") + Ok(AnnotatedLlmResponse::default()) } } #[test] - fn retained_facades_do_not_keep_backing_codec_alive_after_lease_expiry() { - let backing = Arc::new(DropProbeCodec); + fn retained_facades_forward_while_active_and_expire_with_their_lease() { + let backing = Arc::new(LeaseProbeCodec); let backing_probe = Arc::downgrade(&backing); let request_codec: Arc = backing.clone(); let response_codec: Arc = backing.clone(); @@ -278,6 +278,19 @@ mod tests { drop(leased_context); drop(context); + let request = LlmRequest { + headers: serde_json::Map::new(), + content: Json::Null, + }; + let annotated = retained_request.decode(&request).unwrap(); + assert_eq!( + retained_request.encode(&annotated, &request).unwrap(), + request + ); + assert_eq!( + retained_response.decode_response(&Json::Null).unwrap(), + AnnotatedLlmResponse::default() + ); assert!(backing_probe.upgrade().is_some()); drop(guard); @@ -292,10 +305,7 @@ mod tests { LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) ); assert!(matches!( - retained_request.decode(&LlmRequest { - headers: serde_json::Map::new(), - content: Json::Null, - }), + retained_request.decode(&request), Err(FlowError::InvalidArgument(message)) if message == INACTIVE_EXECUTION_CODEC_ERROR )); diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index 3aac06e8d..1bd373b76 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -6809,6 +6809,24 @@ fn native_execution_context_is_directional_and_streaming_omits_response_codec() assert!(response.codec_id.is_null()); assert!(response.codec.is_null()); }); + + let allocation_failure = LlmExecutionContext::new( + LlmSanitizeRequestContext::with_identity(LlmCodecIdentity::Runtime("request.v1".into())), + Some(LlmSanitizeResponseContext::with_identity( + LlmCodecIdentity::Runtime("response.v1".into()), + )), + ); + let live_before = native_string_live_allocations(); + fail_native_string_allocation_after(1); + let error = NativeLlmExecutionContextBridge::new(&allocation_failure, None, None) + .err() + .expect("response codec ID allocation should fail"); + assert!( + error + .to_string() + .contains("failed to allocate native LLM codec ID") + ); + assert_eq!(native_string_live_allocations(), live_before); } #[test] From fa453acbaa37ce2bea4c1137b0923803f39ac295 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 29 Sep 2026 17:43:50 -0700 Subject: [PATCH 13/22] refactor: reuse released codec context types Signed-off-by: Alex Fournier --- crates/core/src/api/runtime.rs | 10 ++--- crates/core/src/api/runtime/callbacks.rs | 22 ++++------- .../src/api/runtime/llm_execution_context.rs | 38 +++++++++---------- crates/core/src/plugin/dynamic/native.rs | 10 ++--- crates/ffi/nemo_relay.h | 21 +++------- crates/ffi/src/callable.rs | 17 +++------ crates/node/plugin.d.ts | 2 - crates/node/root-types.d.ts | 14 ++----- crates/plugin/src/lib.rs | 12 ------ crates/plugin/tests/typed_callbacks.rs | 12 +++--- crates/python/src/py_types/core.rs | 4 +- crates/python/src/py_types/mod.rs | 8 ---- crates/worker/src/lib.rs | 28 ++++++-------- go/nemo_relay/callbacks.go | 18 +++------ python/nemo_relay/__init__.py | 4 -- python/nemo_relay/_native.pyi | 11 ++---- .../plugin/src/nemo_relay_plugin/__init__.py | 12 ++---- python/plugin/src/nemo_relay_plugin/_api.py | 21 ++++------ 18 files changed, 93 insertions(+), 171 deletions(-) diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index 9f29d4315..a86713a57 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -14,11 +14,11 @@ pub mod subscriber_dispatcher; pub use callbacks::{ BuiltinLlmCodec, ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmCollectorFn, LlmConditionalFn, LlmExecutionFn, - LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmRequestCodecContext, - LlmRequestInterceptFn, LlmResponseCodecContext, LlmSanitizeRequestContext, - LlmSanitizeRequestFn, LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, - LlmStreamExecutionNextFn, LlmStreamInner, ToolConditionalFn, ToolExecutionContext, - ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmRequestInterceptFn, + LlmSanitizeRequestContext, LlmSanitizeRequestFn, LlmSanitizeResponseContext, + LlmSanitizeResponseFn, LlmStreamExecutionFn, LlmStreamExecutionNextFn, LlmStreamInner, + ToolConditionalFn, ToolExecutionContext, ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, + ToolSanitizeFn, }; #[doc(hidden)] pub use continuation_context::MiddlewareContinuationContext; diff --git a/crates/core/src/api/runtime/callbacks.rs b/crates/core/src/api/runtime/callbacks.rs index 1044e7410..642827dc7 100644 --- a/crates/core/src/api/runtime/callbacks.rs +++ b/crates/core/src/api/runtime/callbacks.rs @@ -252,22 +252,22 @@ pub(crate) type ToolExecutionOutcomeNextFn = Arc< /// The context distinguishes no codec, Relay built-ins, runtime-registered /// codecs, and active codecs with no stable identity. #[derive(Clone, Default)] -pub struct LlmRequestCodecContext { +pub struct LlmSanitizeRequestContext { /// Identity of the codec active for this payload direction. codec: LlmCodecIdentity, request_codec: Option>, } -impl std::fmt::Debug for LlmRequestCodecContext { +impl std::fmt::Debug for LlmSanitizeRequestContext { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter - .debug_struct("LlmRequestCodecContext") + .debug_struct("LlmSanitizeRequestContext") .field("codec", &self.codec) .finish_non_exhaustive() } } -impl LlmRequestCodecContext { +impl LlmSanitizeRequestContext { /// Construct a context that carries only a codec identity. /// /// Identity-only contexts do not carry a codec handle, so @@ -313,22 +313,22 @@ impl LlmRequestCodecContext { /// The context distinguishes no codec, Relay built-ins, runtime-registered /// codecs, and active codecs with no stable identity. #[derive(Clone, Default)] -pub struct LlmResponseCodecContext { +pub struct LlmSanitizeResponseContext { /// Identity of the codec active for this payload direction. codec: LlmCodecIdentity, response_codec: Option>, } -impl std::fmt::Debug for LlmResponseCodecContext { +impl std::fmt::Debug for LlmSanitizeResponseContext { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter - .debug_struct("LlmResponseCodecContext") + .debug_struct("LlmSanitizeResponseContext") .field("codec", &self.codec) .finish_non_exhaustive() } } -impl LlmResponseCodecContext { +impl LlmSanitizeResponseContext { /// Construct a context that carries only a codec identity. /// /// Identity-only contexts do not carry a codec handle, so @@ -369,12 +369,6 @@ impl LlmResponseCodecContext { } } -/// Backward-compatible name for request codec context supplied to sanitizers. -pub type LlmSanitizeRequestContext = LlmRequestCodecContext; - -/// Backward-compatible name for response codec context supplied to sanitizers. -pub type LlmSanitizeResponseContext = LlmResponseCodecContext; - /// Sanitize an LLM request before the runtime records it. /// /// LLM request sanitizers affect the serialized request payload emitted on diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index bd03de6f9..8e6c7278e 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -6,7 +6,7 @@ use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Weak}; -use super::callbacks::{LlmRequestCodecContext, LlmResponseCodecContext}; +use super::callbacks::{LlmSanitizeRequestContext, LlmSanitizeResponseContext}; use crate::api::llm::LlmRequest; use crate::codec::request::AnnotatedLlmRequest; use crate::codec::response::AnnotatedLlmResponse; @@ -123,16 +123,16 @@ impl LlmResponseCodec for RevocableResponseCodec { /// after expiry. #[derive(Clone, Debug, Default)] pub struct LlmExecutionContext { - request_codec: LlmRequestCodecContext, - response_codec: Option, + request_codec: LlmSanitizeRequestContext, + response_codec: Option, } impl LlmExecutionContext { /// Construct an execution context from its directional codec contexts. #[must_use] pub fn new( - request_codec: LlmRequestCodecContext, - response_codec: Option, + request_codec: LlmSanitizeRequestContext, + response_codec: Option, ) -> Self { Self { request_codec, @@ -146,8 +146,8 @@ impl LlmExecutionContext { response_codec: &Option>, ) -> Self { Self::new( - LlmRequestCodecContext::for_request_codec(request_codec), - Some(LlmResponseCodecContext::for_response_codec( + LlmSanitizeRequestContext::for_request_codec(request_codec), + Some(LlmSanitizeResponseContext::for_response_codec( response_codec.clone(), )), ) @@ -156,7 +156,7 @@ impl LlmExecutionContext { /// Construct the context for a streaming managed execution. pub(crate) fn for_streaming_codec(request_codec: Option>) -> Self { Self::new( - LlmRequestCodecContext::for_request_codec(request_codec), + LlmSanitizeRequestContext::for_request_codec(request_codec), None, ) } @@ -171,31 +171,31 @@ impl LlmExecutionContext { let gate = Arc::new(ExecutionCodecGate::new()); let leased_request_codec = self.request_codec.resolve_codec(); let request_codec = match leased_request_codec.as_ref() { - Some(codec) => { - LlmRequestCodecContext::for_request_codec(Some(Arc::new(RevocableRequestCodec { + Some(codec) => LlmSanitizeRequestContext::for_request_codec(Some(Arc::new( + RevocableRequestCodec { codec: Arc::downgrade(codec), identity: self.request_codec.codec().clone(), gate: Arc::clone(&gate), - }))) - } - None => LlmRequestCodecContext::with_identity(self.request_codec.codec().clone()), + }, + ))), + None => LlmSanitizeRequestContext::with_identity(self.request_codec.codec().clone()), }; let leased_response_codec = self .response_codec .as_ref() - .and_then(LlmResponseCodecContext::resolve_codec); + .and_then(LlmSanitizeResponseContext::resolve_codec); let response_codec = self.response_codec .as_ref() .map(|context| match leased_response_codec.as_ref() { - Some(codec) => LlmResponseCodecContext::for_response_codec(Some(Arc::new( + Some(codec) => LlmSanitizeResponseContext::for_response_codec(Some(Arc::new( RevocableResponseCodec { codec: Arc::downgrade(codec), identity: context.codec().clone(), gate: Arc::clone(&gate), }, ))), - None => LlmResponseCodecContext::with_identity(context.codec().clone()), + None => LlmSanitizeResponseContext::with_identity(context.codec().clone()), }); ( @@ -210,7 +210,7 @@ impl LlmExecutionContext { /// Return the request-direction codec identity and revocable capability. #[must_use] - pub fn request_codec(&self) -> &LlmRequestCodecContext { + pub fn request_codec(&self) -> &LlmSanitizeRequestContext { &self.request_codec } @@ -219,7 +219,7 @@ impl LlmExecutionContext { /// Streaming execution returns `None` because Relay does not expose a /// completed-response codec for individual stream chunks. #[must_use] - pub fn response_codec(&self) -> Option<&LlmResponseCodecContext> { + pub fn response_codec(&self) -> Option<&LlmSanitizeResponseContext> { self.response_codec.as_ref() } } @@ -273,7 +273,7 @@ mod tests { let retained_request = leased_context.request_codec().resolve_codec().unwrap(); let retained_response = leased_context .response_codec() - .and_then(LlmResponseCodecContext::resolve_codec) + .and_then(LlmSanitizeResponseContext::resolve_codec) .unwrap(); drop(leased_context); drop(context); diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 1fe534178..c584460e8 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -70,9 +70,9 @@ use nemo_relay_plugin::{ NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, - NemoRelayNativeLlmExecutionRequestContext, NemoRelayNativeLlmExecutionResponseContext, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, - NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, + NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestCodecContext, + NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, + NemoRelayNativeLlmResponseCodecContext, NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, NemoRelayNativePluginContext, @@ -628,7 +628,7 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { &self, callback: impl FnOnce(NemoRelayNativeLlmExecutionContext) -> T, ) -> T { - let request_codec = NemoRelayNativeLlmExecutionRequestContext { + let request_codec = NemoRelayNativeLlmRequestCodecContext { codec_kind: self.request_kind, codec_id: self .request_id @@ -639,7 +639,7 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { }; let response_codec = self.response_kind - .map(|codec_kind| NemoRelayNativeLlmExecutionResponseContext { + .map(|codec_kind| NemoRelayNativeLlmResponseCodecContext { codec_kind, codec_id: self .response_id diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 6cf20a83a..0b9c5b91f 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -299,8 +299,9 @@ typedef char *(*NemoRelayCodecEncodeFn)(void *user_data, const struct FfiLLMRequest *original_request); /** - * Codec identity supplied to an LLM sanitizer. `codec_id` is null for - * `None` and `Opaque`, and is valid only for the duration of the callback. + * Request codec context shared by LLM sanitizer and execution callbacks. + * `codec_id` is null for `None` and `Opaque`, and is valid only for the + * duration of the callback. */ typedef struct NemoRelayLlmSanitizeRequestContext { /** @@ -328,7 +329,7 @@ typedef struct FfiLLMRequest *(*NemoRelayLlmSanitizeRequestCb)(void *user_data, struct NemoRelayLlmSanitizeRequestContext context); /** - * Directional codec context supplied to an LLM response sanitizer. + * Response codec context shared by LLM sanitizer and execution callbacks. */ typedef struct NemoRelayLlmSanitizeResponseContext { /** @@ -381,16 +382,6 @@ typedef NemoRelayStatus (*NemoRelayLlmRequestInterceptCb)(void *user_data, const char *annotated_json, char **out_outcome_json); -/** - * General name for request codec context used by execution intercepts. - */ -typedef struct NemoRelayLlmSanitizeRequestContext NemoRelayLlmRequestCodecContext; - -/** - * General name for response codec context used by execution intercepts. - */ -typedef struct NemoRelayLlmSanitizeResponseContext NemoRelayLlmResponseCodecContext; - /** * Directional codec context supplied to an LLM execution intercept. * @@ -403,11 +394,11 @@ typedef struct NemoRelayLlmExecutionContext { /** * Active request codec identity and capability. */ - NemoRelayLlmRequestCodecContext request_codec; + struct NemoRelayLlmSanitizeRequestContext request_codec; /** * Active unary-response codec context, or null for streaming execution. */ - const NemoRelayLlmResponseCodecContext *response_codec; + const struct NemoRelayLlmSanitizeResponseContext *response_codec; } NemoRelayLlmExecutionContext; /** diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 7c32c129c..7b89273ef 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -177,8 +177,9 @@ pub enum NemoRelayLlmSanitizeCodecKind { Opaque = 3, } -/// Codec identity supplied to an LLM sanitizer. `codec_id` is null for -/// `None` and `Opaque`, and is valid only for the duration of the callback. +/// Request codec context shared by LLM sanitizer and execution callbacks. +/// `codec_id` is null for `None` and `Opaque`, and is valid only for the +/// duration of the callback. #[repr(C)] #[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmSanitizeRequestContext { @@ -190,7 +191,7 @@ pub struct NemoRelayLlmSanitizeRequestContext { pub codec: *const crate::types::FfiLlmSanitizeRequestCodec, } -/// Directional codec context supplied to an LLM response sanitizer. +/// Response codec context shared by LLM sanitizer and execution callbacks. #[repr(C)] #[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmSanitizeResponseContext { @@ -202,12 +203,6 @@ pub struct NemoRelayLlmSanitizeResponseContext { pub codec: *const crate::types::FfiLlmSanitizeResponseCodec, } -/// General name for request codec context used by execution intercepts. -pub type NemoRelayLlmRequestCodecContext = NemoRelayLlmSanitizeRequestContext; - -/// General name for response codec context used by execution intercepts. -pub type NemoRelayLlmResponseCodecContext = NemoRelayLlmSanitizeResponseContext; - /// Directional codec context supplied to an LLM execution intercept. /// /// `request_codec` is always present. `response_codec` is non-null for unary @@ -218,9 +213,9 @@ pub type NemoRelayLlmResponseCodecContext = NemoRelayLlmSanitizeResponseContext; #[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmExecutionContext { /// Active request codec identity and capability. - pub request_codec: NemoRelayLlmRequestCodecContext, + pub request_codec: NemoRelayLlmSanitizeRequestContext, /// Active unary-response codec context, or null for streaming execution. - pub response_codec: *const NemoRelayLlmResponseCodecContext, + pub response_codec: *const NemoRelayLlmSanitizeResponseContext, } /// LLM request sanitizer. It receives the request first and its codec context diff --git a/crates/node/plugin.d.ts b/crates/node/plugin.d.ts index f4a52648b..86d0bb26f 100644 --- a/crates/node/plugin.d.ts +++ b/crates/node/plugin.d.ts @@ -29,9 +29,7 @@ export type { LlmOptimizationModelTransition, LlmOptimizationTokenImpact, LlmOptimizationTokens, - LlmRequestCodecContext, LlmRequestInterceptOutcome, - LlmResponseCodecContext, LlmSanitizeRequestContext, LlmSanitizeResponseContext, } from './index'; diff --git a/crates/node/root-types.d.ts b/crates/node/root-types.d.ts index e635d0d04..b5bd121e6 100644 --- a/crates/node/root-types.d.ts +++ b/crates/node/root-types.d.ts @@ -11,32 +11,26 @@ export type LlmCodecIdentity = | { kind: 'runtime'; id: string } | { kind: 'opaque' }; -/** Codec context available while an LLM request is sanitized. */ +/** Request codec context shared by sanitizer and execution callbacks. */ export interface LlmSanitizeRequestContext { codec: LlmCodecIdentity; /** Resolve the active codec for this callback. Do not retain the result after the callback returns. */ resolveCodec(): import('./typed').LlmCodec | null; } -/** Codec context available while an LLM response is sanitized. */ +/** Response codec context shared by sanitizer and execution callbacks. */ export interface LlmSanitizeResponseContext { codec: LlmCodecIdentity; /** Resolve the active codec for this callback. Do not retain the result after the callback returns. */ resolveCodec(): import('./typed').LlmResponseCodec | null; } -/** General name for request codec context used outside sanitizer callbacks. */ -export type LlmRequestCodecContext = LlmSanitizeRequestContext; - -/** General name for response codec context used outside sanitizer callbacks. */ -export type LlmResponseCodecContext = LlmSanitizeResponseContext; - /** Codec capabilities for one managed LLM execution intercept invocation. */ export interface LlmExecutionContext { /** Request codec identity plus optional decode and encode capability. */ - requestCodec: LlmRequestCodecContext; + requestCodec: LlmSanitizeRequestContext; /** Unary response codec identity plus optional decode capability; `null` for streaming execution. */ - responseCodec: LlmResponseCodecContext | null; + responseCodec: LlmSanitizeResponseContext | null; } /** Schema tag attached to an opaque optimization contribution payload. */ diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index ff51526d7..eac749dd1 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -128,12 +128,6 @@ impl<'a> LlmExecutionContext<'a> { } } -/// Previous name for request codec context on typed execution intercepts. -pub type LlmExecutionRequestContext<'a> = LlmRequestCodecContext<'a>; - -/// Previous name for response codec context on typed execution intercepts. -pub type LlmExecutionResponseContext<'a> = LlmResponseCodecContext<'a>; - /// Status codes returned by stable native ABI functions. #[repr(i32)] #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -252,12 +246,6 @@ pub struct NemoRelayNativeLlmResponseCodecContext { pub codec: *const NemoRelayNativeLlmResponseCodec, } -/// Previous native name for request codec context on execution intercepts. -pub type NemoRelayNativeLlmExecutionRequestContext = NemoRelayNativeLlmRequestCodecContext; - -/// Previous native name for response codec context on execution intercepts. -pub type NemoRelayNativeLlmExecutionResponseContext = NemoRelayNativeLlmResponseCodecContext; - /// Codec context passed to native LLM execution intercept callbacks. #[repr(C)] #[derive(Debug, Clone, Copy)] diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index ecf72bf62..e549d353b 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -39,9 +39,9 @@ use nemo_relay_plugin::{ NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, - NemoRelayNativeLlmExecutionRequestContext, NemoRelayNativeLlmExecutionResponseContext, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestInterceptCb, - NemoRelayNativeLlmResponseCodec, NemoRelayNativeLlmSanitizeRequestCb, + NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestCodecContext, + NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, + NemoRelayNativeLlmResponseCodecContext, NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, NemoRelayNativePluginContext, @@ -3108,7 +3108,7 @@ fn invoke_async_llm_execution_registration( } struct TestExecutionContext { - response: Option>, + response: Option>, context: NemoRelayNativeLlmExecutionContext, } @@ -3116,14 +3116,14 @@ impl TestExecutionContext { fn new(unary: bool) -> Self { let request_codec = NonNull::::dangling().as_ptr(); let response = unary.then(|| { - Box::new(NemoRelayNativeLlmExecutionResponseContext { + Box::new(NemoRelayNativeLlmResponseCodecContext { codec_kind: NemoRelayNativeLlmCodecKind::Opaque, codec_id: ptr::null(), codec: NonNull::::dangling().as_ptr(), }) }); let context = NemoRelayNativeLlmExecutionContext { - request_codec: NemoRelayNativeLlmExecutionRequestContext { + request_codec: NemoRelayNativeLlmRequestCodecContext { codec_kind: NemoRelayNativeLlmCodecKind::Opaque, codec_id: ptr::null(), codec: request_codec, diff --git a/crates/python/src/py_types/core.rs b/crates/python/src/py_types/core.rs index d514d506e..cdacfc2e9 100644 --- a/crates/python/src/py_types/core.rs +++ b/crates/python/src/py_types/core.rs @@ -57,7 +57,7 @@ impl PyLlmCodecIdentity { } } -/// Structured per-call context delivered to LLM request sanitizer callbacks. +/// Per-call request codec context shared by sanitizer and execution callbacks. #[pyclass(name = "LlmSanitizeRequestContext", frozen)] pub struct PyLlmSanitizeRequestContext { pub(crate) inner: LlmSanitizeRequestContext, @@ -81,7 +81,7 @@ impl PyLlmSanitizeRequestContext { } } -/// Structured per-call context delivered to LLM response sanitizer callbacks. +/// Per-call response codec context shared by sanitizer and execution callbacks. #[pyclass(name = "LlmSanitizeResponseContext", frozen)] pub struct PyLlmSanitizeResponseContext { pub(crate) inner: LlmSanitizeResponseContext, diff --git a/crates/python/src/py_types/mod.rs b/crates/python/src/py_types/mod.rs index c6dc34be1..6e56ec7f2 100644 --- a/crates/python/src/py_types/mod.rs +++ b/crates/python/src/py_types/mod.rs @@ -156,14 +156,6 @@ fn register_llm_types(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; - m.add( - "LlmRequestCodecContext", - m.getattr("LlmSanitizeRequestContext")?, - )?; - m.add( - "LlmResponseCodecContext", - m.getattr("LlmSanitizeResponseContext")?, - )?; m.add_class::()?; m.add_class::()?; m.add_class::()?; diff --git a/crates/worker/src/lib.rs b/crates/worker/src/lib.rs index f6d438417..f87cf7a55 100644 --- a/crates/worker/src/lib.rs +++ b/crates/worker/src/lib.rs @@ -297,7 +297,7 @@ impl ToolExecutionContext { /// Active codec identity and capability for an LLM request. #[derive(Clone)] -pub struct LlmRequestCodecContext { +pub struct LlmSanitizeRequestContext { /// Identity of the active codec. pub codec: LlmCodecIdentity, runtime: Option, @@ -307,7 +307,7 @@ pub struct LlmRequestCodecContext { /// Active codec identity and capability for an LLM response. #[derive(Clone)] -pub struct LlmResponseCodecContext { +pub struct LlmSanitizeResponseContext { /// Identity of the active codec. pub codec: LlmCodecIdentity, runtime: Option, @@ -315,7 +315,7 @@ pub struct LlmResponseCodecContext { invocation_id: Option, } -impl LlmRequestCodecContext { +impl LlmSanitizeRequestContext { /// Resolves the active request codec for this callback. #[must_use] pub fn resolve_codec(&self) -> Option { @@ -327,7 +327,7 @@ impl LlmRequestCodecContext { } } -impl LlmResponseCodecContext { +impl LlmSanitizeResponseContext { /// Resolves the active response codec for this callback. #[must_use] pub fn resolve_codec(&self) -> Option { @@ -339,12 +339,6 @@ impl LlmResponseCodecContext { } } -/// Backward-compatible name for request codec context supplied to sanitizers. -pub type LlmSanitizeRequestContext = LlmRequestCodecContext; - -/// Backward-compatible name for response codec context supplied to sanitizers. -pub type LlmSanitizeResponseContext = LlmResponseCodecContext; - /// Invocation-scoped proxy for the active LLM request codec. #[derive(Clone)] pub struct WorkerRequestCodec { @@ -401,8 +395,8 @@ impl WorkerResponseCodec { /// not change these identities; codec operations reject incompatible payloads. #[derive(Clone)] pub struct LlmExecutionContext { - request_codec: LlmRequestCodecContext, - response_codec: Option, + request_codec: LlmSanitizeRequestContext, + response_codec: Option, } impl std::fmt::Debug for LlmExecutionContext { @@ -421,7 +415,7 @@ impl std::fmt::Debug for LlmExecutionContext { impl LlmExecutionContext { /// Request codec identity and invocation-scoped operations. #[must_use] - pub fn request_codec(&self) -> &LlmRequestCodecContext { + pub fn request_codec(&self) -> &LlmSanitizeRequestContext { &self.request_codec } @@ -430,7 +424,7 @@ impl LlmExecutionContext { /// Streaming execution returns `None` because Relay response codecs decode /// completed provider responses, not individual stream chunks. #[must_use] - pub fn response_codec(&self) -> Option<&LlmResponseCodecContext> { + pub fn response_codec(&self) -> Option<&LlmSanitizeResponseContext> { self.response_codec.as_ref() } } @@ -2939,12 +2933,12 @@ impl LlmPayload { let response_codec = context .response .as_ref() - .map(|response| -> Result { + .map(|response| -> Result { let identity = require_execution_field( response.codec.as_ref(), "response codec identity is missing", )?; - Ok(LlmResponseCodecContext { + Ok(LlmSanitizeResponseContext { codec: codec_identity_from_proto(Some(identity)), runtime: Some(runtime.clone()), codec_capability_id: response.codec_capability_id.clone(), @@ -2958,7 +2952,7 @@ impl LlmPayload { )); } Ok(LlmExecutionContext { - request_codec: LlmRequestCodecContext { + request_codec: LlmSanitizeRequestContext { codec: codec_identity_from_proto(Some(request_identity)), runtime: Some(runtime.clone()), codec_capability_id: request.codec_capability_id.clone(), diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index 0a70fea7d..c56bd22c2 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -233,7 +233,8 @@ type LLMCodec struct { CodecID *string } -// LLMSanitizeRequestContext provides request codec context for one sanitizer call. +// LLMSanitizeRequestContext provides request codec context to sanitizer and +// execution callbacks. type LLMSanitizeRequestContext struct { Codec LLMCodec resolved *LLMRequestSanitizeCodec @@ -244,26 +245,19 @@ func (context LLMSanitizeRequestContext) ResolveCodec() *LLMRequestSanitizeCodec return context.resolved } -// LLMSanitizeResponseContext provides response codec context for one sanitizer call. +// LLMSanitizeResponseContext provides response codec context to sanitizer and +// execution callbacks. type LLMSanitizeResponseContext struct { Codec LLMCodec resolved *LLMResponseSanitizeCodec } -// LLMRequestCodecContext is the general name for request codec context used by -// execution intercepts. The sanitizer-specific name remains source-compatible. -type LLMRequestCodecContext = LLMSanitizeRequestContext - -// LLMResponseCodecContext is the general name for response codec context used -// by execution intercepts. The sanitizer-specific name remains source-compatible. -type LLMResponseCodecContext = LLMSanitizeResponseContext - // LLMExecutionContext provides invocation-scoped codec access to an LLM // execution intercept. RequestCodec is always present. ResponseCodec is // available for unary execution and nil for streaming execution. type LLMExecutionContext struct { - RequestCodec LLMRequestCodecContext - ResponseCodec *LLMResponseCodecContext + RequestCodec LLMSanitizeRequestContext + ResponseCodec *LLMSanitizeResponseContext } // ResolveCodec returns the active callback-scoped response codec, if any. diff --git a/python/nemo_relay/__init__.py b/python/nemo_relay/__init__.py index a49ba65f9..75a55d51f 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -104,9 +104,7 @@ async def main(): LlmExecutionContext, LLMHandle, LLMRequest, - LlmRequestCodecContext, LLMRequestInterceptOutcome, - LlmResponseCodecContext, LlmSanitizeRequestCodec, LlmSanitizeRequestContext, LlmSanitizeResponseCodec, @@ -845,8 +843,6 @@ def worker() -> None: "LlmSanitizeResponseGuardrail", "LlmCodecIdentity", "LlmExecutionContext", - "LlmRequestCodecContext", - "LlmResponseCodecContext", "LlmSanitizeRequestContext", "LlmSanitizeResponseContext", "LlmSanitizeRequestCodec", diff --git a/python/nemo_relay/_native.pyi b/python/nemo_relay/_native.pyi index c5434f5d9..63fb7b402 100644 --- a/python/nemo_relay/_native.pyi +++ b/python/nemo_relay/_native.pyi @@ -130,27 +130,24 @@ class LlmExecutionContext: """Codec capabilities for one managed LLM execution intercept invocation.""" @property - def request_codec(self) -> LlmRequestCodecContext: ... + def request_codec(self) -> LlmSanitizeRequestContext: ... @property - def response_codec(self) -> LlmResponseCodecContext | None: ... + def response_codec(self) -> LlmSanitizeResponseContext | None: ... class LlmSanitizeRequestContext: - """Per-call context passed to an LLM request sanitizer callback.""" + """Request codec context shared by sanitizer and execution callbacks.""" @property def codec(self) -> LlmCodecIdentity: ... def resolve_codec(self) -> LlmSanitizeRequestCodec | None: ... class LlmSanitizeResponseContext: - """Per-call context passed to an LLM response sanitizer callback.""" + """Response codec context shared by sanitizer and execution callbacks.""" @property def codec(self) -> LlmCodecIdentity: ... def resolve_codec(self) -> LlmSanitizeResponseCodec | None: ... -LlmRequestCodecContext = LlmSanitizeRequestContext -LlmResponseCodecContext = LlmSanitizeResponseContext - class LlmSanitizeRequestCodec: def decode(self, request: LLMRequest) -> AnnotatedLLMRequest: ... def encode(self, annotated: AnnotatedLLMRequest, original: LLMRequest) -> LLMRequest: ... diff --git a/python/plugin/src/nemo_relay_plugin/__init__.py b/python/plugin/src/nemo_relay_plugin/__init__.py index 32e8bf9e2..2789f87e4 100644 --- a/python/plugin/src/nemo_relay_plugin/__init__.py +++ b/python/plugin/src/nemo_relay_plugin/__init__.py @@ -27,10 +27,10 @@ EventSanitizeFields: Mutable event observability fields. LlmRequest: A Relay LLM request represented as a JSON object. LlmCodecIdentity: Typed discriminator for the active LLM codec. - LlmRequestCodecContext: Request-direction codec identity and operations. - LlmResponseCodecContext: Response-direction codec identity and operations. - LlmSanitizeRequestContext: Per-call context supplied to an LLM request sanitizer. - LlmSanitizeResponseContext: Per-call context supplied to an LLM response sanitizer. + LlmSanitizeRequestContext: Request codec context shared by sanitizer and + execution callbacks. + LlmSanitizeResponseContext: Response codec context shared by sanitizer and + execution callbacks. LlmExecutionContext: Invocation-scoped codec context supplied to an LLM execution intercept. WorkerRequestCodec: Invocation-scoped async proxy for an active request codec. @@ -117,9 +117,7 @@ LlmOptimizationTokens, LlmRequest, LlmRequestCallback, - LlmRequestCodecContext, LlmRequestInterceptOutcome, - LlmResponseCodecContext, LlmSanitizeRequestCallback, LlmSanitizeRequestContext, LlmSanitizeResponseCallback, @@ -173,8 +171,6 @@ "LlmCodecIdentity", "LlmExecutionCallback", "LlmExecutionContext", - "LlmRequestCodecContext", - "LlmResponseCodecContext", "LogSeverity", "MetricKind", "MetricMeasurement", diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index 0ae5cfc19..4f081dc4a 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -192,7 +192,7 @@ def get(self, code: str) -> RuntimeDiagnostic | None: @dataclass(frozen=True) class LlmSanitizeRequestContext: - """Structured per-call context provided to an LLM request sanitizer.""" + """Request codec context shared by sanitizer and execution callbacks.""" codec: LlmCodecIdentity _runtime: "PluginRuntime | None" = field(default=None, repr=False, compare=False) @@ -210,7 +210,7 @@ def resolve_codec(self) -> "WorkerRequestCodec | None": @dataclass(frozen=True) class LlmSanitizeResponseContext: - """Structured per-call context provided to an LLM response sanitizer.""" + """Response codec context shared by sanitizer and execution callbacks.""" codec: LlmCodecIdentity _runtime: "PluginRuntime | None" = field(default=None, repr=False, compare=False) @@ -226,13 +226,6 @@ def resolve_codec(self) -> "WorkerResponseCodec | None": return WorkerResponseCodec(self._runtime, self._capability_id, self._invocation_id) -# General names for contexts that are also supplied to execution interceptors. -# The original class objects remain canonical at runtime for compatibility with -# repr, pickling, and code that inspects ``__name__``. -LlmRequestCodecContext = LlmSanitizeRequestContext -LlmResponseCodecContext = LlmSanitizeResponseContext - - @dataclass(frozen=True) class WorkerRequestCodec: """Invocation-scoped async proxy for an active request codec.""" @@ -274,8 +267,8 @@ class LlmExecutionContext: not select a new codec; incompatible codec operations fail. """ - request_codec: LlmRequestCodecContext - response_codec: LlmResponseCodecContext | None + request_codec: LlmSanitizeRequestContext + response_codec: LlmSanitizeResponseContext | None def _llm_codec_identity(invocation: pb.LlmInvocation) -> LlmCodecIdentity: @@ -318,7 +311,7 @@ def _llm_execution_context( raise WorkerSdkError("malformed LLM execution codec context: request codec identity is missing") request_id = context.request.codec_capability_id if context.request.HasField("codec_capability_id") else None - request_context = LlmRequestCodecContext( + request_context = LlmSanitizeRequestContext( codec=_codec_identity( context.request.codec.kind, context.request.codec.id if context.request.codec.HasField("id") else None, @@ -327,12 +320,12 @@ def _llm_execution_context( _capability_id=request_id, _invocation_id=invocation_id, ) - response_context: LlmResponseCodecContext | None = None + response_context: LlmSanitizeResponseContext | None = None if context.HasField("response"): if not context.response.HasField("codec"): raise WorkerSdkError("malformed LLM execution codec context: response codec identity is missing") response_id = context.response.codec_capability_id if context.response.HasField("codec_capability_id") else None - response_context = LlmResponseCodecContext( + response_context = LlmSanitizeResponseContext( codec=_codec_identity( context.response.codec.kind, context.response.codec.id if context.response.codec.HasField("id") else None, From 861ad360c331fa92cf327760eb8f62d4154ffe6b Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 29 Sep 2026 18:43:00 -0700 Subject: [PATCH 14/22] refactor: simplify execution codec documentation Signed-off-by: Alex Fournier --- .../src/api/runtime/llm_execution_context.rs | 33 +++++++------------ 1 file changed, 11 insertions(+), 22 deletions(-) diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index 8e6c7278e..417a23912 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -14,10 +14,8 @@ use crate::codec::traits::{LlmCodec, LlmResponseCodec}; use crate::error::{FlowError, Result}; use crate::json::Json; -const INACTIVE_EXECUTION_CODEC_ERROR: &str = "LLM execution codec capability is no longer active"; - fn inactive_execution_codec_error() -> FlowError { - FlowError::InvalidArgument(INACTIVE_EXECUTION_CODEC_ERROR.into()) + FlowError::InvalidArgument("LLM execution codec capability is no longer active".into()) } fn upgrade_active_codec(codec: &Weak, gate: &ExecutionCodecGate) -> Result> { @@ -104,23 +102,16 @@ impl LlmResponseCodec for RevocableResponseCodec { } } -/// Active request and response codec context for one managed LLM execution. +/// Codec access for one LLM execution. /// -/// The request direction is always present and distinguishes an invocation -/// with no request codec from an invocation with a built-in, runtime, or opaque -/// codec. Unary execution also carries a response direction. Streaming -/// execution deliberately leaves [`Self::response_codec`] unavailable because -/// Relay's response codecs operate on complete provider responses rather than -/// individual stream chunks. +/// Request codec information is always present. Response codec access is +/// available only for non-streaming calls because response codecs expect a +/// complete response, not individual chunks. /// -/// The codecs are fixed when the managed invocation is created. Rewriting a -/// payload does not select another codec; decoding or encoding an incompatible -/// wire representation fails rather than inferring a different format. -/// Resolved codec capabilities are valid only for the callback that received -/// this context. Unary capabilities expire when that callback settles; -/// streaming request capabilities remain valid until its returned stream ends -/// or closes. Retained capabilities return [`FlowError::InvalidArgument`] -/// after expiry. +/// Relay chooses the codecs before interceptors run. Changing the payload does +/// not select a different codec. Codec handles expire when the interceptor +/// finishes; a streaming request handle remains valid until its returned stream +/// ends or closes. Later use returns [`FlowError::InvalidArgument`]. #[derive(Clone, Debug, Default)] pub struct LlmExecutionContext { request_codec: LlmSanitizeRequestContext, @@ -306,13 +297,11 @@ mod tests { ); assert!(matches!( retained_request.decode(&request), - Err(FlowError::InvalidArgument(message)) - if message == INACTIVE_EXECUTION_CODEC_ERROR + Err(FlowError::InvalidArgument(_)) )); assert!(matches!( retained_response.decode_response(&Json::Null), - Err(FlowError::InvalidArgument(message)) - if message == INACTIVE_EXECUTION_CODEC_ERROR + Err(FlowError::InvalidArgument(_)) )); } } From aca721dbc4969305123847b00b35c4439c5fc810 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Tue, 29 Sep 2026 21:28:17 -0700 Subject: [PATCH 15/22] fix: harden execution codec context ownership Signed-off-by: Alex Fournier --- .../src/api/runtime/llm_execution_context.rs | 104 ++-------- crates/core/src/api/runtime/state.rs | 79 ++------ crates/core/src/plugin/dynamic/native.rs | 189 +++++++++--------- crates/core/src/plugin/dynamic/worker.rs | 142 ++++++++----- .../tests/integration/middleware_tests.rs | 20 +- .../tests/integration/worker_plugin_tests.rs | 9 + .../core/tests/unit/dynamic_worker_tests.rs | 45 ++++- crates/core/tests/unit/llm_api_tests.rs | 35 ++-- .../tests/unit/llm_execution_context_tests.rs | 116 +++++++++++ crates/plugin/README.md | 4 +- crates/plugin/src/async_sdk.rs | 90 +++++---- crates/plugin/src/lib.rs | 107 +++++----- crates/plugin/tests/typed_callbacks.rs | 107 ++++------ crates/worker-proto/README.md | 14 +- crates/worker/README.md | 18 +- docs/about-nemo-relay/concepts/codecs.mdx | 15 ++ docs/about-nemo-relay/concepts/middleware.mdx | 6 + docs/about-nemo-relay/release-notes/index.mdx | 13 +- docs/build-plugins/about.mdx | 19 +- .../language-binding/register-behavior.mdx | 4 +- docs/build-plugins/native/about.mdx | 6 +- .../native/native-abi-reference.mdx | 33 +-- docs/build-plugins/native/wrap-execution.mdx | 16 +- .../package-discoverable-plugins.mdx | 8 +- docs/build-plugins/plugin-context.mdx | 2 +- docs/build-plugins/workers/about.mdx | 14 +- .../workers/grpc-v1-protocol.mdx | 26 +-- .../workers/middleware-and-continuations.mdx | 10 +- docs/build-plugins/workers/python.mdx | 6 +- docs/build-plugins/workers/rust.mdx | 4 +- docs/reference/migration-guides.mdx | 12 +- examples/python-grpc-worker-plugin/README.md | 6 +- examples/rust-grpc-worker-plugin/README.md | 9 +- examples/rust-native-plugin/README.md | 8 +- python/plugin/README.md | 10 +- skills/nemo-relay-plugin-build/SKILL.md | 20 +- 36 files changed, 740 insertions(+), 586 deletions(-) create mode 100644 crates/core/tests/unit/llm_execution_context_tests.rs diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index 417a23912..08fe1b39c 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -100,6 +100,11 @@ impl LlmResponseCodec for RevocableResponseCodec { fn decode_response(&self, response: &Json) -> Result { upgrade_active_codec(&self.codec, &self.gate)?.decode_response(response) } + + fn allows_estimated_cost(&self, response: &Json) -> bool { + upgrade_active_codec(&self.codec, &self.gate) + .is_ok_and(|codec| codec.allows_estimated_cost(response)) + } } /// Codec access for one LLM execution. @@ -112,6 +117,8 @@ impl LlmResponseCodec for RevocableResponseCodec { /// not select a different codec. Codec handles expire when the interceptor /// finishes; a streaming request handle remains valid until its returned stream /// ends or closes. Later use returns [`FlowError::InvalidArgument`]. +/// Public construction supports direct callback tests; only Relay can attach +/// active, revocable codec handles. #[derive(Clone, Debug, Default)] pub struct LlmExecutionContext { request_codec: LlmSanitizeRequestContext, @@ -119,7 +126,7 @@ pub struct LlmExecutionContext { } impl LlmExecutionContext { - /// Construct an execution context from its directional codec contexts. + /// Construct an execution context from request and optional response codec context. #[must_use] pub fn new( request_codec: LlmSanitizeRequestContext, @@ -131,7 +138,7 @@ impl LlmExecutionContext { } } - /// Construct the context for a unary managed execution. + /// Construct the context for a non-streaming managed execution. pub(crate) fn for_unary_codecs( request_codec: Option>, response_codec: &Option>, @@ -205,7 +212,7 @@ impl LlmExecutionContext { &self.request_codec } - /// Return the unary response-direction codec identity and revocable capability. + /// Return the completed-response codec identity and revocable capability. /// /// Streaming execution returns `None` because Relay does not expose a /// completed-response codec for individual stream chunks. @@ -216,92 +223,5 @@ impl LlmExecutionContext { } #[cfg(test)] -mod tests { - use super::*; - use crate::api::runtime::{BuiltinLlmCodec, LlmCodecIdentity}; - - struct LeaseProbeCodec; - - impl LlmCodec for LeaseProbeCodec { - fn codec_identity(&self) -> LlmCodecIdentity { - LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) - } - - fn decode(&self, _request: &LlmRequest) -> Result { - Ok(AnnotatedLlmRequest::default()) - } - - fn encode( - &self, - _annotated: &AnnotatedLlmRequest, - original: &LlmRequest, - ) -> Result { - Ok(original.clone()) - } - } - - impl LlmResponseCodec for LeaseProbeCodec { - fn codec_identity(&self) -> LlmCodecIdentity { - LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) - } - - fn decode_response(&self, _response: &Json) -> Result { - Ok(AnnotatedLlmResponse::default()) - } - } - - #[test] - fn retained_facades_forward_while_active_and_expire_with_their_lease() { - let backing = Arc::new(LeaseProbeCodec); - let backing_probe = Arc::downgrade(&backing); - let request_codec: Arc = backing.clone(); - let response_codec: Arc = backing.clone(); - let context = - LlmExecutionContext::for_unary_codecs(Some(request_codec), &Some(response_codec)); - drop(backing); - - let (leased_context, guard) = context.lease(); - let retained_request = leased_context.request_codec().resolve_codec().unwrap(); - let retained_response = leased_context - .response_codec() - .and_then(LlmSanitizeResponseContext::resolve_codec) - .unwrap(); - drop(leased_context); - drop(context); - - let request = LlmRequest { - headers: serde_json::Map::new(), - content: Json::Null, - }; - let annotated = retained_request.decode(&request).unwrap(); - assert_eq!( - retained_request.encode(&annotated, &request).unwrap(), - request - ); - assert_eq!( - retained_response.decode_response(&Json::Null).unwrap(), - AnnotatedLlmResponse::default() - ); - assert!(backing_probe.upgrade().is_some()); - - drop(guard); - - assert!(backing_probe.upgrade().is_none()); - assert_eq!( - retained_request.codec_identity(), - LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) - ); - assert_eq!( - retained_response.codec_identity(), - LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) - ); - assert!(matches!( - retained_request.decode(&request), - Err(FlowError::InvalidArgument(_)) - )); - assert!(matches!( - retained_response.decode_response(&Json::Null), - Err(FlowError::InvalidArgument(_)) - )); - } -} +#[path = "../../../tests/unit/llm_execution_context_tests.rs"] +mod tests; diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index fa5d6b284..4d8ff1118 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -64,14 +64,10 @@ use chrono::{Duration, Utc}; use serde_json::json; use uuid::Uuid; -struct ContinuationGuardedLlmStream { +struct ExecutionGuardedLlmStream { inner: LlmJsonStream, - guard: Option, -} - -struct CodecGuardedLlmStream { - inner: LlmJsonStream, - guard: Option, + continuation_guard: Option, + codec_guard: Option, } struct ContextualizedLlmStream { @@ -120,23 +116,27 @@ pub(crate) fn contextualize_stream( }) } -impl Stream for ContinuationGuardedLlmStream { +impl Stream for ExecutionGuardedLlmStream { type Item = crate::error::Result; fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { let this = self.get_mut(); let result = Pin::new(&mut this.inner).poll_next(cx); if matches!(&result, Poll::Ready(None)) { - this.guard.take(); + this.continuation_guard.take(); + this.codec_guard.take(); + } else if matches!(&result, Poll::Ready(Some(Err(_)))) { + this.codec_guard.take(); } result } } -impl LlmStreamInner for ContinuationGuardedLlmStream { +impl LlmStreamInner for ExecutionGuardedLlmStream { fn terminalize(self: Pin<&mut Self>) { let this = self.get_mut(); - this.guard.take(); + this.codec_guard.take(); + this.continuation_guard.take(); this.inner.terminalize(); } @@ -145,60 +145,24 @@ impl LlmStreamInner for ContinuationGuardedLlmStream { ) -> Pin> + Send + '_>> { Box::pin(async move { let this = self.get_mut(); - let guard = this.guard.take(); + let continuation_guard = this.continuation_guard.take(); let result = this.inner.close().await; - drop(guard); + drop(continuation_guard); + this.codec_guard.take(); result }) } } -fn guard_stream_continuation( +fn guard_execution_stream( stream: LlmJsonStream, - guard: MiddlewareContinuationGuard, + continuation_guard: MiddlewareContinuationGuard, + codec_guard: LlmExecutionCodecLeaseGuard, ) -> LlmJsonStream { - LlmJsonStream::from_closeable(ContinuationGuardedLlmStream { - inner: stream, - guard: Some(guard), - }) -} - -impl Stream for CodecGuardedLlmStream { - type Item = crate::error::Result; - - fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - let this = self.get_mut(); - let result = Pin::new(&mut this.inner).poll_next(cx); - if matches!(&result, Poll::Ready(None) | Poll::Ready(Some(Err(_)))) { - this.guard.take(); - } - result - } -} - -impl LlmStreamInner for CodecGuardedLlmStream { - fn terminalize(self: Pin<&mut Self>) { - let this = self.get_mut(); - this.guard.take(); - this.inner.terminalize(); - } - - fn close( - self: Pin<&mut Self>, - ) -> Pin> + Send + '_>> { - Box::pin(async move { - let this = self.get_mut(); - let result = this.inner.close().await; - this.guard.take(); - result - }) - } -} - -fn guard_stream_codec(stream: LlmJsonStream, guard: LlmExecutionCodecLeaseGuard) -> LlmJsonStream { - LlmJsonStream::from_closeable(CodecGuardedLlmStream { + LlmJsonStream::from_closeable(ExecutionGuardedLlmStream { inner: stream, - guard: Some(guard), + continuation_guard: Some(continuation_guard), + codec_guard: Some(codec_guard), }) } @@ -1947,8 +1911,7 @@ impl NemoRelayContextState { }); let result = callable(¤t_name, request, current_context, raw_next).await; result.map(|stream| { - let stream = guard_stream_continuation(stream, continuation_guard); - guard_stream_codec(stream, codec_guard) + guard_execution_stream(stream, continuation_guard, codec_guard) }) }) }); diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index c584460e8..376b456b5 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -14,7 +14,7 @@ use std::future::Future; use std::panic::{AssertUnwindSafe, catch_unwind}; use std::path::{Path, PathBuf}; use std::pin::Pin; -use std::ptr; +use std::ptr::{self, NonNull}; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, OnceLock, Weak}; use std::task::{Context, Poll}; @@ -578,15 +578,39 @@ struct NativeHostString(Vec); struct NativeHostLlmRequestCodec(Arc); struct NativeHostLlmResponseCodec(Arc); +struct OwnedNativeString { + ptr: NonNull, +} + +impl OwnedNativeString { + fn new(ptr: *mut NemoRelayNativeString) -> FlowResult { + Ok(Self { + ptr: NonNull::new(ptr).ok_or_else(|| { + FlowError::Internal("native string allocation returned null".into()) + })?, + }) + } + + fn as_ptr(&self) -> *const NemoRelayNativeString { + self.ptr.as_ptr() + } +} + +impl Drop for OwnedNativeString { + fn drop(&mut self) { + unsafe { native_string_free(self.ptr.as_ptr()) }; + } +} + /// Borrows the host codec handles and owns the native strings exposed to one /// native LLM execution callback. struct NativeLlmExecutionContextBridge<'a> { request_codec: Option<&'a NativeHostLlmRequestCodec>, request_kind: NemoRelayNativeLlmCodecKind, - request_id: Option, + request_id: Option, response_codec: Option<&'a NativeHostLlmResponseCodec>, response_kind: Option, - response_id: Option, + response_id: Option, } impl<'a> NativeLlmExecutionContextBridge<'a> { @@ -597,19 +621,11 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { ) -> FlowResult { let (request_kind, request_id) = native_llm_codec_identity(context.request_codec().codec())?; - let request_id = request_id.map(|value| value as usize); + let request_id = request_id.map(OwnedNativeString::new).transpose()?; let (response_kind, response_id) = if let Some(response) = context.response_codec() { - let (kind, id) = match native_llm_codec_identity(response.codec()) { - Ok(value) => value, - Err(error) => { - if let Some(request_id) = request_id { - unsafe { native_string_free(request_id as *mut NemoRelayNativeString) }; - } - return Err(error); - } - }; - (Some(kind), id.map(|value| value as usize)) + let (kind, id) = native_llm_codec_identity(response.codec())?; + (Some(kind), id.map(OwnedNativeString::new).transpose()?) } else { (None, None) }; @@ -632,7 +648,8 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { codec_kind: self.request_kind, codec_id: self .request_id - .map_or(ptr::null(), |value| value as *const NemoRelayNativeString), + .as_ref() + .map_or(ptr::null(), OwnedNativeString::as_ptr), codec: self .request_codec .map_or(ptr::null(), |value| std::ptr::from_ref(value).cast()), @@ -643,7 +660,8 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { codec_kind, codec_id: self .response_id - .map_or(ptr::null(), |value| value as *const NemoRelayNativeString), + .as_ref() + .map_or(ptr::null(), OwnedNativeString::as_ptr), codec: self .response_codec .map_or(ptr::null(), |value| std::ptr::from_ref(value).cast()), @@ -657,17 +675,6 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { } } -impl Drop for NativeLlmExecutionContextBridge<'_> { - fn drop(&mut self) { - if let Some(request_id) = self.request_id.take() { - unsafe { native_string_free(request_id as *mut NemoRelayNativeString) }; - } - if let Some(response_id) = self.response_id.take() { - unsafe { native_string_free(response_id as *mut NemoRelayNativeString) }; - } - } -} - struct NativeHostScopeHandle(ScopeHandle); struct NativeHostScopeStack(ScopeStackHandle); @@ -2179,75 +2186,77 @@ async fn invoke_native_async_callback_inner( before_settlement_lock: None, _callback_user_data: Some(user_data.clone()), }); - let native_context = match &callback { - NativeAsyncCallback::LlmExecution { context, .. } => { - Some(NativeLlmExecutionContextBridge::new( - context, - completion.request_codec.as_ref(), - completion.response_codec.as_ref(), - )) - } - NativeAsyncCallback::Middleware(_) => None, - } - .transpose(); - let native_context = match native_context { - Ok(context) => context, - Err(error) => { - unsafe { native_string_free(invocation as *mut NemoRelayNativeString) }; - return Err(error); - } - }; let mut wait = NativeAsyncWait { completion: Arc::clone(&completion), receiver, completed: false, }; - 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::with_completion_owner( - inner, - runtime, - Some(user_data.clone()), - &completion, - ), - )) as usize), - (None, None) => None, - _ => unreachable!("runtime is present exactly for native async intercepts"), - }; - // ABI v3 exposes a thread-stack capture operation. Mirror the effective - // task-local stack into that slot only while entering plugin code so the - // SDK can capture it before moving the future to its own executor. - let previous_thread_stack = capture_thread_scope_stack(); - sync_thread_scope_stack(current_scope_stack()); - let state = catch_unwind(AssertUnwindSafe(|| match callback { - NativeAsyncCallback::Middleware(callback) => unsafe { - callback( - user_data.ptr, - invocation as *const NemoRelayNativeString, - next_ref - .map(|next| next as *const NemoRelayNativeAsyncNext) - .unwrap_or(ptr::null()), - completion_ref as *const NemoRelayNativeAsyncCompletion, - ) - }, - NativeAsyncCallback::LlmExecution { callback, .. } => native_context - .as_ref() - .expect("LLM execution callbacks always build a native context") - .with_native_context(|context| unsafe { + let (state, completion_ref) = { + let native_context = match &callback { + NativeAsyncCallback::LlmExecution { context, .. } => { + Some(NativeLlmExecutionContextBridge::new( + context, + completion.request_codec.as_ref(), + completion.response_codec.as_ref(), + )) + } + NativeAsyncCallback::Middleware(_) => None, + } + .transpose(); + let native_context = match native_context { + Ok(context) => context, + Err(error) => { + unsafe { native_string_free(invocation as *mut NemoRelayNativeString) }; + return Err(error); + } + }; + 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::with_completion_owner( + inner, + runtime, + Some(user_data.clone()), + &completion, + ), + )) as usize), + (None, None) => None, + _ => unreachable!("runtime is present exactly for native async intercepts"), + }; + // ABI v3 exposes a thread-stack capture operation. Mirror the effective + // task-local stack into that slot only while entering plugin code so the + // SDK can capture it before moving the future to its own executor. + let previous_thread_stack = capture_thread_scope_stack(); + sync_thread_scope_stack(current_scope_stack()); + let state = catch_unwind(AssertUnwindSafe(|| match callback { + NativeAsyncCallback::Middleware(callback) => unsafe { callback( user_data.ptr, invocation as *const NemoRelayNativeString, - std::ptr::from_ref(&context), next_ref .map(|next| next as *const NemoRelayNativeAsyncNext) .unwrap_or(ptr::null()), completion_ref as *const NemoRelayNativeAsyncCompletion, ) - }), - })); - restore_thread_scope_stack(previous_thread_stack); - drop(native_context); + }, + NativeAsyncCallback::LlmExecution { callback, .. } => native_context + .as_ref() + .expect("LLM execution callbacks always build a native context") + .with_native_context(|context| unsafe { + callback( + user_data.ptr, + invocation as *const NemoRelayNativeString, + std::ptr::from_ref(&context), + next_ref + .map(|next| next as *const NemoRelayNativeAsyncNext) + .unwrap_or(ptr::null()), + completion_ref as *const NemoRelayNativeAsyncCompletion, + ) + }), + })); + restore_thread_scope_stack(previous_thread_stack); + (state, completion_ref) + }; let state = match state { Ok(state) => state, Err(_) => { @@ -3120,7 +3129,7 @@ impl Drop for NativeCallbackTaskGuard { } } -/// Invokes a unary continuation with an independent per-call result callback. +/// Invokes a non-streaming continuation with an independent per-call result callback. unsafe extern "C" fn native_async_next_invoke_result( next: *const NemoRelayNativeAsyncNext, invocation_json: *const NemoRelayNativeString, @@ -3159,7 +3168,7 @@ unsafe extern "C" fn native_async_next_invoke_result( } NativeAsyncNextInner::LlmStream(_) => { set_native_last_error( - "stream continuations require async_next_invoke_stream; unary result callbacks cannot buffer a stream", + "stream continuations require async_next_invoke_stream; non-streaming result callbacks cannot buffer a stream", ); return NemoRelayStatus::InvalidArg; } @@ -4040,16 +4049,16 @@ fn wrap_native_incremental_llm_stream_execution_with_user_data( before_settlement_lock: None, _callback_user_data: Some(user_data.clone()), }); - let native_context = NativeLlmExecutionContextBridge::new( - &context, - stream.request_codec.as_ref(), - None, - )?; let output = NativeAsyncStreamReceiver { receiver, stream: Arc::clone(&stream), }; let state = { + let native_context = NativeLlmExecutionContextBridge::new( + &context, + stream.request_codec.as_ref(), + None, + )?; let invocation = native_string_from_json(&serde_json::json!({"name": name, "request": request})) .ok_or_else(|| { diff --git a/crates/core/src/plugin/dynamic/worker.rs b/crates/core/src/plugin/dynamic/worker.rs index 8520fa6e9..771925482 100644 --- a/crates/core/src/plugin/dynamic/worker.rs +++ b/crates/core/src/plugin/dynamic/worker.rs @@ -1651,21 +1651,37 @@ impl WorkerInvocationGuard { } } - fn cancel(&mut self, reason: impl Into) { + fn cancellation( + &mut self, + reason: impl Into, + ) -> Option<(PluginWorkerClient, CancelInvocationRequest)> { if !self.cancel_on_drop { - return; + return None; } self.cancel_on_drop = false; - let mut client = self.client.clone(); - let request = CancelInvocationRequest { - activation_id: self.activation_id.clone(), - invocation_id: self.invocation_id.clone(), - auth_token: self.auth_token.clone(), - reason: reason.into(), - }; - self.runtime.spawn(async move { + Some(( + self.client.clone(), + CancelInvocationRequest { + activation_id: self.activation_id.clone(), + invocation_id: self.invocation_id.clone(), + auth_token: self.auth_token.clone(), + reason: reason.into(), + }, + )) + } + + fn cancel(&mut self, reason: impl Into) { + if let Some((mut client, request)) = self.cancellation(reason) { + self.runtime.spawn(async move { + let _ = worker_rpc(client.cancel_invocation(worker_rpc_request(request))).await; + }); + } + } + + async fn cancel_and_wait(&mut self, reason: impl Into) { + if let Some((mut client, request)) = self.cancellation(reason) { let _ = worker_rpc(client.cancel_invocation(worker_rpc_request(request))).await; - }); + } } fn finish(&mut self) { @@ -2068,8 +2084,11 @@ impl WorkerPluginCallback { None, )), ); - let codec_capabilities = - self.attach_llm_execution_codec_context(&mut invoke, &execution_context); + let codec_capabilities = self.attach_llm_execution_codec_context( + &mut invoke, + &execution_context, + WorkerLlmExecutionMode::CompleteResponse, + ); let _codec_capabilities = self.cleanup_after_setup_error(&invoke, codec_capabilities)?; json_from_invoke_response(self.invoke_async(invoke).await?) } @@ -2096,8 +2115,11 @@ impl WorkerPluginCallback { None, )), ); - let codec_capabilities = - self.attach_llm_execution_codec_context(&mut invoke, &execution_context); + let codec_capabilities = self.attach_llm_execution_codec_context( + &mut invoke, + &execution_context, + WorkerLlmExecutionMode::Streaming, + ); let codec_capabilities = self.cleanup_after_setup_error(&invoke, codec_capabilities)?; let mut client = self.client.clone(); let mut guard = WorkerInvocationGuard::new(self, &invoke); @@ -2110,7 +2132,7 @@ impl WorkerPluginCallback { let result = tokio::select! { result = worker_rpc(client.invoke_stream(worker_rpc_request(invoke))) => result, _ = tx.closed() => { - guard.cancel("host stopped consuming the worker stream"); + guard.cancel_and_wait("host stopped consuming the worker stream").await; guard.finish(); return; } @@ -2123,7 +2145,7 @@ impl WorkerPluginCallback { let item = tokio::select! { item = stream.next() => item, _ = tx.closed() => { - guard.cancel("host stopped consuming the worker stream"); + guard.cancel_and_wait("host stopped consuming the worker stream").await; break; } }; @@ -2138,7 +2160,7 @@ impl WorkerPluginCallback { }; let terminal = result.is_err(); if tx.send(result).await.is_err() { - guard.cancel("host stopped consuming the worker stream"); + guard.cancel_and_wait("host stopped consuming the worker stream").await; break; } if terminal { @@ -2152,13 +2174,13 @@ impl WorkerPluginCallback { } else { "worker stream transport failed" }; - guard.cancel(reason); let _ = tx .send(Err(worker_status_to_flow( "worker stream invoke failed", err, ))) .await; + guard.cancel_and_wait(reason).await; } } guard.finish(); @@ -2178,38 +2200,52 @@ impl WorkerPluginCallback { &self, invoke: &mut InvokeRequest, context: &LlmExecutionContext, - ) -> FlowResult> { - let mut guards = Vec::with_capacity(2); + mode: WorkerLlmExecutionMode, + ) -> FlowResult { + match (mode, context.response_codec()) { + (WorkerLlmExecutionMode::CompleteResponse, None) => { + return Err(FlowError::InvalidArgument( + "complete-response execution requires a response codec context".into(), + )); + } + (WorkerLlmExecutionMode::Streaming, Some(_)) => { + return Err(FlowError::InvalidArgument( + "streaming execution cannot use a response codec context".into(), + )); + } + _ => {} + } - let mut request = ProtoLlmSanitizeRequestContext { + let request_guard = context + .request_codec() + .resolve_codec() + .map(|codec| { + self.host_state + .issue_request_codec(&invoke.invocation_id, codec) + }) + .transpose()?; + let request = ProtoLlmSanitizeRequestContext { codec: Some(codec_identity_to_proto(context.request_codec().codec())), - codec_capability_id: None, + codec_capability_id: request_guard.as_ref().map(|guard| guard.id().into()), }; - if let Some(codec) = context.request_codec().resolve_codec() { - let capability = self - .host_state - .issue_request_codec(&invoke.invocation_id, codec)?; - request.codec_capability_id = Some(capability.id().into()); - guards.push(capability); - } - let response = context - .response_codec() - .map(|response_context| -> FlowResult<_> { - let mut response = ProtoLlmSanitizeResponseContext { + let (response, response_guard) = match context.response_codec() { + Some(response_context) => { + let guard = response_context + .resolve_codec() + .map(|codec| { + self.host_state + .issue_response_codec(&invoke.invocation_id, codec) + }) + .transpose()?; + let response = ProtoLlmSanitizeResponseContext { codec: Some(codec_identity_to_proto(response_context.codec())), - codec_capability_id: None, + codec_capability_id: guard.as_ref().map(|guard| guard.id().into()), }; - if let Some(codec) = response_context.resolve_codec() { - let capability = self - .host_state - .issue_response_codec(&invoke.invocation_id, codec)?; - response.codec_capability_id = Some(capability.id().into()); - guards.push(capability); - } - Ok(response) - }) - .transpose()?; + (Some(response), guard) + } + None => (None, None), + }; let Some(invoke_request_payload::Payload::Llm(llm)) = invoke.payload.as_mut() else { unreachable!("LLM execution invocation must have an LLM payload"); @@ -2218,7 +2254,10 @@ impl WorkerPluginCallback { request: Some(request), response, })); - Ok(guards) + Ok(WorkerCodecCapabilityGuards { + _request: request_guard, + _response: response_guard, + }) } fn cleanup_after_setup_error( @@ -2482,6 +2521,17 @@ struct WorkerCodecCapabilityGuard { capability_id: String, } +struct WorkerCodecCapabilityGuards { + _request: Option, + _response: Option, +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum WorkerLlmExecutionMode { + CompleteResponse, + Streaming, +} + impl WorkerCodecCapabilityGuard { fn id(&self) -> &str { &self.capability_id diff --git a/crates/core/tests/integration/middleware_tests.rs b/crates/core/tests/integration/middleware_tests.rs index df0ca2404..e988fa3c7 100644 --- a/crates/core/tests/integration/middleware_tests.rs +++ b/crates/core/tests/integration/middleware_tests.rs @@ -72,10 +72,12 @@ use nemo_relay::api::tool::{ ToolExecutionInterceptOutcome, ToolExecutionResult, tool_call, tool_call_end, tool_call_execute, tool_conditional_execution, tool_request_intercepts, }; +use nemo_relay::codec::openai_chat::OpenAIChatCodec; use nemo_relay::codec::optimization::{ LlmOptimizationContribution, LlmOptimizationEvidenceQuality, LlmOptimizationTokenImpact, LlmOptimizationTokens, }; +use nemo_relay::codec::traits::LlmCodec; use nemo_relay::error::FlowError; use nemo_relay::json::Json; use nemo_relay::observability::OpenTelemetryType; @@ -2325,15 +2327,22 @@ async fn stream_next_is_revoked_when_the_managed_stream_terminalizes_with_an_err let request = LlmRequest { headers: serde_json::Map::new(), - content: json!({"prompt": "terminal-error"}), + content: json!({ + "model": "test-model", + "messages": [{"role": "user", "content": "terminal-error"}] + }), }; let upstream_error_next = Arc::new(Mutex::new(None::)); + let upstream_error_codec = Arc::new(Mutex::new(None::>)); let captured_upstream_error_next = Arc::clone(&upstream_error_next); + let captured_upstream_error_codec = Arc::clone(&upstream_error_codec); register_llm_stream_execution_intercept( "upstream_error_stream_next", 1, - Arc::new(move |_name, request, _context, next| { + Arc::new(move |_name, request, context, next| { *captured_upstream_error_next.lock().unwrap() = Some(next.clone()); + *captured_upstream_error_codec.lock().unwrap() = + context.request_codec().resolve_codec(); next(request) }), ) @@ -2354,6 +2363,7 @@ async fn stream_next_is_revoked_when_the_managed_stream_terminalizes_with_an_err })) .collector(Box::new(|_| Ok(()))) .finalizer(Box::new(|| json!({}))) + .codec(Arc::new(OpenAIChatCodec)) .build(), ) .await @@ -2369,6 +2379,12 @@ async fn stream_next_is_revoked_when_the_managed_stream_terminalizes_with_an_err FlowError::InvalidArgument(message) if message == "execution continuation is no longer active" )); + let codec = upstream_error_codec.lock().unwrap().take().unwrap(); + assert!(matches!( + codec.decode(&request), + Err(FlowError::InvalidArgument(message)) + if message == "LLM execution codec capability is no longer active" + )); assert_eq!(upstream_provider_calls.load(Ordering::Acquire), 1); deregister_llm_stream_execution_intercept("upstream_error_stream_next").unwrap(); diff --git a/crates/core/tests/integration/worker_plugin_tests.rs b/crates/core/tests/integration/worker_plugin_tests.rs index b2345dc4f..78a89096d 100644 --- a/crates/core/tests/integration/worker_plugin_tests.rs +++ b/crates/core/tests/integration/worker_plugin_tests.rs @@ -1462,6 +1462,15 @@ fn worker_llm_execution_context_requires_zero_ten_compatibility() { let (_manifest_dir, manifest_ref) = write_manifest_with_relay(fixture.binary_path(), ">=0.9,<1.0"); + let activation = load_worker_plugins([WorkerPluginLoadSpec { + plugin_id: "fixture_worker".into(), + manifest_ref: manifest_ref.to_string_lossy().into_owned(), + environment_ref: None, + config: Map::from_iter([("event_metadata_injector_only".into(), json!(true))]), + }]) + .expect("the 0.10 floor applies only to workers that register LLM execution intercepts"); + activation.clear(); + let error = match load_worker_plugins([WorkerPluginLoadSpec { plugin_id: "fixture_worker".into(), manifest_ref: manifest_ref.to_string_lossy().into_owned(), diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index 4e4ce3f8f..0a536922e 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -1265,6 +1265,43 @@ async fn llm_worker_execution_codec_context_is_required_and_ephemeral() { ); } +#[tokio::test(flavor = "multi_thread")] +async fn worker_execution_rejects_mismatched_response_codec_context() { + enable_operational_logs(); + let (callback, _shutdown) = fake_callback_service(|_| InvokeResponse { + result: Some(InvokeResult::Empty(EmptyResult {})), + }) + .await; + let result = callback + .invoke_llm_stream_execution( + "invalid-stream-context", + "model", + valid_llm_request(), + openai_execution_codec_context(), + Arc::new(|_| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })), + ) + .await; + let Err(error) = result else { + panic!("streaming response codec context must be rejected"); + }; + + assert!(matches!(error, FlowError::InvalidArgument(_))); + + let error = callback + .invoke_llm_execution( + "invalid-complete-response-context", + "model", + valid_llm_request(), + openai_stream_execution_codec_context(), + Arc::new(|_| Box::pin(async { Ok(json!({})) })), + ) + .await + .expect_err("complete-response execution must require response codec context"); + assert!(matches!(error, FlowError::InvalidArgument(_))); + assert!(callback.host_state.continuations.lock().unwrap().is_empty()); + assert!(callback.host_state.scope_stacks.lock().unwrap().is_empty()); +} + #[tokio::test(flavor = "multi_thread")] async fn cancelling_worker_execution_expires_context_and_continuation_state() { enable_operational_logs(); @@ -1898,12 +1935,18 @@ async fn closing_worker_stream_waits_for_cancellation_and_codec_cleanup() { .await .expect("explicit close must wait for worker stream cleanup"); + let cancellation = fixture + .cancel_rx + .try_recv() + .expect("close must wait for worker cancellation"); + assert_eq!(cancellation.invocation_id, fixture.invocation_id); + assert!(cancellation.reason.contains("stopped consuming")); assert_request_codec_expired( &fixture.callback.host_state, &fixture.request_id, &fixture.invocation_id, ); - assert_worker_stream_cancelled_and_cleaned(&mut fixture).await; + assert_worker_stream_dropped_and_cleaned(&mut fixture).await; } #[tokio::test(flavor = "multi_thread")] diff --git a/crates/core/tests/unit/llm_api_tests.rs b/crates/core/tests/unit/llm_api_tests.rs index 9919cd17b..3443de023 100644 --- a/crates/core/tests/unit/llm_api_tests.rs +++ b/crates/core/tests/unit/llm_api_tests.rs @@ -124,6 +124,19 @@ fn multi_turn_annotation() -> Arc { Arc::new(OpenAIChatCodec.decode(&multi_turn_request()).unwrap()) } +fn execution_response(content: &str) -> Json { + json!({ + "id": "chatcmpl-execution-context", + "object": "chat.completion", + "model": "demo", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": content}, + "finish_reason": "stop" + }] + }) +} + fn assert_openai_execution_context(context: &LlmExecutionContext) { assert_eq!( context.request_codec().codec(), @@ -368,16 +381,7 @@ fn managed_execution_codec_context_decodes_encodes_and_decodes_response() { json!("rewritten by interceptor") ); assert_eq!(request.content["provider_only"], json!({"preserved": true})); - Ok(json!({ - "id": "chatcmpl-test", - "object": "chat.completion", - "model": "demo", - "choices": [{ - "index": 0, - "message": {"role": "assistant", "content": "accepted"}, - "finish_reason": "stop" - }] - })) + Ok(execution_response("accepted")) }) })) .codec(Arc::new(OpenAIChatCodec)) @@ -429,16 +433,7 @@ fn unary_execution_codec_facades_expire_after_interceptor_settlement() { ) .unwrap(); - let response = json!({ - "id": "chatcmpl-expiry", - "object": "chat.completion", - "model": "demo", - "choices": [{ - "index": 0, - "message": {"role": "assistant", "content": "accepted"}, - "finish_reason": "stop" - }] - }); + let response = execution_response("accepted"); let provider_response = response.clone(); tokio::runtime::Runtime::new().unwrap().block_on(async { let actual = execute_openai_call( diff --git a/crates/core/tests/unit/llm_execution_context_tests.rs b/crates/core/tests/unit/llm_execution_context_tests.rs new file mode 100644 index 000000000..d33830299 --- /dev/null +++ b/crates/core/tests/unit/llm_execution_context_tests.rs @@ -0,0 +1,116 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use super::*; +use crate::api::runtime::{BuiltinLlmCodec, LlmCodecIdentity}; + +struct LeaseProbeCodec { + allows_estimated_cost: bool, +} + +impl LlmCodec for LeaseProbeCodec { + fn codec_identity(&self) -> LlmCodecIdentity { + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + } + + fn decode(&self, _request: &LlmRequest) -> Result { + Ok(AnnotatedLlmRequest::default()) + } + + fn encode( + &self, + _annotated: &AnnotatedLlmRequest, + original: &LlmRequest, + ) -> Result { + Ok(original.clone()) + } +} + +impl LlmResponseCodec for LeaseProbeCodec { + fn codec_identity(&self) -> LlmCodecIdentity { + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + } + + fn allows_estimated_cost(&self, _response: &Json) -> bool { + self.allows_estimated_cost + } + + fn decode_response(&self, _response: &Json) -> Result { + Ok(AnnotatedLlmResponse::default()) + } +} + +#[test] +fn retained_codecs_work_while_active_and_expire_with_their_lease() { + let backing = Arc::new(LeaseProbeCodec { + allows_estimated_cost: false, + }); + let backing_probe = Arc::downgrade(&backing); + let request_codec: Arc = backing.clone(); + let response_codec: Arc = backing.clone(); + let context = LlmExecutionContext::for_unary_codecs(Some(request_codec), &Some(response_codec)); + drop(backing); + + let (leased_context, guard) = context.lease(); + let retained_request = leased_context.request_codec().resolve_codec().unwrap(); + let retained_response = leased_context + .response_codec() + .and_then(LlmSanitizeResponseContext::resolve_codec) + .unwrap(); + drop(leased_context); + drop(context); + + let request = LlmRequest { + headers: serde_json::Map::new(), + content: Json::Null, + }; + let annotated = retained_request.decode(&request).unwrap(); + assert_eq!( + retained_request.encode(&annotated, &request).unwrap(), + request + ); + assert_eq!( + retained_response.decode_response(&Json::Null).unwrap(), + AnnotatedLlmResponse::default() + ); + assert!(!retained_response.allows_estimated_cost(&Json::Null)); + assert!(backing_probe.upgrade().is_some()); + + drop(guard); + + assert!(backing_probe.upgrade().is_none()); + assert_eq!( + retained_request.codec_identity(), + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert_eq!( + retained_response.codec_identity(), + LlmCodecIdentity::BuiltIn(BuiltinLlmCodec::OpenAiChat) + ); + assert!(!retained_response.allows_estimated_cost(&Json::Null)); + assert!(matches!( + retained_request.decode(&request), + Err(FlowError::InvalidArgument(_)) + )); + assert!(matches!( + retained_response.decode_response(&Json::Null), + Err(FlowError::InvalidArgument(_)) + )); +} + +#[test] +fn estimated_cost_defaults_to_false_after_expiry() { + let response_codec: Arc = Arc::new(LeaseProbeCodec { + allows_estimated_cost: true, + }); + let context = LlmExecutionContext::for_unary_codecs(None, &Some(response_codec)); + let (leased_context, guard) = context.lease(); + let retained = leased_context + .response_codec() + .and_then(LlmSanitizeResponseContext::resolve_codec) + .unwrap(); + + assert!(retained.allows_estimated_cost(&Json::Null)); + drop(guard); + assert!(!retained.allows_estimated_cost(&Json::Null)); +} diff --git a/crates/plugin/README.md b/crates/plugin/README.md index dcda060cc..c4a60784a 100644 --- a/crates/plugin/README.md +++ b/crates/plugin/README.md @@ -32,7 +32,7 @@ the dynamic-library boundary on the stable C-compatible ABI. | `PluginContext` | Installs component-owned subscribers, guardrails, intercepts, continuations, and streams. | | `PluginRuntime` | Emits marks and manages Relay-owned scopes and scope stacks through typed host helpers. | | `nemo_relay_plugin!` | Exports the one versioned native entry point used by the loader. | -| Native ABI v7 | Keeps C-compatible host and plugin tables behind the safe Rust interface. Relay 0.10 adds directional codec context for LLM execution callbacks and requires native plugins to rebuild against that layout. | +| Native ABI v7 | Keeps C-compatible host and plugin tables behind the safe Rust interface. Relay 0.10 adds request and response codec context to LLM execution callbacks and requires native plugins to rebuild against that layout. | | Typed async middleware | Drives guardrails, sanitizers, and intercepts on a per-component SDK-owned Tokio executor. Subscribers and raw ABI registrations remain synchronous. | | Async continuations and streams | `ToolNext`, `LlmNext`, and `LlmStreamNext` support repeated or concurrent downstream calls. Streaming LLM continuations use a pull-based host handle. | | Tool results | `ToolNext` returns `ToolExecutionResult`, which keeps an application result and optional annotation together. | @@ -105,7 +105,7 @@ table. Under Relay 0.9, typed async plugins that did not use this registration could retain `compat.relay = ">=0.8.0,<1.0"`. Relay 0.10 advances the internal table to ABI v7 and makes -`LlmExecutionContext` part of every unary and streaming LLM execution callback. +`LlmExecutionContext` part of every non-streaming and streaming LLM execution callback. Because this changes callback layouts, the 0.10 host rejects v2-v6 tables. Rebuild every plugin with the 0.10 v7 SDK and set `compat.relay = ">=0.10.0,<1.0"`. The authored manifest contract remains diff --git a/crates/plugin/src/async_sdk.rs b/crates/plugin/src/async_sdk.rs index 968921cc2..264ef7588 100644 --- a/crates/plugin/src/async_sdk.rs +++ b/crates/plugin/src/async_sdk.rs @@ -244,18 +244,18 @@ impl CompletionRef { fn execution_request_context( self, + host: Arc, codec: LlmCodecIdentity, resolved: bool, - ) -> Result> { + ) -> Result { let resolved = if resolved { - let status = unsafe { (self.host.0.async_completion_retain)(self.raw) }; + let status = unsafe { (host.async_completion_retain)(self.raw) }; status_result(status, "retain native async completion capability")?; Some(LlmExecutionRequestCodec { owner: LlmExecutionRequestCodecOwner::Completion { - host: self.host.0, + host, completion: self.raw, }, - _lifetime: PhantomData, }) } else { None @@ -265,16 +265,16 @@ impl CompletionRef { fn execution_response_context( self, + host: Arc, codec: LlmCodecIdentity, resolved: bool, - ) -> Result> { + ) -> Result { let resolved = if resolved { - let status = unsafe { (self.host.0.async_completion_retain)(self.raw) }; + let status = unsafe { (host.async_completion_retain)(self.raw) }; status_result(status, "retain native async completion capability")?; Some(LlmExecutionResponseCodec { - host: self.host.0, + host, completion: self.raw, - _lifetime: PhantomData, }) } else { None @@ -283,9 +283,8 @@ impl CompletionRef { } } -#[derive(Clone, Copy)] struct StreamRef { - host: HostV7, + host: Arc, raw: *const NemoRelayNativeAsyncStream, } @@ -296,16 +295,15 @@ impl StreamRef { self, codec: LlmCodecIdentity, resolved: bool, - ) -> Result> { + ) -> Result { let resolved = if resolved { - let status = unsafe { (self.host.0.async_stream_retain)(self.raw) }; + let status = unsafe { (self.host.async_stream_retain)(self.raw) }; status_result(status, "retain native async stream capability")?; Some(LlmExecutionRequestCodec { owner: LlmExecutionRequestCodecOwner::Stream { - host: self.host.0, + host: self.host, stream: self.raw, }, - _lifetime: PhantomData, }) } else { None @@ -611,17 +609,13 @@ unsafe extern "C" fn unary_next_callback( } type UnaryFuture = Pin> + Send>>; -type UnaryAdapter = dyn Fn( - Json, - Option>, - Option>, - CompletionRef, - ) -> UnaryFuture +type UnaryAdapter = dyn Fn(Json, Option, Option>, CompletionRef) -> UnaryFuture + Send + Sync; struct UnaryCallbackState { host: HostV4, + execution_host: Option>, executor: Arc, adapter: Box, } @@ -676,7 +670,18 @@ unsafe fn unary_trampoline_impl( let invocation = read_json_value(&state.host.0.v3.v1, invocation_json, "async invocation") .map_err(|status| format!("invalid async invocation: {status:?}")); let context = (!context.is_null()) - .then(|| llm_execution_context_from_completion(completion_ref, unsafe { &*context })) + .then(|| { + llm_execution_context_from_completion( + completion_ref, + Arc::clone( + state + .execution_host + .as_ref() + .expect("LLM execution callbacks have a codec host"), + ), + unsafe { &*context }, + ) + }) .transpose(); let binding = ScopePollBinding::capture(state.host.0.v3.v1); let future = catch_unwind(AssertUnwindSafe(|| match (invocation, context) { @@ -956,11 +961,11 @@ struct EventMetadataInvocation { } type StreamFuture = Pin> + Send>>; -type StreamAdapter = - dyn Fn(Json, LlmExecutionContext<'static>, LlmStreamNext) -> StreamFuture + Send + Sync; +type StreamAdapter = dyn Fn(Json, LlmExecutionContext, LlmStreamNext) -> StreamFuture + Send + Sync; struct StreamCallbackState { host: HostV7, + codec_host: Arc, executor: Arc, adapter: Box, } @@ -1076,7 +1081,7 @@ unsafe extern "C" fn stream_trampoline( .and_then(|context| { llm_stream_execution_context_from_native( StreamRef { - host: state.host, + host: Arc::clone(&state.codec_host), raw: stream, }, context, @@ -1220,18 +1225,20 @@ fn execution_codec_identity( fn llm_execution_context_from_completion( completion: CompletionRef, + host: Arc, context: &NemoRelayNativeLlmExecutionContext, -) -> Result> { +) -> Result { let request = context.request_codec; - let host = &completion.host.0.v3.v1; let request_codec = completion.execution_request_context( - execution_codec_identity(host, request.codec_kind, request.codec_id)?, + Arc::clone(&host), + execution_codec_identity(&host.v3.v1, request.codec_kind, request.codec_id)?, !request.codec.is_null(), )?; let response_codec = unsafe { context.response_codec.as_ref() } .map(|response| { completion.execution_response_context( - execution_codec_identity(host, response.codec_kind, response.codec_id)?, + Arc::clone(&host), + execution_codec_identity(&host.v3.v1, response.codec_kind, response.codec_id)?, !response.codec.is_null(), ) }) @@ -1245,16 +1252,17 @@ fn llm_execution_context_from_completion( fn llm_stream_execution_context_from_native( stream: StreamRef, context: &NemoRelayNativeLlmExecutionContext, -) -> Result> { +) -> Result { if !context.response_codec.is_null() { return Err("native LLM stream execution context exposed a response codec".into()); } let request = context.request_codec; - let host = &stream.host.0.v6.v5.v4.v3.v1; - let request_codec = stream.execution_request_context( - execution_codec_identity(host, request.codec_kind, request.codec_id)?, - !request.codec.is_null(), + let codec = execution_codec_identity( + &stream.host.v6.v5.v4.v3.v1, + request.codec_kind, + request.codec_id, )?; + let request_codec = stream.execution_request_context(codec, !request.codec.is_null())?; Ok(LlmExecutionContext { request_codec, response_codec: None, @@ -1294,6 +1302,7 @@ impl PluginContext<'_> { ) -> Result<()> { let state = Box::into_raw(Box::new(UnaryCallbackState { host: self.host_v4()?, + execution_host: None, executor: Arc::clone(&self.executor), adapter, })); @@ -1325,8 +1334,10 @@ impl PluginContext<'_> { priority: i32, adapter: Box, ) -> Result<()> { + let host = self.host_v4()?; let state = Box::into_raw(Box::new(UnaryCallbackState { - host: self.host_v4()?, + host, + execution_host: Some(Arc::new(host.0)), executor: Arc::clone(&self.executor), adapter, })); @@ -1760,10 +1771,7 @@ impl PluginContext<'_> { callback: F, ) -> Result<()> where - F: Fn(String, LlmRequest, LlmExecutionContext<'static>, LlmNext) -> Fut - + Send - + Sync - + 'static, + F: Fn(String, LlmRequest, LlmExecutionContext, LlmNext) -> Fut + Send + Sync + 'static, Fut: Future> + Send + 'static, { let callback = Arc::new(callback); @@ -1794,15 +1802,17 @@ impl PluginContext<'_> { callback: F, ) -> Result<()> where - F: Fn(String, LlmRequest, LlmExecutionContext<'static>, LlmStreamNext) -> Fut + F: Fn(String, LlmRequest, LlmExecutionContext, LlmStreamNext) -> Fut + Send + Sync + 'static, Fut: Future> + Send + 'static, { let callback = Arc::new(callback); + let host = self.host_v7()?; let state = Box::into_raw(Box::new(StreamCallbackState { - host: self.host_v7()?, + host, + codec_host: Arc::new(host.0), executor: Arc::clone(&self.executor), adapter: Box::new(move |value, context, next| { let callback = Arc::clone(&callback); diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index eac749dd1..d4084871a 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -92,38 +92,38 @@ unsafe impl Send for LlmSanitizeResponseContext<'_> {} /// Per-call codec context delivered to an LLM execution intercept. /// -/// Request codec context is always available. Unary execution also supplies a +/// Request codec context is always available. Non-streaming execution also supplies a /// response codec context, while streaming execution leaves it unavailable /// until Relay has a completed-response streaming codec contract. -pub struct LlmExecutionContext<'a> { - request_codec: LlmRequestCodecContext<'a>, - response_codec: Option>, +pub struct LlmExecutionContext { + request_codec: LlmRequestCodecContext, + response_codec: Option, } /// Request codec context for one LLM execution intercept invocation. -pub struct LlmRequestCodecContext<'a> { +pub struct LlmRequestCodecContext { /// Identity of the active request codec. pub codec: LlmCodecIdentity, - resolved: Option>, + resolved: Option, } -/// Response codec context for one unary LLM execution intercept invocation. -pub struct LlmResponseCodecContext<'a> { +/// Response codec context for one non-streaming LLM execution intercept invocation. +pub struct LlmResponseCodecContext { /// Identity of the active response codec. pub codec: LlmCodecIdentity, - resolved: Option>, + resolved: Option, } -impl<'a> LlmExecutionContext<'a> { +impl LlmExecutionContext { /// Return the active request codec context. #[must_use] - pub fn request_codec(&self) -> &LlmRequestCodecContext<'a> { + pub fn request_codec(&self) -> &LlmRequestCodecContext { &self.request_codec } - /// Return the unary response codec context, or `None` for streaming execution. + /// Return the completed-response codec context, or `None` for streaming execution. #[must_use] - pub fn response_codec(&self) -> Option<&LlmResponseCodecContext<'a>> { + pub fn response_codec(&self) -> Option<&LlmResponseCodecContext> { self.response_codec.as_ref() } } @@ -165,7 +165,10 @@ pub struct NemoRelayNativeString { _marker: PhantomData<(*mut u8, PhantomPinned)>, } -/// Opaque callback-scoped request codec capability owned by the host. +/// Opaque host-owned request codec capability. +/// +/// Its valid lifetime is defined by the callback that receives it. See +/// [`NemoRelayNativeLlmExecutionContext`] for the streaming exception. #[repr(C)] pub struct NemoRelayNativeLlmRequestCodec { _private: [u8; 0], @@ -234,7 +237,7 @@ pub struct NemoRelayNativeLlmRequestCodecContext { pub codec: *const NemoRelayNativeLlmRequestCodec, } -/// Response codec context passed to a native unary LLM execution intercept. +/// Response codec context passed to a native non-streaming LLM execution intercept. #[repr(C)] #[derive(Debug, Clone, Copy)] pub struct NemoRelayNativeLlmResponseCodecContext { @@ -247,12 +250,19 @@ pub struct NemoRelayNativeLlmResponseCodecContext { } /// Codec context passed to native LLM execution intercept callbacks. +/// +/// The context and codec IDs are borrowed for the callback. A synchronous +/// streaming callback may retain the request codec pointer only in the state +/// of its returned [`NemoRelayNativeLlmStreamV1`]; it remains valid until Relay +/// invokes that stream's drop callback. Other raw callbacks must not retain +/// codec pointers. Asynchronous streaming callbacks must instead retain the +/// output stream and use the v7 stream codec functions. #[repr(C)] #[derive(Debug, Clone, Copy)] pub struct NemoRelayNativeLlmExecutionContext { /// Request codec context, always present. pub request_codec: NemoRelayNativeLlmRequestCodecContext, - /// Unary response codec context, or null for streaming execution. + /// Completed-response codec context, or null for streaming execution. pub response_codec: *const NemoRelayNativeLlmResponseCodecContext, } @@ -351,49 +361,48 @@ impl LlmSanitizeResponseCodec<'_> { enum LlmExecutionRequestCodecOwner { Completion { - host: NemoRelayNativeHostApiV4, + host: Arc, completion: *const NemoRelayNativeAsyncCompletion, }, Stream { - host: NemoRelayNativeHostApiV7, + host: Arc, stream: *const NemoRelayNativeAsyncStream, }, } /// Invocation-lifetime request codec facade for an LLM execution intercept. -pub struct LlmExecutionRequestCodec<'a> { +pub struct LlmExecutionRequestCodec { owner: LlmExecutionRequestCodecOwner, - _lifetime: PhantomData<&'a NemoRelayNativeLlmRequestCodec>, } // SAFETY: construction retains the completion or stream that owns the codec. // Host operations are thread-safe and reject calls after that owner settles. -unsafe impl Send for LlmExecutionRequestCodec<'_> {} -unsafe impl Sync for LlmExecutionRequestCodec<'_> {} +unsafe impl Send for LlmExecutionRequestCodec {} +unsafe impl Sync for LlmExecutionRequestCodec {} -impl Drop for LlmExecutionRequestCodec<'_> { +impl Drop for LlmExecutionRequestCodec { fn drop(&mut self) { - match self.owner { + match &self.owner { LlmExecutionRequestCodecOwner::Completion { host, completion } => unsafe { - (host.v3.async_completion_release)(completion) + (host.v3.async_completion_release)(*completion) }, LlmExecutionRequestCodecOwner::Stream { host, stream } => unsafe { - (host.v6.v5.v4.v3.async_stream_release)(stream) + (host.v6.v5.v4.v3.async_stream_release)(*stream) }, } } } -impl LlmExecutionRequestCodec<'_> { +impl LlmExecutionRequestCodec { /// Decode an opaque request into Relay's normalized request model. pub fn decode(&self, request: &LlmRequest) -> Result { - match self.owner { + match &self.owner { LlmExecutionRequestCodecOwner::Completion { host, completion } => { native_codec_call(&host.v3.v1, |out| unsafe { let request = HostString::from_json(&host.v3.v1, request) .ok_or_else(|| "failed to serialize LLM request".to_string())?; let status = (host.async_completion_llm_request_codec_decode)( - completion, + *completion, request.as_ptr(), out, ); @@ -404,8 +413,11 @@ impl LlmExecutionRequestCodec<'_> { native_codec_call(&host.v6.v5.v4.v3.v1, |out| unsafe { let request = HostString::from_json(&host.v6.v5.v4.v3.v1, request) .ok_or_else(|| "failed to serialize LLM request".to_string())?; - let status = - (host.async_stream_llm_request_codec_decode)(stream, request.as_ptr(), out); + let status = (host.async_stream_llm_request_codec_decode)( + *stream, + request.as_ptr(), + out, + ); codec_status(&host.v6.v5.v4.v3.v1, status) }) } @@ -418,7 +430,7 @@ impl LlmExecutionRequestCodec<'_> { annotated: &AnnotatedLlmRequest, original: &LlmRequest, ) -> Result { - match self.owner { + match &self.owner { LlmExecutionRequestCodecOwner::Completion { host, completion } => { native_codec_call(&host.v3.v1, |out| unsafe { let annotated = HostString::from_json(&host.v3.v1, annotated) @@ -426,7 +438,7 @@ impl LlmExecutionRequestCodec<'_> { let original = HostString::from_json(&host.v3.v1, original) .ok_or_else(|| "failed to serialize original request".to_string())?; let status = (host.async_completion_llm_request_codec_encode)( - completion, + *completion, annotated.as_ptr(), original.as_ptr(), out, @@ -441,7 +453,7 @@ impl LlmExecutionRequestCodec<'_> { let original = HostString::from_json(&host.v6.v5.v4.v3.v1, original) .ok_or_else(|| "failed to serialize original request".to_string())?; let status = (host.async_stream_llm_request_codec_encode)( - stream, + *stream, annotated.as_ptr(), original.as_ptr(), out, @@ -453,25 +465,24 @@ impl LlmExecutionRequestCodec<'_> { } } -/// Invocation-lifetime response codec facade for a unary LLM execution intercept. -pub struct LlmExecutionResponseCodec<'a> { - host: NemoRelayNativeHostApiV4, +/// Response codec for one non-streaming LLM execution intercept. +pub struct LlmExecutionResponseCodec { + host: Arc, completion: *const NemoRelayNativeAsyncCompletion, - _lifetime: PhantomData<&'a NemoRelayNativeLlmResponseCodec>, } // SAFETY: construction retains the completion that owns the codec. Host // operations are thread-safe and reject calls after that completion settles. -unsafe impl Send for LlmExecutionResponseCodec<'_> {} -unsafe impl Sync for LlmExecutionResponseCodec<'_> {} +unsafe impl Send for LlmExecutionResponseCodec {} +unsafe impl Sync for LlmExecutionResponseCodec {} -impl Drop for LlmExecutionResponseCodec<'_> { +impl Drop for LlmExecutionResponseCodec { fn drop(&mut self) { unsafe { (self.host.v3.async_completion_release)(self.completion) }; } } -impl LlmExecutionResponseCodec<'_> { +impl LlmExecutionResponseCodec { /// Decode an opaque response into Relay's normalized response model. pub fn decode(&self, response: &Json) -> Result { native_codec_call(&self.host.v3.v1, |out| unsafe { @@ -487,18 +498,18 @@ impl LlmExecutionResponseCodec<'_> { } } -impl LlmRequestCodecContext<'_> { +impl LlmRequestCodecContext { /// Resolve the active request codec capability. #[must_use] - pub fn resolve_codec(&self) -> Option<&LlmExecutionRequestCodec<'_>> { + pub fn resolve_codec(&self) -> Option<&LlmExecutionRequestCodec> { self.resolved.as_ref() } } -impl LlmResponseCodecContext<'_> { +impl LlmResponseCodecContext { /// Resolve the active response codec capability. #[must_use] - pub fn resolve_codec(&self) -> Option<&LlmExecutionResponseCodec<'_>> { + pub fn resolve_codec(&self) -> Option<&LlmExecutionResponseCodec> { self.resolved.as_ref() } } @@ -1217,7 +1228,7 @@ pub type NemoRelayNativeAsyncNextStreamCb = unsafe extern "C" fn( done: bool, ) -> bool; -/// Receives one completion from a unary execution-continuation invocation. +/// Receives one completion from a non-streaming execution-continuation invocation. /// /// Exactly one of `value_json` and `error` is non-null. The callback owns its /// `user_data` and is invoked exactly once after a successful @@ -1403,7 +1414,7 @@ pub struct NemoRelayNativeHostApiV3 { free_fn: NemoRelayNativeFreeFn, ) -> NemoRelayStatus, - /// Invokes a unary execution continuation with an independent result sink. + /// Invokes a non-streaming execution continuation with an independent result sink. /// /// Unlike the legacy completion-coupled `async_next_invoke`, this hook may /// be called repeatedly or concurrently with distinct `user_data`. For a diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index e549d353b..060d83d54 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -1860,6 +1860,14 @@ fn test_host() -> NemoRelayNativeHostApiV1 { } } +fn wait_until(message: &str, mut condition: impl FnMut() -> bool) { + let deadline = Instant::now() + Duration::from_secs(5); + while !condition() { + assert!(Instant::now() < deadline, "{message}"); + std::thread::yield_now(); + } +} + #[derive(Debug)] struct MockAsyncCompletion { settled: Mutex>>, @@ -1899,11 +1907,9 @@ impl MockAsyncCompletion { } fn wait_for_release(&self) { - let deadline = Instant::now() + Duration::from_secs(5); - while self.releases.load(Ordering::SeqCst) == 0 { - assert!(Instant::now() < deadline, "completion was not released"); - std::thread::yield_now(); - } + wait_until("completion was not released", || { + self.releases.load(Ordering::SeqCst) != 0 + }); } } @@ -1971,11 +1977,9 @@ impl MockAsyncOutput { } fn wait_for_release(&self) { - let deadline = Instant::now() + Duration::from_secs(5); - while self.releases.load(Ordering::SeqCst) == 0 { - assert!(Instant::now() < deadline, "async output was not released"); - std::thread::yield_now(); - } + wait_until("async output was not released", || { + self.releases.load(Ordering::SeqCst) != 0 + }); } } @@ -4882,14 +4886,9 @@ fn typed_async_unary_execution_codecs_expire_after_completion_settles() { .is_err() ); drop(context); - let deadline = Instant::now() + Duration::from_secs(5); - while completion.releases.load(Ordering::SeqCst) < 3 { - assert!( - Instant::now() < deadline, - "retained codec facades were not released" - ); - std::thread::yield_now(); - } + wait_until("retained codec facades were not released", || { + completion.releases.load(Ordering::SeqCst) >= 3 + }); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; } @@ -4962,14 +4961,9 @@ fn typed_async_stream_execution_codec_expires_after_stream_finishes() { .is_err() ); drop(context); - let deadline = Instant::now() + Duration::from_secs(5); - while output.releases.load(Ordering::SeqCst) < 2 { - assert!( - Instant::now() < deadline, - "retained stream codec facade was not released" - ); - std::thread::yield_now(); - } + wait_until("retained stream codec facade was not released", || { + output.releases.load(Ordering::SeqCst) >= 2 + }); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; } @@ -5012,14 +5006,10 @@ fn typed_async_llm_sanitize_context_decodes_oci_genai_builtin_identity() { // The context's retained codec capability is released on the SDK executor // after the result completion is delivered, so poll instead of asserting // immediately. - let deadline = Instant::now() + Duration::from_secs(5); - while live_host_strings() != 0 { - assert!( - Instant::now() < deadline, - "host strings were not released after the sanitize invocation" - ); - std::thread::yield_now(); - } + wait_until( + "host strings were not released after the sanitize invocation", + || live_host_strings() == 0, + ); } #[test] @@ -5062,14 +5052,10 @@ fn typed_async_llm_sanitize_context_decodes_all_builtin_identities() { ); unsafe { registration.free() }; - let deadline = Instant::now() + Duration::from_secs(5); - while live_host_strings() != 0 { - assert!( - Instant::now() < deadline, - "host strings were not released after the sanitize invocation" - ); - std::thread::yield_now(); - } + wait_until( + "host strings were not released after the sanitize invocation", + || live_host_strings() == 0, + ); } } @@ -5340,14 +5326,9 @@ fn typed_async_cancellation_drops_future_and_releases_owned_handles() { Ok(NemoRelayNativeAsyncCallbackState::Pending) ); completion.wait_for_release(); - let deadline = Instant::now() + Duration::from_secs(5); - while future_drops.load(Ordering::SeqCst) == 0 || next.releases.load(Ordering::SeqCst) == 0 { - assert!( - Instant::now() < deadline, - "cancelled state was not reclaimed" - ); - std::thread::yield_now(); - } + wait_until("cancelled state was not reclaimed", || { + future_drops.load(Ordering::SeqCst) != 0 && next.releases.load(Ordering::SeqCst) != 0 + }); assert!(completion.settled.lock().unwrap().is_none()); assert_eq!(completion.releases.load(Ordering::SeqCst), 1); assert_eq!(future_drops.load(Ordering::SeqCst), 1); @@ -5389,20 +5370,14 @@ fn typed_async_cancellation_while_awaiting_reclaims_future() { ) }; unsafe { (host.v3.v1.string_free)(invocation) }; - let deadline = Instant::now() + Duration::from_secs(5); - while !started.load(Ordering::SeqCst) { - assert!(Instant::now() < deadline, "callback future never started"); - std::thread::yield_now(); - } + wait_until("callback future never started", || { + started.load(Ordering::SeqCst) + }); completion.cancelled.store(true, Ordering::SeqCst); completion.wait_for_release(); - while future_drops.load(Ordering::SeqCst) == 0 { - assert!( - Instant::now() < deadline, - "cancelled future was not dropped" - ); - std::thread::yield_now(); - } + wait_until("cancelled future was not dropped", || { + future_drops.load(Ordering::SeqCst) != 0 + }); assert!(completion.settled.lock().unwrap().is_none()); assert_eq!(completion.releases.load(Ordering::SeqCst), 1); assert_eq!(future_drops.load(Ordering::SeqCst), 1); @@ -5559,11 +5534,9 @@ fn typed_async_stream_cancellation_while_polling_releases_output() { ) }; unsafe { (host_v1.string_free)(invocation) }; - let deadline = Instant::now() + Duration::from_secs(5); - while !started.load(Ordering::SeqCst) { - assert!(Instant::now() < deadline, "returned stream was not polled"); - std::thread::yield_now(); - } + wait_until("returned stream was not polled", || { + started.load(Ordering::SeqCst) + }); output.cancelled.store(true, Ordering::SeqCst); output.wait_for_release(); assert_eq!(next.releases.load(Ordering::SeqCst), 1); diff --git a/crates/worker-proto/README.md b/crates/worker-proto/README.md index 2b38e45a0..5ed80ee2f 100644 --- a/crates/worker-proto/README.md +++ b/crates/worker-proto/README.md @@ -30,12 +30,12 @@ for earlier releases must regenerate their bindings, rebuild, and declare `compat.relay` beginning at `0.8.0`. `ToolNext` returns `ToolExecutionResultResponse`, and tool execution intercepts use structural `ToolExecutionInterceptOutcome` messages. -Relay 0.10 adds directional execution codec context to `LlmInvocation` and makes -that context required by the 0.10 worker SDK's unary and streaming LLM execution -callbacks. The protobuf field is additive, but the SDK callback contract is a -release-level source break. Rebuild workers, regenerate custom bindings, and use -a `compat.relay` lower bound of `0.10.0`. The protocol identifier remains -`grpc-v1`. +Relay 0.10 adds execution codec context to `LlmInvocation`. The protobuf field +is additive. Workers that register non-streaming or streaming LLM execution +callbacks must update their callback signatures, rebuild with the 0.10 SDK, and +set `compat.relay` to begin at `0.10.0`. Custom workers that register either +surface must regenerate their bindings to read the field. The protocol +identifier remains `grpc-v1`. ## Protocol Surface @@ -48,7 +48,7 @@ a `compat.relay` lower bound of `0.10.0`. The protocol identifier remains | Tool results | `ToolNext` returns `ToolExecutionResultResponse`, and `ToolExecutionInterceptResult` returns `ToolExecutionInterceptOutcome`. Both preserve the application result and optional annotation. Intercept outcomes also include ordered pending marks. These fields use lossless protobuf `JsonValue` wrappers rather than `google.protobuf.Value`. | | Mark options | `EmitMarkRequest.data_schema` carries a `nemo.relay.DataSchema@1` envelope, `severity` carries the log severity, and `category` carries an optional semantic mark category. Omitting these fields preserves legacy behavior. | | Runtime diagnostics | Authenticated `GetRuntimeDiagnostics` returns a bounded active-host `{ code, message, count }` snapshot. Older hosts return gRPC `UNIMPLEMENTED`. | -| LLM execution codec context | `LlmInvocation.execution_codec_context` carries request identity and an invocation-scoped request capability. Unary execution also carries response identity and a response capability; streaming execution omits the response direction. | +| LLM execution codec context | `LlmInvocation.execution_codec_context` carries request identity and request codec access. Non-streaming execution also carries completed-response identity and codec access; streaming execution omits it. | ## Installation diff --git a/crates/worker/README.md b/crates/worker/README.md index e52368b6d..86872e835 100644 --- a/crates/worker/README.md +++ b/crates/worker/README.md @@ -26,23 +26,21 @@ for an earlier release must rebuild with this SDK and declare `compat.relay` beg `ToolExecutionResult`, which keeps an optional opaque annotation beside the application result. -Relay 0.10 makes LLM execution codec context part of every unary and streaming -execution callback. Rebuild workers with the 0.10 SDK and raise their -`compat.relay` lower bound to `0.10.0`. The callback receives -`LlmExecutionContext` immediately before `next`. Its request direction reports -the active codec and may resolve request decode/encode operations. Unary -callbacks also receive response decode context; streaming callbacks deliberately -receive no response codec because Relay does not decode incomplete chunks. The -wire protocol remains named `grpc-v1`. +Relay 0.10 adds `LlmExecutionContext` to non-streaming and streaming LLM +execution callbacks. Workers that register either surface must rebuild with the +0.10 SDK and set `compat.relay` to begin at `0.10.0`. The context provides +request decode and encode operations when Relay selected a codec. Non-streaming +callbacks can also decode the completed response; streaming callbacks cannot +decode incomplete chunks. The wire protocol remains `grpc-v1`. ## Authoring Surface | Surface | Role | |---|---| | `WorkerPlugin` | Defines plugin identity, validation, registration, and multiple-component behavior in the worker process. | -| `PluginContext` | Installs typed handlers for all 16 supported registration surfaces. | +| `PluginContext` | Installs typed handlers for all 17 supported registration surfaces. | | `PluginRuntime` and continuations | Emit marks, manage scopes, and call the remaining tool or LLM execution chain through the authenticated host service. | -| `LlmExecutionContext` | Reports request and unary-response codec identity and exposes invocation-scoped codec operations when Relay resolved a codec. | +| `LlmExecutionContext` | Reports the selected request and completed-response codecs and provides codec operations while the callback is active. | | Canonical tool results | Preserve application results and opaque annotations across tool callbacks and continuations. | | `serve_plugin` | Starts the Tokio gRPC server from the activation identity, local endpoints, and token supplied by Relay. | diff --git a/docs/about-nemo-relay/concepts/codecs.mdx b/docs/about-nemo-relay/concepts/codecs.mdx index 931d5f8ae..8894130e6 100644 --- a/docs/about-nemo-relay/concepts/codecs.mdx +++ b/docs/about-nemo-relay/concepts/codecs.mdx @@ -83,6 +83,21 @@ Response decoding improves observability and downstream consistency. It does not automatically change the value returned to the application unless a separate typed value codec also does so. +### Codec Access in Execution Intercepts + +LLM execution intercepts receive `LlmExecutionContext`, which exposes the codec +Relay selected before the execution chain starts. + +| Call Type | Available Codec Operations | Valid Until | +|---|---|---| +| Non-streaming | Decode and encode the request; decode the completed response | The interceptor finishes | +| Streaming | Decode and encode the request | The returned stream ends or closes | + +The request context is always present and reports `none` when Relay did not +select a codec. Runtime and opaque codecs remain usable when Relay has a codec +implementation. Rewriting a request does not select another codec, so the +replacement must remain compatible with the selected codec. + The built-in request codecs are lossless patch codecs. For an unchanged annotation, the following identity holds at the JSON-value level: diff --git a/docs/about-nemo-relay/concepts/middleware.mdx b/docs/about-nemo-relay/concepts/middleware.mdx index c62fd31ab..40689846f 100644 --- a/docs/about-nemo-relay/concepts/middleware.mdx +++ b/docs/about-nemo-relay/concepts/middleware.mdx @@ -72,6 +72,12 @@ stream extends that active lifetime until it closes, so a lazy stream adapter can call `next` while it is being consumed. A stream successfully returned by streaming `next` keeps its ordinary stream lifetime. +LLM execution intercepts also receive `LlmExecutionContext`. Non-streaming +calls expose request and completed-response codec access. Streaming calls expose +request codec access only. Relay selects the codecs before the execution chain +starts. Refer to [Provider Codecs](/about-nemo-relay/concepts/codecs#codec-access-in-execution-intercepts) +for the operations and lifetimes. + A tool execution continuation returns `ToolExecutionResult`, which contains the application-owned `result` and an optional opaque `annotation`. A forwarding intercept must preserve both fields in its `ToolExecutionInterceptOutcome` or diff --git a/docs/about-nemo-relay/release-notes/index.mdx b/docs/about-nemo-relay/release-notes/index.mdx index 64b940aed..705d3be82 100644 --- a/docs/about-nemo-relay/release-notes/index.mdx +++ b/docs/about-nemo-relay/release-notes/index.mdx @@ -35,19 +35,20 @@ changes, compatibility updates, and fixed known issues. ### LLM Execution Codec Context -Relay now passes the selected request and unary-response codec context to every +Relay now passes the selected request and completed-response codec context to every LLM execution interceptor. Policies can use Relay's host-owned codec to inspect -or safely rewrite the final provider request and decode a completed unary +or safely rewrite the final provider request and decode a completed response without importing or copying Relay's codec implementations. Streaming interceptors receive request codec access only. **Breaking change:** Callback signatures change across Rust, Python, Node.js, Go, C, native plugins, and Rust and Python gRPC workers. Python language-binding streaming and public C callbacks also gain the logical LLM -name. Native plugins must rebuild for the internal ABI v7 layout, and affected workers -must regenerate their protobuf bindings and rebuild. Authored compatibility -labels remain `native_api = "1"` and `grpc-v1`; plugin manifests must use a -Relay range that begins at 0.10 or otherwise excludes 0.9. Refer to the +name. Every native plugin must rebuild for the internal ABI v7 layout. Workers +that register an LLM execution intercept must update their callback signatures +and rebuild. Authored compatibility labels remain `native_api = "1"` and +`grpc-v1`; affected plugin manifests must use a Relay range that begins at 0.10 +or otherwise excludes 0.9. Refer to the [Migration Guides](/reference/migration-guides) for callback shapes and upgrade steps. diff --git a/docs/build-plugins/about.mdx b/docs/build-plugins/about.mdx index 352434c3d..fac8c1deb 100644 --- a/docs/build-plugins/about.mdx +++ b/docs/build-plugins/about.mdx @@ -75,14 +75,17 @@ same pair plus Relay-owned pending marks. Native and worker plugins built for an release must rebuild for this contract. Workers retain the `grpc-v1` name and protobuf package; their tool-result fields, not their protocol identity, changed. -Relay 0.10 makes codec context part of every LLM execution callback in language -bindings, native plugins, and worker SDKs. This is a source break. Native and worker -plugins must rebuild and constrain `compat.relay` to `>=0.10.0`; their authored -`native_api = "1"` and `grpc-v1` labels do not change. - -The codec capability belongs to that execution callback. Unary capabilities expire when -the callback settles; streaming request capabilities remain valid until the returned -stream completes or closes. Retaining a codec object does not extend either lifetime. +Relay 0.10 adds codec context to every LLM execution callback in language +bindings, native plugins, and worker SDKs. Every native plugin must rebuild for +ABI v7. Workers that register a non-streaming or streaming LLM execution +intercept must update their callback signatures, rebuild, and set +`compat.relay` to begin at `0.10.0`. The authored `native_api = "1"` and +`grpc-v1` labels do not change. + +Codec access belongs to that execution callback. For a non-streaming call, it +expires when the callback finishes. For a streaming call, request codec access +remains valid until the returned stream completes or closes. Retaining a codec +object does not extend either lifetime. ## Match Behavior to a Plugin diff --git a/docs/build-plugins/language-binding/register-behavior.mdx b/docs/build-plugins/language-binding/register-behavior.mdx index 7ef7e8636..79dc01ac9 100644 --- a/docs/build-plugins/language-binding/register-behavior.mdx +++ b/docs/build-plugins/language-binding/register-behavior.mdx @@ -349,8 +349,8 @@ ctx.register_llm_stream_execution_intercept( -Execution codec objects are scoped to the callback that received them. Unary -capabilities expire when that callback settles. A streaming request capability remains +Execution codec objects are scoped to the callback that received them. Non-streaming +capabilities expire when that callback finishes. A streaming request capability remains usable while the returned stream is live and expires on completion, error, close, or drop. Retaining a codec object past that boundary produces an error; it does not extend the invocation lifetime. diff --git a/docs/build-plugins/native/about.mdx b/docs/build-plugins/native/about.mdx index 5a016d63f..e369a9f8d 100644 --- a/docs/build-plugins/native/about.mdx +++ b/docs/build-plugins/native/about.mdx @@ -33,7 +33,7 @@ The checked example registers a tool execution intercept, so it declares that lower bound because an older Relay cannot load the ABI v7 function table, even when the component itself does not register an LLM execution intercept. -Relay 0.10 adds ABI v7 with directional codec context on every LLM +Relay 0.10 adds ABI v7 with request and response codec context on every LLM execution callback. Rebuild all native plugins for this release; the authored `native_api = "1"` label does not change. @@ -52,7 +52,7 @@ attributed to a particular plugin. ## What the SDK Owns The Rust SDK exports the stable entry symbol, converts host-owned JSON handles into -typed DTOs, registers all 16 plugin surfaces, and drives async middleware on one +typed DTOs, registers the supported subscriber and middleware surfaces, and drives async middleware on one SDK-owned multi-thread Tokio runtime per configured component. A plugin can set a default executor size and accept a positive `executor.worker_threads` component override. The default is two workers; change it only after measuring queued async work @@ -153,7 +153,7 @@ Follow these pages in order to build, activate, exercise, and remove the native including subscribers and all three event sanitizer surfaces. 3. Add policy and request rewriting with [Control Requests](/build-plugins/native/control-requests), preserving annotations and making priority and `break_chain` explicit. -4. Add tool, unary model, and lazy stream wrappers with [Wrap Execution](/build-plugins/native/wrap-execution). +4. Add tool, non-streaming model, and lazy stream wrappers with [Wrap Execution](/build-plugins/native/wrap-execution). 5. Verify marks, scopes, isolated stacks, cleanup, and executor control with [Runtime Events and Scopes](/build-plugins/native/runtime-events-and-scopes). 6. Consult [Native ABI Reference](/build-plugins/native/native-abi-reference) only when diff --git a/docs/build-plugins/native/native-abi-reference.mdx b/docs/build-plugins/native/native-abi-reference.mdx index aa84eac0b..5ebf47dd4 100644 --- a/docs/build-plugins/native/native-abi-reference.mdx +++ b/docs/build-plugins/native/native-abi-reference.mdx @@ -29,10 +29,10 @@ layout would be unsafe. Rebuild every native plugin for Relay 0.10 and raise its `compat.relay` lower bound to `0.10.0`. The authored manifest label remains `compat.native_api = "1"`. -ABI v7 adds directional request and unary-response codec context to raw, typed, -and asynchronous LLM execution callbacks. The generic asynchronous middleware -callback remains unchanged; ABI v7 appends an execution-specific unary registration, -and the existing execution-specific stream callback gains the context parameter. +ABI v7 adds request and completed-response codec context to raw, typed, and +asynchronous LLM execution callbacks. The generic asynchronous middleware +callback remains unchanged; ABI v7 appends an execution-specific non-streaming +registration, and the existing stream callback gains the context parameter. ABI v6 adds a host-routed operational logging function. Native plugins pass a level, optional target, message, and optional JSON-object fields; Relay applies its operational logging policy and preserves those fields in structured JSONL output. @@ -50,11 +50,11 @@ function signatures and field order are defined by the public | Table Level | Operations | |---|---| | Frozen v1/v2 prefix | Version and struct-size negotiation; host version; string allocation, access, and release; thread-local error reporting; callback-scoped LLM request decode and encode plus response decode; subscriber, five tool, six LLM, and three event-sanitizer registrations; current scope, scope push and pop, mark emission, isolated stack creation and release, thread-stack set, capture, and restore, captured-binding release, active-stack inspection, and scoped binding. | -| Frozen v3 extension | Completion resolve, reject, cancellation inspection, and release; one-shot completion-coupled continuation invocation; continuation release; generic async middleware registration; bounded output-stream push, finish, reject, cancellation inspection, and release; downstream stream invocation; stream-middleware registration; and repeated or concurrent unary continuation invocation with independent result callbacks. | +| Frozen v3 extension | Completion resolve, reject, cancellation inspection, and release; one-shot completion-coupled continuation invocation; continuation release; generic async middleware registration; bounded output-stream push, finish, reject, cancellation inspection, and release; downstream stream invocation; stream-middleware registration; and repeated or concurrent non-streaming continuation invocation with independent result callbacks. | | Frozen v4 extension | Completion-scoped LLM request decode and encode plus response decode; pull-based downstream LLM stream open, pull, cancel, and release; completion retain for typed codec facades; output-stream backpressure inspection; extended mark emission; runtime diagnostics; activation-owned runtime capability creation, retain, and release; global runtime-registration discovery; owned conditional middleware guardrail registration and deregistration; and activation-owned and runtime-discovered callback gate registration. The callback registration slots are appended after the original constant-reason slots. | | v5 extension | Context-carrying tool execution intercept registration. The callback receives one JSON object containing `tool_name`, `args`, and `tool_call_id`. | | v6 extension | Host-routed operational logging with structured fields. | -| Current v7 extension | Directional codec context for raw and typed LLM execution callbacks, an execution-specific asynchronous unary registration, context on the asynchronous stream-execution callback, and stream retain plus stream-scoped request decode and encode operations. | +| Current v7 extension | Codec context for raw and typed LLM execution callbacks, an execution-specific asynchronous non-streaming registration, context on the asynchronous stream-execution callback, and stream retain plus stream-scoped request decode and encode operations. | The prefix and descriptor layout are explicit. A plugin fills the descriptor with its stable kind, component multiplicity, opaque state, callbacks, and destructor. The host @@ -164,14 +164,14 @@ or binding, and restoration. `PluginContext::register_async_middleware_raw` registers non-stream middleware other than LLM execution. ABI v7 uses the execution-specific -`plugin_context_register_async_llm_execution_intercept` entry for asynchronous unary +`plugin_context_register_async_llm_execution_intercept` entry for asynchronous non-streaming LLM execution callbacks. Both paths can settle later. Return `Complete` only after resolving or rejecting the completion inside the callback. Return `Pending` only after retaining it. A retained completion must settle exactly once and then be released. Release every async `next` reference after its last use. -`async_next_invoke_result` supports repeated or concurrent unary continuation calls with +`async_next_invoke_result` supports repeated or concurrent non-streaming continuation calls with independent result callbacks. The older completion-coupled `async_next_invoke` is one-shot because the continuation result settles the middleware completion. Settle the owner only after every started continuation call has finished. When the owner settles or @@ -231,17 +231,18 @@ request or application response. Never retain a raw handle or resolved typed fac after the sanitizer callback ends. LLM execution context uses the same identity and opaque codec operations. Its -request direction is always present and may resolve request decode/encode. Unary +request direction is always present and may resolve request decode/encode. Non-streaming execution also carries an optional response direction with response decode. Streaming execution sets the response direction to unavailable because the response codecs operate on completed provider responses, not individual chunks. -Raw execution-context handles are borrowed for the owning callback. Typed unary -facades retain their completion and typed streaming request facades retain their -output stream, so they remain memory-safe while the callback future or returned -stream owns them. Codec calls fail after the completion or stream settles. Raw -plugins that need request codec access after an asynchronous stream callback -returns must retain the stream and use the v7 stream-scoped decode and encode -operations; release that stream reference after the last call. +Raw execution-context handles are borrowed for the callback. A synchronous raw stream +callback may store its request handle only in the state of the returned stream; Relay +keeps it valid until that stream's drop callback. Typed non-streaming facades retain +their completion and typed streaming request facades retain their output stream, so +they remain memory-safe while the callback future or returned stream owns them. Codec +calls fail after the completion or stream settles. Asynchronous raw stream callbacks +must retain the output stream and use the v7 stream-scoped decode and encode operations; +release that stream reference after the last call. ## Unload Ordering diff --git a/docs/build-plugins/native/wrap-execution.mdx b/docs/build-plugins/native/wrap-execution.mdx index a628b61c6..9dae61bbe 100644 --- a/docs/build-plugins/native/wrap-execution.mdx +++ b/docs/build-plugins/native/wrap-execution.mdx @@ -1,6 +1,6 @@ --- title: "Wrap Execution" -description: "Use native tool, unary LLM, and streaming continuations correctly." +description: "Use native tool, non-streaming LLM, and streaming continuations correctly." position: 24 --- {/* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. @@ -8,7 +8,7 @@ SPDX-License-Identifier: Apache-2.0 */} [Execution intercepts](/about-nemo-relay/concepts/middleware) receive the real request and a continuation representing the rest -of the call path. The example's `execution` group registers a tool wrapper, a unary LLM +of the call path. The example's `execution` group registers a tool wrapper, a non-streaming LLM wrapper, and an LLM stream wrapper at the configured priority. ## Tool and Unary Results @@ -20,7 +20,7 @@ example's pending mark. The conversion preserves both the application `result` and any opaque `annotation`. Relay owns the pending mark: it is emitted in the managed lifecycle and does not appear in the application-visible result. -The unary LLM wrapper has a deliberately smaller contract and returns provider-response +The non-streaming LLM wrapper has a deliberately smaller contract and returns provider-response JSON. LLM pending marks, annotations, and optimization contributions belong to the request-intercept outcome shown in [Control Requests](/build-plugins/native/control-requests), not the execution result. @@ -56,7 +56,7 @@ reads `tool_result.result["answer"]`. Relay separately emits callback returns an outcome instead of returning a plain JSON value. The continuation is reusable. A deliberate configuration or request flag in the example -can invoke unary `next` twice concurrently, await both responses, and select one. This +can invoke non-streaming `next` twice concurrently, await both responses, and select one. This demonstrates the API while preserving the cost of repetition. A repeated tool can perform its side effect twice, and a repeated model call can incur provider cost twice. Production plugins need an idempotency @@ -137,15 +137,15 @@ or cancellation. ## Verify Execution Behavior -Use the following procedure to verify unary, repeated, and streaming continuation +Use the following procedure to verify non-streaming, repeated, and streaming continuation behavior: 1. Activate the example with `execution.enabled = true`, priority 30, and `emit_pending_marks = true`. -2. Execute a tool and a unary model call. Confirm each downstream callback runs once, +2. Execute a tool and a non-streaming model call. Confirm each downstream callback runs once, the application receives only its expected result, and an additional pending mark is emitted under the managed call scope. -3. Enable the example's repeated-continuation input and confirm two downstream unary +3. Enable the example's repeated-continuation input and confirm two downstream non-streaming invocations can overlap. Verify the selected result and accounting explicitly. 4. Consume a three-chunk LLM stream one item at a time. Confirm the first transformed chunk arrives before the downstream stream completes. @@ -154,6 +154,6 @@ behavior: 6. Clear the component and repeat the calls to prove that no wrapper or pending mark remains registered. -Success means unary and stream continuations preserve scope, errors, and cancellation, +Success means non-streaming and stream continuations preserve scope, errors, and cancellation, while Relay-owned tool marks and LLM request accounting remain separate from application results. diff --git a/docs/build-plugins/package-discoverable-plugins.mdx b/docs/build-plugins/package-discoverable-plugins.mdx index 4343b0b76..0b3469abb 100644 --- a/docs/build-plugins/package-discoverable-plugins.mdx +++ b/docs/build-plugins/package-discoverable-plugins.mdx @@ -143,10 +143,10 @@ directory Relay installs into its managed environment. The digest still covers o declared artifact, so regenerate it whenever `worker.py` changes. A compiled worker has no managed Python environment and starts the executable named by `load.entrypoint`. -Relay 0.8 keeps `grpc-v1` as the worker protocol identifier, but changes its tool-result -messages. Both worker manifests therefore begin their Relay compatibility range at -`0.8.0`. Rebuild SDK workers and regenerate bindings in a custom worker before packaging; -the unchanged protocol name does not make an earlier worker wire-compatible. +Both worker examples register LLM execution intercepts, so their Relay +compatibility range begins at `0.10.0`. A worker that does not register those +surfaces does not need a 0.10 lower bound solely for the additive execution +codec context field. The worker protocol remains `grpc-v1`. ## Package and Register the Artifact diff --git a/docs/build-plugins/plugin-context.mdx b/docs/build-plugins/plugin-context.mdx index 49c8d1105..db1f0274f 100644 --- a/docs/build-plugins/plugin-context.mdx +++ b/docs/build-plugins/plugin-context.mdx @@ -30,7 +30,7 @@ language-binding, native typed, and worker plugins. | LLM | Response sanitizer | Sanitize response observability with the structured LLM context and its codec handle. | | LLM | Conditional guardrail | Allow or block real model execution from the provider request. The current callback contract does not receive a separate model name or annotation argument. | | LLM | Request intercept | Return the complete request-intercept outcome, preserving or deliberately replacing annotations. | -| LLM | Execution intercept | Wrap unary execution through a continuation and return provider-response JSON. | +| LLM | Execution intercept | Wrap non-streaming execution through a continuation and return provider-response JSON. | | LLM | Stream execution intercept | Transform response chunks while preserving order, cancellation, and error behavior. Native and worker SDKs expose a lazy stream; the Node.js language binding currently supplies the downstream chunks as an array. | Sanitize guardrails change emitted observability payloads only. They do not rewrite the diff --git a/docs/build-plugins/workers/about.mdx b/docs/build-plugins/workers/about.mdx index 8e7d516e1..159e5ea54 100644 --- a/docs/build-plugins/workers/about.mdx +++ b/docs/build-plugins/workers/about.mdx @@ -18,12 +18,14 @@ SDK worker, regenerate custom protobuf bindings, and declare `compat.relay` begi `0.8.0`. The protocol remains named `grpc-v1`; this is a changed tool-result contract, not a new protocol family. -Relay 0.10 adds `LlmExecutionContext` before the continuation in every unary and -streaming LLM execution callback. Rebuild every SDK worker, regenerate custom -protobuf bindings, and use a `compat.relay` lower bound of `0.10.0`. The request -direction can expose decode/encode operations; unary callbacks can also decode -the complete response. Streaming callbacks do not receive a response codec. -The protocol remains named `grpc-v1`. +Relay 0.10 adds `LlmExecutionContext` before the continuation in non-streaming +and streaming LLM execution callbacks. Workers that register either surface +must update their callback signatures, rebuild with the 0.10 SDK, and set +`compat.relay` to begin at `0.10.0`. Custom workers that register either surface +must regenerate their bindings to read the field. The context can decode and +encode the request; non-streaming callbacks can also decode the completed +response. Streaming callbacks do not receive a response codec. The protocol +remains `grpc-v1`. Workers can add a data schema, log severity, and semantic category to a mark, emit validated metric measurements, and read the current host-level runtime diagnostics. These diagnostics are diff --git a/docs/build-plugins/workers/grpc-v1-protocol.mdx b/docs/build-plugins/workers/grpc-v1-protocol.mdx index 50933d0e6..089a4570f 100644 --- a/docs/build-plugins/workers/grpc-v1-protocol.mdx +++ b/docs/build-plugins/workers/grpc-v1-protocol.mdx @@ -18,9 +18,11 @@ worker cannot decode the current `ToolNext` response or tool-execution outcome. Relay 0.10 retains those protocol identifiers and adds execution codec context to `LlmInvocation`. The field is protobuf-compatible, but the 0.10 Rust and Python -SDKs intentionally change every LLM execution callback to receive -`LlmExecutionContext` before its continuation. Rebuild workers, regenerate custom -bindings, and declare `compat.relay` beginning at `0.10.0`. +SDKs change non-streaming and streaming LLM execution callbacks to receive +`LlmExecutionContext` before the continuation. Workers that register either +surface must update their callback signatures, rebuild, and declare +`compat.relay` beginning at `0.10.0`. Custom workers that register either +surface must regenerate their bindings to read the field. The protocol consists of the worker-facing service implemented by the plugin process and the host-runtime service implemented by Relay. These are the current service @@ -65,7 +67,7 @@ service RelayHostRuntime { | `Health` | Confirms the authenticated activation, protocol, plugin identity, and current worker readiness. | | `Validate` | Receives component config in a JSON envelope and returns encoded diagnostics or a structured worker error. It must not register behavior. | | `Register` | Receives valid component config and returns owned registrations with local name, surface, priority, and `break_chain`. | -| `Invoke` | Dispatches one subscriber, sanitizer, guardrail, request intercept, or unary execution callback and returns the surface-appropriate result. | +| `Invoke` | Dispatches one subscriber, sanitizer, guardrail, request intercept, or non-streaming execution callback and returns the surface-appropriate result. | | `InvokeStream` | Dispatches an LLM stream execution intercept and emits incremental value or error chunks. | | `CancelInvocation` | Cooperatively cancels one active invocation ID and reports whether cancellation was accepted. Unknown, completed, and already-cancelled IDs receive a negative acknowledgment. | | `Shutdown` | Stops new work for an activation and begins orderly worker termination with a reason. | @@ -136,8 +138,8 @@ An invocation names the activation, invocation, registration, surface, optional continuation, captured scope, and token. Its payload is exactly one event, tool invocation, LLM invocation, or conditional middleware invocation. LLM sanitizer invocations additionally carry codec identity and an opaque invocation-scoped codec -capability. LLM execution invocations carry separate request and unary-response codec -context. Streaming execution carries only the request direction. +capability. LLM execution invocations carry separate request and completed-response +codec context. Streaming execution carries only the request context. For `EVENT_METADATA_INJECTOR`, Relay sends the immutable Event snapshot in `InvokeRequest.event`. The worker returns an `InvokeResponse.json` object containing @@ -267,10 +269,10 @@ invocation ends. For an LLM execution invocation, `execution_codec_context.request` is required. Its identity is present even when no codec resolves, and its capability ID is -present only when host codec operations are available. Unary execution also -includes `response`; streaming execution omits it entirely because Relay's -response codecs require a complete provider response. The worker SDK hides -both opaque capability IDs behind `resolve_codec()` proxies. +present only when host codec operations are available. Non-streaming execution +also includes `response`; streaming execution omits it because Relay's response +codecs require a complete provider response. The worker SDK hides both opaque +capability IDs behind `resolve_codec()` proxies. ## Host-Runtime Service @@ -284,7 +286,7 @@ both opaque capability IDs behind `resolve_codec()` proxies. | `CreateScopeStack` | Allocates an isolated stack and returns its opaque ID. | | `DropScopeStack` | Releases an isolated stack owned by the activation. | | `ToolNext` | Executes a tool continuation with JSON arguments and captured scope, returning a structural tool result or worker error. | -| `LlmNext` | Executes a unary LLM continuation with a request and captured scope. | +| `LlmNext` | Executes a non-streaming LLM continuation with a request and captured scope. | | `LlmStreamNext` | Executes a streaming LLM continuation and returns incremental chunks. | | `DecodeLlmCodecRequest` | Uses the invocation-scoped capability to decode a request into its annotated representation. | | `EncodeLlmCodecRequest` | Applies an annotated request to the original provider envelope. | @@ -416,6 +418,6 @@ process: only during explicit package removal. Successful protocol verification covers authentication failures, envelope schema -failures, every registration and result variant, repeated unary and incremental stream +failures, every registration and result variant, repeated non-streaming and incremental stream continuations, cancellation before and during callbacks, codec capability expiry, host runtime ownership errors, health, and orderly plus forced shutdown. diff --git a/docs/build-plugins/workers/middleware-and-continuations.mdx b/docs/build-plugins/workers/middleware-and-continuations.mdx index 2d5446729..d4ece4e5a 100644 --- a/docs/build-plugins/workers/middleware-and-continuations.mdx +++ b/docs/build-plugins/workers/middleware-and-continuations.mdx @@ -72,7 +72,7 @@ context.register_event_metadata_injector( -Relay sends the immutable Event snapshot through the existing unary `Invoke` RPC. It +Relay sends the immutable Event snapshot through the existing non-streaming `Invoke` RPC. It validates and merges accepted additions before Event sanitizers run. A callback error omits that callback's additions without dropping the Event. Stopping the worker removes the registration. @@ -279,7 +279,7 @@ originating tool call. `ToolNext`, `LlmNext`, and `LlmStreamNext` are host proxies identified by an opaque continuation ID. Calling one issues a host-runtime RPC under the scope snapshot captured -for that worker invocation. A callback can call a unary proxy zero, one, or multiple +for that worker invocation. A callback can call a non-streaming proxy zero, one, or multiple times, including concurrently when the SDK type permits cloning. Each call can repeat side effects, provider charges, events, and downstream middleware. @@ -293,7 +293,7 @@ pending marks. The example uses one ordinary wrapper and one explicitly requested concurrent path. It does not retry implicitly. Tool pending marks remain in the tool execution outcome; LLM annotations, pending marks, and optimization contributions remain in the request -intercept outcome. The unary execution callback returns only provider-response JSON. +intercept outcome. The non-streaming execution callback returns only provider-response JSON. `LlmStreamNext` returns a remote stream. The worker transforms each chunk as it arrives and yields immediately. If the host cancels the invocation or the consumer abandons the @@ -465,7 +465,7 @@ instead of implying that individual chunks can be decoded as completed responses. Codec proxies are invocation-scoped and must not be retained after the callback or returned stream finishes. -The second unary result is awaited even though the first response is selected. That +The second non-streaming result is awaited even though the first response is selected. That prevents an unobserved continuation from outliving the worker callback. In the stream case, each error remains an error item and no chunks are requested before the consumer polls the mapped stream. @@ -479,7 +479,7 @@ Use the following procedure to verify equivalent behavior across the two worker that only observability values change. 3. Block configured tool and model names, rewrite allowed requests, and preserve an annotated LLM request through later request intercepts and the managed start event. -4. Call each unary continuation once, then use the explicit repeated path to call it +4. Call each non-streaming continuation once, then use the explicit repeated path to call it twice concurrently and account for both downstream invocations. 5. Consume a transformed stream incrementally, then cancel a second stream and confirm worker cleanup. diff --git a/docs/build-plugins/workers/python.mdx b/docs/build-plugins/workers/python.mdx index 950b6be8e..02af98367 100644 --- a/docs/build-plugins/workers/python.mdx +++ b/docs/build-plugins/workers/python.mdx @@ -8,7 +8,7 @@ SPDX-License-Identifier: Apache-2.0 */} The `examples/python-grpc-worker-plugin` package uses the 0.10.0 `nemo-relay-plugin` SDK and the shared documentation configuration. Its worker registers -all 16 [surfaces](/build-plugins/workers/middleware-and-continuations), accepts +all 17 [surfaces](/build-plugins/workers/middleware-and-continuations), accepts synchronous and asynchronous callback forms where the Python SDK permits them, and relies on the SDK for protobuf stubs, the authenticated server, and cooperative task cancellation. @@ -184,14 +184,14 @@ start the worker: ``` The activation report should identify `examples.python_grpc_worker`, and the worker - handshake should advertise all 16 supported surfaces. + handshake should advertise all 17 supported surfaces. ## Verify Behavior and Clean Up Use the following procedure to verify each callback family and clean up the managed environment: -1. Exercise one allowed and one blocked tool, one allowed and one blocked model, a unary +1. Exercise one allowed and one blocked tool, one allowed and one blocked model, a non-streaming LLM continuation, and a multi-chunk stream. Confirm configured headers, sanitized event fields, preserved annotations, pending marks, and lazy chunk transformation. 2. Cancel a long-running async callback and abandon a worker stream. Confirm the SDK diff --git a/docs/build-plugins/workers/rust.mdx b/docs/build-plugins/workers/rust.mdx index 629e0fe67..9db408250 100644 --- a/docs/build-plugins/workers/rust.mdx +++ b/docs/build-plugins/workers/rust.mdx @@ -116,7 +116,7 @@ Use the following procedure to register the manifest and inspect worker activati 2. Start Relay from the repository root. Inspect the activation report and confirm the handshake reports plugin identity, SDK and runtime metadata, - `grpc-v1`, multiple-component behavior, and all 16 surfaces. + `grpc-v1`, multiple-component behavior, and all 17 surfaces. ```bash nemo-relay --bind 127.0.0.1:4040 @@ -202,7 +202,7 @@ async fn main() -> Result<()> { Use the following procedure to verify the registered callbacks and stop the worker cleanly: -1. Exercise allowed and blocked calls, unary and streaming continuations, codec decode and +1. Exercise allowed and blocked calls, non-streaming and streaming continuations, codec decode and encode proxies, pending marks, optimization contributions, nested and isolated scopes, and cancellation. 2. Disable and remove the component only after in-flight invocations have settled. diff --git a/docs/reference/migration-guides.mdx b/docs/reference/migration-guides.mdx index 5a71d70b1..81efe3516 100644 --- a/docs/reference/migration-guides.mdx +++ b/docs/reference/migration-guides.mdx @@ -13,25 +13,25 @@ upgrade actions as they are identified during the 0.10 development cycle. ### Update LLM Execution Intercepts -Relay 0.10 changes every unary and streaming LLM execution-intercept callback. +Relay 0.10 changes every non-streaming and streaming LLM execution-intercept callback. Update callbacks as follows: | Surface | Relay 0.9 | Relay 0.10 | |---|---|---| -| Rust, Python unary, typed native plugin, Rust worker, Python worker | `(name, request, next)` | `(name, request, context, next)` | +| Rust, Python non-streaming, typed native plugin, Rust worker, Python worker | `(name, request, next)` | `(name, request, context, next)` | | Python language-binding streaming | `(request, next)` | `(name, request, context, next)` | | Node.js, Go | `(request, next)` | `(request, context, next)` | | C | `(user_data, request, next, next_ctx)` | `(user_data, name, request, context, next, next_ctx)` | | Raw native callback | `(..., name, request, next, ...)` | `(..., name, request, context, next, ...)` | The request direction reports the selected codec and exposes decode and encode -operations when Relay resolved one. Unary execution also exposes response +operations when Relay resolved one. Non-streaming execution also exposes response identity and decode. Streaming execution has no response codec because chunks are not complete provider responses. -Codec access is invocation-scoped. Do not cache the context or a resolved codec: -unary access expires when the callback settles, and streaming request access -expires when the returned stream closes. +Codec access is limited to the current call. Do not cache the context or a +resolved codec: non-streaming access expires when the callback finishes, and +streaming request access expires when the returned stream closes. This is a source and binary compatibility break for execution-intercept users: diff --git a/examples/python-grpc-worker-plugin/README.md b/examples/python-grpc-worker-plugin/README.md index 5c50c6aca..8e730cab6 100644 --- a/examples/python-grpc-worker-plugin/README.md +++ b/examples/python-grpc-worker-plugin/README.md @@ -12,9 +12,9 @@ invocation-scoped codec proxies, transforms streams lazily, and cleans up marks, scopes, isolated stacks, and cancelled tasks. The worker targets Relay 0.10 while retaining the `grpc-v1` protocol name. Its tool -continuation returns `ToolExecutionResult`. Its LLM execution callbacks receive -directional codec context before the continuation; unary execution can decode a complete -response, while streaming execution exposes request codec operations only. +continuation returns `ToolExecutionResult`. Its LLM execution callbacks receive codec +context before the continuation; non-streaming execution can decode a completed response, +while streaming execution exposes request codec operations only. Run the example's own test project from this directory: diff --git a/examples/rust-grpc-worker-plugin/README.md b/examples/rust-grpc-worker-plugin/README.md index bb4aa8883..c258a5e8e 100644 --- a/examples/rust-grpc-worker-plugin/README.md +++ b/examples/rust-grpc-worker-plugin/README.md @@ -10,10 +10,11 @@ guide. It validates the shared documentation configuration, registers every safe `grpc-v1` surface, exercises continuations and lazy streams, uses invocation-scoped codecs, and demonstrates marks and scope-stack cleanup. -The worker targets Relay 0.10 while retaining the `grpc-v1` protocol name. Unary and -streaming LLM execution callbacks receive directional codec context before their -continuation. Streaming deliberately has no response codec because chunks are not -complete provider responses. +The worker targets Relay 0.10 while retaining the `grpc-v1` protocol name. +Non-streaming and streaming LLM execution callbacks receive request codec +context before their continuation. Non-streaming callbacks also receive +completed-response codec context; streaming callbacks do not because chunks +are not complete provider responses. Run `cargo test` and `cargo build` from this directory. The configuration and schema tests are order-independent. The lifecycle test builds a fresh worker, diff --git a/examples/rust-native-plugin/README.md b/examples/rust-native-plugin/README.md index 83d00ec7a..6922b0dac 100644 --- a/examples/rust-native-plugin/README.md +++ b/examples/rust-native-plugin/README.md @@ -11,10 +11,10 @@ helpers live in separate source modules. Together they register the subscriber, all three event sanitizers, five tool surfaces, and six LLM surfaces exposed by the current typed 0.10.0 SDK. -Relay 0.10 uses native ABI v7. Every LLM execution callback receives directional codec -context before its continuation; streaming execution exposes request codec operations -but no response decoder. The manifest continues to declare `native_api = "1"`, and its -Relay lower bound is `0.10.0` because Relay 0.9 uses the v5 callback layout. +Relay 0.10 uses native ABI v7. Every LLM execution callback receives codec context +before its continuation; streaming execution exposes request codec operations but no +response decoder. The manifest continues to declare `native_api = "1"`, and its Relay +lower bound is `0.10.0` because Relay 0.9 uses the earlier v5/v6 callback layouts. Run the focused tests and build the shared library from this directory. The configuration tests isolate validation and schema contracts. The lifecycle test diff --git a/python/plugin/README.md b/python/plugin/README.md index bb86e2a65..f1ee42674 100644 --- a/python/plugin/README.md +++ b/python/plugin/README.md @@ -27,16 +27,22 @@ and declare `compat.relay` beginning at `0.8.0`. `ToolNext.call()` returns `ToolExecutionResult`, preserving opaque annotation metadata independently of the application result JSON. +Relay 0.10 adds `LlmExecutionContext` to non-streaming and streaming LLM +execution callbacks. Workers that register either surface must update their +callback signatures, rebuild with the 0.10 SDK, and set `compat.relay` to begin +at `0.10.0`. The protocol remains `grpc-v1`. + ## Authoring Surface The following rows describe the plugin authoring surfaces available through this SDK. | Surface | Role | |---|---| -| `WorkerPlugin` and `PluginContext` | Define validation and install all 16 worker-owned subscriber and middleware registrations. | +| `WorkerPlugin` and `PluginContext` | Define validation and install all 17 worker-owned subscriber and middleware registrations. | | `serve_plugin` | Starts an AsyncIO gRPC server from the Relay-managed environment and authenticated local activation endpoints. | | Typed runtime helpers | Share JSON, event, scope, middleware, continuation, and diagnostic contracts with the Relay host. | | Canonical tool results | Preserve application results and opaque annotations across tool callbacks and continuations. | +| `LlmExecutionContext` | Provides the selected request codec and, for non-streaming calls, the completed-response codec while the callback is active. | | Generated transport bindings | Ship private protobuf bindings in the wheel, so installation does not require `protoc` or `grpcio-tools`. | The worker process isolates Python dependencies and crashes from Relay while preserving @@ -170,7 +176,7 @@ an application-level RPC admission limit. ## Invocation Cancellation -Relay assigns every unary and streaming callback an invocation ID. The host +Relay assigns every non-streaming and streaming callback an invocation ID. The host sends `CancelInvocation` when its managed caller is cancelled, its worker RPC times out, or it stops consuming a worker-backed stream. The SDK cancels the matching `asyncio.Task` and reports a structured `worker.cancelled` result. diff --git a/skills/nemo-relay-plugin-build/SKILL.md b/skills/nemo-relay-plugin-build/SKILL.md index 84a490978..c1da5e585 100644 --- a/skills/nemo-relay-plugin-build/SKILL.md +++ b/skills/nemo-relay-plugin-build/SKILL.md @@ -155,18 +155,25 @@ is needed. ## Native Version Compatibility -Choose the native callback model from the target Relay version; do not present -the 0.8 SDK as source-compatible with 0.7: +Choose the native callback model from the target Relay version: - **Relay 0.7:** Keep typed Rust middleware callbacks synchronous. Use the raw native ABI v3 completion-based registration path only when asynchronous work is required, and constrain the manifest to `compat.relay = ">=0.7,<0.8"`. - **Relay 0.8:** Return futures through the typed Rust SDK and constrain - the manifest to `compat.relay = ">=0.8.0,<1.0"`. The SDK runs typed middleware + the manifest to `compat.relay = ">=0.8.0,<0.10"`. The SDK runs typed middleware on an SDK-owned Tokio executor; subscribers and raw synchronous ABI registrations remain synchronous. Do not block executor workers. Scope context does not automatically propagate to tasks created with `tokio::spawn`, and teardown must stop new callbacks and drain accepted work before unload. +- **Relay 0.9:** Use the 0.9 SDK and constrain the manifest to + `compat.relay = ">=0.9.0,<0.10"`. Context-aware tool execution callbacks and + host-routed plugin logging are not available in earlier native layouts. +- **Relay 0.10:** Rebuild every native plugin with the ABI v7 SDK and constrain + the manifest to `compat.relay = ">=0.10.0,<1.0"`. LLM execution callbacks + receive request codec context; non-streaming callbacks also receive + completed-response codec context. The authored `compat.native_api = "1"` + value does not change. The 0.8 typed SDK lets native components configure its executor with a positive `executor.worker_threads` value. When their manifest exposes a @@ -191,8 +198,7 @@ this skill. diagnostic. - Do not treat native process isolation as a security boundary, or claim that a worker is sandboxed. -- Do not apply the 0.8 typed-async callback contract to Relay 0.7, or the 0.7 - raw completion contract to typed 0.8 middleware. +- Do not mix callback layouts from different Relay releases. - Do not block the 0.8 SDK-owned executor. ## Validation Checklist @@ -211,9 +217,7 @@ this skill. JSON Schema and declare `config_schema` in `relay-plugin.toml`. - [ ] Dynamic manifests validate their lane-specific kind, load contract, SemVer compatibility, integrity, and disabled-record behavior. -- [ ] Native callback APIs and compatibility constraints match the target Relay - version: synchronous typed callbacks and raw ABI v3 completion work for - 0.7; typed async SDK middleware and executor configuration for 0.8. +- [ ] Native callback APIs and `compat.relay` match the target Relay release. - [ ] Async native paths cover cancellation, settlement, and drain-before- unload behavior. From 9f27dc4657821a3b11cb62da9bd110619b1d21a1 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 30 Sep 2026 15:09:14 -0700 Subject: [PATCH 16/22] refactor(ffi): remove unused codec context derives Signed-off-by: Alex Fournier --- crates/ffi/src/callable.rs | 3 --- 1 file changed, 3 deletions(-) diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index 7b89273ef..fcbcd3dd3 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -181,7 +181,6 @@ pub enum NemoRelayLlmSanitizeCodecKind { /// `codec_id` is null for `None` and `Opaque`, and is valid only for the /// duration of the callback. #[repr(C)] -#[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmSanitizeRequestContext { /// Kind of active codec identity. pub codec_kind: NemoRelayLlmSanitizeCodecKind, @@ -193,7 +192,6 @@ pub struct NemoRelayLlmSanitizeRequestContext { /// Response codec context shared by LLM sanitizer and execution callbacks. #[repr(C)] -#[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmSanitizeResponseContext { /// Kind of active codec identity. pub codec_kind: NemoRelayLlmSanitizeCodecKind, @@ -210,7 +208,6 @@ pub struct NemoRelayLlmSanitizeResponseContext { /// response to decode. Pointers reachable from this value are borrowed and /// valid only until the intercept callback returns. #[repr(C)] -#[derive(Debug, Clone, Copy)] pub struct NemoRelayLlmExecutionContext { /// Active request codec identity and capability. pub request_codec: NemoRelayLlmSanitizeRequestContext, From 57ff8405536feaff7f096c0d1e6207bb13c0fa8a Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 30 Sep 2026 15:22:45 -0700 Subject: [PATCH 17/22] chore: remove plugin skill update Signed-off-by: Alex Fournier --- skills/nemo-relay-plugin-build/SKILL.md | 20 ++++++++------------ 1 file changed, 8 insertions(+), 12 deletions(-) diff --git a/skills/nemo-relay-plugin-build/SKILL.md b/skills/nemo-relay-plugin-build/SKILL.md index c1da5e585..84a490978 100644 --- a/skills/nemo-relay-plugin-build/SKILL.md +++ b/skills/nemo-relay-plugin-build/SKILL.md @@ -155,25 +155,18 @@ is needed. ## Native Version Compatibility -Choose the native callback model from the target Relay version: +Choose the native callback model from the target Relay version; do not present +the 0.8 SDK as source-compatible with 0.7: - **Relay 0.7:** Keep typed Rust middleware callbacks synchronous. Use the raw native ABI v3 completion-based registration path only when asynchronous work is required, and constrain the manifest to `compat.relay = ">=0.7,<0.8"`. - **Relay 0.8:** Return futures through the typed Rust SDK and constrain - the manifest to `compat.relay = ">=0.8.0,<0.10"`. The SDK runs typed middleware + the manifest to `compat.relay = ">=0.8.0,<1.0"`. The SDK runs typed middleware on an SDK-owned Tokio executor; subscribers and raw synchronous ABI registrations remain synchronous. Do not block executor workers. Scope context does not automatically propagate to tasks created with `tokio::spawn`, and teardown must stop new callbacks and drain accepted work before unload. -- **Relay 0.9:** Use the 0.9 SDK and constrain the manifest to - `compat.relay = ">=0.9.0,<0.10"`. Context-aware tool execution callbacks and - host-routed plugin logging are not available in earlier native layouts. -- **Relay 0.10:** Rebuild every native plugin with the ABI v7 SDK and constrain - the manifest to `compat.relay = ">=0.10.0,<1.0"`. LLM execution callbacks - receive request codec context; non-streaming callbacks also receive - completed-response codec context. The authored `compat.native_api = "1"` - value does not change. The 0.8 typed SDK lets native components configure its executor with a positive `executor.worker_threads` value. When their manifest exposes a @@ -198,7 +191,8 @@ this skill. diagnostic. - Do not treat native process isolation as a security boundary, or claim that a worker is sandboxed. -- Do not mix callback layouts from different Relay releases. +- Do not apply the 0.8 typed-async callback contract to Relay 0.7, or the 0.7 + raw completion contract to typed 0.8 middleware. - Do not block the 0.8 SDK-owned executor. ## Validation Checklist @@ -217,7 +211,9 @@ this skill. JSON Schema and declare `config_schema` in `relay-plugin.toml`. - [ ] Dynamic manifests validate their lane-specific kind, load contract, SemVer compatibility, integrity, and disabled-record behavior. -- [ ] Native callback APIs and `compat.relay` match the target Relay release. +- [ ] Native callback APIs and compatibility constraints match the target Relay + version: synchronous typed callbacks and raw ABI v3 completion work for + 0.7; typed async SDK middleware and executor configuration for 0.8. - [ ] Async native paths cover cancellation, settlement, and drain-before- unload behavior. From 6e30b5f47756f730adb5cbe395e91df871a95696 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 30 Sep 2026 15:34:34 -0700 Subject: [PATCH 18/22] docs: keep codec context updates scoped Signed-off-by: Alex Fournier --- crates/worker/README.md | 2 +- docs/build-plugins/native/about.mdx | 4 ++-- .../native/native-abi-reference.mdx | 4 ++-- docs/build-plugins/native/wrap-execution.mdx | 16 ++++++++-------- docs/build-plugins/plugin-context.mdx | 2 +- docs/build-plugins/workers/grpc-v1-protocol.mdx | 6 +++--- .../workers/middleware-and-continuations.mdx | 10 +++++----- docs/build-plugins/workers/python.mdx | 6 +++--- docs/build-plugins/workers/rust.mdx | 4 ++-- python/plugin/README.md | 4 ++-- 10 files changed, 29 insertions(+), 29 deletions(-) diff --git a/crates/worker/README.md b/crates/worker/README.md index 86872e835..3dbc3938d 100644 --- a/crates/worker/README.md +++ b/crates/worker/README.md @@ -38,7 +38,7 @@ decode incomplete chunks. The wire protocol remains `grpc-v1`. | Surface | Role | |---|---| | `WorkerPlugin` | Defines plugin identity, validation, registration, and multiple-component behavior in the worker process. | -| `PluginContext` | Installs typed handlers for all 17 supported registration surfaces. | +| `PluginContext` | Installs typed handlers for all 16 supported registration surfaces. | | `PluginRuntime` and continuations | Emit marks, manage scopes, and call the remaining tool or LLM execution chain through the authenticated host service. | | `LlmExecutionContext` | Reports the selected request and completed-response codecs and provides codec operations while the callback is active. | | Canonical tool results | Preserve application results and opaque annotations across tool callbacks and continuations. | diff --git a/docs/build-plugins/native/about.mdx b/docs/build-plugins/native/about.mdx index e369a9f8d..f1367af7c 100644 --- a/docs/build-plugins/native/about.mdx +++ b/docs/build-plugins/native/about.mdx @@ -52,7 +52,7 @@ attributed to a particular plugin. ## What the SDK Owns The Rust SDK exports the stable entry symbol, converts host-owned JSON handles into -typed DTOs, registers the supported subscriber and middleware surfaces, and drives async middleware on one +typed DTOs, registers all 16 plugin surfaces, and drives async middleware on one SDK-owned multi-thread Tokio runtime per configured component. A plugin can set a default executor size and accept a positive `executor.worker_threads` component override. The default is two workers; change it only after measuring queued async work @@ -153,7 +153,7 @@ Follow these pages in order to build, activate, exercise, and remove the native including subscribers and all three event sanitizer surfaces. 3. Add policy and request rewriting with [Control Requests](/build-plugins/native/control-requests), preserving annotations and making priority and `break_chain` explicit. -4. Add tool, non-streaming model, and lazy stream wrappers with [Wrap Execution](/build-plugins/native/wrap-execution). +4. Add tool, unary model, and lazy stream wrappers with [Wrap Execution](/build-plugins/native/wrap-execution). 5. Verify marks, scopes, isolated stacks, cleanup, and executor control with [Runtime Events and Scopes](/build-plugins/native/runtime-events-and-scopes). 6. Consult [Native ABI Reference](/build-plugins/native/native-abi-reference) only when diff --git a/docs/build-plugins/native/native-abi-reference.mdx b/docs/build-plugins/native/native-abi-reference.mdx index 5ebf47dd4..0d3d7b94c 100644 --- a/docs/build-plugins/native/native-abi-reference.mdx +++ b/docs/build-plugins/native/native-abi-reference.mdx @@ -50,7 +50,7 @@ function signatures and field order are defined by the public | Table Level | Operations | |---|---| | Frozen v1/v2 prefix | Version and struct-size negotiation; host version; string allocation, access, and release; thread-local error reporting; callback-scoped LLM request decode and encode plus response decode; subscriber, five tool, six LLM, and three event-sanitizer registrations; current scope, scope push and pop, mark emission, isolated stack creation and release, thread-stack set, capture, and restore, captured-binding release, active-stack inspection, and scoped binding. | -| Frozen v3 extension | Completion resolve, reject, cancellation inspection, and release; one-shot completion-coupled continuation invocation; continuation release; generic async middleware registration; bounded output-stream push, finish, reject, cancellation inspection, and release; downstream stream invocation; stream-middleware registration; and repeated or concurrent non-streaming continuation invocation with independent result callbacks. | +| Frozen v3 extension | Completion resolve, reject, cancellation inspection, and release; one-shot completion-coupled continuation invocation; continuation release; generic async middleware registration; bounded output-stream push, finish, reject, cancellation inspection, and release; downstream stream invocation; stream-middleware registration; and repeated or concurrent unary continuation invocation with independent result callbacks. | | Frozen v4 extension | Completion-scoped LLM request decode and encode plus response decode; pull-based downstream LLM stream open, pull, cancel, and release; completion retain for typed codec facades; output-stream backpressure inspection; extended mark emission; runtime diagnostics; activation-owned runtime capability creation, retain, and release; global runtime-registration discovery; owned conditional middleware guardrail registration and deregistration; and activation-owned and runtime-discovered callback gate registration. The callback registration slots are appended after the original constant-reason slots. | | v5 extension | Context-carrying tool execution intercept registration. The callback receives one JSON object containing `tool_name`, `args`, and `tool_call_id`. | | v6 extension | Host-routed operational logging with structured fields. | @@ -171,7 +171,7 @@ the callback. Return `Pending` only after retaining it. A retained completion mu exactly once and then be released. Release every async `next` reference after its last use. -`async_next_invoke_result` supports repeated or concurrent non-streaming continuation calls with +`async_next_invoke_result` supports repeated or concurrent unary continuation calls with independent result callbacks. The older completion-coupled `async_next_invoke` is one-shot because the continuation result settles the middleware completion. Settle the owner only after every started continuation call has finished. When the owner settles or diff --git a/docs/build-plugins/native/wrap-execution.mdx b/docs/build-plugins/native/wrap-execution.mdx index 9dae61bbe..a628b61c6 100644 --- a/docs/build-plugins/native/wrap-execution.mdx +++ b/docs/build-plugins/native/wrap-execution.mdx @@ -1,6 +1,6 @@ --- title: "Wrap Execution" -description: "Use native tool, non-streaming LLM, and streaming continuations correctly." +description: "Use native tool, unary LLM, and streaming continuations correctly." position: 24 --- {/* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. @@ -8,7 +8,7 @@ SPDX-License-Identifier: Apache-2.0 */} [Execution intercepts](/about-nemo-relay/concepts/middleware) receive the real request and a continuation representing the rest -of the call path. The example's `execution` group registers a tool wrapper, a non-streaming LLM +of the call path. The example's `execution` group registers a tool wrapper, a unary LLM wrapper, and an LLM stream wrapper at the configured priority. ## Tool and Unary Results @@ -20,7 +20,7 @@ example's pending mark. The conversion preserves both the application `result` and any opaque `annotation`. Relay owns the pending mark: it is emitted in the managed lifecycle and does not appear in the application-visible result. -The non-streaming LLM wrapper has a deliberately smaller contract and returns provider-response +The unary LLM wrapper has a deliberately smaller contract and returns provider-response JSON. LLM pending marks, annotations, and optimization contributions belong to the request-intercept outcome shown in [Control Requests](/build-plugins/native/control-requests), not the execution result. @@ -56,7 +56,7 @@ reads `tool_result.result["answer"]`. Relay separately emits callback returns an outcome instead of returning a plain JSON value. The continuation is reusable. A deliberate configuration or request flag in the example -can invoke non-streaming `next` twice concurrently, await both responses, and select one. This +can invoke unary `next` twice concurrently, await both responses, and select one. This demonstrates the API while preserving the cost of repetition. A repeated tool can perform its side effect twice, and a repeated model call can incur provider cost twice. Production plugins need an idempotency @@ -137,15 +137,15 @@ or cancellation. ## Verify Execution Behavior -Use the following procedure to verify non-streaming, repeated, and streaming continuation +Use the following procedure to verify unary, repeated, and streaming continuation behavior: 1. Activate the example with `execution.enabled = true`, priority 30, and `emit_pending_marks = true`. -2. Execute a tool and a non-streaming model call. Confirm each downstream callback runs once, +2. Execute a tool and a unary model call. Confirm each downstream callback runs once, the application receives only its expected result, and an additional pending mark is emitted under the managed call scope. -3. Enable the example's repeated-continuation input and confirm two downstream non-streaming +3. Enable the example's repeated-continuation input and confirm two downstream unary invocations can overlap. Verify the selected result and accounting explicitly. 4. Consume a three-chunk LLM stream one item at a time. Confirm the first transformed chunk arrives before the downstream stream completes. @@ -154,6 +154,6 @@ behavior: 6. Clear the component and repeat the calls to prove that no wrapper or pending mark remains registered. -Success means non-streaming and stream continuations preserve scope, errors, and cancellation, +Success means unary and stream continuations preserve scope, errors, and cancellation, while Relay-owned tool marks and LLM request accounting remain separate from application results. diff --git a/docs/build-plugins/plugin-context.mdx b/docs/build-plugins/plugin-context.mdx index db1f0274f..49c8d1105 100644 --- a/docs/build-plugins/plugin-context.mdx +++ b/docs/build-plugins/plugin-context.mdx @@ -30,7 +30,7 @@ language-binding, native typed, and worker plugins. | LLM | Response sanitizer | Sanitize response observability with the structured LLM context and its codec handle. | | LLM | Conditional guardrail | Allow or block real model execution from the provider request. The current callback contract does not receive a separate model name or annotation argument. | | LLM | Request intercept | Return the complete request-intercept outcome, preserving or deliberately replacing annotations. | -| LLM | Execution intercept | Wrap non-streaming execution through a continuation and return provider-response JSON. | +| LLM | Execution intercept | Wrap unary execution through a continuation and return provider-response JSON. | | LLM | Stream execution intercept | Transform response chunks while preserving order, cancellation, and error behavior. Native and worker SDKs expose a lazy stream; the Node.js language binding currently supplies the downstream chunks as an array. | Sanitize guardrails change emitted observability payloads only. They do not rewrite the diff --git a/docs/build-plugins/workers/grpc-v1-protocol.mdx b/docs/build-plugins/workers/grpc-v1-protocol.mdx index 089a4570f..0d73f62ff 100644 --- a/docs/build-plugins/workers/grpc-v1-protocol.mdx +++ b/docs/build-plugins/workers/grpc-v1-protocol.mdx @@ -67,7 +67,7 @@ service RelayHostRuntime { | `Health` | Confirms the authenticated activation, protocol, plugin identity, and current worker readiness. | | `Validate` | Receives component config in a JSON envelope and returns encoded diagnostics or a structured worker error. It must not register behavior. | | `Register` | Receives valid component config and returns owned registrations with local name, surface, priority, and `break_chain`. | -| `Invoke` | Dispatches one subscriber, sanitizer, guardrail, request intercept, or non-streaming execution callback and returns the surface-appropriate result. | +| `Invoke` | Dispatches one subscriber, sanitizer, guardrail, request intercept, or unary execution callback and returns the surface-appropriate result. | | `InvokeStream` | Dispatches an LLM stream execution intercept and emits incremental value or error chunks. | | `CancelInvocation` | Cooperatively cancels one active invocation ID and reports whether cancellation was accepted. Unknown, completed, and already-cancelled IDs receive a negative acknowledgment. | | `Shutdown` | Stops new work for an activation and begins orderly worker termination with a reason. | @@ -286,7 +286,7 @@ capability IDs behind `resolve_codec()` proxies. | `CreateScopeStack` | Allocates an isolated stack and returns its opaque ID. | | `DropScopeStack` | Releases an isolated stack owned by the activation. | | `ToolNext` | Executes a tool continuation with JSON arguments and captured scope, returning a structural tool result or worker error. | -| `LlmNext` | Executes a non-streaming LLM continuation with a request and captured scope. | +| `LlmNext` | Executes a unary LLM continuation with a request and captured scope. | | `LlmStreamNext` | Executes a streaming LLM continuation and returns incremental chunks. | | `DecodeLlmCodecRequest` | Uses the invocation-scoped capability to decode a request into its annotated representation. | | `EncodeLlmCodecRequest` | Applies an annotated request to the original provider envelope. | @@ -418,6 +418,6 @@ process: only during explicit package removal. Successful protocol verification covers authentication failures, envelope schema -failures, every registration and result variant, repeated non-streaming and incremental stream +failures, every registration and result variant, repeated unary and incremental stream continuations, cancellation before and during callbacks, codec capability expiry, host runtime ownership errors, health, and orderly plus forced shutdown. diff --git a/docs/build-plugins/workers/middleware-and-continuations.mdx b/docs/build-plugins/workers/middleware-and-continuations.mdx index d4ece4e5a..2d5446729 100644 --- a/docs/build-plugins/workers/middleware-and-continuations.mdx +++ b/docs/build-plugins/workers/middleware-and-continuations.mdx @@ -72,7 +72,7 @@ context.register_event_metadata_injector( -Relay sends the immutable Event snapshot through the existing non-streaming `Invoke` RPC. It +Relay sends the immutable Event snapshot through the existing unary `Invoke` RPC. It validates and merges accepted additions before Event sanitizers run. A callback error omits that callback's additions without dropping the Event. Stopping the worker removes the registration. @@ -279,7 +279,7 @@ originating tool call. `ToolNext`, `LlmNext`, and `LlmStreamNext` are host proxies identified by an opaque continuation ID. Calling one issues a host-runtime RPC under the scope snapshot captured -for that worker invocation. A callback can call a non-streaming proxy zero, one, or multiple +for that worker invocation. A callback can call a unary proxy zero, one, or multiple times, including concurrently when the SDK type permits cloning. Each call can repeat side effects, provider charges, events, and downstream middleware. @@ -293,7 +293,7 @@ pending marks. The example uses one ordinary wrapper and one explicitly requested concurrent path. It does not retry implicitly. Tool pending marks remain in the tool execution outcome; LLM annotations, pending marks, and optimization contributions remain in the request -intercept outcome. The non-streaming execution callback returns only provider-response JSON. +intercept outcome. The unary execution callback returns only provider-response JSON. `LlmStreamNext` returns a remote stream. The worker transforms each chunk as it arrives and yields immediately. If the host cancels the invocation or the consumer abandons the @@ -465,7 +465,7 @@ instead of implying that individual chunks can be decoded as completed responses. Codec proxies are invocation-scoped and must not be retained after the callback or returned stream finishes. -The second non-streaming result is awaited even though the first response is selected. That +The second unary result is awaited even though the first response is selected. That prevents an unobserved continuation from outliving the worker callback. In the stream case, each error remains an error item and no chunks are requested before the consumer polls the mapped stream. @@ -479,7 +479,7 @@ Use the following procedure to verify equivalent behavior across the two worker that only observability values change. 3. Block configured tool and model names, rewrite allowed requests, and preserve an annotated LLM request through later request intercepts and the managed start event. -4. Call each non-streaming continuation once, then use the explicit repeated path to call it +4. Call each unary continuation once, then use the explicit repeated path to call it twice concurrently and account for both downstream invocations. 5. Consume a transformed stream incrementally, then cancel a second stream and confirm worker cleanup. diff --git a/docs/build-plugins/workers/python.mdx b/docs/build-plugins/workers/python.mdx index 02af98367..950b6be8e 100644 --- a/docs/build-plugins/workers/python.mdx +++ b/docs/build-plugins/workers/python.mdx @@ -8,7 +8,7 @@ SPDX-License-Identifier: Apache-2.0 */} The `examples/python-grpc-worker-plugin` package uses the 0.10.0 `nemo-relay-plugin` SDK and the shared documentation configuration. Its worker registers -all 17 [surfaces](/build-plugins/workers/middleware-and-continuations), accepts +all 16 [surfaces](/build-plugins/workers/middleware-and-continuations), accepts synchronous and asynchronous callback forms where the Python SDK permits them, and relies on the SDK for protobuf stubs, the authenticated server, and cooperative task cancellation. @@ -184,14 +184,14 @@ start the worker: ``` The activation report should identify `examples.python_grpc_worker`, and the worker - handshake should advertise all 17 supported surfaces. + handshake should advertise all 16 supported surfaces. ## Verify Behavior and Clean Up Use the following procedure to verify each callback family and clean up the managed environment: -1. Exercise one allowed and one blocked tool, one allowed and one blocked model, a non-streaming +1. Exercise one allowed and one blocked tool, one allowed and one blocked model, a unary LLM continuation, and a multi-chunk stream. Confirm configured headers, sanitized event fields, preserved annotations, pending marks, and lazy chunk transformation. 2. Cancel a long-running async callback and abandon a worker stream. Confirm the SDK diff --git a/docs/build-plugins/workers/rust.mdx b/docs/build-plugins/workers/rust.mdx index 9db408250..629e0fe67 100644 --- a/docs/build-plugins/workers/rust.mdx +++ b/docs/build-plugins/workers/rust.mdx @@ -116,7 +116,7 @@ Use the following procedure to register the manifest and inspect worker activati 2. Start Relay from the repository root. Inspect the activation report and confirm the handshake reports plugin identity, SDK and runtime metadata, - `grpc-v1`, multiple-component behavior, and all 17 surfaces. + `grpc-v1`, multiple-component behavior, and all 16 surfaces. ```bash nemo-relay --bind 127.0.0.1:4040 @@ -202,7 +202,7 @@ async fn main() -> Result<()> { Use the following procedure to verify the registered callbacks and stop the worker cleanly: -1. Exercise allowed and blocked calls, non-streaming and streaming continuations, codec decode and +1. Exercise allowed and blocked calls, unary and streaming continuations, codec decode and encode proxies, pending marks, optimization contributions, nested and isolated scopes, and cancellation. 2. Disable and remove the component only after in-flight invocations have settled. diff --git a/python/plugin/README.md b/python/plugin/README.md index f1ee42674..32c9041f0 100644 --- a/python/plugin/README.md +++ b/python/plugin/README.md @@ -38,7 +38,7 @@ The following rows describe the plugin authoring surfaces available through this | Surface | Role | |---|---| -| `WorkerPlugin` and `PluginContext` | Define validation and install all 17 worker-owned subscriber and middleware registrations. | +| `WorkerPlugin` and `PluginContext` | Define validation and install all 16 worker-owned subscriber and middleware registrations. | | `serve_plugin` | Starts an AsyncIO gRPC server from the Relay-managed environment and authenticated local activation endpoints. | | Typed runtime helpers | Share JSON, event, scope, middleware, continuation, and diagnostic contracts with the Relay host. | | Canonical tool results | Preserve application results and opaque annotations across tool callbacks and continuations. | @@ -176,7 +176,7 @@ an application-level RPC admission limit. ## Invocation Cancellation -Relay assigns every non-streaming and streaming callback an invocation ID. The host +Relay assigns every unary and streaming callback an invocation ID. The host sends `CancelInvocation` when its managed caller is cancelled, its worker RPC times out, or it stops consuming a worker-backed stream. The SDK cancels the matching `asyncio.Task` and reports a structured `worker.cancelled` result. From 006995ffb90562c3c712bc4269c999e2d2d7786b Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 30 Sep 2026 15:45:13 -0700 Subject: [PATCH 19/22] refactor: tighten execution codec context ownership Signed-off-by: Alex Fournier --- crates/core/src/api/llm.rs | 5 ++-- .../src/api/runtime/llm_execution_context.rs | 8 +++--- crates/core/src/plugin/dynamic.rs | 26 ++++++++----------- crates/core/src/plugin/dynamic/native.rs | 20 +++++++++++--- .../tests/unit/llm_execution_context_tests.rs | 4 +-- crates/node/tests/llm_tests.mjs | 22 +++++++++------- python/nemo_relay/__init__.pyi | 6 ----- 7 files changed, 50 insertions(+), 41 deletions(-) diff --git a/crates/core/src/api/llm.rs b/crates/core/src/api/llm.rs index 76b673e21..b303302e2 100644 --- a/crates/core/src/api/llm.rs +++ b/crates/core/src/api/llm.rs @@ -1739,7 +1739,8 @@ pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result { ); let execution_name = name.clone(); let event_uuid = handle.uuid; - let execution_context = LlmExecutionContext::for_unary_codecs(request_codec, &response_codec); + let execution_context = + LlmExecutionContext::for_non_streaming(request_codec, response_codec.clone()); let execution = with_active_event_trace_context( event_uuid, Some(active_trace_context), @@ -1972,7 +1973,7 @@ pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Resu let execution_name = name.clone(); let event_uuid = handle.uuid; let stream_started_at = Instant::now(); - let execution_context = LlmExecutionContext::for_streaming_codec(request_codec); + let execution_context = LlmExecutionContext::for_streaming(request_codec); let execution = with_active_event_trace_context( event_uuid, Some(active_trace_context), diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index 08fe1b39c..c1cc55d3f 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -139,20 +139,20 @@ impl LlmExecutionContext { } /// Construct the context for a non-streaming managed execution. - pub(crate) fn for_unary_codecs( + pub(crate) fn for_non_streaming( request_codec: Option>, - response_codec: &Option>, + response_codec: Option>, ) -> Self { Self::new( LlmSanitizeRequestContext::for_request_codec(request_codec), Some(LlmSanitizeResponseContext::for_response_codec( - response_codec.clone(), + response_codec, )), ) } /// Construct the context for a streaming managed execution. - pub(crate) fn for_streaming_codec(request_codec: Option>) -> Self { + pub(crate) fn for_streaming(request_codec: Option>) -> Self { Self::new( LlmSanitizeRequestContext::for_request_codec(request_codec), None, diff --git a/crates/core/src/plugin/dynamic.rs b/crates/core/src/plugin/dynamic.rs index a6b496073..cf5d6ec12 100644 --- a/crates/core/src/plugin/dynamic.rs +++ b/crates/core/src/plugin/dynamic.rs @@ -118,9 +118,7 @@ pub(super) fn validate_annotated_request_consumer_compatibility( relay: &str, plugin_kind: &str, ) -> crate::plugin::Result<()> { - let requirement = VersionReq::parse(relay).map_err(|error| { - PluginError::InvalidConfig(format!("invalid compat.relay version requirement: {error}")) - })?; + let requirement = parse_relay_requirement(relay)?; if requirement.matches(&Version::new(0, 5, u64::MAX)) { return Err(PluginError::InvalidConfig(format!( "dynamic plugin '{plugin_kind}' registers an LLM request intercept and must declare compat.relay = \">=0.6,<1.0\" or another range that excludes Relay 0.5" @@ -133,9 +131,7 @@ pub(super) fn validate_tool_execution_context_compatibility( relay: &str, plugin_kind: &str, ) -> crate::plugin::Result<()> { - let requirement = VersionReq::parse(relay).map_err(|error| { - PluginError::InvalidConfig(format!("invalid compat.relay version requirement: {error}")) - })?; + let requirement = parse_relay_requirement(relay)?; if version_requirement_matches_minor(&requirement, 0, 8) { return Err(PluginError::InvalidConfig(format!( "dynamic plugin '{plugin_kind}' registers a context-aware tool execution intercept and must declare compat.relay = \">=0.9,<1.0\" or another range that excludes Relay 0.8" @@ -149,9 +145,7 @@ pub(super) fn validate_llm_execution_context_compatibility( relay: &str, plugin_kind: &str, ) -> crate::plugin::Result<()> { - let requirement = VersionReq::parse(relay).map_err(|error| { - PluginError::InvalidConfig(format!("invalid compat.relay version requirement: {error}")) - })?; + let requirement = parse_relay_requirement(relay)?; if version_requirement_matches_minor(&requirement, 0, 9) { return Err(PluginError::InvalidConfig(format!( "dynamic plugin '{plugin_kind}' registers an LLM execution intercept and must declare compat.relay = \">=0.10,<1.0\" or another range that excludes Relay 0.9" @@ -164,9 +158,7 @@ pub(super) fn validate_native_abi_compatibility( relay: &str, plugin_kind: &str, ) -> crate::plugin::Result<()> { - let requirement = VersionReq::parse(relay).map_err(|error| { - PluginError::InvalidConfig(format!("invalid compat.relay version requirement: {error}")) - })?; + let requirement = parse_relay_requirement(relay)?; if version_requirement_matches_minor(&requirement, 0, 9) { return Err(PluginError::InvalidConfig(format!( "dynamic native plugin '{plugin_kind}' uses native ABI v7 and must declare compat.relay = \">=0.10,<1.0\" or another range that excludes Relay 0.9" @@ -190,6 +182,12 @@ fn version_requirement_matches_minor(requirement: &VersionReq, major: u64, minor candidate < next_minor && requirement.matches(&candidate) } +fn parse_relay_requirement(relay: &str) -> crate::plugin::Result { + VersionReq::parse(relay).map_err(|error| { + PluginError::InvalidConfig(format!("invalid compat.relay version requirement: {error}")) + }) +} + fn parse_dynamic_plugin_relay_requirement<'a>( relay: Option<&'a str>, plugin_type: &str, @@ -198,9 +196,7 @@ fn parse_dynamic_plugin_relay_requirement<'a>( .map(str::trim) .filter(|value| !value.is_empty()) .ok_or_else(|| PluginError::InvalidConfig("compat.relay is required".into()))?; - let requirement = VersionReq::parse(relay).map_err(|error| { - PluginError::InvalidConfig(format!("invalid compat.relay version requirement: {error}")) - })?; + let requirement = parse_relay_requirement(relay)?; let minimum = Version::new(0, 8, 0); let declares_minimum = requirement .comparators diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index 376b456b5..e6f4d782c 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -414,6 +414,8 @@ fn load_one_native_plugin( library_path.display() )) })?; + // ABI v7 changes native LLM execution callback layouts. Do not negotiate + // v2-v6: plugins compiled against those tables must rebuild before loading. let status = entry(native_host_api(), &mut plugin); if status != NemoRelayStatus::Ok { drop_native_plugin_descriptor(&mut plugin); @@ -583,7 +585,12 @@ struct OwnedNativeString { } impl OwnedNativeString { - fn new(ptr: *mut NemoRelayNativeString) -> FlowResult { + /// Take ownership of a string allocated by the Relay native-string API. + /// + /// # Safety + /// A non-null `ptr` must identify a live, uniquely owned allocation from + /// `native_string_new` that has not already been freed. + unsafe fn from_raw(ptr: *mut NemoRelayNativeString) -> FlowResult { Ok(Self { ptr: NonNull::new(ptr).ok_or_else(|| { FlowError::Internal("native string allocation returned null".into()) @@ -621,11 +628,18 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { ) -> FlowResult { let (request_kind, request_id) = native_llm_codec_identity(context.request_codec().codec())?; - let request_id = request_id.map(OwnedNativeString::new).transpose()?; + let request_id = request_id + // SAFETY: `native_llm_codec_identity` returns a fresh host-owned string. + .map(|ptr| unsafe { OwnedNativeString::from_raw(ptr) }) + .transpose()?; let (response_kind, response_id) = if let Some(response) = context.response_codec() { let (kind, id) = native_llm_codec_identity(response.codec())?; - (Some(kind), id.map(OwnedNativeString::new).transpose()?) + let id = id + // SAFETY: `native_llm_codec_identity` returns a fresh host-owned string. + .map(|ptr| unsafe { OwnedNativeString::from_raw(ptr) }) + .transpose()?; + (Some(kind), id) } else { (None, None) }; diff --git a/crates/core/tests/unit/llm_execution_context_tests.rs b/crates/core/tests/unit/llm_execution_context_tests.rs index d33830299..ab1751187 100644 --- a/crates/core/tests/unit/llm_execution_context_tests.rs +++ b/crates/core/tests/unit/llm_execution_context_tests.rs @@ -48,7 +48,7 @@ fn retained_codecs_work_while_active_and_expire_with_their_lease() { let backing_probe = Arc::downgrade(&backing); let request_codec: Arc = backing.clone(); let response_codec: Arc = backing.clone(); - let context = LlmExecutionContext::for_unary_codecs(Some(request_codec), &Some(response_codec)); + let context = LlmExecutionContext::for_non_streaming(Some(request_codec), Some(response_codec)); drop(backing); let (leased_context, guard) = context.lease(); @@ -103,7 +103,7 @@ fn estimated_cost_defaults_to_false_after_expiry() { let response_codec: Arc = Arc::new(LeaseProbeCodec { allows_estimated_cost: true, }); - let context = LlmExecutionContext::for_unary_codecs(None, &Some(response_codec)); + let context = LlmExecutionContext::for_non_streaming(None, Some(response_codec)); let (leased_context, guard) = context.lease(); let retained = leased_context .response_codec() diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index ef555d12b..7de9ecc0a 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -1421,15 +1421,19 @@ describe('LLM intercepts', () => { const events = []; const observed = {}; registerSubscriber('node_llm_replace_callback_context', (event) => events.push(event)); - registerLlmExecutionIntercept('node_llm_replace_callback_context', 10, async (request, next) => { - observed.before = lib.capturePropagationContext(); - observed.replacement = lib.withScopeStack(stack, () => ({ - context: lib.capturePropagationContext(), - traceparent: lib.captureTraceparent(), - })); - observed.after = lib.capturePropagationContext(); - return next(request); - }); + registerLlmExecutionIntercept( + 'node_llm_replace_callback_context', + 10, + async (request, _context, next) => { + observed.before = lib.capturePropagationContext(); + observed.replacement = lib.withScopeStack(stack, () => ({ + context: lib.capturePropagationContext(), + traceparent: lib.captureTraceparent(), + })); + observed.after = lib.capturePropagationContext(); + return next(request); + }, + ); try { await llmCallExecuteAsync( 'replace_callback_context_llm', diff --git a/python/nemo_relay/__init__.pyi b/python/nemo_relay/__init__.pyi index f203f8fe6..326ce1909 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -71,9 +71,6 @@ from nemo_relay._native import ( from nemo_relay._native import ( LlmExecutionContext as LlmExecutionContext, ) -from nemo_relay._native import ( - LlmRequestCodecContext as LlmRequestCodecContext, -) from nemo_relay._native import ( LLMHandle as LLMHandle, ) @@ -95,9 +92,6 @@ from nemo_relay._native import ( from nemo_relay._native import ( LlmSanitizeResponseContext as LlmSanitizeResponseContext, ) -from nemo_relay._native import ( - LlmResponseCodecContext as LlmResponseCodecContext, -) from nemo_relay._native import ( LogSeverity as LogSeverity, ) From 43ef01da88130ba729f037bda4de742f99a3ebd4 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Wed, 30 Sep 2026 17:21:02 -0700 Subject: [PATCH 20/22] fix: address execution context review feedback Signed-off-by: Alex Fournier --- crates/core/src/api/runtime/state.rs | 2 -- crates/plugin/src/lib.rs | 20 ++++++++++++++++++++ crates/plugin/tests/typed_callbacks.rs | 24 ++++++++++++++---------- docs/build-plugins/native/about.mdx | 10 +++++----- 4 files changed, 39 insertions(+), 17 deletions(-) diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index 4d8ff1118..fdcaf01f2 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -125,8 +125,6 @@ impl Stream for ExecutionGuardedLlmStream { if matches!(&result, Poll::Ready(None)) { this.continuation_guard.take(); this.codec_guard.take(); - } else if matches!(&result, Poll::Ready(Some(Err(_)))) { - this.codec_guard.take(); } result } diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index d4084871a..444ec3af3 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -3201,6 +3201,10 @@ impl<'a> PluginContext<'a> { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT || self.host.struct_size < std::mem::size_of::() { + set_last_error( + self.host, + "LLM execution intercepts require Relay native ABI v7", + ); if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } @@ -3231,6 +3235,10 @@ impl<'a> PluginContext<'a> { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT || self.host.struct_size < std::mem::size_of::() { + set_last_error( + self.host, + "LLM execution intercepts require Relay native ABI v7", + ); if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } @@ -3272,6 +3280,10 @@ impl<'a> PluginContext<'a> { NemoRelayNativeAsyncMiddlewareKind::LlmExecutionIntercept | NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept ) { + set_last_error( + self.host, + "LLM execution intercepts require the ABI v7 registration functions", + ); if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } @@ -3319,6 +3331,10 @@ impl<'a> PluginContext<'a> { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT || self.host.struct_size < std::mem::size_of::() { + set_last_error( + self.host, + "LLM execution intercepts require Relay native ABI v7", + ); if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } @@ -3356,6 +3372,10 @@ impl<'a> PluginContext<'a> { if self.host.abi_version < NEMO_RELAY_NATIVE_ABI_VERSION_LLM_EXECUTION_CONTEXT || self.host.struct_size < std::mem::size_of::() { + set_last_error( + self.host, + "LLM stream execution intercepts require Relay native ABI v7", + ); if let Some(free_fn) = free_fn { unsafe { free_fn(user_data) }; } diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index 060d83d54..fff9e4a10 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -27,6 +27,7 @@ use nemo_relay_plugin::{ LlmStreamNext, LogSeverity, MetricKind, MetricMeasurement, MetricValueType, NEMO_RELAY_NATIVE_ABI_VERSION, NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE, NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, + NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, NativeExecutorConfig, NativePlugin, NemoRelayNativeAsyncCallbackState, NemoRelayNativeAsyncCompletion, NemoRelayNativeAsyncLlmExecutionCb, NemoRelayNativeAsyncLlmStreamOpenCb, @@ -1976,9 +1977,9 @@ impl MockAsyncOutput { std::mem::take(&mut *events) } - fn wait_for_release(&self) { + fn wait_for_releases(&self, expected: usize) { wait_until("async output was not released", || { - self.releases.load(Ordering::SeqCst) != 0 + self.releases.load(Ordering::SeqCst) >= expected }); } } @@ -4804,7 +4805,7 @@ fn typed_async_middleware_registers_and_round_trips_every_surface() { MockOutputEvent::Finished, ] ); - output.wait_for_release(); + output.wait_for_releases(2); assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(pull_stream.releases.load(Ordering::SeqCst), 1); assert_eq!(next.releases.load(Ordering::SeqCst), 3); @@ -5275,7 +5276,7 @@ fn typed_async_continuations_are_concurrent_and_executor_owned() { "host returned neither an LLM stream nor an error".into() )] ); - output.wait_for_release(); + output.wait_for_releases(2); assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; @@ -5538,7 +5539,7 @@ fn typed_async_stream_cancellation_while_polling_releases_output() { started.load(Ordering::SeqCst) }); output.cancelled.store(true, Ordering::SeqCst); - output.wait_for_release(); + output.wait_for_releases(1); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; } @@ -5586,7 +5587,7 @@ fn typed_async_stream_restores_callback_scope_while_polling_returned_stream() { MockOutputEvent::Finished, ] ); - output.wait_for_release(); + output.wait_for_releases(1); assert!(SCOPE_STACK_BINDING_RESTORES.load(Ordering::SeqCst) >= 3); unsafe { registration.free() }; } @@ -5637,7 +5638,7 @@ fn typed_async_stream_rejects_item_errors_and_releases_output() { MockOutputEvent::Rejected("stream item failed".into()), ] ); - output.wait_for_release(); + output.wait_for_releases(2); assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; @@ -5693,7 +5694,7 @@ fn typed_async_stream_rejects_poll_panics_and_releases_output() { MockOutputEvent::Rejected("typed native stream panicked while polling".into()), ] ); - output.wait_for_release(); + output.wait_for_releases(2); assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(next.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; @@ -5742,7 +5743,7 @@ fn typed_async_stream_propagates_downstream_pull_errors() { output.wait_terminal(), vec![MockOutputEvent::Rejected("downstream pull failed".into())] ); - output.wait_for_release(); + output.wait_for_releases(2); assert_eq!(output.releases.load(Ordering::SeqCst), 2); assert_eq!(pull_stream.releases.load(Ordering::SeqCst), 1); assert_eq!(next.releases.load(Ordering::SeqCst), 1); @@ -5790,7 +5791,7 @@ fn typed_async_stream_rejects_missing_continuation() { "native stream middleware requires a continuation".into() )] ); - output.wait_for_release(); + output.wait_for_releases(1); assert_eq!(output.releases.load(Ordering::SeqCst), 1); unsafe { registration.free() }; } @@ -6712,6 +6713,9 @@ fn exported_entry_symbol_rejects_prior_host_versions() { for abi_version in [ NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY, NEMO_RELAY_NATIVE_ABI_VERSION_ASYNC_MIDDLEWARE, + NEMO_RELAY_NATIVE_ABI_VERSION_RUNTIME_CONTROL, + NEMO_RELAY_NATIVE_ABI_VERSION_TOOL_EXECUTION_CONTEXT, + NEMO_RELAY_NATIVE_ABI_VERSION_LOGGING, ] { let mut host = test_host(); host.abi_version = abi_version; diff --git a/docs/build-plugins/native/about.mdx b/docs/build-plugins/native/about.mdx index f1367af7c..356a3da99 100644 --- a/docs/build-plugins/native/about.mdx +++ b/docs/build-plugins/native/about.mdx @@ -41,13 +41,13 @@ Relay 0.8 changed the native API 1 tool-result JSON contract without changing th host-table layout. A tool callback and `ToolNext` continuation return `ToolExecutionResult`, which carries `result` and optional opaque `annotation`; an execution intercept returns that pair plus pending marks. Rebuild native plugins for -that contract even though its negotiated host-table ABI remains v4. Relay 0.9 +that contract even though its negotiated host-table ABI remained v4. Relay 0.9 adds the context-carrying tool execution callback as the v5 host-table extension. -The v4 SDK also lets a plugin add a data schema and log severity to a mark, or emit -validated metric measurements. `PluginRuntime::runtime_diagnostics()` returns the -current host-level diagnostic snapshot. The snapshot is ordered by code and is not -attributed to a particular plugin. +In Relay 0.8, the v4 SDK also let a plugin add a data schema and log severity to a +mark, or emit validated metric measurements. `PluginRuntime::runtime_diagnostics()` +returns the current host-level diagnostic snapshot. The snapshot is ordered by code +and is not attributed to a particular plugin. ## What the SDK Owns From b8d13671aa06d7edde0d3fffe1bf5e33b991a27b Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Thu, 1 Oct 2026 09:40:13 -0700 Subject: [PATCH 21/22] refactor: rename execution codec context aliases Signed-off-by: Alex Fournier --- crates/core/src/api/runtime.rs | 10 ++--- crates/core/src/api/runtime/callbacks.rs | 12 ++++++ .../src/api/runtime/llm_execution_context.rs | 41 ++++++++----------- crates/core/src/plugin/dynamic/native.rs | 8 ++-- .../tests/unit/llm_execution_context_tests.rs | 6 +-- crates/ffi/nemo_relay.h | 14 ++++++- crates/ffi/src/callable.rs | 10 ++++- crates/ffi/tests/integration/api_tests.rs | 6 +-- crates/ffi/tests/unit/api/registry_tests.rs | 2 +- crates/ffi/tests/unit/api_tests.rs | 6 +-- crates/node/plugin.d.ts | 4 ++ crates/node/root-types.d.ts | 10 ++++- crates/node/tests/llm_tests.mjs | 10 ++++- crates/plugin/src/async_sdk.rs | 12 +++--- crates/plugin/src/lib.rs | 24 +++++------ crates/plugin/tests/typed_callbacks.rs | 10 ++--- crates/python/src/py_types/core.rs | 12 ++++-- crates/python/src/py_types/mod.rs | 7 +++- crates/worker/src/lib.rs | 20 +++++---- .../tests/test_worker.py | 8 ++-- go/nemo_relay/callbacks.go | 22 ++++++---- go/nemo_relay/nemo_relay.go | 4 +- go/nemo_relay/plugin.go | 4 +- python/nemo_relay/__init__.py | 4 ++ python/nemo_relay/__init__.pyi | 6 +++ python/nemo_relay/_native.pyi | 19 +++++---- .../plugin/src/nemo_relay_plugin/__init__.py | 14 +++++-- python/plugin/src/nemo_relay_plugin/_api.py | 14 ++++--- 28 files changed, 204 insertions(+), 115 deletions(-) diff --git a/crates/core/src/api/runtime.rs b/crates/core/src/api/runtime.rs index a86713a57..72ff74b1a 100644 --- a/crates/core/src/api/runtime.rs +++ b/crates/core/src/api/runtime.rs @@ -14,11 +14,11 @@ pub mod subscriber_dispatcher; pub use callbacks::{ BuiltinLlmCodec, ConditionalMiddlewareGuardrailFn, EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmCollectorFn, LlmConditionalFn, LlmExecutionFn, - LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmRequestInterceptFn, - LlmSanitizeRequestContext, LlmSanitizeRequestFn, LlmSanitizeResponseContext, - LlmSanitizeResponseFn, LlmStreamExecutionFn, LlmStreamExecutionNextFn, LlmStreamInner, - ToolConditionalFn, ToolExecutionContext, ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, - ToolSanitizeFn, + LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream, LlmRequestContext, LlmRequestInterceptFn, + LlmResponseContext, LlmSanitizeRequestContext, LlmSanitizeRequestFn, + LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionFn, + LlmStreamExecutionNextFn, LlmStreamInner, ToolConditionalFn, ToolExecutionContext, + ToolExecutionFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, }; #[doc(hidden)] pub use continuation_context::MiddlewareContinuationContext; diff --git a/crates/core/src/api/runtime/callbacks.rs b/crates/core/src/api/runtime/callbacks.rs index 642827dc7..02a9374c1 100644 --- a/crates/core/src/api/runtime/callbacks.rs +++ b/crates/core/src/api/runtime/callbacks.rs @@ -258,6 +258,12 @@ pub struct LlmSanitizeRequestContext { request_codec: Option>, } +/// Request codec context exposed to an LLM execution intercept. +/// +/// This is the same context passed to request sanitizers. The alias keeps the +/// execution API independent of sanitizer-specific naming. +pub type LlmRequestContext = LlmSanitizeRequestContext; + impl std::fmt::Debug for LlmSanitizeRequestContext { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter @@ -319,6 +325,12 @@ pub struct LlmSanitizeResponseContext { response_codec: Option>, } +/// Response codec context exposed to a non-streaming LLM execution intercept. +/// +/// This is the same context passed to response sanitizers. The alias keeps the +/// execution API independent of sanitizer-specific naming. +pub type LlmResponseContext = LlmSanitizeResponseContext; + impl std::fmt::Debug for LlmSanitizeResponseContext { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter diff --git a/crates/core/src/api/runtime/llm_execution_context.rs b/crates/core/src/api/runtime/llm_execution_context.rs index c1cc55d3f..80fafa6dc 100644 --- a/crates/core/src/api/runtime/llm_execution_context.rs +++ b/crates/core/src/api/runtime/llm_execution_context.rs @@ -6,7 +6,7 @@ use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Weak}; -use super::callbacks::{LlmSanitizeRequestContext, LlmSanitizeResponseContext}; +use super::callbacks::{LlmRequestContext, LlmResponseContext}; use crate::api::llm::LlmRequest; use crate::codec::request::AnnotatedLlmRequest; use crate::codec::response::AnnotatedLlmResponse; @@ -121,16 +121,16 @@ impl LlmResponseCodec for RevocableResponseCodec { /// active, revocable codec handles. #[derive(Clone, Debug, Default)] pub struct LlmExecutionContext { - request_codec: LlmSanitizeRequestContext, - response_codec: Option, + request_codec: LlmRequestContext, + response_codec: Option, } impl LlmExecutionContext { /// Construct an execution context from request and optional response codec context. #[must_use] pub fn new( - request_codec: LlmSanitizeRequestContext, - response_codec: Option, + request_codec: LlmRequestContext, + response_codec: Option, ) -> Self { Self { request_codec, @@ -144,19 +144,14 @@ impl LlmExecutionContext { response_codec: Option>, ) -> Self { Self::new( - LlmSanitizeRequestContext::for_request_codec(request_codec), - Some(LlmSanitizeResponseContext::for_response_codec( - response_codec, - )), + LlmRequestContext::for_request_codec(request_codec), + Some(LlmResponseContext::for_response_codec(response_codec)), ) } /// Construct the context for a streaming managed execution. pub(crate) fn for_streaming(request_codec: Option>) -> Self { - Self::new( - LlmSanitizeRequestContext::for_request_codec(request_codec), - None, - ) + Self::new(LlmRequestContext::for_request_codec(request_codec), None) } /// Issue revocable codec facades for one execution-intercept invocation. @@ -169,31 +164,31 @@ impl LlmExecutionContext { let gate = Arc::new(ExecutionCodecGate::new()); let leased_request_codec = self.request_codec.resolve_codec(); let request_codec = match leased_request_codec.as_ref() { - Some(codec) => LlmSanitizeRequestContext::for_request_codec(Some(Arc::new( - RevocableRequestCodec { + Some(codec) => { + LlmRequestContext::for_request_codec(Some(Arc::new(RevocableRequestCodec { codec: Arc::downgrade(codec), identity: self.request_codec.codec().clone(), gate: Arc::clone(&gate), - }, - ))), - None => LlmSanitizeRequestContext::with_identity(self.request_codec.codec().clone()), + }))) + } + None => LlmRequestContext::with_identity(self.request_codec.codec().clone()), }; let leased_response_codec = self .response_codec .as_ref() - .and_then(LlmSanitizeResponseContext::resolve_codec); + .and_then(LlmResponseContext::resolve_codec); let response_codec = self.response_codec .as_ref() .map(|context| match leased_response_codec.as_ref() { - Some(codec) => LlmSanitizeResponseContext::for_response_codec(Some(Arc::new( + Some(codec) => LlmResponseContext::for_response_codec(Some(Arc::new( RevocableResponseCodec { codec: Arc::downgrade(codec), identity: context.codec().clone(), gate: Arc::clone(&gate), }, ))), - None => LlmSanitizeResponseContext::with_identity(context.codec().clone()), + None => LlmResponseContext::with_identity(context.codec().clone()), }); ( @@ -208,7 +203,7 @@ impl LlmExecutionContext { /// Return the request-direction codec identity and revocable capability. #[must_use] - pub fn request_codec(&self) -> &LlmSanitizeRequestContext { + pub fn request_codec(&self) -> &LlmRequestContext { &self.request_codec } @@ -217,7 +212,7 @@ impl LlmExecutionContext { /// Streaming execution returns `None` because Relay does not expose a /// completed-response codec for individual stream chunks. #[must_use] - pub fn response_codec(&self) -> Option<&LlmSanitizeResponseContext> { + pub fn response_codec(&self) -> Option<&LlmResponseContext> { self.response_codec.as_ref() } } diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index e6f4d782c..7617087eb 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -70,9 +70,9 @@ use nemo_relay_plugin::{ NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestCodecContext, + NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestContext, NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, - NemoRelayNativeLlmResponseCodecContext, NemoRelayNativeLlmSanitizeRequestCb, + NemoRelayNativeLlmResponseContext, NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, NemoRelayNativePluginContext, @@ -658,7 +658,7 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { &self, callback: impl FnOnce(NemoRelayNativeLlmExecutionContext) -> T, ) -> T { - let request_codec = NemoRelayNativeLlmRequestCodecContext { + let request_codec = NemoRelayNativeLlmRequestContext { codec_kind: self.request_kind, codec_id: self .request_id @@ -670,7 +670,7 @@ impl<'a> NativeLlmExecutionContextBridge<'a> { }; let response_codec = self.response_kind - .map(|codec_kind| NemoRelayNativeLlmResponseCodecContext { + .map(|codec_kind| NemoRelayNativeLlmResponseContext { codec_kind, codec_id: self .response_id diff --git a/crates/core/tests/unit/llm_execution_context_tests.rs b/crates/core/tests/unit/llm_execution_context_tests.rs index ab1751187..017f271b9 100644 --- a/crates/core/tests/unit/llm_execution_context_tests.rs +++ b/crates/core/tests/unit/llm_execution_context_tests.rs @@ -2,7 +2,7 @@ // SPDX-License-Identifier: Apache-2.0 use super::*; -use crate::api::runtime::{BuiltinLlmCodec, LlmCodecIdentity}; +use crate::api::runtime::{BuiltinLlmCodec, LlmCodecIdentity, LlmResponseContext}; struct LeaseProbeCodec { allows_estimated_cost: bool, @@ -55,7 +55,7 @@ fn retained_codecs_work_while_active_and_expire_with_their_lease() { let retained_request = leased_context.request_codec().resolve_codec().unwrap(); let retained_response = leased_context .response_codec() - .and_then(LlmSanitizeResponseContext::resolve_codec) + .and_then(LlmResponseContext::resolve_codec) .unwrap(); drop(leased_context); drop(context); @@ -107,7 +107,7 @@ fn estimated_cost_defaults_to_false_after_expiry() { let (leased_context, guard) = context.lease(); let retained = leased_context .response_codec() - .and_then(LlmSanitizeResponseContext::resolve_codec) + .and_then(LlmResponseContext::resolve_codec) .unwrap(); assert!(retained.allows_estimated_cost(&Json::Null)); diff --git a/crates/ffi/nemo_relay.h b/crates/ffi/nemo_relay.h index 0b9c5b91f..7988f7069 100644 --- a/crates/ffi/nemo_relay.h +++ b/crates/ffi/nemo_relay.h @@ -382,6 +382,16 @@ typedef NemoRelayStatus (*NemoRelayLlmRequestInterceptCb)(void *user_data, const char *annotated_json, char **out_outcome_json); +/** + * Request codec context exposed to an LLM execution intercept. + */ +typedef struct NemoRelayLlmSanitizeRequestContext NemoRelayLlmRequestContext; + +/** + * Response codec context exposed to an LLM execution intercept. + */ +typedef struct NemoRelayLlmSanitizeResponseContext NemoRelayLlmResponseContext; + /** * Directional codec context supplied to an LLM execution intercept. * @@ -394,11 +404,11 @@ typedef struct NemoRelayLlmExecutionContext { /** * Active request codec identity and capability. */ - struct NemoRelayLlmSanitizeRequestContext request_codec; + NemoRelayLlmRequestContext request_codec; /** * Active unary-response codec context, or null for streaming execution. */ - const struct NemoRelayLlmSanitizeResponseContext *response_codec; + const NemoRelayLlmResponseContext *response_codec; } NemoRelayLlmExecutionContext; /** diff --git a/crates/ffi/src/callable.rs b/crates/ffi/src/callable.rs index fcbcd3dd3..e43f1f602 100644 --- a/crates/ffi/src/callable.rs +++ b/crates/ffi/src/callable.rs @@ -201,6 +201,12 @@ pub struct NemoRelayLlmSanitizeResponseContext { pub codec: *const crate::types::FfiLlmSanitizeResponseCodec, } +/// Request codec context exposed to an LLM execution intercept. +pub type NemoRelayLlmRequestContext = NemoRelayLlmSanitizeRequestContext; + +/// Response codec context exposed to an LLM execution intercept. +pub type NemoRelayLlmResponseContext = NemoRelayLlmSanitizeResponseContext; + /// Directional codec context supplied to an LLM execution intercept. /// /// `request_codec` is always present. `response_codec` is non-null for unary @@ -210,9 +216,9 @@ pub struct NemoRelayLlmSanitizeResponseContext { #[repr(C)] pub struct NemoRelayLlmExecutionContext { /// Active request codec identity and capability. - pub request_codec: NemoRelayLlmSanitizeRequestContext, + pub request_codec: NemoRelayLlmRequestContext, /// Active unary-response codec context, or null for streaming execution. - pub response_codec: *const NemoRelayLlmSanitizeResponseContext, + pub response_codec: *const NemoRelayLlmResponseContext, } /// LLM request sanitizer. It receives the request first and its codec context diff --git a/crates/ffi/tests/integration/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index 5175da0b9..e52aa0a60 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -15,9 +15,9 @@ use serde_json::{Value as Json, json}; use uuid::Uuid; use nemo_relay_ffi::callable::{ - NemoRelayLlmExecNextFn, NemoRelayLlmExecutionContext, NemoRelayLlmSanitizeCodecKind, - NemoRelayLlmSanitizeRequestContext, NemoRelayLlmSanitizeResponseContext, - NemoRelayToolExecNextFn, + NemoRelayLlmExecNextFn, NemoRelayLlmExecutionContext, NemoRelayLlmRequestContext, + NemoRelayLlmSanitizeCodecKind, NemoRelayLlmSanitizeRequestContext, + NemoRelayLlmSanitizeResponseContext, NemoRelayToolExecNextFn, }; use nemo_relay_ffi::convert::nemo_relay_string_free; use nemo_relay_ffi::error::{NemoRelayStatus, nemo_relay_last_error, set_last_error}; diff --git a/crates/ffi/tests/unit/api/registry_tests.rs b/crates/ffi/tests/unit/api/registry_tests.rs index 911b8e40b..c3d71a0e7 100644 --- a/crates/ffi/tests/unit/api/registry_tests.rs +++ b/crates/ffi/tests/unit/api/registry_tests.rs @@ -2663,7 +2663,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { intercept_name.as_ptr(), request.as_ptr(), NemoRelayLlmExecutionContext { - request_codec: NemoRelayLlmSanitizeRequestContext { + request_codec: NemoRelayLlmRequestContext { codec_kind: NemoRelayLlmSanitizeCodecKind::None, codec_id: ptr::null(), codec: ptr::null(), diff --git a/crates/ffi/tests/unit/api_tests.rs b/crates/ffi/tests/unit/api_tests.rs index 83b8665dc..02893c96f 100644 --- a/crates/ffi/tests/unit/api_tests.rs +++ b/crates/ffi/tests/unit/api_tests.rs @@ -16,9 +16,9 @@ use serde_json::{Value as Json, json}; use uuid::Uuid; use crate::callable::{ - NemoRelayLlmExecNextFn, NemoRelayLlmExecutionContext, NemoRelayLlmSanitizeCodecKind, - NemoRelayLlmSanitizeRequestContext, NemoRelayLlmSanitizeResponseContext, - NemoRelayToolExecNextFn, + NemoRelayLlmExecNextFn, NemoRelayLlmExecutionContext, NemoRelayLlmRequestContext, + NemoRelayLlmSanitizeCodecKind, NemoRelayLlmSanitizeRequestContext, + NemoRelayLlmSanitizeResponseContext, NemoRelayToolExecNextFn, }; use crate::convert::nemo_relay_string_free; use crate::error::{NemoRelayStatus, nemo_relay_last_error}; diff --git a/crates/node/plugin.d.ts b/crates/node/plugin.d.ts index 86d0bb26f..afd4ef30a 100644 --- a/crates/node/plugin.d.ts +++ b/crates/node/plugin.d.ts @@ -9,6 +9,8 @@ import type { Json, LlmExecutionContext, LlmRequestInterceptOutcome, + LlmRequestContext, + LlmResponseContext, LlmSanitizeRequestContext, LlmSanitizeResponseContext, PendingMarkSpec, @@ -30,6 +32,8 @@ export type { LlmOptimizationTokenImpact, LlmOptimizationTokens, LlmRequestInterceptOutcome, + LlmRequestContext, + LlmResponseContext, LlmSanitizeRequestContext, LlmSanitizeResponseContext, } from './index'; diff --git a/crates/node/root-types.d.ts b/crates/node/root-types.d.ts index b5bd121e6..afab62ec1 100644 --- a/crates/node/root-types.d.ts +++ b/crates/node/root-types.d.ts @@ -25,12 +25,18 @@ export interface LlmSanitizeResponseContext { resolveCodec(): import('./typed').LlmResponseCodec | null; } +/** Request codec context exposed to an LLM execution intercept. */ +export type LlmRequestContext = LlmSanitizeRequestContext; + +/** Response codec context exposed to an LLM execution intercept. */ +export type LlmResponseContext = LlmSanitizeResponseContext; + /** Codec capabilities for one managed LLM execution intercept invocation. */ export interface LlmExecutionContext { /** Request codec identity plus optional decode and encode capability. */ - requestCodec: LlmSanitizeRequestContext; + requestCodec: LlmRequestContext; /** Unary response codec identity plus optional decode capability; `null` for streaming execution. */ - responseCodec: LlmSanitizeResponseContext | null; + responseCodec: LlmResponseContext | null; } /** Schema tag attached to an opaque optimization contribution payload. */ diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index 7de9ecc0a..024a2b807 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -2570,7 +2570,15 @@ describe('LLM intercepts', () => { ); assert.match( declarations, - /export interface LlmExecutionContext \{[\s\S]*?requestCodec: LlmSanitizeRequestContext[\s\S]*?responseCodec: LlmSanitizeResponseContext \| null/, + /export type LlmRequestContext = LlmSanitizeRequestContext/, + ); + assert.match( + declarations, + /export type LlmResponseContext = LlmSanitizeResponseContext/, + ); + assert.match( + declarations, + /export interface LlmExecutionContext \{[\s\S]*?requestCodec: LlmRequestContext[\s\S]*?responseCodec: LlmResponseContext \| null/, ); assert.doesNotMatch(declarations, /registerToolExecutionInterceptV2|scopeRegisterToolExecutionInterceptV2/); }); diff --git a/crates/plugin/src/async_sdk.rs b/crates/plugin/src/async_sdk.rs index 264ef7588..4d7c75fad 100644 --- a/crates/plugin/src/async_sdk.rs +++ b/crates/plugin/src/async_sdk.rs @@ -247,7 +247,7 @@ impl CompletionRef { host: Arc, codec: LlmCodecIdentity, resolved: bool, - ) -> Result { + ) -> Result { let resolved = if resolved { let status = unsafe { (host.async_completion_retain)(self.raw) }; status_result(status, "retain native async completion capability")?; @@ -260,7 +260,7 @@ impl CompletionRef { } else { None }; - Ok(LlmRequestCodecContext { codec, resolved }) + Ok(LlmRequestContext { codec, resolved }) } fn execution_response_context( @@ -268,7 +268,7 @@ impl CompletionRef { host: Arc, codec: LlmCodecIdentity, resolved: bool, - ) -> Result { + ) -> Result { let resolved = if resolved { let status = unsafe { (host.async_completion_retain)(self.raw) }; status_result(status, "retain native async completion capability")?; @@ -279,7 +279,7 @@ impl CompletionRef { } else { None }; - Ok(LlmResponseCodecContext { codec, resolved }) + Ok(LlmResponseContext { codec, resolved }) } } @@ -295,7 +295,7 @@ impl StreamRef { self, codec: LlmCodecIdentity, resolved: bool, - ) -> Result { + ) -> Result { let resolved = if resolved { let status = unsafe { (self.host.async_stream_retain)(self.raw) }; status_result(status, "retain native async stream capability")?; @@ -308,7 +308,7 @@ impl StreamRef { } else { None }; - Ok(LlmRequestCodecContext { codec, resolved }) + Ok(LlmRequestContext { codec, resolved }) } } diff --git a/crates/plugin/src/lib.rs b/crates/plugin/src/lib.rs index 444ec3af3..e0f443039 100644 --- a/crates/plugin/src/lib.rs +++ b/crates/plugin/src/lib.rs @@ -96,19 +96,19 @@ unsafe impl Send for LlmSanitizeResponseContext<'_> {} /// response codec context, while streaming execution leaves it unavailable /// until Relay has a completed-response streaming codec contract. pub struct LlmExecutionContext { - request_codec: LlmRequestCodecContext, - response_codec: Option, + request_codec: LlmRequestContext, + response_codec: Option, } /// Request codec context for one LLM execution intercept invocation. -pub struct LlmRequestCodecContext { +pub struct LlmRequestContext { /// Identity of the active request codec. pub codec: LlmCodecIdentity, resolved: Option, } /// Response codec context for one non-streaming LLM execution intercept invocation. -pub struct LlmResponseCodecContext { +pub struct LlmResponseContext { /// Identity of the active response codec. pub codec: LlmCodecIdentity, resolved: Option, @@ -117,13 +117,13 @@ pub struct LlmResponseCodecContext { impl LlmExecutionContext { /// Return the active request codec context. #[must_use] - pub fn request_codec(&self) -> &LlmRequestCodecContext { + pub fn request_codec(&self) -> &LlmRequestContext { &self.request_codec } /// Return the completed-response codec context, or `None` for streaming execution. #[must_use] - pub fn response_codec(&self) -> Option<&LlmResponseCodecContext> { + pub fn response_codec(&self) -> Option<&LlmResponseContext> { self.response_codec.as_ref() } } @@ -228,7 +228,7 @@ pub struct NemoRelayNativeLlmSanitizeResponseContext { /// Request codec context passed to a native LLM execution intercept. #[repr(C)] #[derive(Debug, Clone, Copy)] -pub struct NemoRelayNativeLlmRequestCodecContext { +pub struct NemoRelayNativeLlmRequestContext { /// Discriminator for the active request codec. pub codec_kind: NemoRelayNativeLlmCodecKind, /// Optional borrowed built-in or runtime codec identifier. @@ -240,7 +240,7 @@ pub struct NemoRelayNativeLlmRequestCodecContext { /// Response codec context passed to a native non-streaming LLM execution intercept. #[repr(C)] #[derive(Debug, Clone, Copy)] -pub struct NemoRelayNativeLlmResponseCodecContext { +pub struct NemoRelayNativeLlmResponseContext { /// Discriminator for the active response codec. pub codec_kind: NemoRelayNativeLlmCodecKind, /// Optional borrowed built-in or runtime codec identifier. @@ -261,9 +261,9 @@ pub struct NemoRelayNativeLlmResponseCodecContext { #[derive(Debug, Clone, Copy)] pub struct NemoRelayNativeLlmExecutionContext { /// Request codec context, always present. - pub request_codec: NemoRelayNativeLlmRequestCodecContext, + pub request_codec: NemoRelayNativeLlmRequestContext, /// Completed-response codec context, or null for streaming execution. - pub response_codec: *const NemoRelayNativeLlmResponseCodecContext, + pub response_codec: *const NemoRelayNativeLlmResponseContext, } /// Safe completion-backed request codec facade for typed native plugins. @@ -498,7 +498,7 @@ impl LlmExecutionResponseCodec { } } -impl LlmRequestCodecContext { +impl LlmRequestContext { /// Resolve the active request codec capability. #[must_use] pub fn resolve_codec(&self) -> Option<&LlmExecutionRequestCodec> { @@ -506,7 +506,7 @@ impl LlmRequestCodecContext { } } -impl LlmResponseCodecContext { +impl LlmResponseContext { /// Resolve the active response codec capability. #[must_use] pub fn resolve_codec(&self) -> Option<&LlmExecutionResponseCodec> { diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index fff9e4a10..cd1ab39d0 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -40,9 +40,9 @@ use nemo_relay_plugin::{ NemoRelayNativeHostApiV5, NemoRelayNativeHostApiV6, NemoRelayNativeHostApiV7, NemoRelayNativeLlmAsyncStream, NemoRelayNativeLlmCodecKind, NemoRelayNativeLlmConditionalCb, NemoRelayNativeLlmExecutionCb, NemoRelayNativeLlmExecutionContext, - NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestCodecContext, + NemoRelayNativeLlmRequestCodec, NemoRelayNativeLlmRequestContext, NemoRelayNativeLlmRequestInterceptCb, NemoRelayNativeLlmResponseCodec, - NemoRelayNativeLlmResponseCodecContext, NemoRelayNativeLlmSanitizeRequestCb, + NemoRelayNativeLlmResponseContext, NemoRelayNativeLlmSanitizeRequestCb, NemoRelayNativeLlmSanitizeRequestContext, NemoRelayNativeLlmSanitizeResponseCb, NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamExecutionCb, NemoRelayNativeLlmStreamV1, NemoRelayNativeLogLevel, NemoRelayNativePluginContext, @@ -3113,7 +3113,7 @@ fn invoke_async_llm_execution_registration( } struct TestExecutionContext { - response: Option>, + response: Option>, context: NemoRelayNativeLlmExecutionContext, } @@ -3121,14 +3121,14 @@ impl TestExecutionContext { fn new(unary: bool) -> Self { let request_codec = NonNull::::dangling().as_ptr(); let response = unary.then(|| { - Box::new(NemoRelayNativeLlmResponseCodecContext { + Box::new(NemoRelayNativeLlmResponseContext { codec_kind: NemoRelayNativeLlmCodecKind::Opaque, codec_id: ptr::null(), codec: NonNull::::dangling().as_ptr(), }) }); let context = NemoRelayNativeLlmExecutionContext { - request_codec: NemoRelayNativeLlmRequestCodecContext { + request_codec: NemoRelayNativeLlmRequestContext { codec_kind: NemoRelayNativeLlmCodecKind::Opaque, codec_id: ptr::null(), codec: request_codec, diff --git a/crates/python/src/py_types/core.rs b/crates/python/src/py_types/core.rs index cdacfc2e9..2ef4f9e5c 100644 --- a/crates/python/src/py_types/core.rs +++ b/crates/python/src/py_types/core.rs @@ -63,6 +63,8 @@ pub struct PyLlmSanitizeRequestContext { pub(crate) inner: LlmSanitizeRequestContext, } +pub(crate) type PyLlmRequestContext = PyLlmSanitizeRequestContext; + #[pymethods] impl PyLlmSanitizeRequestContext { /// The active codec identity for this request or response payload. @@ -87,6 +89,8 @@ pub struct PyLlmSanitizeResponseContext { pub(crate) inner: LlmSanitizeResponseContext, } +pub(crate) type PyLlmResponseContext = PyLlmSanitizeResponseContext; + #[pymethods] impl PyLlmSanitizeResponseContext { /// The active codec identity for this request or response payload. @@ -115,8 +119,8 @@ pub struct PyLlmExecutionContext { impl PyLlmExecutionContext { /// Request codec identity and optional decode/encode capability. #[getter] - fn request_codec(&self) -> PyLlmSanitizeRequestContext { - PyLlmSanitizeRequestContext { + fn request_codec(&self) -> PyLlmRequestContext { + PyLlmRequestContext { inner: self.inner.request_codec().clone(), } } @@ -126,11 +130,11 @@ impl PyLlmExecutionContext { /// Streaming execution returns ``None`` because Relay does not have a /// complete-response codec contract for response chunks. #[getter] - fn response_codec(&self) -> Option { + fn response_codec(&self) -> Option { self.inner .response_codec() .cloned() - .map(|inner| PyLlmSanitizeResponseContext { inner }) + .map(|inner| PyLlmResponseContext { inner }) } } diff --git a/crates/python/src/py_types/mod.rs b/crates/python/src/py_types/mod.rs index 6e56ec7f2..c127421b7 100644 --- a/crates/python/src/py_types/mod.rs +++ b/crates/python/src/py_types/mod.rs @@ -153,9 +153,14 @@ fn register_runtime_types(m: &Bound<'_, PyModule>) -> PyResult<()> { fn register_llm_types(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; - m.add_class::()?; m.add_class::()?; m.add_class::()?; + m.add("LlmRequestContext", m.getattr("LlmSanitizeRequestContext")?)?; + m.add( + "LlmResponseContext", + m.getattr("LlmSanitizeResponseContext")?, + )?; + m.add_class::()?; m.add_class::()?; m.add_class::()?; m.add_class::()?; diff --git a/crates/worker/src/lib.rs b/crates/worker/src/lib.rs index f87cf7a55..da513f4f1 100644 --- a/crates/worker/src/lib.rs +++ b/crates/worker/src/lib.rs @@ -315,6 +315,12 @@ pub struct LlmSanitizeResponseContext { invocation_id: Option, } +/// Request codec context supplied to an LLM execution intercept. +pub type LlmRequestContext = LlmSanitizeRequestContext; + +/// Response codec context supplied to an LLM execution intercept. +pub type LlmResponseContext = LlmSanitizeResponseContext; + impl LlmSanitizeRequestContext { /// Resolves the active request codec for this callback. #[must_use] @@ -395,8 +401,8 @@ impl WorkerResponseCodec { /// not change these identities; codec operations reject incompatible payloads. #[derive(Clone)] pub struct LlmExecutionContext { - request_codec: LlmSanitizeRequestContext, - response_codec: Option, + request_codec: LlmRequestContext, + response_codec: Option, } impl std::fmt::Debug for LlmExecutionContext { @@ -415,7 +421,7 @@ impl std::fmt::Debug for LlmExecutionContext { impl LlmExecutionContext { /// Request codec identity and invocation-scoped operations. #[must_use] - pub fn request_codec(&self) -> &LlmSanitizeRequestContext { + pub fn request_codec(&self) -> &LlmRequestContext { &self.request_codec } @@ -424,7 +430,7 @@ impl LlmExecutionContext { /// Streaming execution returns `None` because Relay response codecs decode /// completed provider responses, not individual stream chunks. #[must_use] - pub fn response_codec(&self) -> Option<&LlmSanitizeResponseContext> { + pub fn response_codec(&self) -> Option<&LlmResponseContext> { self.response_codec.as_ref() } } @@ -2933,12 +2939,12 @@ impl LlmPayload { let response_codec = context .response .as_ref() - .map(|response| -> Result { + .map(|response| -> Result { let identity = require_execution_field( response.codec.as_ref(), "response codec identity is missing", )?; - Ok(LlmSanitizeResponseContext { + Ok(LlmResponseContext { codec: codec_identity_from_proto(Some(identity)), runtime: Some(runtime.clone()), codec_capability_id: response.codec_capability_id.clone(), @@ -2952,7 +2958,7 @@ impl LlmPayload { )); } Ok(LlmExecutionContext { - request_codec: LlmSanitizeRequestContext { + request_codec: LlmRequestContext { codec: codec_identity_from_proto(Some(request_identity)), runtime: Some(runtime.clone()), codec_capability_id: request.codec_capability_id.clone(), diff --git a/examples/python-grpc-worker-plugin/tests/test_worker.py b/examples/python-grpc-worker-plugin/tests/test_worker.py index c21996fc2..7e96971e1 100644 --- a/examples/python-grpc-worker-plugin/tests/test_worker.py +++ b/examples/python-grpc-worker-plugin/tests/test_worker.py @@ -30,8 +30,8 @@ from nemo_relay_plugin import ( # noqa: E402 LlmCodecIdentity, LlmExecutionContext, - LlmSanitizeRequestContext, - LlmSanitizeResponseContext, + LlmRequestContext, + LlmResponseContext, PluginContext, PluginRuntime, ToolExecutionContext, @@ -98,8 +98,8 @@ def callback(context: MagicMock, method: str, name: str | None = None) -> Any: def execution_context(*, streaming: bool = False) -> LlmExecutionContext: - request = LlmSanitizeRequestContext(LlmCodecIdentity("none")) - response = None if streaming else LlmSanitizeResponseContext(LlmCodecIdentity("none")) + request = LlmRequestContext(LlmCodecIdentity("none")) + response = None if streaming else LlmResponseContext(LlmCodecIdentity("none")) return LlmExecutionContext(request_codec=request, response_codec=response) diff --git a/go/nemo_relay/callbacks.go b/go/nemo_relay/callbacks.go index c56bd22c2..c61ab939e 100644 --- a/go/nemo_relay/callbacks.go +++ b/go/nemo_relay/callbacks.go @@ -38,9 +38,11 @@ typedef struct NemoRelayLlmSanitizeResponseContext { const char* codec_id; const FfiLlmSanitizeResponseCodec* codec; } NemoRelayLlmSanitizeResponseContext; +typedef NemoRelayLlmSanitizeRequestContext NemoRelayLlmRequestContext; +typedef NemoRelayLlmSanitizeResponseContext NemoRelayLlmResponseContext; typedef struct NemoRelayLlmExecutionContext { - NemoRelayLlmSanitizeRequestContext request_codec; - const NemoRelayLlmSanitizeResponseContext* response_codec; + NemoRelayLlmRequestContext request_codec; + const NemoRelayLlmResponseContext* response_codec; } NemoRelayLlmExecutionContext; typedef void (*NemoRelayFreeFn)(void* user_data); @@ -233,31 +235,35 @@ type LLMCodec struct { CodecID *string } -// LLMSanitizeRequestContext provides request codec context to sanitizer and -// execution callbacks. +// LLMSanitizeRequestContext provides request codec context to sanitizer callbacks. type LLMSanitizeRequestContext struct { Codec LLMCodec resolved *LLMRequestSanitizeCodec } +// LLMRequestContext is the request codec context exposed to execution intercepts. +type LLMRequestContext = LLMSanitizeRequestContext + // ResolveCodec returns the active callback-scoped request codec, if any. func (context LLMSanitizeRequestContext) ResolveCodec() *LLMRequestSanitizeCodec { return context.resolved } -// LLMSanitizeResponseContext provides response codec context to sanitizer and -// execution callbacks. +// LLMSanitizeResponseContext provides response codec context to sanitizer callbacks. type LLMSanitizeResponseContext struct { Codec LLMCodec resolved *LLMResponseSanitizeCodec } +// LLMResponseContext is the response codec context exposed to execution intercepts. +type LLMResponseContext = LLMSanitizeResponseContext + // LLMExecutionContext provides invocation-scoped codec access to an LLM // execution intercept. RequestCodec is always present. ResponseCodec is // available for unary execution and nil for streaming execution. type LLMExecutionContext struct { - RequestCodec LLMSanitizeRequestContext - ResponseCodec *LLMSanitizeResponseContext + RequestCodec LLMRequestContext + ResponseCodec *LLMResponseContext } // ResolveCodec returns the active callback-scoped response codec, if any. diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index 553278239..afdb4cdf5 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -42,7 +42,9 @@ typedef struct FfiLlmSanitizeRequestCodec FfiLlmSanitizeRequestCodec; typedef struct FfiLlmSanitizeResponseCodec FfiLlmSanitizeResponseCodec; typedef struct NemoRelayLlmSanitizeRequestContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeRequestCodec* codec; } NemoRelayLlmSanitizeRequestContext; typedef struct NemoRelayLlmSanitizeResponseContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeResponseCodec* codec; } NemoRelayLlmSanitizeResponseContext; -typedef struct NemoRelayLlmExecutionContext { NemoRelayLlmSanitizeRequestContext request_codec; const NemoRelayLlmSanitizeResponseContext* response_codec; } NemoRelayLlmExecutionContext; +typedef NemoRelayLlmSanitizeRequestContext NemoRelayLlmRequestContext; +typedef NemoRelayLlmSanitizeResponseContext NemoRelayLlmResponseContext; +typedef struct NemoRelayLlmExecutionContext { NemoRelayLlmRequestContext request_codec; const NemoRelayLlmResponseContext* response_codec; } NemoRelayLlmExecutionContext; typedef void (*NemoRelayFreeFn)(void* user_data); diff --git a/go/nemo_relay/plugin.go b/go/nemo_relay/plugin.go index ed733ce0d..55761c7cd 100644 --- a/go/nemo_relay/plugin.go +++ b/go/nemo_relay/plugin.go @@ -14,7 +14,9 @@ typedef struct FfiLlmSanitizeRequestCodec FfiLlmSanitizeRequestCodec; typedef struct FfiLlmSanitizeResponseCodec FfiLlmSanitizeResponseCodec; typedef struct NemoRelayLlmSanitizeRequestContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeRequestCodec* codec; } NemoRelayLlmSanitizeRequestContext; typedef struct NemoRelayLlmSanitizeResponseContext { uint32_t codec_kind; const char* codec_id; const FfiLlmSanitizeResponseCodec* codec; } NemoRelayLlmSanitizeResponseContext; -typedef struct NemoRelayLlmExecutionContext { NemoRelayLlmSanitizeRequestContext request_codec; const NemoRelayLlmSanitizeResponseContext* response_codec; } NemoRelayLlmExecutionContext; +typedef NemoRelayLlmSanitizeRequestContext NemoRelayLlmRequestContext; +typedef NemoRelayLlmSanitizeResponseContext NemoRelayLlmResponseContext; +typedef struct NemoRelayLlmExecutionContext { NemoRelayLlmRequestContext request_codec; const NemoRelayLlmResponseContext* response_codec; } NemoRelayLlmExecutionContext; typedef void (*NemoRelayFreeFn)(void* user_data); typedef char* (*NemoRelayPluginValidateCb)(void* user_data, const char* plugin_config_json); diff --git a/python/nemo_relay/__init__.py b/python/nemo_relay/__init__.py index 75a55d51f..794ec0884 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -106,8 +106,10 @@ async def main(): LLMRequest, LLMRequestInterceptOutcome, LlmSanitizeRequestCodec, + LlmRequestContext, LlmSanitizeRequestContext, LlmSanitizeResponseCodec, + LlmResponseContext, LlmSanitizeResponseContext, LogSeverity, MarkEvent, @@ -843,6 +845,8 @@ def worker() -> None: "LlmSanitizeResponseGuardrail", "LlmCodecIdentity", "LlmExecutionContext", + "LlmRequestContext", + "LlmResponseContext", "LlmSanitizeRequestContext", "LlmSanitizeResponseContext", "LlmSanitizeRequestCodec", diff --git a/python/nemo_relay/__init__.pyi b/python/nemo_relay/__init__.pyi index 326ce1909..8842f67b3 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -83,12 +83,18 @@ from nemo_relay._native import ( from nemo_relay._native import ( LlmSanitizeRequestCodec as LlmSanitizeRequestCodec, ) +from nemo_relay._native import ( + LlmRequestContext as LlmRequestContext, +) from nemo_relay._native import ( LlmSanitizeRequestContext as LlmSanitizeRequestContext, ) from nemo_relay._native import ( LlmSanitizeResponseCodec as LlmSanitizeResponseCodec, ) +from nemo_relay._native import ( + LlmResponseContext as LlmResponseContext, +) from nemo_relay._native import ( LlmSanitizeResponseContext as LlmSanitizeResponseContext, ) diff --git a/python/nemo_relay/_native.pyi b/python/nemo_relay/_native.pyi index 63fb7b402..b8d5f35a1 100644 --- a/python/nemo_relay/_native.pyi +++ b/python/nemo_relay/_native.pyi @@ -126,14 +126,6 @@ class LlmCodecIdentity: @property def id(self) -> str | None: ... -class LlmExecutionContext: - """Codec capabilities for one managed LLM execution intercept invocation.""" - - @property - def request_codec(self) -> LlmSanitizeRequestContext: ... - @property - def response_codec(self) -> LlmSanitizeResponseContext | None: ... - class LlmSanitizeRequestContext: """Request codec context shared by sanitizer and execution callbacks.""" @@ -148,6 +140,17 @@ class LlmSanitizeResponseContext: def codec(self) -> LlmCodecIdentity: ... def resolve_codec(self) -> LlmSanitizeResponseCodec | None: ... +LlmRequestContext: TypeAlias = LlmSanitizeRequestContext +LlmResponseContext: TypeAlias = LlmSanitizeResponseContext + +class LlmExecutionContext: + """Codec capabilities for one managed LLM execution intercept invocation.""" + + @property + def request_codec(self) -> LlmRequestContext: ... + @property + def response_codec(self) -> LlmResponseContext | None: ... + class LlmSanitizeRequestCodec: def decode(self, request: LLMRequest) -> AnnotatedLLMRequest: ... def encode(self, annotated: AnnotatedLLMRequest, original: LLMRequest) -> LLMRequest: ... diff --git a/python/plugin/src/nemo_relay_plugin/__init__.py b/python/plugin/src/nemo_relay_plugin/__init__.py index 2789f87e4..a938e4098 100644 --- a/python/plugin/src/nemo_relay_plugin/__init__.py +++ b/python/plugin/src/nemo_relay_plugin/__init__.py @@ -27,10 +27,12 @@ EventSanitizeFields: Mutable event observability fields. LlmRequest: A Relay LLM request represented as a JSON object. LlmCodecIdentity: Typed discriminator for the active LLM codec. - LlmSanitizeRequestContext: Request codec context shared by sanitizer and - execution callbacks. - LlmSanitizeResponseContext: Response codec context shared by sanitizer and - execution callbacks. + LlmRequestContext: Request codec context supplied to an LLM execution + intercept. + LlmResponseContext: Response codec context supplied to an LLM execution + intercept. + LlmSanitizeRequestContext: Request codec context supplied to a sanitizer. + LlmSanitizeResponseContext: Response codec context supplied to a sanitizer. LlmExecutionContext: Invocation-scoped codec context supplied to an LLM execution intercept. WorkerRequestCodec: Invocation-scoped async proxy for an active request codec. @@ -117,11 +119,13 @@ LlmOptimizationTokens, LlmRequest, LlmRequestCallback, + LlmRequestContext, LlmRequestInterceptOutcome, LlmSanitizeRequestCallback, LlmSanitizeRequestContext, LlmSanitizeResponseCallback, LlmSanitizeResponseContext, + LlmResponseContext, LlmStreamExecutionCallback, LlmStreamNext, LogSeverity, @@ -184,8 +188,10 @@ "LlmOptimizationTokens", "LlmNext", "LlmRequest", + "LlmRequestContext", "LlmSanitizeRequestContext", "LlmSanitizeResponseContext", + "LlmResponseContext", "LlmRequestCallback", "LlmRequestInterceptOutcome", "LlmSanitizeRequestCallback", diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index 4f081dc4a..402558570 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -226,6 +226,10 @@ def resolve_codec(self) -> "WorkerResponseCodec | None": return WorkerResponseCodec(self._runtime, self._capability_id, self._invocation_id) +LlmRequestContext: TypeAlias = LlmSanitizeRequestContext +LlmResponseContext: TypeAlias = LlmSanitizeResponseContext + + @dataclass(frozen=True) class WorkerRequestCodec: """Invocation-scoped async proxy for an active request codec.""" @@ -267,8 +271,8 @@ class LlmExecutionContext: not select a new codec; incompatible codec operations fail. """ - request_codec: LlmSanitizeRequestContext - response_codec: LlmSanitizeResponseContext | None + request_codec: LlmRequestContext + response_codec: LlmResponseContext | None def _llm_codec_identity(invocation: pb.LlmInvocation) -> LlmCodecIdentity: @@ -311,7 +315,7 @@ def _llm_execution_context( raise WorkerSdkError("malformed LLM execution codec context: request codec identity is missing") request_id = context.request.codec_capability_id if context.request.HasField("codec_capability_id") else None - request_context = LlmSanitizeRequestContext( + request_context = LlmRequestContext( codec=_codec_identity( context.request.codec.kind, context.request.codec.id if context.request.codec.HasField("id") else None, @@ -320,12 +324,12 @@ def _llm_execution_context( _capability_id=request_id, _invocation_id=invocation_id, ) - response_context: LlmSanitizeResponseContext | None = None + response_context: LlmResponseContext | None = None if context.HasField("response"): if not context.response.HasField("codec"): raise WorkerSdkError("malformed LLM execution codec context: response codec identity is missing") response_id = context.response.codec_capability_id if context.response.HasField("codec_capability_id") else None - response_context = LlmSanitizeResponseContext( + response_context = LlmResponseContext( codec=_codec_identity( context.response.codec.kind, context.response.codec.id if context.response.codec.HasField("id") else None, From 182bc681d69beaf25fcf1a46b6a4f2c02832cfe5 Mon Sep 17 00:00:00 2001 From: Alex Fournier Date: Thu, 1 Oct 2026 09:49:24 -0700 Subject: [PATCH 22/22] style: sort Python context imports Signed-off-by: Alex Fournier --- python/nemo_relay/__init__.py | 4 ++-- python/nemo_relay/__init__.pyi | 10 +++++----- python/plugin/src/nemo_relay_plugin/__init__.py | 2 +- 3 files changed, 8 insertions(+), 8 deletions(-) diff --git a/python/nemo_relay/__init__.py b/python/nemo_relay/__init__.py index 794ec0884..8d3e86262 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -104,12 +104,12 @@ async def main(): LlmExecutionContext, LLMHandle, LLMRequest, + LlmRequestContext, LLMRequestInterceptOutcome, + LlmResponseContext, LlmSanitizeRequestCodec, - LlmRequestContext, LlmSanitizeRequestContext, LlmSanitizeResponseCodec, - LlmResponseContext, LlmSanitizeResponseContext, LogSeverity, MarkEvent, diff --git a/python/nemo_relay/__init__.pyi b/python/nemo_relay/__init__.pyi index 8842f67b3..c9cc3a0ca 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -77,14 +77,17 @@ from nemo_relay._native import ( from nemo_relay._native import ( LLMRequest as LLMRequest, ) +from nemo_relay._native import ( + LlmRequestContext as LlmRequestContext, +) from nemo_relay._native import ( LLMRequestInterceptOutcome as LLMRequestInterceptOutcome, ) from nemo_relay._native import ( - LlmSanitizeRequestCodec as LlmSanitizeRequestCodec, + LlmResponseContext as LlmResponseContext, ) from nemo_relay._native import ( - LlmRequestContext as LlmRequestContext, + LlmSanitizeRequestCodec as LlmSanitizeRequestCodec, ) from nemo_relay._native import ( LlmSanitizeRequestContext as LlmSanitizeRequestContext, @@ -92,9 +95,6 @@ from nemo_relay._native import ( from nemo_relay._native import ( LlmSanitizeResponseCodec as LlmSanitizeResponseCodec, ) -from nemo_relay._native import ( - LlmResponseContext as LlmResponseContext, -) from nemo_relay._native import ( LlmSanitizeResponseContext as LlmSanitizeResponseContext, ) diff --git a/python/plugin/src/nemo_relay_plugin/__init__.py b/python/plugin/src/nemo_relay_plugin/__init__.py index a938e4098..56e7d71cf 100644 --- a/python/plugin/src/nemo_relay_plugin/__init__.py +++ b/python/plugin/src/nemo_relay_plugin/__init__.py @@ -121,11 +121,11 @@ LlmRequestCallback, LlmRequestContext, LlmRequestInterceptOutcome, + LlmResponseContext, LlmSanitizeRequestCallback, LlmSanitizeRequestContext, LlmSanitizeResponseCallback, LlmSanitizeResponseContext, - LlmResponseContext, LlmStreamExecutionCallback, LlmStreamNext, LogSeverity,