1use serde_json::json;
5
6use crate::api::runtime::NemoRelayContextState;
7use crate::api::runtime::ToolExecutionNextFn;
8use crate::api::runtime::current_scope_stack;
9use crate::api::runtime::global_context;
10use crate::api::scope::event;
11use crate::api::scope::{EmitMarkEventParams, ScopeHandle};
12use crate::api::shared::{
13 ensure_runtime_owner, metadata_with_otel_status, resolve_parent_uuid,
14 snapshot_event_subscribers,
15};
16use crate::error::{FlowError, Result};
17use crate::json::Json;
18use bitflags::bitflags;
19use chrono::{DateTime, Utc};
20use serde::{Deserialize, Serialize};
21use typed_builder::TypedBuilder;
22use uuid::Uuid;
23
24bitflags! {
25 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
27 pub struct ToolAttributes: u32 {
28 const REMOTE = 0b01;
30 }
31}
32
33#[derive(Debug, Clone, Serialize, Deserialize, TypedBuilder)]
35#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
36pub struct ToolHandle {
37 #[builder(default = Uuid::now_v7())]
39 pub uuid: Uuid,
40 #[builder(default = Utc::now())]
42 pub started_at: DateTime<Utc>,
43 #[builder(setter(into))]
45 pub name: String,
46 #[builder(default)]
48 pub data: Option<Json>,
49 #[builder(default)]
51 pub metadata: Option<Json>,
52 #[builder(default = ToolAttributes::empty())]
54 pub attributes: ToolAttributes,
55 #[builder(default)]
57 pub parent_uuid: Option<Uuid>,
58 #[builder(default, setter(into))]
60 pub tool_call_id: Option<String>,
61}
62
63#[derive(Debug, Clone, TypedBuilder)]
65#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
66pub struct CreateToolHandleParams<'a> {
67 pub name: &'a str,
69 #[builder(default)]
71 pub parent_uuid: Option<uuid::Uuid>,
72 #[builder(default = ToolAttributes::empty())]
74 pub attributes: ToolAttributes,
75 #[builder(default)]
77 pub data: Option<Json>,
78 #[builder(default)]
80 pub metadata: Option<Json>,
81 #[builder(default, setter(into))]
83 pub tool_call_id: Option<String>,
84 #[builder(default)]
87 pub timestamp: Option<DateTime<Utc>>,
88}
89
90#[derive(Debug, Clone, TypedBuilder)]
92#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
93pub struct EndToolHandleParams<'a> {
94 pub handle: &'a ToolHandle,
96 #[builder(default)]
98 pub data: Option<Json>,
99 #[builder(default)]
101 pub metadata: Option<Json>,
102 #[builder(default)]
106 pub timestamp: Option<DateTime<Utc>>,
107}
108
109#[derive(TypedBuilder)]
111#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
112pub struct ToolCallParams<'a> {
113 pub name: &'a str,
115 pub args: Json,
117 #[builder(default)]
119 pub parent: Option<&'a ScopeHandle>,
120 #[builder(default = ToolAttributes::empty())]
122 pub attributes: ToolAttributes,
123 #[builder(default)]
126 pub data: Option<Json>,
127 #[builder(default)]
129 pub metadata: Option<Json>,
130 #[builder(default, setter(into))]
132 pub tool_call_id: Option<String>,
133 #[builder(default)]
136 pub timestamp: Option<DateTime<Utc>>,
137}
138
139#[derive(TypedBuilder)]
141#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
142pub struct ToolCallExecuteParams {
143 #[builder(setter(into))]
145 pub name: String,
146 pub args: Json,
148 pub func: ToolExecutionNextFn,
150 #[builder(default)]
152 pub parent: Option<ScopeHandle>,
153 #[builder(default = ToolAttributes::empty())]
155 pub attributes: ToolAttributes,
156 #[builder(default)]
159 pub data: Option<Json>,
160 #[builder(default)]
162 pub metadata: Option<Json>,
163}
164
165#[derive(TypedBuilder)]
167#[builder(field_defaults(setter(strip_option(ignore_invalid, fallback_suffix = "_opt"))))]
168pub struct ToolCallEndParams<'a> {
169 pub handle: &'a ToolHandle,
171 pub result: Json,
173 #[builder(default)]
176 pub data: Option<Json>,
177 #[builder(default)]
179 pub metadata: Option<Json>,
180 #[builder(default)]
184 pub timestamp: Option<DateTime<Utc>>,
185}
186
187pub fn tool_call(params: ToolCallParams<'_>) -> Result<ToolHandle> {
215 ensure_runtime_owner()?;
216 let parent_uuid = resolve_parent_uuid(params.parent);
217 let (handle, event, subscribers) = {
218 let scope_stack = current_scope_stack();
219 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
220 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
221 ®istries.tool_sanitize_request_guardrails
222 });
223 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
224 let subscribers = snapshot_event_subscribers(scope_subscribers)?;
225 let context = global_context();
226 let state = context
227 .read()
228 .map_err(|error| FlowError::Internal(error.to_string()))?;
229
230 let sanitized_args =
231 state.tool_sanitize_request_chain(params.name, params.args, &scope_locals);
232 let handle_params = CreateToolHandleParams::builder()
233 .name(params.name)
234 .parent_uuid_opt(parent_uuid)
235 .attributes(params.attributes)
236 .data_opt(params.data)
237 .metadata_opt(params.metadata)
238 .tool_call_id_opt(params.tool_call_id)
239 .timestamp_opt(params.timestamp)
240 .build();
241 let handle = state.create_tool_handle(handle_params);
242 let event = state.build_tool_start_event(&handle, Some(sanitized_args));
243 (handle, event, subscribers)
244 };
245 NemoRelayContextState::emit_event(&event, &subscribers);
246 Ok(handle)
247}
248
249pub fn tool_call_end(params: ToolCallEndParams<'_>) -> Result<()> {
276 ensure_runtime_owner()?;
277 let (event, subscribers) = {
278 let scope_stack = current_scope_stack();
279 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
280 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
281 ®istries.tool_sanitize_response_guardrails
282 });
283 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
284 let subscribers = snapshot_event_subscribers(scope_subscribers)?;
285 let context = global_context();
286 let state = context
287 .read()
288 .map_err(|error| FlowError::Internal(error.to_string()))?;
289
290 let sanitized_result =
291 state.tool_sanitize_response_chain(¶ms.handle.name, params.result, &scope_locals);
292 let data = if sanitized_result.is_null() {
293 params.data
294 } else {
295 Some(sanitized_result)
296 };
297 let event = state.build_tool_end_event(
298 EndToolHandleParams::builder()
299 .handle(params.handle)
300 .data_opt(data)
301 .metadata_opt(params.metadata)
302 .timestamp_opt(params.timestamp)
303 .build(),
304 );
305 (event, subscribers)
306 };
307 NemoRelayContextState::emit_event(&event, &subscribers);
308 Ok(())
309}
310
311fn emit_tool_end_without_output(handle: &ToolHandle, metadata: Option<Json>) -> Result<()> {
312 ensure_runtime_owner()?;
313 let (event, subscribers) = {
314 let scope_stack = current_scope_stack();
315 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
316 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
317 let subscribers = snapshot_event_subscribers(scope_subscribers)?;
318 let context = global_context();
319 let state = context
320 .read()
321 .map_err(|error| FlowError::Internal(error.to_string()))?;
322 let event = state.end_tool_handle(handle, handle.data.clone(), metadata);
323 (event, subscribers)
324 };
325 NemoRelayContextState::emit_event(&event, &subscribers);
326 Ok(())
327}
328
329pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result<Json> {
358 let ToolCallExecuteParams {
359 name,
360 args,
361 func,
362 parent,
363 attributes,
364 data,
365 metadata,
366 } = params;
367 ensure_runtime_owner()?;
368 {
369 let (entries, subscribers, parent_uuid, guardrail_metadata) = {
370 let scope_stack = current_scope_stack();
371 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
372 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
373 ®istries.tool_conditional_execution_guardrails
374 });
375 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
376 let context = global_context();
377 let state = context
378 .read()
379 .map_err(|error| FlowError::Internal(error.to_string()))?;
380 let entries = state.tool_conditional_execution_entries(&scope_locals);
381 let subscribers = state.collect_event_subscribers(&scope_subscribers);
382 (
383 entries,
384 subscribers,
385 resolve_parent_uuid(parent.as_ref()),
386 metadata.clone(),
387 )
388 };
389 if let Some(error) = NemoRelayContextState::tool_conditional_execution_snapshot_chain(
390 &name,
391 &args,
392 &entries,
393 &subscribers,
394 parent_uuid,
395 guardrail_metadata,
396 )? {
397 let mut rejection_data = json!({});
398 if let Some(object) = rejection_data.as_object_mut() {
399 object.insert("rejected".into(), json!(true));
400 object.insert("rejection_reason".into(), json!(&error));
401 }
402 let _ = event(
403 EmitMarkEventParams::builder()
404 .name(&name)
405 .parent_opt(parent.as_ref())
406 .data(rejection_data)
407 .metadata_opt(metadata.clone())
408 .build(),
409 );
410 return Err(FlowError::GuardrailRejected(error));
411 }
412 }
413
414 let intercepted_args = {
415 let scope_stack = current_scope_stack();
416 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
417 let scope_locals = scope_guard
418 .collect_scope_local_registries(|registries| ®istries.tool_request_intercepts);
419 let context = global_context();
420 let state = context
421 .read()
422 .map_err(|error| FlowError::Internal(error.to_string()))?;
423 state.tool_request_intercepts_chain(&name, args, &scope_locals)?
424 };
425
426 let handle = tool_call(
427 ToolCallParams::builder()
428 .name(name.as_str())
429 .args(intercepted_args.clone())
430 .parent_opt(parent.as_ref())
431 .attributes(attributes)
432 .data_opt(data.clone())
433 .metadata_opt(metadata.clone())
434 .build(),
435 )?;
436
437 let execution = {
438 let scope_stack = current_scope_stack();
439 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
440 let scope_locals = scope_guard
441 .collect_scope_local_registries(|registries| ®istries.tool_execution_intercepts);
442 let context = global_context();
443 let state = context
444 .read()
445 .map_err(|error| FlowError::Internal(error.to_string()))?;
446 state.tool_build_execution_chain(&name, func, &scope_locals)
447 };
448
449 match execution(intercepted_args).await {
450 Ok(result) => {
451 let end_metadata = metadata_with_otel_status(metadata, "OK", None);
452 tool_call_end(
453 ToolCallEndParams::builder()
454 .handle(&handle)
455 .result(result.clone())
456 .data_opt(data)
457 .metadata_opt(end_metadata)
458 .build(),
459 )?;
460 Ok(result)
461 }
462 Err(error) => {
463 let end_metadata =
464 metadata_with_otel_status(metadata, "ERROR", Some(error.to_string()));
465 let _ = emit_tool_end_without_output(&handle, end_metadata);
466 Err(error)
467 }
468 }
469}
470
471pub fn tool_request_intercepts(name: &str, args: Json) -> Result<Json> {
489 ensure_runtime_owner()?;
490 let scope_stack = current_scope_stack();
491 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
492 let scope_locals = scope_guard
493 .collect_scope_local_registries(|registries| ®istries.tool_request_intercepts);
494 let context = global_context();
495 let state = context
496 .read()
497 .map_err(|error| FlowError::Internal(error.to_string()))?;
498 state.tool_request_intercepts_chain(name, args, &scope_locals)
499}
500
501pub fn tool_conditional_execution(name: &str, args: &Json) -> Result<()> {
523 ensure_runtime_owner()?;
524 let (entries, subscribers, parent_uuid) = {
525 let scope_stack = current_scope_stack();
526 let scope_guard = scope_stack.read().expect("scope stack lock poisoned");
527 let scope_locals = scope_guard.collect_scope_local_registries(|registries| {
528 ®istries.tool_conditional_execution_guardrails
529 });
530 let scope_subscribers = scope_guard.collect_scope_local_subscribers();
531 let context = global_context();
532 let state = context
533 .read()
534 .map_err(|error| FlowError::Internal(error.to_string()))?;
535 let entries = state.tool_conditional_execution_entries(&scope_locals);
536 let subscribers = state.collect_event_subscribers(&scope_subscribers);
537 (entries, subscribers, resolve_parent_uuid(None))
538 };
539 if let Some(error) = NemoRelayContextState::tool_conditional_execution_snapshot_chain(
540 name,
541 args,
542 &entries,
543 &subscribers,
544 parent_uuid,
545 None,
546 )? {
547 return Err(FlowError::GuardrailRejected(error));
548 }
549 Ok(())
550}
551
552#[cfg(test)]
553#[path = "../../tests/unit/tool_api_tests.rs"]
554mod tests;