1use crate::config::constants::tools;
19use crate::types::CompactStr;
20use hashbrown::HashMap;
21use std::sync::Arc;
22
23use async_trait::async_trait;
24
25use crate::tools::tool_intent;
26
27use super::tool_handler::{
28 ConfiguredToolSpec, ToolCallError, ToolHandler, ToolInvocation, ToolKind, ToolOutput, ToolPayload, ToolSession,
29 ToolSpec, TurnContext,
30};
31
32#[derive(Clone, Debug)]
34pub struct ToolCall {
35 pub tool_name: String,
37 pub call_id: String,
39 pub payload: ToolPayload,
41}
42
43struct DispatchEntry {
44 canonical_name: String,
45 handler: Arc<dyn ToolHandler>,
46}
47
48pub struct DispatchRegistry {
50 handlers: HashMap<CompactStr, DispatchEntry>,
51}
52
53fn normalize_router_tool_name(tool_name: &str) -> Option<String> {
54 let lowered = tool_name.trim().to_ascii_lowercase();
55 if lowered.is_empty() {
56 return None;
57 }
58
59 let normalized = lowered.replace([' ', '-'], "_").replace(['(', ')', '\'', '"'], "");
60
61 let mapped = match normalized.as_str() {
62 alias if tool_intent::is_command_session_tool(alias) => tools::UNIFIED_EXEC,
63 "exec_code" | "run_code" | "run_command" | "run_command_pty" => tools::UNIFIED_EXEC,
64 "search_text" | "search" | "find" => tools::GREP_FILE,
65 "applypatch" | "apply_patch" => tools::APPLY_PATCH,
66 "create_file" | "new_file" => tools::CREATE_FILE,
67 "delete_file" | "remove_file" => tools::DELETE_FILE,
68 "move_file" | "rename_file" => tools::MOVE_FILE,
69 "copy_file" | "duplicate_file" => tools::COPY_FILE,
70 "search_replace" | "find_replace" => tools::SEARCH_REPLACE,
71 "file_op" | "file_operation" => tools::FILE_OP,
72 tools::READ_FILE => tools::READ_FILE,
73 tools::WRITE_FILE => tools::WRITE_FILE,
74 tools::EDIT_FILE => tools::EDIT_FILE,
75 tools::LIST_FILES => tools::LIST_FILES,
76 _ => normalized.as_str(),
77 };
78
79 if mapped == lowered {
80 None
81 } else {
82 Some(mapped.to_string())
83 }
84}
85
86fn suggest_similar_tool_names(requested_tool_name: &str, handlers: &HashMap<CompactStr, DispatchEntry>) -> Vec<String> {
87 let requested_lower = requested_tool_name.to_ascii_lowercase();
88 let normalized = normalize_router_tool_name(requested_tool_name).unwrap_or_default();
89
90 let mut available: Vec<CompactStr> = handlers.keys().cloned().collect();
91 available.sort_unstable();
92
93 available
94 .into_iter()
95 .filter(|candidate| {
96 let c: &str = candidate;
97 c.contains(&requested_lower)
98 || requested_lower.contains(c)
99 || (!normalized.is_empty() && (c.contains(&*normalized) || normalized.contains(c)))
100 })
101 .take(3)
102 .map(|c| c.to_string())
103 .collect()
104}
105
106impl DispatchRegistry {
107 pub fn new(handlers: HashMap<String, Arc<dyn ToolHandler>>) -> Self {
108 let handlers: HashMap<CompactStr, DispatchEntry> = handlers
109 .into_iter()
110 .map(|(name, handler)| (CompactStr::from(name.clone()), DispatchEntry { canonical_name: name, handler }))
111 .collect();
112 Self { handlers }
113 }
114
115 pub fn handler(&self, name: &str) -> Option<Arc<dyn ToolHandler>> {
116 self.handlers.get(name).map(|entry| entry.handler.clone())
117 }
118
119 pub fn resolve_tool_name(&self, requested_name: &str) -> Result<&str, ToolCallError> {
120 self.resolve_entry(requested_name).map(|entry| entry.canonical_name.as_str())
121 }
122
123 pub async fn dispatch(&self, invocation: ToolInvocation) -> Result<ToolOutput, ToolCallError> {
125 let entry = self.resolve_entry(&invocation.tool_name)?;
126 let handler = &entry.handler;
127
128 if !handler.matches_kind(&invocation.payload) {
129 return Err(ToolCallError::respond(format!(
130 "Tool {} invoked with incompatible payload type",
131 invocation.tool_name
132 )));
133 }
134
135 handler.handle(invocation).await
136 }
137
138 fn resolve_entry(&self, requested_name: &str) -> Result<&DispatchEntry, ToolCallError> {
139 if let Some(entry) = self.handlers.get(requested_name) {
140 return Ok(entry);
141 }
142
143 let normalized_name = normalize_router_tool_name(requested_name);
144 normalized_name
145 .as_deref()
146 .and_then(|candidate| self.handlers.get(candidate))
147 .ok_or_else(|| {
148 let suggested = suggest_similar_tool_names(requested_name, &self.handlers);
149 let normalized_hint = normalized_name
150 .as_deref()
151 .filter(|candidate| *candidate != requested_name)
152 .map(|candidate| format!(" Normalized as '{candidate}'."))
153 .unwrap_or_default();
154 let suggestion_hint = if suggested.is_empty() {
155 String::new()
156 } else {
157 format!(" Did you mean: {}?", suggested.join(", "))
158 };
159 ToolCallError::respond(format!("Unknown tool: {requested_name}.{normalized_hint}{suggestion_hint}"))
160 })
161 }
162}
163
164pub struct DispatchRegistryBuilder {
166 handlers: HashMap<CompactStr, DispatchEntry>,
167 specs: Vec<ConfiguredToolSpec>,
168}
169
170impl Default for DispatchRegistryBuilder {
171 fn default() -> Self {
172 Self::new()
173 }
174}
175
176impl DispatchRegistryBuilder {
177 pub fn new() -> Self {
178 Self { handlers: HashMap::new(), specs: Vec::new() }
179 }
180
181 pub fn push_spec(&mut self, spec: ToolSpec) -> &mut Self {
183 self.push_spec_with_parallel_support(spec, false)
184 }
185
186 pub fn push_spec_with_parallel_support(&mut self, spec: ToolSpec, supports_parallel_tool_calls: bool) -> &mut Self {
188 self.specs.push(ConfiguredToolSpec::new(spec, supports_parallel_tool_calls));
189 self
190 }
191
192 pub fn register_handler(&mut self, name: impl Into<String>, handler: Arc<dyn ToolHandler>) -> &mut Self {
194 let name = name.into();
195 self.register_route(name.clone(), name, handler)
196 }
197
198 pub fn register_route(
200 &mut self,
201 name: impl Into<String>,
202 canonical_name: impl Into<String>,
203 handler: Arc<dyn ToolHandler>,
204 ) -> &mut Self {
205 let name = name.into();
206 let canonical_name = canonical_name.into();
207 let previous = self.handlers.insert(
208 CompactStr::from(&*name),
209 DispatchEntry {
210 canonical_name: canonical_name.clone(),
211 handler: Arc::new(RouteAliasHandler { canonical_name, inner: handler }),
212 },
213 );
214 if previous.is_some() {
215 tracing::warn!("Overwriting handler for tool");
216 }
217 self
218 }
219
220 pub fn register_aliases(&mut self, names: &[&str], handler: Arc<dyn ToolHandler>) -> &mut Self {
222 for name in names {
223 self.register_handler((*name).to_string(), handler.clone());
224 }
225 self
226 }
227
228 pub fn build(self) -> (Vec<ConfiguredToolSpec>, DispatchRegistry) {
230 let registry = DispatchRegistry { handlers: self.handlers };
231 (self.specs, registry)
232 }
233}
234
235pub struct ToolRouter {
242 registry: DispatchRegistry,
243 specs: Vec<ConfiguredToolSpec>,
244}
245
246impl ToolRouter {
247 pub fn from_builder(builder: DispatchRegistryBuilder) -> Self {
249 let (specs, registry) = builder.build();
250 Self { registry, specs }
251 }
252
253 pub fn specs(&self) -> Vec<ToolSpec> {
255 self.specs.iter().map(|c| c.spec.clone()).collect()
256 }
257
258 pub fn configured_specs(&self) -> &[ConfiguredToolSpec] {
260 &self.specs
261 }
262
263 pub fn tool_supports_parallel(&self, tool_name: &str) -> bool {
265 self.specs
266 .iter()
267 .filter(|c| c.supports_parallel_tool_calls)
268 .any(|c| c.spec.name() == tool_name)
269 }
270
271 pub fn resolve_tool_name(&self, tool_name: &str) -> Result<&str, ToolCallError> {
273 self.registry.resolve_tool_name(tool_name)
274 }
275
276 pub fn build_tool_call(
280 name: String,
281 call_id: String,
282 arguments: String,
283 mcp_prefix: Option<&str>,
284 ) -> Result<ToolCall, ToolCallError> {
285 if let Some(prefix) = mcp_prefix
287 && name.starts_with(prefix)
288 {
289 let parts: Vec<&str> = name.splitn(2, '/').collect();
290 if parts.len() == 2 {
291 return Ok(ToolCall {
292 tool_name: name.clone(),
293 call_id,
294 payload: ToolPayload::Mcp {
295 arguments: Some(serde_json::from_str(&arguments).unwrap_or_default()),
296 },
297 });
298 }
299 }
300
301 Ok(ToolCall {
303 tool_name: name,
304 call_id,
305 payload: ToolPayload::Function { arguments },
306 })
307 }
308
309 pub async fn dispatch_tool_call(
311 &self,
312 session: Arc<dyn ToolSession>,
313 turn: Arc<TurnContext>,
314 call: ToolCall,
315 ) -> Result<ToolOutput, ToolCallError> {
316 let invocation = ToolInvocation {
317 session,
318 turn,
319 tracker: None,
320 call_id: call.call_id,
321 tool_name: call.tool_name,
322 payload: call.payload,
323 };
324
325 self.registry.dispatch(invocation).await
326 }
327
328 #[cold]
330 pub fn failure_response(_call_id: String, error: ToolCallError) -> ToolOutput {
331 ToolOutput::error(error.to_string())
332 }
333}
334
335struct RouteAliasHandler {
336 canonical_name: String,
337 inner: Arc<dyn ToolHandler>,
338}
339
340#[async_trait]
341impl ToolHandler for RouteAliasHandler {
342 fn kind(&self) -> ToolKind {
343 self.inner.kind()
344 }
345
346 fn matches_kind(&self, payload: &ToolPayload) -> bool {
347 self.inner.matches_kind(payload)
348 }
349
350 async fn is_mutating(&self, invocation: &ToolInvocation) -> bool {
351 self.inner.is_mutating(invocation).await
352 }
353
354 async fn handle(&self, mut invocation: ToolInvocation) -> Result<ToolOutput, ToolCallError> {
355 invocation.tool_name = self.canonical_name.clone();
356 self.inner.handle(invocation).await
357 }
358}
359
360#[async_trait]
362pub trait ToolRouterProvider: Send + Sync {
363 async fn get_tool_router(&self) -> Arc<ToolRouter>;
365}
366
367#[cfg(test)]
368mod tests {
369 use super::super::tool_handler::{ResponsesApiTool, ToolKind};
370 use super::*;
371 use serde_json::json;
372
373 struct MockHandler;
374
375 #[async_trait]
376 impl ToolHandler for MockHandler {
377 fn kind(&self) -> ToolKind {
378 ToolKind::Function
379 }
380
381 async fn handle(&self, invocation: ToolInvocation) -> Result<ToolOutput, ToolCallError> {
382 Ok(ToolOutput::simple(format!("Handled: {}", invocation.tool_name)))
383 }
384 }
385
386 #[test]
387 fn test_build_tool_call_function() {
388 let call = ToolRouter::build_tool_call(
389 "test_tool".to_string(),
390 "call-1".to_string(),
391 r#"{"arg": "value"}"#.to_string(),
392 None,
393 )
394 .unwrap();
395
396 assert_eq!(call.tool_name, "test_tool");
397 assert_eq!(call.call_id, "call-1");
398 assert!(matches!(call.payload, ToolPayload::Function { .. }));
399 }
400
401 #[test]
402 fn test_build_tool_call_mcp() {
403 let call = ToolRouter::build_tool_call(
404 "mcp_server/do_thing".to_string(),
405 "call-2".to_string(),
406 r#"{"arg": "value"}"#.to_string(),
407 Some("mcp_server"),
408 )
409 .unwrap();
410
411 assert_eq!(call.tool_name, "mcp_server/do_thing");
412 assert!(matches!(call.payload, ToolPayload::Mcp { arguments: Some(_) }));
413 }
414
415 #[test]
416 fn test_registry_builder() {
417 let handler = Arc::new(MockHandler);
418 let spec = ToolSpec::Function(ResponsesApiTool {
419 name: "test_tool".to_string(),
420 description: "A test tool".to_string(),
421 parameters: json!({"type": "object"}),
422 strict: false,
423 });
424
425 let mut builder = DispatchRegistryBuilder::new();
426 builder
427 .push_spec_with_parallel_support(spec, true)
428 .register_handler("test_tool", handler);
429
430 let (specs, registry) = builder.build();
431
432 assert_eq!(specs.len(), 1);
433 assert!(specs[0].supports_parallel_tool_calls);
434 assert!(registry.handler("test_tool").is_some());
435 }
436
437 #[test]
438 fn test_router_parallel_support() {
439 let handler = Arc::new(MockHandler);
440 let spec = ToolSpec::Function(ResponsesApiTool {
441 name: "parallel_tool".to_string(),
442 description: "Supports parallel".to_string(),
443 parameters: json!({"type": "object"}),
444 strict: false,
445 });
446
447 let mut builder = DispatchRegistryBuilder::new();
448 builder
449 .push_spec_with_parallel_support(spec, true)
450 .register_handler("parallel_tool", handler);
451
452 let router = ToolRouter::from_builder(builder);
453
454 assert!(router.tool_supports_parallel("parallel_tool"));
455 assert!(!router.tool_supports_parallel("nonexistent"));
456 }
457
458 #[test]
459 fn test_normalize_router_tool_name_exec_code_label() {
460 assert_eq!(normalize_router_tool_name("Exec code").as_deref(), Some(tools::UNIFIED_EXEC));
461 assert_eq!(normalize_router_tool_name("run command (PTY)").as_deref(), Some(tools::UNIFIED_EXEC));
462 assert_eq!(normalize_router_tool_name("bash").as_deref(), Some(tools::UNIFIED_EXEC));
463 assert_eq!(normalize_router_tool_name("container.exec").as_deref(), Some(tools::UNIFIED_EXEC));
464 }
465
466 #[test]
467 fn test_normalize_router_tool_name_apply_patch_variants() {
468 assert_eq!(normalize_router_tool_name("apply_patch").as_deref(), None);
470 assert_eq!(normalize_router_tool_name("Apply_Patch").as_deref(), None);
472 assert_eq!(normalize_router_tool_name("apply-patch").as_deref(), Some("apply_patch"));
474 assert_eq!(normalize_router_tool_name("applypatch").as_deref(), Some("apply_patch"));
476 }
477
478 #[test]
479 fn test_normalize_router_tool_name_file_ops() {
480 assert_eq!(normalize_router_tool_name("Create File").as_deref(), Some("create_file"));
482 assert_eq!(normalize_router_tool_name("delete-file").as_deref(), Some("delete_file"));
484 assert_eq!(normalize_router_tool_name("move file").as_deref(), Some("move_file"));
486 assert_eq!(normalize_router_tool_name("copy_file").as_deref(), None);
488 }
489
490 #[test]
491 fn test_normalize_router_tool_name_search_variants() {
492 assert_eq!(normalize_router_tool_name("search_text").as_deref(), Some("grep_file"));
494 assert_eq!(normalize_router_tool_name("Search").as_deref(), Some("grep_file"));
496 assert_eq!(normalize_router_tool_name("find").as_deref(), Some("grep_file"));
498 }
499
500 #[test]
501 fn test_suggest_similar_tool_names_uses_normalized_form() {
502 let mut handlers = HashMap::new();
503 handlers.insert(
504 CompactStr::from(tools::UNIFIED_EXEC),
505 DispatchEntry {
506 canonical_name: tools::UNIFIED_EXEC.to_string(),
507 handler: Arc::new(MockHandler) as Arc<dyn ToolHandler>,
508 },
509 );
510
511 let suggestions = suggest_similar_tool_names("Exec code", &handlers);
512 assert_eq!(suggestions, vec![tools::UNIFIED_EXEC]);
513 }
514}