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, Event, MarkEvent, PendingMarkSpec};
use crate::api::runtime::NemoRelayContextState;
use crate::api::runtime::current_scope_stack;
use crate::api::runtime::global_context;
use crate::api::runtime::{
EventSubscriberFn, LlmCollectorFn, LlmExecutionNextFn, LlmFinalizerFn, LlmJsonStream,
LlmStreamExecutionNextFn,
};
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, snapshot_event_subscribers,
};
use crate::codec::request::AnnotatedLlmRequest;
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::{LlmAttributes, LlmRequest, LlmRequestInterceptOutcome};
#[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>,
}
#[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 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 = current_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 = 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_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 sanitized_request =
NemoRelayContextState::llm_sanitize_request_snapshot_chain(request.clone(), &entries);
let 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 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)
};
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,
));
NemoRelayContextState::emit_event(&event, subscribers);
}
Ok(())
}
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 = 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_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 decode_error = None;
let annotated_response = match annotated_response {
Some(annotated_response) => Some(annotated_response),
None => match (response_codec.as_ref(), data.as_ref()) {
(Some(codec), Some(response)) => match codec.decode_response(response) {
Ok(mut decoded) => {
if behavior.attach_estimated_cost {
attach_estimated_cost_for_provider(&mut decoded, Some(&handle.name));
}
Some(Arc::new(decoded))
}
Err(error) => {
decode_error = Some(error);
None
}
},
_ => None,
},
};
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(),
)
};
NemoRelayContextState::emit_event(&event, &subscribers);
if let Some(error) = decode_error
&& behavior.response_codec_errors_fatal
{
Err(error)
} else {
Ok(())
}
}
fn emit_llm_end_without_output(
handle: &LlmHandle,
metadata: Option<Json>,
lifecycle_subscribers: Option<&[EventSubscriberFn]>,
) -> Result<()> {
ensure_runtime_owner()?;
let (event, subscribers) = {
let scope_stack = current_scope_stack();
let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
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 event = state.end_llm_handle(handle, handle.data.clone(), metadata, None);
(event, subscribers)
};
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 (intercepted_request, annotated_request, pending_marks) =
run_request_intercepts_with_codec(&name, request, codec)?;
let 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(),
)?;
let lifecycle_subscribers = {
let scope_stack = current_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)?;
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(&name, func, &scope_locals)
};
match execution(intercepted_request).await {
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 (intercepted_request, annotated_request, pending_marks) =
run_request_intercepts_with_codec(&name, request, codec)?;
let 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(),
)?;
let lifecycle_subscribers = {
let scope_stack = current_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)?;
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(&name, func, &scope_locals)
};
match execution(intercepted_request).await {
Ok(raw_stream) => {
let wrapper = LlmStreamWrapper::new_managed(
raw_stream,
handle,
collector,
finalizer,
metadata,
response_codec,
lifecycle_subscribers,
);
Ok(Box::pin(wrapper) as LlmJsonStream)
}
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;