Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions crates/core/src/api/tool.rs
Original file line number Diff line number Diff line change
Expand Up @@ -191,6 +191,9 @@ pub struct ToolCallExecuteParams {
/// Optional JSON metadata recorded on emitted events.
#[builder(default)]
pub metadata: Option<Json>,
/// Optional provider-specific correlation identifier.
#[builder(default, setter(into))]
pub tool_call_id: Option<String>,
}

/// Builder parameters for [`tool_call_end`].
Expand Down Expand Up @@ -721,6 +724,8 @@ impl Drop for ManagedToolCompletion {
/// - `data`: Optional application payload stored on the managed tool handle.
/// It may be used on failure end events that have no output payload.
/// - `metadata`: Optional JSON metadata recorded on emitted events.
/// - `tool_call_id`: Optional provider-specific correlation identifier recorded
/// on both managed lifecycle events.
///
/// # Returns
/// A [`Result`] containing the application result and optional opaque
Expand All @@ -743,6 +748,7 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result<ToolExec
attributes,
data,
metadata,
tool_call_id,
} = params;
ensure_runtime_owner()?;
{
Expand Down Expand Up @@ -819,6 +825,7 @@ pub async fn tool_call_execute(params: ToolCallExecuteParams) -> Result<ToolExec
.attributes(attributes)
.data_opt(data.clone())
.metadata_opt(metadata.clone())
.tool_call_id_opt(tool_call_id)
.build(),
)
.await?;
Expand Down
244 changes: 241 additions & 3 deletions crates/core/tests/integration/middleware_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -171,6 +171,25 @@ fn captured_events_snapshot(events: &Arc<Mutex<Vec<Event>>>) -> Vec<Event> {
events.lock().unwrap().clone()
}

fn assert_tool_lifecycle_correlation(
events: &[Event],
name: &str,
expected_tool_call_id: Option<&str>,
) {
let lifecycle = events
.iter()
.filter(|event| event.name() == name && event.scope_type() == Some(ScopeType::Tool))
.collect::<Vec<_>>();
assert_eq!(lifecycle.len(), 2, "expected one managed Start/End pair");
let start = lifecycle[0];
let end = lifecycle[1];
assert_eq!(start.scope_category(), Some(ScopeCategory::Start));
assert_eq!(end.scope_category(), Some(ScopeCategory::End));
assert_eq!(start.uuid(), end.uuid());
assert_eq!(start.tool_call_id(), expected_tool_call_id);
assert_eq!(end.tool_call_id(), expected_tool_call_id);
}

fn assert_middleware_callback_locks_are_free() {
let scope_stack = current_scope_stack();
assert!(
Expand Down Expand Up @@ -523,6 +542,87 @@ async fn test_no_break_chain_runs_all_intercepts() {
// Execution Intercepts (Middleware Chain) Tests
// =========================================================================

#[tokio::test]
async fn managed_tool_call_id_correlates_success_and_absent_lifecycles() {
let _lock = TEST_MUTEX.lock().unwrap();
reset_global();
setup_isolated_thread();

let events = Arc::new(Mutex::new(Vec::<Event>::new()));
let captured = events.clone();
register_subscriber(
"managed_tool_call_id_success_observer",
Arc::new(move |event| captured.lock().unwrap().push(event.clone())),
)
.unwrap();

let result = tool_call_execute(
nemo_relay::api::tool::ToolCallExecuteParams::builder()
.name("managed-correlated-tool")
.args(json!({"value": 42}))
.func(Arc::new(|args| Box::pin(async move { Ok(args.into()) })))
.tool_call_id("call-managed-42")
.build(),
)
.await
.unwrap();
assert_eq!(result.result, json!({"value": 42}));

tool_call_execute(
nemo_relay::api::tool::ToolCallExecuteParams::builder()
.name("managed-uncorrelated-tool")
.args(json!({}))
.func(Arc::new(|args| Box::pin(async move { Ok(args.into()) })))
.build(),
)
.await
.unwrap();

let captured = captured_events_snapshot(&events);
assert_tool_lifecycle_correlation(
&captured,
"managed-correlated-tool",
Some("call-managed-42"),
);
assert_tool_lifecycle_correlation(&captured, "managed-uncorrelated-tool", None);

deregister_subscriber("managed_tool_call_id_success_observer").unwrap();
}

#[tokio::test]
async fn managed_tool_call_id_correlates_callback_error_lifecycle() {
let _lock = TEST_MUTEX.lock().unwrap();
reset_global();
setup_isolated_thread();

let events = Arc::new(Mutex::new(Vec::<Event>::new()));
let captured = events.clone();
register_subscriber(
"managed_tool_call_id_error_observer",
Arc::new(move |event| captured.lock().unwrap().push(event.clone())),
)
.unwrap();

let error = tool_call_execute(
nemo_relay::api::tool::ToolCallExecuteParams::builder()
.name("managed-error-tool")
.args(json!({}))
.func(Arc::new(|_args| {
Box::pin(async { Err(FlowError::Internal("callback failed".into())) })
}))
.tool_call_id("call-managed-error")
.build(),
)
.await
.unwrap_err();
assert!(matches!(error, FlowError::Internal(message) if message == "callback failed"));

let captured = captured_events_snapshot(&events);
assert_tool_lifecycle_correlation(&captured, "managed-error-tool", Some("call-managed-error"));

deregister_subscriber("managed_tool_call_id_error_observer").unwrap();
}

/// Register an execution intercept that calls next().
/// Verify the original callable is invoked.
#[tokio::test]
Expand Down Expand Up @@ -1239,6 +1339,123 @@ async fn test_tool_execution_outcome_marks_follow_end_with_tool_parentage() {
deregister_subscriber("tool_outcome_mark_observer").unwrap();
}

#[tokio::test]
async fn managed_tool_call_id_projects_to_atif_and_otel() {
let _lock = TEST_MUTEX.lock().unwrap();
reset_global();
let agent = setup_isolated_scope("managed-tool-exporter-correlation");

let atif = AtifExporter::new(
"managed-tool-exporter-correlation".to_string(),
AtifAgentInfo {
name: "test-agent".to_string(),
version: "1.0.0".to_string(),
model_name: None,
tool_definitions: None,
extra: None,
},
);
register_subscriber("managed_tool_exporter_atif", atif.subscriber()).unwrap();

let otel_exporter = InMemorySpanExporterBuilder::new().build();
let otel_provider = SdkTracerProvider::builder()
.with_simple_exporter(otel_exporter.clone())
.build();
let otel =
OpenTelemetrySubscriber::from_tracer_provider(otel_provider, "managed-tool-exporter-otel");
register_subscriber("managed_tool_exporter_otel", otel.subscriber()).unwrap();

let tool_call_id = "call-managed-exporter-42";
llm_call_execute(
LlmCallExecuteParams::builder()
.name("model-call")
.request(LlmRequest {
headers: serde_json::Map::new(),
content: json!({
"model": "test-model",
"messages": [{"role": "user", "content": "weather in NYC"}],
}),
})
.func(Arc::new(move |_| {
Box::pin(async move {
Ok(json!({
"model": "test-model",
"choices": [{
"message": {
"role": "assistant",
"content": null,
"tool_calls": [{
"id": tool_call_id,
"type": "function",
"function": {
"name": "get_weather",
"arguments": "{\\\"city\\\":\\\"NYC\\\"}",
},
}],
},
}],
}))
})
}))
.build(),
)
.await
.unwrap();

tool_call_execute(
nemo_relay::api::tool::ToolCallExecuteParams::builder()
.name("get_weather")
.args(json!({"city": "NYC"}))
.tool_call_id(tool_call_id)
.func(Arc::new(|args| {
Box::pin(async move { Ok(ToolExecutionResult::new(args)) })
}))
.build(),
)
.await
.unwrap();

pop_scope(
nemo_relay::api::scope::PopScopeParams::builder()
.handle_uuid(&agent.uuid)
.build(),
)
.unwrap();
flush_subscribers().unwrap();

let trajectory = atif.export().unwrap();
let agent_step = trajectory
.steps
.iter()
.find(|step| step.source == "agent")
.expect("managed LLM event must create an agent step");
assert_eq!(
agent_step.tool_calls.as_ref().unwrap()[0].tool_call_id,
tool_call_id
);
assert_eq!(
agent_step.observation.as_ref().unwrap().results[0]
.source_call_id
.as_deref(),
Some(tool_call_id)
);

otel.force_flush().unwrap();
let tool_span = otel_exporter
.get_finished_spans()
.unwrap()
.into_iter()
.find(|span| span.name.as_ref() == "get_weather")
.expect("managed tool span must be exported");
assert!(tool_span.attributes.iter().any(|attribute| {
attribute.key.as_str() == "nemo_relay.tool_call_id"
&& attribute.value.to_string() == tool_call_id
}));

deregister_subscriber("managed_tool_exporter_atif").unwrap();
deregister_subscriber("managed_tool_exporter_otel").unwrap();
}

#[tokio::test]
async fn test_managed_tool_pending_marks_project_through_trace_exporters_only() {
let _lock = TEST_MUTEX.lock().unwrap();
Expand Down Expand Up @@ -2376,6 +2593,7 @@ async fn dropping_pending_tool_execution_closes_the_managed_lifecycle() {
.name("cancelled-tool")
.args(json!({}))
.func(Arc::new(|args| Box::pin(async move { Ok(args.into()) })))
.tool_call_id("call-managed-cancelled")
.build(),
));
tokio::select! {
Expand All @@ -2385,6 +2603,9 @@ async fn dropping_pending_tool_execution_closes_the_managed_lifecycle() {

assert_flush_waits_for_pending_completion(|| drop(execution));

let captured = captured_events_snapshot(&events);
assert_tool_lifecycle_correlation(&captured, "cancelled-tool", Some("call-managed-cancelled"));

let lifecycle = events
.lock()
.unwrap()
Expand All @@ -2393,8 +2614,8 @@ async fn dropping_pending_tool_execution_closes_the_managed_lifecycle() {
.filter_map(Event::scope_category)
.collect::<Vec<_>>();
assert_eq!(lifecycle, [ScopeCategory::Start, ScopeCategory::End]);
let events = events.lock().unwrap();
let cancelled_end = events
let event_log = events.lock().unwrap();
let cancelled_end = event_log
.iter()
.find(|event| {
event.name() == "cancelled-tool" && event.scope_category() == Some(ScopeCategory::End)
Expand All @@ -2406,7 +2627,7 @@ async fn dropping_pending_tool_execution_closes_the_managed_lifecycle() {
.as_ref()
.is_none_or(|value| value.is_null())
}));
drop(events);
drop(event_log);

deregister_tool_execution_intercept("pending_tool_execution").unwrap();
deregister_subscriber("cancelled_tool_lifecycle").unwrap();
Expand Down Expand Up @@ -2714,6 +2935,14 @@ async fn test_conditional_guardrail_rejects() {
reset_global();
setup_isolated_thread();

let events = Arc::new(Mutex::new(Vec::<Event>::new()));
let captured = events.clone();
register_subscriber(
"managed_tool_call_id_rejection_observer",
Arc::new(move |event| captured.lock().unwrap().push(event.clone())),
)
.unwrap();

register_tool_conditional_execution_guardrail(
"rejector",
1,
Expand All @@ -2728,6 +2957,7 @@ async fn test_conditional_guardrail_rejects() {
.name("tool")
.args(json!({}))
.func(func)
.tool_call_id("call-managed-rejected")
.build(),
)
.await;
Expand All @@ -2740,8 +2970,16 @@ async fn test_conditional_guardrail_rejects() {
other => panic!("Expected GuardrailRejected, got: {:?}", other),
}

let captured = captured_events_snapshot(&events);
assert!(
captured
.iter()
.all(|event| { event.name() != "tool" || event.scope_type() != Some(ScopeType::Tool) })
);

// Cleanup
deregister_tool_conditional_execution_guardrail("rejector").unwrap();
deregister_subscriber("managed_tool_call_id_rejection_observer").unwrap();
}

/// Register a conditional guardrail that allows (returns None). Execution proceeds.
Expand Down
Loading
Loading