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::{EventSubscriberFn, ToolExecutionNextFn};
use crate::api::scope::event;
use crate::api::scope::{EmitMarkEventParams, ScopeHandle};
use crate::api::shared::{
ensure_runtime_owner, metadata_with_otel_status, resolve_parent_uuid, sanitize_event,
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};
#[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>>,
}
#[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> {
let (handle, _) = tool_call_with_subscriber_snapshot(params)?;
Ok(handle)
}
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 handled_skill_loads = params
.metadata
.as_ref()
.and_then(Json::as_object)
.and_then(|metadata| metadata.get(skill_load::HANDLED_METADATA_KEY))
.and_then(Json::as_bool)
.is_some_and(|handled| handled);
let skill_loads = if handled_skill_loads {
Vec::new()
} else if let Some(skill_loads) = skill_load::precomputed(params.metadata.as_ref()) {
skill_loads
} else {
skill_load::detect(params.name, ¶ms.args)
};
let sanitized_args = NemoRelayContextState::tool_sanitize_request_snapshot_chain(
params.name,
params.args,
&entries,
);
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, Some(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)
};
let marks = marks
.into_iter()
.filter_map(sanitize_event)
.collect::<Vec<_>>();
if let Some(event) = sanitize_event(event) {
NemoRelayContextState::emit_event(&event, &subscribers);
}
for mark in marks {
NemoRelayContextState::emit_event(&mark, &subscribers);
}
Ok((handle, subscribers))
}
pub fn tool_call_end(params: ToolCallEndParams<'_>) -> Result<()> {
tool_call_end_with_pending_marks(params, Vec::new(), None)
}
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,
);
let data = if sanitized_result.is_null() {
params.data
} else {
Some(sanitized_result)
};
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,
))
})
.filter_map(sanitize_event)
.collect::<Vec<_>>();
if let Some(event) = sanitize_event(event) {
NemoRelayContextState::emit_event(&event, subscribers);
}
for mark in marks {
NemoRelayContextState::emit_event(&mark, subscribers);
}
Ok(())
}
fn emit_tool_end_without_output(
handle: &ToolHandle,
metadata: Option<Json>,
lifecycle_subscribers: &[EventSubscriberFn],
) -> 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)
};
if let Some(event) = sanitize_event(event) {
NemoRelayContextState::emit_event(&event, lifecycle_subscribers);
}
Ok(())
}
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,
)? {
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,
)?;
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(),
)?;
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(&name, func, &scope_locals)
};
match execution(intercepted_args).await {
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),
)?;
Ok(result)
}
Err(error) => {
let end_metadata =
metadata_with_otel_status(metadata, "ERROR", Some(error.to_string()));
let _ = emit_tool_end_without_output(&handle, end_metadata, &lifecycle_subscribers);
Err(error)
}
}
}
pub 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)
}
pub 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,
)? {
return Err(FlowError::GuardrailRejected(error));
}
Ok(())
}
#[cfg(test)]
#[path = "../../tests/unit/tool_api_tests.rs"]
mod tests;