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::runtime::NemoRelayContextState;
use crate::api::runtime::global_context;
use crate::api::runtime::{
EventSubscriberFn, LlmCollectorFn, LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream,
LlmStreamExecutionNextFn,
};
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_status,
resolve_parent_uuid, run_request_intercepts_with_codec_and_recorder,
sanitize_event_with_scope_stack, 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,
};
#[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 { .. }),
)
}
fn emit_llm_start(
handle: &LlmHandle,
request: &LlmRequest,
annotated_request: Option<Arc<AnnotatedLlmRequest>>,
request_codec: Option<&dyn LlmCodec>,
) -> Result<()> {
ensure_runtime_owner()?;
let subscribers = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
snapshot_event_subscribers(scope_guard.collect_scope_local_subscribers())?
};
emit_llm_start_with_subscribers(
handle,
request,
annotated_request,
request_codec,
&subscribers,
)
}
fn emit_llm_start_with_subscribers(
handle: &LlmHandle,
request: &LlmRequest,
annotated_request: Option<Arc<AnnotatedLlmRequest>>,
request_codec: Option<&dyn LlmCodec>,
subscribers: &[EventSubscriberFn],
) -> Result<()> {
ensure_runtime_owner()?;
let entries = {
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)
};
let mut sanitized_request =
NemoRelayContextState::llm_sanitize_request_snapshot_chain(request.clone(), &entries);
let mut annotated_request = match request_codec {
Some(codec)
if sanitized_request.headers != request.headers
|| sanitized_request.content != request.content =>
{
codec.decode(&sanitized_request).ok().map(Arc::new)
}
_ => annotated_request,
};
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 !agent_is_fresh {
project_llm_request_to_current_user_turn(
&mut sanitized_request,
&mut annotated_request,
request_codec,
);
}
let input = serde_json::to_value(&sanitized_request).unwrap_or(Json::Null);
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_llm_start_event(handle, Some(input), annotated_request)
};
if let Some(event) = sanitize_event_with_scope_stack(event, scope_stack) {
NemoRelayContextState::emit_event(&event, subscribers);
}
Ok(())
}
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,
));
if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) {
NemoRelayContextState::emit_event(&event, subscribers);
}
}
Ok(())
}
pub(crate) fn emit_optimization_marks(handle: &LlmHandle, subscribers: &[EventSubscriberFn]) {
emit_optimization_marks_with(
handle,
subscribers,
|event| sanitize_event_with_scope_stack(event, handle.captured_scope_stack()),
|event, subscribers| NemoRelayContextState::try_emit_event(event, subscribers),
);
}
fn emit_optimization_marks_with(
handle: &LlmHandle,
subscribers: &[EventSubscriberFn],
mut sanitize: impl FnMut(Event) -> Option<Event>,
mut enqueue: impl FnMut(&Event, &[EventSubscriberFn]) -> bool,
) {
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 offset = contribution.sequence.unwrap_or(0).saturating_add(2);
let offset = i64::try_from(offset).unwrap_or(i64::MAX);
let request_ordered_timestamp = handle.started_at + TimeDelta::microseconds(offset);
let timestamp = recorded_at.max(request_ordered_timestamp);
let data = serde_json::to_value(&contribution).unwrap_or(Json::Null);
let event = Event::Mark(MarkEvent::new(
BaseEvent::builder()
.name("nemo_relay.llm.optimization")
.parent_uuid(handle.uuid)
.timestamp(timestamp)
.data(data)
.data_schema(DataSchema {
name: "nemo.relay.llm_optimization_contribution".to_string(),
version: "1".to_string(),
})
.build(),
Some(EventCategory::custom()),
Some(
CategoryProfile::builder()
.subtype("nemo_relay.llm.optimization")
.build(),
),
));
let Some(event) = sanitize(event) else {
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)?;
emit_llm_start(&handle, params.request, params.annotated_request, None)?;
Ok(handle)
}
#[derive(Clone, Copy)]
struct LlmCallEndBehavior {
response_codec_errors_fatal: bool,
attach_estimated_cost: bool,
}
pub fn llm_call_end(params: LlmCallEndParams<'_>) -> Result<()> {
llm_call_end_with_behavior(
params,
LlmCallEndBehavior {
response_codec_errors_fatal: true,
attach_estimated_cost: false,
},
None,
)
}
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)
};
let sanitized_response =
NemoRelayContextState::llm_sanitize_response_snapshot_chain(response, &entries);
let data = if sanitized_response.is_null() {
data
} else {
Some(sanitized_response)
};
let (mut annotated_response, decode_error) = resolve_llm_end_annotation(
annotated_response,
response_codec,
data.as_ref(),
&behavior,
&handle.name,
);
handle.optimization_recorder.close_for_finalization(None);
emit_optimization_marks(handle, &subscribers);
let pricing = crate::codec::response::active_pricing_resolver();
let summary = finalize_optimization_summary(
&handle.optimization_recorder,
annotated_response.as_mut(),
handle.model_name.as_deref(),
&pricing,
);
if annotated_response.is_none()
&& let Some(summary) = summary
{
annotated_response = Some(AnnotatedLlmResponse {
optimization_summary: Some(summary),
..AnnotatedLlmResponse::default()
});
}
let annotated_response = annotated_response.map(Arc::new);
let event = {
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(data)
.metadata_opt(end_metadata)
.annotated_response_opt(annotated_response)
.timestamp_opt(timestamp)
.build(),
)
};
if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) {
NemoRelayContextState::emit_event(&event, &subscribers);
}
if let Some(error) = decode_error
&& behavior.response_codec_errors_fatal
{
Err(error)
} else {
Ok(())
}
}
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)),
}
}
fn emit_llm_end_without_output(
handle: &LlmHandle,
metadata: Option<Json>,
lifecycle_subscribers: Option<&[EventSubscriberFn]>,
) -> Result<()> {
ensure_runtime_owner()?;
let subscribers = {
let scope_stack = handle.captured_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
match lifecycle_subscribers {
Some(subscribers) => subscribers.to_vec(),
None => snapshot_event_subscribers(scope_subscribers)?,
}
};
handle.optimization_recorder.close_for_finalization(None);
emit_optimization_marks(handle, &subscribers);
let pricing = crate::codec::response::active_pricing_resolver();
let annotated_response = finalize_optimization_summary(
&handle.optimization_recorder,
None,
handle.model_name.as_deref(),
&pricing,
)
.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, handle.data.clone(), metadata, annotated_response)
};
if let Some(event) = sanitize_event_with_scope_stack(event, handle.captured_scope_stack()) {
NemoRelayContextState::emit_event(&event, &subscribers);
}
Ok(())
}
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,
)? {
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?;
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.as_deref(),
&lifecycle_subscribers,
)?;
emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers)?;
handle
.optimization_recorder
.record_all(optimization_contributions);
emit_optimization_marks(&handle, &lifecycle_subscribers);
let execution_name = name.clone();
let execution =
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),
)?;
Ok(response)
}
Err(error) => {
let end_metadata =
metadata_with_otel_status(metadata, "ERROR", Some(error.to_string()));
let _ =
emit_llm_end_without_output(&handle, end_metadata, Some(&lifecycle_subscribers));
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,
)? {
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?;
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.as_deref(),
&lifecycle_subscribers,
)?;
emit_pending_request_marks(&handle, pending_marks, &lifecycle_subscribers)?;
handle
.optimization_recorder
.record_all(optimization_contributions);
emit_optimization_marks(&handle, &lifecycle_subscribers);
let execution_name = name.clone();
let execution =
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)
};
execution(intercepted_request).await
})
.await;
match execution {
Ok(raw_stream) => {
let wrapper = LlmStreamWrapper::new_managed(
raw_stream,
handle,
collector,
finalizer,
metadata,
response_codec,
lifecycle_subscribers,
);
Ok(LlmJsonStream::from_closeable(wrapper))
}
Err(error) => {
let end_metadata =
metadata_with_otel_status(metadata, "ERROR", Some(error.to_string()));
let _ =
emit_llm_end_without_output(&handle, end_metadata, Some(&lifecycle_subscribers));
Err(error)
}
}
}
pub 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,
)?;
inject_dynamo_session_ids(&mut outcome.request);
Ok(outcome)
}
pub 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,
)? {
return Err(FlowError::GuardrailRejected(error));
}
Ok(())
}
#[cfg(test)]
#[path = "../../tests/unit/llm_api_tests.rs"]
mod tests;