use serde_json::json;
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::subscriber_dispatcher::{
PendingPublication, dispatch_sanitized_event, dispatch_transformed_event,
register_pending_publication,
};
use crate::api::runtime::{
EventSubscriberFn, ScopeStackHandle, ToolExecutionNextFn, with_active_event_uuid,
};
use crate::api::scope::event;
use crate::api::scope::{EmitMarkEventParams, ScopeHandle};
use crate::api::shared::{
ensure_runtime_owner, metadata_with_otel_error, metadata_with_otel_status, resolve_parent_uuid,
snapshot_event_sanitizers, snapshot_event_subscribers,
};
use crate::api::skill_load;
use crate::error::{FlowError, Result};
use crate::json::Json;
use chrono::{DateTime, TimeDelta, Utc};
use serde::{Deserialize, Serialize};
use typed_builder::TypedBuilder;
use uuid::Uuid;
pub use nemo_relay_types::api::tool::{ToolAttributes, ToolExecutionInterceptOutcome};
fn queue_sanitized_event(event: Event, subscribers: &[EventSubscriberFn]) -> bool {
let scope_stack = current_scope_stack();
queue_sanitized_event_with_scope_stack(event, subscribers, scope_stack)
}
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)
}
#[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct ToolHandle {
#[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 = ToolAttributes::empty())]
pub attributes: ToolAttributes,
#[builder(default)]
pub parent_uuid: Option<Uuid>,
#[builder(default, setter(into))]
pub tool_call_id: Option<String>,
}
#[derive(Debug, Clone, TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct CreateToolHandleParams<'a> {
pub name: &'a str,
#[builder(default)]
pub parent_uuid: Option<uuid::Uuid>,
#[builder(default = ToolAttributes::empty())]
pub attributes: ToolAttributes,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default, setter(into))]
pub tool_call_id: Option<String>,
#[builder(default)]
pub timestamp: Option<DateTime<Utc>>,
}
fn resolve_skill_loads(
name: &str,
args: &Json,
metadata: Option<&Json>,
) -> Vec<skill_load::SkillLoad> {
let already_handled = metadata
.and_then(Json::as_object)
.and_then(|metadata| metadata.get(skill_load::HANDLED_METADATA_KEY))
.and_then(Json::as_bool)
.unwrap_or(false);
if already_handled {
Vec::new()
} else if let Some(skill_loads) = skill_load::precomputed(metadata) {
skill_loads
} else {
skill_load::detect(name, args)
}
}
#[derive(Debug, Clone, TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct EndToolHandleParams<'a> {
pub handle: &'a ToolHandle,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default)]
pub timestamp: Option<DateTime<Utc>>,
}
#[derive(TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct ToolCallParams<'a> {
pub name: &'a str,
pub args: Json,
#[builder(default)]
pub parent: Option<&'a ScopeHandle>,
#[builder(default = ToolAttributes::empty())]
pub attributes: ToolAttributes,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default, setter(into))]
pub tool_call_id: Option<String>,
#[builder(default)]
pub timestamp: Option<DateTime<Utc>>,
}
#[derive(TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct ToolCallExecuteParams {
#[builder(setter(into))]
pub name: String,
pub args: Json,
pub func: ToolExecutionNextFn,
#[builder(default)]
pub parent: Option<ScopeHandle>,
#[builder(default = ToolAttributes::empty())]
pub attributes: ToolAttributes,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
}
#[derive(TypedBuilder)]
#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
pub struct ToolCallEndParams<'a> {
pub handle: &'a ToolHandle,
pub result: Json,
#[builder(default)]
pub data: Option<Json>,
#[builder(default)]
pub metadata: Option<Json>,
#[builder(default)]
pub timestamp: Option<DateTime<Utc>>,
}
pub fn tool_call(params: ToolCallParams<'_>) -> Result<ToolHandle> {
ensure_runtime_owner()?;
let scope_stack = current_scope_stack();
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.tool_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()))?;
(
state.tool_sanitize_request_entries(&scope_locals),
subscribers,
)
};
let skill_loads = resolve_skill_loads(params.name, ¶ms.args, params.metadata.as_ref());
let raw_args = params.args;
let (handle, event, marks) = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let handle = state.create_tool_handle(
CreateToolHandleParams::builder()
.name(params.name)
.parent_uuid_opt(resolve_parent_uuid(params.parent))
.attributes(params.attributes)
.data_opt(params.data)
.metadata_opt(params.metadata)
.tool_call_id_opt(params.tool_call_id)
.timestamp_opt(params.timestamp)
.build(),
);
let event = state.build_tool_start_event(&handle, None);
let marks = skill_loads
.into_iter()
.map(|skill_load| {
state.create_event(MarkEvent::new(
BaseEvent::builder()
.name("skill.load")
.parent_uuid(handle.uuid)
.timestamp(handle.started_at)
.data(json!({"skill_name": skill_load.name}))
.metadata(json!({
"skill_load_source": <&str>::from(skill_load.source),
"tool_name": handle.name,
}))
.build(),
None,
None,
))
})
.collect::<Vec<_>>();
(handle, event, marks)
};
let tool_name = handle.name.clone();
let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
dispatch_transformed_event(
event,
Box::new(move |mut event| {
Box::pin(async move {
let sanitized = NemoRelayContextState::tool_sanitize_request_snapshot_chain(
&tool_name, raw_args, &entries,
)
.await;
let mut fields = event.sanitize_fields();
fields.data = sanitized;
event.apply_sanitize_fields(fields);
event
})
}),
event_sanitizers,
&subscribers,
scope_stack.clone(),
);
for mark in marks {
let sanitizers = snapshot_event_sanitizers(&mark, &scope_stack).unwrap_or_default();
dispatch_sanitized_event(mark, sanitizers, &subscribers, scope_stack.clone());
}
Ok(handle)
}
async fn tool_call_with_subscriber_snapshot(
params: ToolCallParams<'_>,
) -> Result<(ToolHandle, Vec<EventSubscriberFn>)> {
ensure_runtime_owner()?;
let parent_uuid = resolve_parent_uuid(params.parent);
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.tool_sanitize_request_guardrails
});
let scope_subscribers = scope_guard.collect_scope_local_subscribers();
let subscribers = snapshot_event_subscribers(scope_subscribers)?;
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let entries = state.tool_sanitize_request_entries(&scope_locals);
(entries, subscribers)
};
let skill_loads = resolve_skill_loads(params.name, ¶ms.args, params.metadata.as_ref());
let sanitized_args = NemoRelayContextState::tool_sanitize_request_snapshot_chain(
params.name,
params.args,
&entries,
)
.await;
let (handle, event, marks) = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
let handle_params = CreateToolHandleParams::builder()
.name(params.name)
.parent_uuid_opt(parent_uuid)
.attributes(params.attributes)
.data_opt(params.data)
.metadata_opt(params.metadata)
.tool_call_id_opt(params.tool_call_id)
.timestamp_opt(params.timestamp)
.build();
let handle = state.create_tool_handle(handle_params);
let event = state.build_tool_start_event(&handle, sanitized_args);
let marks = skill_loads
.into_iter()
.map(|skill_load| {
state.create_event(MarkEvent::new(
BaseEvent::builder()
.name("skill.load")
.parent_uuid(handle.uuid)
.timestamp(handle.started_at)
.data(json!({"skill_name": skill_load.name}))
.metadata(json!({
"skill_load_source": <&str>::from(skill_load.source),
"tool_name": handle.name,
}))
.build(),
None,
None,
))
})
.collect::<Vec<_>>();
(handle, event, marks)
};
queue_sanitized_event(event, &subscribers);
for mark in marks {
queue_sanitized_event(mark, &subscribers);
}
Ok((handle, subscribers))
}
pub fn tool_call_end(params: ToolCallEndParams<'_>) -> Result<()> {
ensure_runtime_owner()?;
let scope_stack = current_scope_stack();
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.tool_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.tool_sanitize_response_entries(&scope_locals),
subscribers,
)
};
let result = params.result;
let fallback = params.data;
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_tool_end_event(
EndToolHandleParams::builder()
.handle(params.handle)
.data(Json::Null)
.metadata_opt(params.metadata)
.timestamp_opt(params.timestamp)
.build(),
)
};
let tool_name = params.handle.name.clone();
let event_sanitizers = snapshot_event_sanitizers(&event, &scope_stack).unwrap_or_default();
dispatch_transformed_event(
event,
Box::new(move |mut event| {
Box::pin(async move {
let sanitized = NemoRelayContextState::tool_sanitize_response_snapshot_chain(
&tool_name, result, &entries,
)
.await;
let mut fields = event.sanitize_fields();
fields.data = sanitized.and_then(|value| {
if value.is_null() {
fallback
} else {
Some(value)
}
});
event.apply_sanitize_fields(fields);
event
})
}),
event_sanitizers,
&subscribers,
scope_stack,
);
Ok(())
}
async fn tool_call_end_with_pending_marks(
params: ToolCallEndParams<'_>,
pending_marks: Vec<PendingMarkSpec>,
lifecycle_subscribers: Option<&[EventSubscriberFn]>,
) -> Result<()> {
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.tool_sanitize_response_guardrails
});
let subscribers = if lifecycle_subscribers.is_some() {
Vec::new()
} else {
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.tool_sanitize_response_entries(&scope_locals);
(entries, subscribers)
};
let subscribers = lifecycle_subscribers.unwrap_or(&subscribers);
let sanitized_result = NemoRelayContextState::tool_sanitize_response_snapshot_chain(
¶ms.handle.name,
params.result,
&entries,
)
.await;
let data = sanitized_result.and_then(|value| {
if value.is_null() {
params.data
} else {
Some(value)
}
});
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.build_tool_end_event(
EndToolHandleParams::builder()
.handle(params.handle)
.data_opt(data)
.metadata_opt(params.metadata)
.timestamp_opt(params.timestamp)
.build(),
)
};
let marks = pending_marks
.into_iter()
.enumerate()
.map(|(index, mark)| {
let timestamp = *event.timestamp()
+ TimeDelta::microseconds(i64::try_from(index).unwrap_or_default() + 1);
Event::Mark(MarkEvent::new(
BaseEvent::builder()
.name(mark.name)
.parent_uuid(params.handle.uuid)
.timestamp(timestamp)
.data_opt(mark.data)
.metadata_opt(mark.metadata)
.build(),
mark.category,
mark.category_profile,
))
})
.collect::<Vec<_>>();
queue_sanitized_event(event, subscribers);
for mark in marks {
queue_sanitized_event(mark, subscribers);
}
Ok(())
}
fn emit_tool_end_without_output(
handle: &ToolHandle,
metadata: Option<Json>,
lifecycle_subscribers: &[EventSubscriberFn],
scope_stack: ScopeStackHandle,
) -> Result<()> {
ensure_runtime_owner()?;
let event = {
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.end_tool_handle(handle, handle.data.clone(), metadata)
};
queue_sanitized_event_with_scope_stack(event, lifecycle_subscribers, scope_stack);
Ok(())
}
struct ManagedToolCompletion {
handle: Option<ToolHandle>,
metadata: Option<Json>,
subscribers: Vec<EventSubscriberFn>,
scope_stack: ScopeStackHandle,
pending_publication: Option<PendingPublication>,
}
impl ManagedToolCompletion {
fn new(
handle: &ToolHandle,
metadata: Option<Json>,
subscribers: &[EventSubscriberFn],
scope_stack: ScopeStackHandle,
) -> Self {
Self {
handle: Some(handle.clone()),
metadata,
subscribers: subscribers.to_vec(),
scope_stack,
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 ManagedToolCompletion {
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("tool execution cancelled".into()),
);
let _ = emit_tool_end_without_output(
&handle,
metadata,
&self.subscribers,
self.scope_stack.clone(),
);
drop(pending_publication);
}
}
pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result<Json> {
let ToolCallExecuteParams {
name,
args,
func,
parent,
attributes,
data,
metadata,
} = 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.tool_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.tool_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::tool_conditional_execution_snapshot_chain(
&name,
&args,
&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 intercept_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.tool_request_intercepts);
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.tool_request_intercept_entries(&scope_locals)
};
let intercepted_args = NemoRelayContextState::tool_request_intercepts_snapshot_chain(
&name,
args,
&intercept_entries,
)
.await?;
let (handle, lifecycle_subscribers) = tool_call_with_subscriber_snapshot(
ToolCallParams::builder()
.name(name.as_str())
.args(intercepted_args.clone())
.parent_opt(parent.as_ref())
.attributes(attributes)
.data_opt(data.clone())
.metadata_opt(metadata.clone())
.build(),
)
.await?;
let lifecycle_scope_stack = current_scope_stack();
let mut completion = ManagedToolCompletion::new(
&handle,
metadata.clone(),
&lifecycle_subscribers,
lifecycle_scope_stack.clone(),
);
let execution_name = name.clone();
let execution = with_active_event_uuid(handle.uuid, 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.tool_execution_intercepts);
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.tool_build_execution_chain(&execution_name, func, &scope_locals)
};
execution(intercepted_args).await
})
.await;
match execution {
Ok(outcome) => {
let ToolExecutionInterceptOutcome {
result,
pending_marks,
} = outcome;
let end_metadata = metadata_with_otel_status(metadata, "OK", None);
tool_call_end_with_pending_marks(
ToolCallEndParams::builder()
.handle(&handle)
.result(result.clone())
.data_opt(data)
.metadata_opt(end_metadata)
.build(),
pending_marks,
Some(&lifecycle_subscribers),
)
.await?;
completion.disarm();
Ok(result)
}
Err(error) => {
let end_metadata = metadata_with_otel_error(metadata, &error);
let _ = emit_tool_end_without_output(
&handle,
end_metadata,
&lifecycle_subscribers,
lifecycle_scope_stack,
);
completion.disarm();
Err(error)
}
}
}
pub async fn tool_request_intercepts(name: &str, args: Json) -> Result<Json> {
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.tool_request_intercepts);
let context = global_context();
let state = context
.read()
.map_err(|error| FlowError::Internal(error.to_string()))?;
state.tool_request_intercept_entries(&scope_locals)
};
NemoRelayContextState::tool_request_intercepts_snapshot_chain(name, args, &entries).await
}
pub async fn tool_conditional_execution(name: &str, args: &Json) -> 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.tool_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.tool_conditional_execution_entries(&scope_locals);
let subscribers = state.collect_event_subscribers(&scope_subscribers);
(entries, subscribers, resolve_parent_uuid(None))
};
if let Some(error) = NemoRelayContextState::tool_conditional_execution_snapshot_chain(
name,
args,
&entries,
&subscribers,
parent_uuid,
None,
)
.await?
{
return Err(FlowError::GuardrailRejected(error));
}
Ok(())
}
#[cfg(test)]
#[path = "../../tests/unit/tool_api_tests.rs"]
mod tests;