use std::future::Future;
use std::sync::Arc;
use chrono::{DateTime, TimeDelta, Utc};
use serde::{Deserialize, Serialize};
use serde_json::json;
use typed_builder::TypedBuilder;
use uuid::Uuid;
use crate::api::event::{
BaseEvent, CategoryProfile, DataSchema, Event, EventCategory, MarkEvent, PendingMarkSpec,
};
use crate::api::optimization::{
LlmOptimizationRecorder, finalize_optimization_summary, scope_llm_optimization_recorder,
};
use crate::api::registry::RuntimeRegistrationKind;
#[cfg(test)]
use crate::api::runtime::LlmCodecIdentity;
use crate::api::runtime::NemoRelayContextState;
use crate::api::runtime::global_context;
use crate::api::runtime::state::contextualize_stream;
use crate::api::runtime::subscriber_dispatcher::{
EventTransformFn, PendingPublication, dispatch_reserved_sanitized_event,
dispatch_sanitized_event, dispatch_transformed_event, register_pending_publication,
};
use crate::api::runtime::{
EventSubscriberFn, LlmCollectorFn, LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream,
LlmSanitizeRequestContext, LlmSanitizeResponseContext, LlmStreamExecutionNextFn,
MiddlewareContinuationContext, with_active_event_uuid,
};
use crate::api::runtime::{ScopeStackHandle, capture_traceparent, current_scope_stack};
use crate::api::scope::event;
use crate::api::scope::{EmitMarkEventParams, ScopeHandle, metadata_with_log_severity};
use crate::api::shared::{
ensure_runtime_owner, inject_dynamo_session_ids, inject_traceparent, inject_traceparent_value,
metadata_with_otel_error, metadata_with_otel_status, resolve_parent_uuid,
run_request_intercepts_with_codec_and_recorder, snapshot_event_sanitizers,
snapshot_event_subscribers,
};
use crate::codec::request::{AnnotatedLlmRequest, Message};
use crate::codec::response::{AnnotatedLlmResponse, attach_estimated_cost_for_provider};
use crate::codec::traits::{LlmCodec, LlmResponseCodec};
use crate::error::{FlowError, Result};
use crate::json::Json;
use crate::stream::LlmStreamWrapper;
pub use nemo_relay_types::api::llm::{
LLM_REQUEST_INTERCEPT_OUTCOME_SCHEMA, LlmAttributes, LlmRequest, LlmRequestInterceptOutcome,
};
const OBSERVABILITY_CREDENTIAL_HEADERS: [&str; 7] = [
"authorization",
"proxy-authorization",
"cookie",
"x-api-key",
"api-key",
"anthropic-api-key",
"x-goog-api-key",
];
fn queue_sanitized_event_with_scope_stack(
event: Event,
subscribers: &[EventSubscriberFn],
scope_stack: &ScopeStackHandle,
) -> bool {
let sanitizers = snapshot_event_sanitizers(&event, scope_stack).unwrap_or_default();
dispatch_sanitized_event(event, sanitizers, subscribers, scope_stack.clone())
}
#[derive(Clone)]
struct CapturedLlmScopeStack(ScopeStackHandle);
impl Default for CapturedLlmScopeStack {
fn default() -> Self {
Self(current_scope_stack())
}
}
impl std::fmt::Debug for CapturedLlmScopeStack {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("CapturedLlmScopeStack(..)")
}
}
#[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct LlmHandle {
#[builder(default = Uuid::now_v7())]
pub uuid: Uuid,
#[builder(default = Utc::now())]
pub started_at: DateTime<Utc>,
#[builder(setter(into))]
pub name: String,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default = LlmAttributes::empty())]
pub attributes: LlmAttributes,
#[builder(default)]
pub parent_uuid: Option<Uuid>,
#[builder(default, setter(into))]
pub model_name: Option<String>,
#[serde(skip, default)]
#[builder(default)]
pub optimization_recorder: LlmOptimizationRecorder,
#[serde(skip, default)]
#[builder(setter(skip), default)]
captured_scope_stack: CapturedLlmScopeStack,
}
impl LlmHandle {
pub(crate) fn captured_scope_stack(&self) -> &ScopeStackHandle {
&self.captured_scope_stack.0
}
}
#[derive(Debug, Clone, TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct CreateLlmHandleParams<'a> {
pub name: &'a str,
#[builder(default)]
pub uuid: Option<Uuid>,
#[builder(default)]
pub parent_uuid: Option<uuid::Uuid>,
#[builder(default = LlmAttributes::empty())]
pub attributes: LlmAttributes,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default, setter(into))]
pub model_name: Option<String>,
#[builder(default)]
pub timestamp: Option<DateTime<Utc>>,
}
#[derive(Clone, TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct EndLlmHandleParams<'a> {
pub handle: &'a LlmHandle,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default)]
pub annotated_response: Option<Arc<AnnotatedLlmResponse>>,
#[builder(default)]
pub timestamp: Option<DateTime<Utc>>,
}
#[derive(TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct LlmCallParams<'a> {
pub name: &'a str,
pub request: &'a LlmRequest,
#[builder(default)]
pub parent: Option<&'a ScopeHandle>,
#[builder(default = LlmAttributes::empty())]
pub attributes: LlmAttributes,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default, setter(into))]
pub model_name: Option<String>,
#[builder(default)]
pub annotated_request: Option<Arc<AnnotatedLlmRequest>>,
#[builder(default)]
pub timestamp: Option<DateTime<Utc>>,
}
#[derive(TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct LlmCallExecuteParams {
#[builder(setter(into))]
pub name: String,
pub request: LlmRequest,
pub func: LlmExecutionNextFn,
#[builder(default)]
pub parent: Option<ScopeHandle>,
#[builder(default = LlmAttributes::empty())]
pub attributes: LlmAttributes,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default, setter(into))]
pub model_name: Option<String>,
#[builder(default)]
pub codec: Option<Arc<dyn LlmCodec>>,
#[builder(default)]
pub response_codec: Option<Arc<dyn LlmResponseCodec>>,
}
#[derive(TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct LlmStreamCallExecuteParams {
#[builder(setter(into))]
pub name: String,
pub request: LlmRequest,
pub func: LlmStreamExecutionNextFn,
pub collector: LlmCollectorFn,
pub finalizer: LlmFinalizerFn,
#[builder(default)]
pub parent: Option<ScopeHandle>,
#[builder(default = LlmAttributes::empty())]
pub attributes: LlmAttributes,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default, setter(into))]
pub model_name: Option<String>,
#[builder(default)]
pub codec: Option<Arc<dyn LlmCodec>>,
#[builder(default)]
pub response_codec: Option<Arc<dyn LlmResponseCodec>>,
}
#[derive(TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct LlmCallEndParams<'a> {
pub handle: &'a LlmHandle,
pub response: Json,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default)]
pub annotated_response: Option<Arc<AnnotatedLlmResponse>>,
#[builder(default)]
pub response_codec: Option<Arc<dyn LlmResponseCodec>>,
#[builder(default)]
pub timestamp: Option<DateTime<Utc>>,
}
fn create_llm_handle(params: CreateLlmHandleParams<'_>) -> Result<LlmHandle> {
ensure_runtime_owner()?;
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
Ok(state.create_llm_handle(params))
}
fn request_turn_projection_needed<T>(
items: &[T],
is_user: &impl Fn(&T) -> bool,
is_instruction: &impl Fn(&T) -> bool,
) -> bool {
let Some(last_index) = items.len().checked_sub(1) else {
return false;
};
match items.iter().rposition(is_user) {
Some(start) => items[..start].iter().any(|item| !is_instruction(item)),
None => items
.iter()
.enumerate()
.any(|(index, item)| index != last_index && !is_instruction(item)),
}
}
fn retain_current_request_turn<T>(
items: &mut Vec<T>,
is_user: impl Fn(&T) -> bool,
is_instruction: impl Fn(&T) -> bool,
) -> bool {
if !request_turn_projection_needed(items, &is_user, &is_instruction) {
return false;
}
let last_index = items.len() - 1;
let Some(start) = items.iter().rposition(is_user) else {
let mut index = 0;
items.retain(|item| {
let retain = index == last_index || is_instruction(item);
index += 1;
retain
});
return true;
};
let mut current_turn = items.split_off(start);
items.retain(is_instruction);
items.append(&mut current_turn);
true
}
fn project_llm_request_to_current_user_turn(
request: &mut LlmRequest,
annotated_request: &mut Option<Arc<AnnotatedLlmRequest>>,
request_codec: Option<&dyn LlmCodec>,
) {
let Some(annotation) = annotated_request.as_mut() else {
return;
};
if !request_turn_projection_needed(
&annotation.messages,
&|message| matches!(message, Message::User { .. }),
&|message| matches!(message, Message::System { .. }),
) {
return;
}
let original_annotation = request_codec.map(|_| Arc::clone(annotation));
let projected = limit_annotated_request_history_to_current_user_turn(Arc::make_mut(annotation));
debug_assert!(projected);
if let Some(codec) = request_codec {
match codec.encode(annotation, request) {
Ok(encoded) => *request = encoded,
Err(_) => {
log::warn!(
target: "nemo_relay.observability",
event = "projection_failed",
projection = "llm_current_turn",
recovery = "preserve_full_history";
"LLM request projection failed; preserving full event history"
);
*annotation = original_annotation
.expect("codec-backed projection should preserve the original annotation")
}
}
}
}
fn limit_annotated_request_history_to_current_user_turn(
annotated_request: &mut AnnotatedLlmRequest,
) -> bool {
retain_current_request_turn(
&mut annotated_request.messages,
|message| matches!(message, Message::User { .. }),
|message| matches!(message, Message::System { .. }),
)
}
#[cfg(test)]
async fn emit_llm_start_with_subscribers(
handle: &LlmHandle,
request: &LlmRequest,
annotated_request: Option<Arc<AnnotatedLlmRequest>>,
request_codec: Option<Arc<dyn LlmCodec>>,
subscribers: &[EventSubscriberFn],
) -> Result<()> {
ensure_runtime_owner()?;
let (entries, full_payloads_enabled) = {
let scope_stack = handle.captured_scope_stack();
let scope_locals = scope_stack
.read()
.expect("scope stack lock poisoned")
.snapshot_scope_local_registries(|registries| {
®istries.llm_sanitize_request_guardrails
});
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[RuntimeRegistrationKind::LlmSanitizeRequestGuardrail]);
(
state.llm_sanitize_request_entries(&scope_local_refs),
state.observability_full_payloads_enabled,
)
};
let observable_request = remove_observability_credential_headers(request.clone());
let mut sanitized_request = NemoRelayContextState::llm_sanitize_request_snapshot_chain(
observable_request.clone(),
LlmSanitizeRequestContext::for_request_codec(request_codec.clone()),
&entries,
)
.await;
let request_changed = sanitized_request
.as_ref()
.is_some_and(|sanitized_request| sanitized_request != &observable_request);
let mut annotated_request = match (sanitized_request.as_ref(), request_codec.as_deref()) {
(Some(sanitized_request), Some(codec)) if request_changed => {
codec.decode(sanitized_request).ok().map(Arc::new)
}
(Some(_), _) if !request_changed => annotated_request,
(None, _) => None,
(Some(_), _) => None,
};
let scope_stack = handle.captured_scope_stack();
let agent_is_fresh = {
let mut scope_guard = scope_stack.write().expect("scope stack lock poisoned");
scope_guard.take_agent_freshness(handle.parent_uuid)
};
if !full_payloads_enabled
&& !agent_is_fresh
&& let Some(sanitized_request) = sanitized_request.as_mut()
{
project_llm_request_to_current_user_turn(
sanitized_request,
&mut annotated_request,
request_codec.as_deref(),
);
}
let input = sanitized_request
.as_ref()
.and_then(|sanitized_request| serde_json::to_value(sanitized_request).ok());
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_llm_start_event(handle, input, annotated_request)
};
queue_sanitized_event_with_scope_stack(event, subscribers, scope_stack);
Ok(())
}
fn queue_llm_start_with_subscribers(
handle: &LlmHandle,
request: &LlmRequest,
annotated_request: Option<Arc<AnnotatedLlmRequest>>,
request_codec: Option<Arc<dyn LlmCodec>>,
subscribers: &[EventSubscriberFn],
) -> Result<()> {
ensure_runtime_owner()?;
let scope_stack = handle.captured_scope_stack().clone();
let scope_locals = {
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_sanitize_request_guardrails
})
};
let (entries, full_payloads_enabled) = {
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[RuntimeRegistrationKind::LlmSanitizeRequestGuardrail]);
(
state.llm_sanitize_request_entries(&scope_local_refs),
state.observability_full_payloads_enabled,
)
};
let agent_is_fresh = {
let mut scope_guard = scope_stack.write().expect("scope stack lock poisoned");
scope_guard.take_agent_freshness(handle.parent_uuid)
};
let observable_request = remove_observability_credential_headers(request.clone());
let queued_handle = handle.clone();
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_llm_start_event(handle, None, None)
};
let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
dispatch_transformed_event(
event,
Box::new(move |event| {
Box::pin(async move {
let mut sanitized_request =
NemoRelayContextState::llm_sanitize_request_snapshot_chain(
observable_request.clone(),
LlmSanitizeRequestContext::for_request_codec(request_codec.clone()),
&entries,
)
.await;
let request_changed = sanitized_request
.as_ref()
.is_some_and(|sanitized| sanitized != &observable_request);
let mut annotation = match (sanitized_request.as_ref(), request_codec.as_deref()) {
(Some(sanitized), Some(codec)) if request_changed => {
codec.decode(sanitized).ok().map(Arc::new)
}
(Some(_), _) if !request_changed => annotated_request,
_ => None,
};
if !full_payloads_enabled
&& !agent_is_fresh
&& let Some(sanitized_request) = sanitized_request.as_mut()
{
project_llm_request_to_current_user_turn(
sanitized_request,
&mut annotation,
request_codec.as_deref(),
);
}
let input = sanitized_request
.as_ref()
.and_then(|request| serde_json::to_value(request).ok());
global_context()
.read()
.map(|state| state.build_llm_start_event(&queued_handle, input, annotation))
.unwrap_or(event)
})
}),
event_sanitizers,
subscribers,
scope_stack,
);
Ok(())
}
fn remove_observability_credential_headers(mut request: LlmRequest) -> LlmRequest {
request.headers.retain(|name, _| {
!OBSERVABILITY_CREDENTIAL_HEADERS
.iter()
.any(|credential_header| name.eq_ignore_ascii_case(credential_header))
});
request
}
#[cfg(test)]
fn emit_llm_start(
handle: &LlmHandle,
request: &LlmRequest,
annotated_request: Option<Arc<AnnotatedLlmRequest>>,
request_codec: Option<Arc<dyn LlmCodec>>,
) -> Result<()> {
let subscribers = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?
};
crate::api::runtime::subscriber_dispatcher::block_on_sanitizer_future(
emit_llm_start_with_subscribers(
handle,
request,
annotated_request,
request_codec,
&subscribers,
),
)
.map_err(FlowError::Internal)?
}
async fn emit_pending_request_marks(
handle: &LlmHandle,
marks: Vec<PendingMarkSpec>,
subscribers: &[EventSubscriberFn],
) -> Result<()> {
if marks.is_empty() {
return Ok(());
}
ensure_runtime_owner()?;
let timestamp = handle.started_at + TimeDelta::microseconds(1);
for (index, mark) in marks.into_iter().enumerate() {
let metadata = match metadata_with_log_severity(mark.metadata, mark.severity) {
Ok(metadata) => metadata,
Err(error) => {
let llm_uuid = handle.uuid.to_string();
log::warn!(
target: "nemo_relay.observability",
event = "llm_pending_mark_dropped",
llm_name = handle.name.as_str(),
llm_uuid = llm_uuid.as_str(),
pending_mark_index = index,
pending_mark_name = mark.name.as_str();
"LLM pending mark was dropped because its severity metadata is invalid: {error}"
);
continue;
}
};
let event = Event::Mark(MarkEvent::new(
BaseEvent::builder()
.name(mark.name)
.parent_uuid(handle.uuid)
.timestamp(timestamp)
.data_opt(mark.data)
.data_schema_opt(mark.data_schema)
.metadata_opt(metadata)
.build(),
mark.category,
mark.category_profile,
));
queue_sanitized_event_with_scope_stack(event, subscribers, handle.captured_scope_stack());
}
Ok(())
}
pub(crate) async fn emit_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) {
emit_optimization_marks_with_async(
handle,
subscribers,
|event| async { Some(event) },
|event, subscribers| {
queue_sanitized_event_with_scope_stack(
event.clone(),
subscribers,
handle.captured_scope_stack(),
)
},
)
.await;
}
pub(crate) async fn emit_reserved_optimization_marks(
handle: &LlmHandle,
subscribers: &[EventSubscriberFn],
) {
emit_optimization_marks_with_async(
handle,
subscribers,
|event| async { Some(event) },
|event, subscribers| {
let sanitizers =
snapshot_event_sanitizers(event, handle.captured_scope_stack()).unwrap_or_default();
dispatch_reserved_sanitized_event(
event.clone(),
sanitizers,
subscribers,
handle.captured_scope_stack().clone(),
)
},
)
.await;
}
fn enqueue_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) {
let contributions = handle.optimization_recorder.unemitted_with_timestamps();
if contributions.is_empty() || ensure_runtime_owner().is_err() {
return;
}
let scope_stack = handle.captured_scope_stack().clone();
for (contribution, recorded_at) in contributions {
let event = optimization_mark_event(handle, &contribution, recorded_at);
let sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
if dispatch_sanitized_event(event, sanitizers, subscribers, scope_stack.clone()) {
handle.optimization_recorder.mark_emitted(1);
} else {
break;
}
}
}
async fn emit_optimization_marks_with_async<F, Fut>(
handle: &LlmHandle,
subscribers: &[EventSubscriberFn],
mut sanitize: F,
mut enqueue: impl FnMut(&Event, &[EventSubscriberFn]) -> bool,
) where
F: FnMut(Event) -> Fut,
Fut: Future<Output = Option<Event>>,
{
let contributions = handle.optimization_recorder.unemitted_with_timestamps();
if contributions.is_empty() {
return;
}
if ensure_runtime_owner().is_err() {
log::warn!(
target: "nemo_relay.observability",
event = "optimization_marks_skipped",
reason = "runtime_owner_unavailable",
contribution_count = contributions.len();
"LLM optimization marks were skipped"
);
return;
}
for (contribution, recorded_at) in contributions {
let event = optimization_mark_event(handle, &contribution, recorded_at);
let Some(event) = sanitize(event).await else {
break;
};
if enqueue(&event, subscribers) {
handle.optimization_recorder.mark_emitted(1);
} else {
break;
}
}
}
fn optimization_mark_event(
handle: &LlmHandle,
contribution: &crate::codec::optimization::LlmOptimizationContribution,
recorded_at: DateTime<Utc>,
) -> Event {
let offset = contribution.sequence.unwrap_or(0).saturating_add(2);
let offset = i64::try_from(offset).unwrap_or(i64::MAX);
let request_ordered_timestamp = handle.started_at + TimeDelta::microseconds(offset);
Event::Mark(MarkEvent::new(
BaseEvent::builder()
.name("nemo_relay.llm.optimization")
.parent_uuid(handle.uuid)
.timestamp(recorded_at.max(request_ordered_timestamp))
.data(serde_json::to_value(contribution).unwrap_or(Json::Null))
.data_schema(DataSchema {
name: "nemo.relay.llm_optimization_contribution".to_string(),
version: "1".to_string(),
})
.build(),
Some(EventCategory::custom()),
Some(
CategoryProfile::builder()
.subtype("nemo_relay.llm.optimization")
.build(),
),
))
}
#[cfg(test)]
fn emit_optimization_marks_with<F>(
handle: &LlmHandle,
subscribers: &[EventSubscriberFn],
mut sanitize: F,
mut enqueue: impl FnMut(&Event, &[EventSubscriberFn]) -> bool,
) where
F: FnMut(Event) -> Option<Event>,
{
let contributions = handle.optimization_recorder.unemitted_with_timestamps();
if contributions.is_empty() || ensure_runtime_owner().is_err() {
return;
}
for (contribution, recorded_at) in contributions {
let event = optimization_mark_event(handle, &contribution, recorded_at);
let Some(event) = sanitize(event) else {
break;
};
if enqueue(&event, subscribers) {
handle.optimization_recorder.mark_emitted(1);
} else {
break;
}
}
}
pub fn llm_call(params: LlmCallParams<'_>) -> Result<LlmHandle> {
let handle_params = CreateLlmHandleParams::builder()
.name(params.name)
.parent_uuid_opt(resolve_parent_uuid(params.parent))
.attributes(params.attributes)
.data_opt(params.data)
.metadata_opt(params.metadata)
.model_name_opt(params.model_name)
.timestamp_opt(params.timestamp)
.build();
let handle = create_llm_handle(handle_params)?;
let scope_stack = handle.captured_scope_stack().clone();
let (scope_locals, scope_subscribers, agent_is_fresh) = {
let mut scope_guard = scope_stack
.write()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let scope_locals = scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_sanitize_request_guardrails
});
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
let agent_is_fresh = scope_guard.take_agent_freshness(handle.parent_uuid);
(scope_locals, scope_subscribers, agent_is_fresh)
};
let (entries, subscribers, full_payloads_enabled) = {
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let subscribers = snapshot_event_subscribers(scope_subscribers)?;
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[RuntimeRegistrationKind::LlmSanitizeRequestGuardrail]);
let entries = state.llm_sanitize_request_entries(&scope_local_refs);
let full_payloads_enabled = state.observability_full_payloads_enabled;
(entries, subscribers, full_payloads_enabled)
};
let request = remove_observability_credential_headers(params.request.clone());
let annotated_request = params.annotated_request;
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_llm_start_event(&handle, None, None)
};
let queued_handle = handle.clone();
let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
dispatch_transformed_event(
event,
Box::new(move |event| {
Box::pin(async move {
let mut sanitized_request =
NemoRelayContextState::llm_sanitize_request_snapshot_chain(
request.clone(),
LlmSanitizeRequestContext::default(),
&entries,
)
.await;
let request_changed = sanitized_request
.as_ref()
.is_some_and(|sanitized| sanitized != &request);
let mut annotation = if sanitized_request.is_none() || request_changed {
None
} else {
annotated_request
};
if !full_payloads_enabled
&& !agent_is_fresh
&& let Some(sanitized_request) = sanitized_request.as_mut()
{
project_llm_request_to_current_user_turn(
sanitized_request,
&mut annotation,
None,
);
}
let input = sanitized_request
.as_ref()
.and_then(|request| serde_json::to_value(request).ok());
let context = global_context();
match context.read() {
Ok(state) => state.build_llm_start_event(&queued_handle, input, annotation),
Err(_) => event,
}
})
}),
event_sanitizers,
&subscribers,
scope_stack,
);
Ok(handle)
}
#[derive(Clone, Copy)]
struct LlmCallEndBehavior {
attach_estimated_cost: bool,
}
struct LlmEndPayload {
data: Option<Json>,
annotated_response: Option<Arc<AnnotatedLlmResponse>>,
decode_error: Option<FlowError>,
}
fn queue_llm_end_event(
event: Event,
transform: EventTransformFn,
subscribers: &[EventSubscriberFn],
scope_stack: ScopeStackHandle,
) {
let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
dispatch_transformed_event(event, transform, event_sanitizers, subscribers, scope_stack);
}
async fn build_llm_end_payload(
handle: &LlmHandle,
response: Json,
fallback_data: Option<Json>,
annotated_response: Option<Arc<AnnotatedLlmResponse>>,
response_codec: Option<Arc<dyn LlmResponseCodec>>,
entries: &[crate::api::registry::Guardrail<crate::api::runtime::LlmSanitizeResponseFn>],
behavior: LlmCallEndBehavior,
) -> LlmEndPayload {
let response_was_null_without_fallback = response.is_null() && fallback_data.is_none();
let response = if response.is_null() {
fallback_data.unwrap_or(response)
} else {
response
};
let sanitized_response = NemoRelayContextState::llm_sanitize_response_snapshot_chain(
response.clone(),
LlmSanitizeResponseContext::for_response_codec(response_codec.clone()),
entries,
)
.await;
let response_changed = sanitized_response
.as_ref()
.is_some_and(|sanitized_response| sanitized_response != &response);
let data = match sanitized_response {
Some(response) if response_was_null_without_fallback && response.is_null() => None,
response => response,
};
let annotation_omitted = data.as_ref().is_none_or(Json::is_null);
let (mut annotated_response, decode_error) = if annotation_omitted {
(None, None)
} else {
resolve_llm_end_annotation(
(!response_changed).then_some(annotated_response).flatten(),
response_codec,
data.as_ref(),
&behavior,
&handle.name,
)
};
let pricing = crate::codec::response::active_pricing_resolver();
let summary = finalize_optimization_summary(
&handle.optimization_recorder,
annotated_response.as_mut(),
handle.model_name.as_deref(),
&pricing,
);
if !annotation_omitted
&& annotated_response.is_none()
&& let Some(summary) = summary
{
annotated_response = Some(AnnotatedLlmResponse {
optimization_summary: Some(summary),
..AnnotatedLlmResponse::default()
});
}
LlmEndPayload {
data,
annotated_response: annotated_response.map(Arc::new),
decode_error,
}
}
pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> {
ensure_runtime_owner()?;
let scope_stack = params.handle.captured_scope_stack().clone();
let (scope_locals, scope_subscribers) = {
let scope_guard = scope_stack
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
(
scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_sanitize_response_guardrails
}),
scope_guard.collect_scope_local_subscribers(),
)
};
let (entries, subscribers) = {
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let subscribers = snapshot_event_subscribers(scope_subscribers)?;
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[RuntimeRegistrationKind::LlmSanitizeResponseGuardrail]);
(
state.llm_sanitize_response_entries(&scope_local_refs),
subscribers,
)
};
let response = params.response;
let fallback_data = params.data;
let handle = params.handle.clone();
let metadata = params.metadata;
let timestamp = params.timestamp;
let annotated_response = params.annotated_response;
let response_codec = params.response_codec;
handle.optimization_recorder.close_for_finalization(None);
enqueue_optimization_marks(&handle, &subscribers);
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_llm_end_event(
EndLlmHandleParams::builder()
.handle(&handle)
.data(Json::Null)
.metadata_opt(metadata.clone())
.annotated_response_opt(annotated_response.clone())
.timestamp_opt(timestamp)
.build(),
)
};
let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
dispatch_transformed_event(
event,
Box::new(move |event| {
Box::pin(async move {
let payload = build_llm_end_payload(
&handle,
response,
fallback_data,
annotated_response,
response_codec,
&entries,
LlmCallEndBehavior {
attach_estimated_cost: true,
},
)
.await;
if let Some(error) = payload.decode_error {
log::error!(
target: "nemo_relay.runtime",
event = "manual_llm_response_codec_failed";
"Manual LLM response annotation failed during queued publication: {error}"
);
}
let context = global_context();
let Ok(state) = context.read() else {
return event;
};
let end_metadata = metadata_with_otel_status(metadata, "OK", None);
state.build_llm_end_event(
EndLlmHandleParams::builder()
.handle(&handle)
.data_opt(payload.data)
.metadata_opt(end_metadata)
.annotated_response_opt(payload.annotated_response)
.timestamp_opt(timestamp)
.build(),
)
})
}),
event_sanitizers,
&subscribers,
scope_stack,
);
Ok(())
}
async fn llm_call_end_with_behavior(
params: LlmCallEndParams<'_>,
behavior: LlmCallEndBehavior,
lifecycle_subscribers: Option<&[EventSubscriberFn]>,
) -> Result<()> {
let LlmCallEndParams {
handle,
response,
data,
metadata,
annotated_response,
response_codec,
timestamp,
} = params;
let timestamp = timestamp.unwrap_or_else(Utc::now);
ensure_runtime_owner()?;
let (scope_locals, scope_subscribers) = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
(
scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_sanitize_response_guardrails
}),
scope_guard.collect_scope_local_subscribers(),
)
};
let (entries, subscribers) = {
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let subscribers = match lifecycle_subscribers {
Some(subscribers) => subscribers.to_vec(),
None => snapshot_event_subscribers(scope_subscribers)?,
};
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[RuntimeRegistrationKind::LlmSanitizeResponseGuardrail]);
let entries = state.llm_sanitize_response_entries(&scope_local_refs);
(entries, subscribers)
};
handle.optimization_recorder.close_for_finalization(None);
enqueue_optimization_marks(handle, &subscribers);
let queued_handle = handle.clone();
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let end_metadata = metadata_with_otel_status(metadata.clone(), "OK", None);
state.build_llm_end_event(
EndLlmHandleParams::builder()
.handle(handle)
.data(Json::Null)
.metadata_opt(end_metadata)
.timestamp(timestamp)
.build(),
)
};
let scope_stack = handle.captured_scope_stack().clone();
queue_llm_end_event(
event,
Box::new(move |event| {
Box::pin(async move {
let payload = build_llm_end_payload(
&queued_handle,
response,
data,
annotated_response,
response_codec,
&entries,
behavior,
)
.await;
if let Some(error) = payload.decode_error {
log::error!(
target: "nemo_relay.runtime",
event = "managed_llm_response_codec_failed";
"Managed LLM response annotation failed during queued publication: {error}"
);
}
let context = global_context();
let Ok(state) = context.read() else {
return event;
};
let end_metadata = metadata_with_otel_status(metadata, "OK", None);
state.build_llm_end_event(
EndLlmHandleParams::builder()
.handle(&queued_handle)
.data_opt(payload.data)
.metadata_opt(end_metadata)
.annotated_response_opt(payload.annotated_response)
.timestamp(timestamp)
.build(),
)
})
}),
&subscribers,
scope_stack,
);
Ok(())
}
#[cfg(test)]
fn sanitize_context_for_request_codec(codec: Option<&dyn LlmCodec>) -> LlmSanitizeRequestContext {
LlmSanitizeRequestContext::with_identity(
codec.map_or(LlmCodecIdentity::None, LlmCodec::codec_identity),
)
}
#[cfg(test)]
pub(crate) fn sanitize_context_for_response_codec(
codec: Option<&dyn LlmResponseCodec>,
) -> LlmSanitizeResponseContext {
LlmSanitizeResponseContext::with_identity(
codec.map_or(LlmCodecIdentity::None, LlmResponseCodec::codec_identity),
)
}
fn resolve_llm_end_annotation(
annotated_response: Option<Arc<AnnotatedLlmResponse>>,
response_codec: Option<Arc<dyn LlmResponseCodec>>,
data: Option<&Json>,
behavior: &LlmCallEndBehavior,
provider_name: &str,
) -> (Option<AnnotatedLlmResponse>, Option<FlowError>) {
if let Some(annotated_response) = annotated_response {
let mut annotated_response = (*annotated_response).clone();
if behavior.attach_estimated_cost {
attach_estimated_cost_for_provider(&mut annotated_response, Some(provider_name));
}
return (Some(annotated_response), None);
}
let (Some(codec), Some(response)) = (response_codec, data) else {
return (None, None);
};
match codec.decode_response(response) {
Ok(mut decoded) => {
if behavior.attach_estimated_cost {
attach_estimated_cost_for_provider(&mut decoded, Some(provider_name));
}
(Some(decoded), None)
}
Err(error) => (None, Some(error)),
}
}
async fn emit_llm_end_without_output(
handle: &LlmHandle,
metadata: Option<Json>,
response_codec: Option<Arc<dyn LlmResponseCodec>>,
lifecycle_subscribers: Option<&[EventSubscriberFn]>,
) -> Result<()> {
ensure_runtime_owner()?;
let timestamp = Utc::now();
let (scope_locals, scope_subscribers) = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
(
scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_sanitize_response_guardrails
}),
scope_guard.collect_scope_local_subscribers(),
)
};
let (entries, subscribers) = {
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let subscribers = match lifecycle_subscribers {
Some(subscribers) => subscribers.to_vec(),
None => snapshot_event_subscribers(scope_subscribers)?,
};
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[RuntimeRegistrationKind::LlmSanitizeResponseGuardrail]);
let entries = state.llm_sanitize_response_entries(&scope_local_refs);
(entries, subscribers)
};
handle.optimization_recorder.close_for_finalization(None);
enqueue_optimization_marks(handle, &subscribers);
let queued_handle = handle.clone();
let fallback_data = handle.data.clone();
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_llm_end_event(
EndLlmHandleParams::builder()
.handle(handle)
.data(Json::Null)
.metadata_opt(metadata.clone())
.timestamp(timestamp)
.build(),
)
};
let scope_stack = handle.captured_scope_stack().clone();
queue_llm_end_event(
event,
Box::new(move |event| {
Box::pin(async move {
let had_fallback_data = fallback_data.is_some();
let data = match fallback_data {
Some(data) => {
NemoRelayContextState::llm_sanitize_response_snapshot_chain(
data,
LlmSanitizeResponseContext::for_response_codec(response_codec),
&entries,
)
.await
}
None => None,
};
let annotation_omitted = (had_fallback_data && data.is_none())
|| data.as_ref().is_some_and(Json::is_null);
let pricing = crate::codec::response::active_pricing_resolver();
let annotated_response = (!annotation_omitted)
.then(|| {
finalize_optimization_summary(
&queued_handle.optimization_recorder,
None,
queued_handle.model_name.as_deref(),
&pricing,
)
})
.flatten()
.map(|summary| {
Arc::new(AnnotatedLlmResponse {
optimization_summary: Some(summary),
..AnnotatedLlmResponse::default()
})
});
global_context()
.read()
.map(|state| {
state.build_llm_end_event(
EndLlmHandleParams::builder()
.handle(&queued_handle)
.data_opt(data)
.metadata_opt(metadata)
.annotated_response_opt(annotated_response)
.timestamp(timestamp)
.build(),
)
})
.unwrap_or(event)
})
}),
&subscribers,
scope_stack,
);
Ok(())
}
struct ManagedLlmCompletion {
handle: Option<LlmHandle>,
metadata: Option<Json>,
response_codec: Option<Arc<dyn LlmResponseCodec>>,
subscribers: Vec<EventSubscriberFn>,
pending_publication: Option<PendingPublication>,
}
impl ManagedLlmCompletion {
fn new(
handle: &LlmHandle,
metadata: Option<Json>,
response_codec: Option<Arc<dyn LlmResponseCodec>>,
subscribers: &[EventSubscriberFn],
) -> Self {
Self {
handle: Some(handle.clone()),
metadata,
response_codec,
subscribers: subscribers.to_vec(),
pending_publication: (!subscribers.is_empty())
.then(register_pending_publication)
.flatten(),
}
}
fn disarm(&mut self) {
self.handle = None;
drop(self.pending_publication.take());
}
}
impl Drop for ManagedLlmCompletion {
fn drop(&mut self) {
let pending_publication = self.pending_publication.take();
let Some(handle) = self.handle.take() else {
return;
};
let metadata = metadata_with_otel_status(
self.metadata.take(),
"ERROR",
Some("LLM execution cancelled".into()),
);
let scope_stack = handle.captured_scope_stack().clone();
let scope_locals = scope_stack.read().ok().map(|scope_guard| {
scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_sanitize_response_guardrails
})
});
let entries = match scope_locals {
Some(scope_locals) => {
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
global_context()
.read()
.ok()
.map(|state| {
state.registry_snapshot(&[
RuntimeRegistrationKind::LlmSanitizeResponseGuardrail,
])
})
.map(|state| state.llm_sanitize_response_entries(&scope_local_refs))
}
None => None,
};
handle
.optimization_recorder
.close_for_finalization(Some("execution_cancelled"));
enqueue_optimization_marks(&handle, &self.subscribers);
let event = global_context()
.read()
.ok()
.map(|state| state.end_llm_handle(&handle, None, metadata.clone(), None));
let Some(event) = event else {
return;
};
let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
let response_codec = self.response_codec.take();
let subscribers = std::mem::take(&mut self.subscribers);
let fallback_data = handle.data.clone();
dispatch_transformed_event(
event,
Box::new(move |event| {
Box::pin(async move {
let (Some(data), Some(entries)) = (fallback_data, entries) else {
return event;
};
let data = NemoRelayContextState::llm_sanitize_response_snapshot_chain(
data,
LlmSanitizeResponseContext::for_response_codec(response_codec),
&entries,
)
.await;
let annotation_omitted = data.as_ref().is_none_or(Json::is_null);
let annotated_response = (!annotation_omitted)
.then(|| {
let pricing = crate::codec::response::active_pricing_resolver();
finalize_optimization_summary(
&handle.optimization_recorder,
None,
handle.model_name.as_deref(),
&pricing,
)
})
.flatten()
.map(|summary| {
Arc::new(AnnotatedLlmResponse {
optimization_summary: Some(summary),
..AnnotatedLlmResponse::default()
})
});
global_context()
.read()
.map(|state| {
state.end_llm_handle(&handle, data, metadata, annotated_response)
})
.unwrap_or(event)
})
}),
event_sanitizers,
&subscribers,
scope_stack,
);
drop(pending_publication);
}
}
pub async fn llm_call_execute(params: LlmCallExecuteParams) -> Result<Json> {
let LlmCallExecuteParams {
name,
request,
func,
parent,
attributes,
data,
metadata,
model_name,
codec,
response_codec,
} = params;
ensure_runtime_owner()?;
{
let (entries, subscribers, parent_uuid, guardrail_metadata) = {
let scope_stack = current_scope_stack();
let (scope_locals, scope_subscribers) = {
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
(
scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_conditional_execution_guardrails
}),
scope_guard.collect_scope_local_subscribers(),
)
};
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[
RuntimeRegistrationKind::LlmConditionalExecutionGuardrail,
RuntimeRegistrationKind::Subscriber,
]);
let entries = state.llm_conditional_execution_entries(&scope_local_refs);
let subscribers = state.collect_event_subscribers(&scope_subscribers);
(
entries,
subscribers,
resolve_parent_uuid(parent.as_ref()),
metadata.clone(),
)
};
if let Some(error) = NemoRelayContextState::llm_conditional_execution_snapshot_chain(
&request,
&entries,
&subscribers,
parent_uuid,
guardrail_metadata,
)
.await?
{
let mut rejection_data = json!({});
if let Some(object) = rejection_data.as_object_mut() {
object.insert("rejected".into(), json!(true));
object.insert("rejection_reason".into(), json!(&error));
}
let _ = event(
EmitMarkEventParams::builder()
.name(&name)
.parent_opt(parent.as_ref())
.data(rejection_data)
.metadata_opt(metadata.clone())
.build(),
);
return Err(FlowError::GuardrailRejected(error));
}
}
let request_codec = codec.clone();
let llm_uuid = Uuid::now_v7();
let optimization_recorder = LlmOptimizationRecorder::default();
let (mut intercepted_request, annotated_request, pending_marks, optimization_contributions) =
scope_llm_optimization_recorder(optimization_recorder.clone(), async {
run_request_intercepts_with_codec_and_recorder(
&name,
request,
codec,
&optimization_recorder,
)
.await
})
.await?;
let mut handle = create_llm_handle(
CreateLlmHandleParams::builder()
.name(name.as_str())
.uuid(llm_uuid)
.parent_uuid_opt(resolve_parent_uuid(parent.as_ref()))
.attributes(attributes)
.data_opt(data.clone())
.metadata_opt(metadata.clone())
.model_name_opt(model_name)
.build(),
)?;
handle.optimization_recorder = optimization_recorder;
let lifecycle_subscribers = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?
};
let observability_request = intercepted_request.clone();
inject_traceparent(&mut intercepted_request, handle.uuid)?;
queue_llm_start_with_subscribers(
&handle,
&observability_request,
annotated_request.clone(),
request_codec.clone(),
&lifecycle_subscribers,
)?;
emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers).await?;
handle
.optimization_recorder
.record_all(optimization_contributions);
emit_optimization_marks(&handle, &lifecycle_subscribers).await;
let mut completion = ManagedLlmCompletion::new(
&handle,
metadata.clone(),
response_codec.clone(),
&lifecycle_subscribers,
);
let execution_name = name.clone();
let event_uuid = handle.uuid;
let execution = with_active_event_uuid(
event_uuid,
scope_llm_optimization_recorder(handle.optimization_recorder.clone(), async move {
let execution = {
let scope_stack = current_scope_stack();
let scope_locals = scope_stack
.read()
.expect("scope stack lock poisoned")
.snapshot_scope_local_registries(|registries| {
®istries.llm_execution_intercepts
});
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[RuntimeRegistrationKind::LlmExecutionIntercept]);
state.llm_build_execution_chain(&execution_name, func, &scope_local_refs)
};
execution(intercepted_request).await
}),
)
.await;
match execution {
Ok(response) => {
llm_call_end_with_behavior(
LlmCallEndParams::builder()
.handle(&handle)
.response(response.clone())
.data_opt(data)
.metadata_opt(metadata)
.response_codec_opt(response_codec)
.build(),
LlmCallEndBehavior {
attach_estimated_cost: true,
},
Some(&lifecycle_subscribers),
)
.await?;
completion.disarm();
Ok(response)
}
Err(error) => {
let end_metadata = metadata_with_otel_error(metadata, &error);
let _ = emit_llm_end_without_output(
&handle,
end_metadata,
response_codec,
Some(&lifecycle_subscribers),
)
.await;
completion.disarm();
Err(error)
}
}
}
pub async fn llm_stream_call_execute(params: LlmStreamCallExecuteParams) -> Result<LlmJsonStream> {
let LlmStreamCallExecuteParams {
name,
request,
func,
collector,
finalizer,
parent,
attributes,
data,
metadata,
model_name,
codec,
response_codec,
} = params;
ensure_runtime_owner()?;
{
let (entries, subscribers, parent_uuid, guardrail_metadata) = {
let scope_stack = current_scope_stack();
let (scope_locals, scope_subscribers) = {
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
(
scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_conditional_execution_guardrails
}),
scope_guard.collect_scope_local_subscribers(),
)
};
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[
RuntimeRegistrationKind::LlmConditionalExecutionGuardrail,
RuntimeRegistrationKind::Subscriber,
]);
let entries = state.llm_conditional_execution_entries(&scope_local_refs);
let subscribers = state.collect_event_subscribers(&scope_subscribers);
(
entries,
subscribers,
resolve_parent_uuid(parent.as_ref()),
metadata.clone(),
)
};
if let Some(error) = NemoRelayContextState::llm_conditional_execution_snapshot_chain(
&request,
&entries,
&subscribers,
parent_uuid,
guardrail_metadata,
)
.await?
{
let mut rejection_data = json!({});
if let Some(object) = rejection_data.as_object_mut() {
object.insert("rejected".into(), json!(true));
object.insert("rejection_reason".into(), json!(&error));
}
let _ = event(
EmitMarkEventParams::builder()
.name(&name)
.parent_opt(parent.as_ref())
.data(rejection_data)
.metadata_opt(metadata.clone())
.build(),
);
return Err(FlowError::GuardrailRejected(error));
}
}
let request_codec = codec.clone();
let llm_uuid = Uuid::now_v7();
let optimization_recorder = LlmOptimizationRecorder::default();
let (mut intercepted_request, annotated_request, pending_marks, optimization_contributions) =
scope_llm_optimization_recorder(optimization_recorder.clone(), async {
run_request_intercepts_with_codec_and_recorder(
&name,
request,
codec,
&optimization_recorder,
)
.await
})
.await?;
let mut handle = create_llm_handle(
CreateLlmHandleParams::builder()
.name(name.as_str())
.uuid(llm_uuid)
.parent_uuid_opt(resolve_parent_uuid(parent.as_ref()))
.attributes(attributes)
.data_opt(data.clone())
.metadata_opt(metadata.clone())
.model_name_opt(model_name)
.build(),
)?;
handle.optimization_recorder = optimization_recorder;
let lifecycle_subscribers = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?
};
let observability_request = intercepted_request.clone();
inject_traceparent(&mut intercepted_request, handle.uuid)?;
queue_llm_start_with_subscribers(
&handle,
&observability_request,
annotated_request,
request_codec.clone(),
&lifecycle_subscribers,
)?;
emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers).await?;
handle
.optimization_recorder
.record_all(optimization_contributions);
emit_optimization_marks(&handle, &lifecycle_subscribers).await;
let mut completion = ManagedLlmCompletion::new(
&handle,
metadata.clone(),
response_codec.clone(),
&lifecycle_subscribers,
);
let execution_name = name.clone();
let event_uuid = handle.uuid;
let execution = with_active_event_uuid(
event_uuid,
scope_llm_optimization_recorder(handle.optimization_recorder.clone(), async move {
let execution = {
let scope_stack = current_scope_stack();
let scope_locals = scope_stack
.read()
.expect("scope stack lock poisoned")
.snapshot_scope_local_registries(|registries| {
®istries.llm_stream_execution_intercepts
});
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let context = global_context();
let state = context
.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)
};
let execution_context = MiddlewareContinuationContext::capture();
execution(intercepted_request)
.await
.map(|stream| contextualize_stream(stream, execution_context))
}),
)
.await;
match execution {
Ok(raw_stream) => {
let wrapper = LlmStreamWrapper::new_managed(
raw_stream,
handle,
collector,
finalizer,
metadata,
response_codec,
lifecycle_subscribers,
);
completion.disarm();
Ok(LlmJsonStream::from_closeable(wrapper))
}
Err(error) => {
let end_metadata = metadata_with_otel_error(metadata, &error);
let _ = emit_llm_end_without_output(
&handle,
end_metadata,
response_codec,
Some(&lifecycle_subscribers),
)
.await;
completion.disarm();
Err(error)
}
}
}
pub async fn llm_request_intercepts(
name: &str,
request: LlmRequest,
) -> Result<LlmRequestInterceptOutcome> {
ensure_runtime_owner()?;
let entries = {
let scope_stack = current_scope_stack();
let scope_locals = scope_stack
.read()
.expect("scope stack lock poisoned")
.snapshot_scope_local_registries(|registries| ®istries.llm_request_intercepts);
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[RuntimeRegistrationKind::LlmRequestIntercept]);
state.llm_request_intercept_entries(&scope_local_refs)
};
let mut outcome = NemoRelayContextState::llm_request_intercepts_snapshot_chain(
name, request, None, &entries, false,
)
.await?;
inject_dynamo_session_ids(&mut outcome.request);
if let Ok(traceparent) = capture_traceparent() {
inject_traceparent_value(&mut outcome.request, traceparent);
}
Ok(outcome)
}
pub async fn llm_conditional_execution(request: &LlmRequest) -> Result<()> {
ensure_runtime_owner()?;
let (entries, subscribers, parent_uuid) = {
let scope_stack = current_scope_stack();
let (scope_locals, scope_subscribers) = {
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
(
scope_guard.snapshot_scope_local_registries(|registries| {
®istries.llm_conditional_execution_guardrails
}),
scope_guard.collect_scope_local_subscribers(),
)
};
let scope_local_refs = scope_locals.iter().collect::<Vec<_>>();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?
.registry_snapshot(&[
RuntimeRegistrationKind::LlmConditionalExecutionGuardrail,
RuntimeRegistrationKind::Subscriber,
]);
let entries = state.llm_conditional_execution_entries(&scope_local_refs);
let subscribers = state.collect_event_subscribers(&scope_subscribers);
(entries, subscribers, resolve_parent_uuid(None))
};
if let Some(error) = NemoRelayContextState::llm_conditional_execution_snapshot_chain(
request,
&entries,
&subscribers,
parent_uuid,
None,
)
.await?
{
return Err(FlowError::GuardrailRejected(error));
}
Ok(())
}
#[cfg(test)]
#[path = "../../tests/unit/llm_api_tests.rs"]
mod tests;