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,
};
#[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::{
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, current_scope_stack};
use crate::api::scope::event;
use crate::api::scope::{EmitMarkEventParams, ScopeHandle};
use crate::api::shared::{
ensure_runtime_owner, inject_dynamo_session_ids, 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 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 { .. }),
)
}
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_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_sanitize_request_guardrails
});
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
(
state.llm_sanitize_request_entries(&scope_locals),
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 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 mark in marks {
let event = Event::Mark(MarkEvent::new(
BaseEvent::builder()
.name(mark.name)
.parent_uuid(handle.uuid)
.timestamp(timestamp)
.data_opt(mark.data)
.metadata_opt(mark.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 (entries, subscribers, agent_is_fresh, full_payloads_enabled) = {
let mut scope_guard = scope_stack
.write()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_sanitize_request_guardrails
});
let subscribers =
snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?;
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let entries = state.llm_sanitize_request_entries(&scope_locals);
let full_payloads_enabled = state.observability_full_payloads_enabled;
drop(state);
let agent_is_fresh = scope_guard.take_agent_freshness(handle.parent_uuid);
(entries, subscribers, agent_is_fresh, 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 {
response_codec_errors_fatal: bool,
attach_estimated_cost: bool,
}
struct LlmEndPayload {
data: Option<Json>,
annotated_response: Option<Arc<AnnotatedLlmResponse>>,
decode_error: Option<FlowError>,
}
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 (entries, subscribers) = {
let scope_guard = scope_stack
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_sanitize_response_guardrails
});
let subscribers =
snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?;
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
(
state.llm_sanitize_response_entries(&scope_locals),
subscribers,
)
};
let response = params.response;
let fallback_data = params.data;
let handle = params.handle.clone();
let metadata = params.metadata;
let timestamp = params.timestamp;
let annotated_response = params.annotated_response;
let response_codec = params.response_codec;
handle.optimization_recorder.close_for_finalization(None);
enqueue_optimization_marks(&handle, &subscribers);
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_llm_end_event(
EndLlmHandleParams::builder()
.handle(&handle)
.data(Json::Null)
.metadata_opt(metadata.clone())
.annotated_response_opt(annotated_response.clone())
.timestamp_opt(timestamp)
.build(),
)
};
let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
dispatch_transformed_event(
event,
Box::new(move |event| {
Box::pin(async move {
let payload = build_llm_end_payload(
&handle,
response,
fallback_data,
annotated_response,
response_codec,
&entries,
LlmCallEndBehavior {
response_codec_errors_fatal: false,
attach_estimated_cost: false,
},
)
.await;
if let Some(error) = payload.decode_error {
log::error!(
target: "nemo_relay.runtime",
event = "manual_llm_response_codec_failed";
"Manual LLM response annotation failed during queued publication: {error}"
);
}
let context = global_context();
let Ok(state) = context.read() else {
return event;
};
let end_metadata = metadata_with_otel_status(metadata, "OK", None);
state.build_llm_end_event(
EndLlmHandleParams::builder()
.handle(&handle)
.data_opt(payload.data)
.metadata_opt(end_metadata)
.annotated_response_opt(payload.annotated_response)
.timestamp_opt(timestamp)
.build(),
)
})
}),
event_sanitizers,
&subscribers,
scope_stack,
);
Ok(())
}
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;
ensure_runtime_owner()?;
let (entries, subscribers) = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_sanitize_response_guardrails
});
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
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()))?;
let entries = state.llm_sanitize_response_entries(&scope_locals);
(entries, subscribers)
};
handle.optimization_recorder.close_for_finalization(None);
emit_optimization_marks(handle, &subscribers).await;
let payload = build_llm_end_payload(
handle,
response,
data,
annotated_response,
response_codec,
&entries,
behavior,
)
.await;
let 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, "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(),
)
};
queue_sanitized_event_with_scope_stack(event, &subscribers, handle.captured_scope_stack());
if let Some(error) = payload.decode_error
&& behavior.response_codec_errors_fatal
{
Err(error)
} else {
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 {
return (Some((*annotated_response).clone()), 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 (entries, subscribers) = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_sanitize_response_guardrails
});
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
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()))?;
let entries = state.llm_sanitize_response_entries(&scope_locals);
(entries, subscribers)
};
let had_fallback_data = handle.data.is_some();
let data = if let Some(data) = handle.data.clone() {
NemoRelayContextState::llm_sanitize_response_snapshot_chain(
data,
LlmSanitizeResponseContext::for_response_codec(response_codec),
&entries,
)
.await
} else {
None
};
let annotation_omitted =
(had_fallback_data && data.is_none()) || data.as_ref().is_some_and(Json::is_null);
handle.optimization_recorder.close_for_finalization(None);
emit_optimization_marks(handle, &subscribers).await;
let pricing = crate::codec::response::active_pricing_resolver();
let annotated_response = (!annotation_omitted)
.then(|| {
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()
})
});
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.end_llm_handle(handle, data, metadata, annotated_response)
};
queue_sanitized_event_with_scope_stack(event, &subscribers, handle.captured_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 entries = match scope_stack.read() {
Ok(scope_guard) => {
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_sanitize_response_guardrails
});
global_context()
.read()
.map(|state| state.llm_sanitize_response_entries(&scope_locals))
.unwrap_or_default()
}
Err(_) => Vec::new(),
};
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) = fallback_data 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_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_conditional_execution_guardrails
});
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let entries = state.llm_conditional_execution_entries(&scope_locals);
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 optimization_recorder = LlmOptimizationRecorder::default();
let (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())
.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())?
};
emit_llm_start_with_subscribers(
&handle,
&intercepted_request,
annotated_request.clone(),
request_codec.clone(),
&lifecycle_subscribers,
)
.await?;
emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers).await?;
handle
.optimization_recorder
.record_all(optimization_contributions);
emit_optimization_marks(&handle, &lifecycle_subscribers).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_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_execution_intercepts
});
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.llm_build_execution_chain(&execution_name, func, &scope_locals)
};
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 {
response_codec_errors_fatal: false,
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_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_conditional_execution_guardrails
});
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let entries = state.llm_conditional_execution_entries(&scope_locals);
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 optimization_recorder = LlmOptimizationRecorder::default();
let (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())
.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())?
};
emit_llm_start_with_subscribers(
&handle,
&intercepted_request,
annotated_request,
request_codec.clone(),
&lifecycle_subscribers,
)
.await?;
emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers).await?;
handle
.optimization_recorder
.record_all(optimization_contributions);
emit_optimization_marks(&handle, &lifecycle_subscribers).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_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_stream_execution_intercepts
});
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.llm_stream_build_execution_chain(&execution_name, func, &scope_locals)
};
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_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard
.collect_scope_local_registries(|registries| ®istries.llm_request_intercepts);
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.llm_request_intercept_entries(&scope_locals)
};
let mut outcome = NemoRelayContextState::llm_request_intercepts_snapshot_chain(
name, request, None, &entries, false,
)
.await?;
inject_dynamo_session_ids(&mut outcome.request);
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_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
®istries.llm_conditional_execution_guardrails
});
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let entries = state.llm_conditional_execution_entries(&scope_locals);
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;