1use parking_lot::RwLock;
2use std::collections::HashMap;
3use std::sync::Arc;
4use std::sync::atomic::{AtomicU64, Ordering};
5
6use ai_agents_core::{LLMProvider, Tool, ToolInfo, ToolSafetyMetadata};
7
8use super::ToolError;
9use super::provider::{ProviderHealth, ToolProvider, ToolProviderError};
10use super::types::{
11 CommandRunner, CommandRunnerSlot, DiagnosticsProvider, DiagnosticsProviderSlot,
12 FileVersionStore, QuestionHandler, QuestionHandlerSlot, TodoItem, TodoStore, ToolAliases,
13 UnavailableCommandRunner, UnavailableDiagnosticsProvider, UnavailableWebSearchProvider,
14 WebSearchProvider, WebSearchProviderSlot,
15};
16
17#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
19#[serde(rename_all = "snake_case")]
20pub enum ToolSchemaPromptMode {
21 #[default]
23 Full,
24 Compact,
26}
27
28#[derive(Debug, Clone)]
30pub struct ToolIdentity {
31 pub requested_name: String,
33 pub canonical_id: String,
35 pub display_name: String,
37 pub provider_id: Option<String>,
39}
40
41#[derive(Clone)]
43pub struct ResolvedTool {
44 pub identity: ToolIdentity,
46 pub tool: Arc<dyn Tool>,
48}
49
50#[derive(Clone)]
51enum ToolRef {
52 Builtin(Arc<dyn Tool>),
53 Provider {
54 provider_id: String,
55 tool: Arc<dyn Tool>,
56 },
57}
58
59pub struct ToolRegistry {
61 builtin_tools: RwLock<HashMap<String, Arc<dyn Tool>>>,
62
63 providers: RwLock<HashMap<String, Arc<dyn ToolProvider>>>,
64
65 tool_index: RwLock<HashMap<String, ToolRef>>,
66
67 alias_index: RwLock<HashMap<String, String>>,
68
69 display_name_index: RwLock<HashMap<String, String>>,
70
71 builtin_aliases: RwLock<HashMap<String, ToolAliases>>,
72
73 question_handler: QuestionHandlerSlot,
74
75 diagnostics_provider: DiagnosticsProviderSlot,
76
77 command_runner: CommandRunnerSlot,
78
79 todo_store: TodoStore,
80
81 file_versions: FileVersionStore,
82
83 web_fetch_extractor: Arc<RwLock<Option<Arc<dyn LLMProvider>>>>,
84
85 web_search_provider: WebSearchProviderSlot,
86
87 registry_version: AtomicU64,
88}
89
90impl ToolRegistry {
91 pub fn new() -> Self {
93 Self {
94 builtin_tools: RwLock::new(HashMap::new()),
95 providers: RwLock::new(HashMap::new()),
96 tool_index: RwLock::new(HashMap::new()),
97 alias_index: RwLock::new(HashMap::new()),
98 display_name_index: RwLock::new(HashMap::new()),
99 builtin_aliases: RwLock::new(HashMap::new()),
100 question_handler: Arc::new(RwLock::new(None)),
101 diagnostics_provider: Arc::new(RwLock::new(Arc::new(UnavailableDiagnosticsProvider))),
102 command_runner: Arc::new(RwLock::new(Arc::new(UnavailableCommandRunner))),
103 todo_store: TodoStore::default(),
104 file_versions: FileVersionStore::default(),
105 web_fetch_extractor: Arc::new(RwLock::new(None)),
106 web_search_provider: Arc::new(RwLock::new(Arc::new(UnavailableWebSearchProvider))),
107 registry_version: AtomicU64::new(1),
108 }
109 }
110
111 fn bump_version(&self) {
112 self.registry_version.fetch_add(1, Ordering::SeqCst);
113 }
114
115 pub fn version(&self) -> u64 {
117 self.registry_version.load(Ordering::SeqCst)
118 }
119
120 fn normalize_key(value: &str) -> String {
121 value.trim().to_lowercase()
122 }
123
124 fn insert_unique_index(index: &mut HashMap<String, String>, key: String, tool_id: &str) {
125 match index.get(&key) {
126 None => {
127 index.insert(key, tool_id.to_string());
128 }
129 Some(existing) if existing == tool_id => {}
130 Some(_) => {
131 index.remove(&key);
132 }
133 }
134 }
135
136 pub fn register(&mut self, tool: Arc<dyn Tool>) -> Result<(), ToolError> {
138 let id = tool.id().to_string();
139
140 let mut builtin_tools = self.builtin_tools.write();
141 let mut tool_index = self.tool_index.write();
142 let mut display_name_index = self.display_name_index.write();
143
144 if builtin_tools.contains_key(&id) || tool_index.contains_key(&id) {
145 return Err(ToolError::Duplicate(id));
146 }
147
148 Self::insert_unique_index(
149 &mut display_name_index,
150 Self::normalize_key(tool.name()),
151 &id,
152 );
153 tool_index.insert(id.clone(), ToolRef::Builtin(tool.clone()));
154 builtin_tools.insert(id, tool);
155 self.bump_version();
156 Ok(())
157 }
158
159 pub fn get(&self, id_or_alias: &str) -> Option<Arc<dyn Tool>> {
160 self.resolve(id_or_alias).map(|resolved| resolved.tool)
161 }
162
163 pub fn canonical_id(&self, id_or_alias: &str) -> Option<String> {
165 self.resolve(id_or_alias)
166 .map(|resolved| resolved.identity.canonical_id)
167 }
168
169 pub fn safety_metadata(&self, id_or_alias: &str) -> Option<ToolSafetyMetadata> {
171 self.resolve(id_or_alias)
172 .map(|resolved| resolved.tool.safety_metadata())
173 }
174
175 pub fn question_handler_slot(&self) -> QuestionHandlerSlot {
177 Arc::clone(&self.question_handler)
178 }
179
180 pub fn set_question_handler(&self, handler: Option<Arc<dyn QuestionHandler>>) {
182 *self.question_handler.write() = handler;
183 }
184
185 pub fn diagnostics_provider_slot(&self) -> DiagnosticsProviderSlot {
187 Arc::clone(&self.diagnostics_provider)
188 }
189
190 pub fn set_diagnostics_provider(&self, provider: Arc<dyn DiagnosticsProvider>) {
192 *self.diagnostics_provider.write() = provider;
193 }
194
195 pub fn diagnostics_available(&self) -> bool {
197 self.diagnostics_provider.read().is_available()
198 }
199
200 pub fn command_runner_slot(&self) -> CommandRunnerSlot {
202 Arc::clone(&self.command_runner)
203 }
204
205 pub fn set_command_runner(&self, runner: Arc<dyn CommandRunner>) {
207 *self.command_runner.write() = runner;
208 }
209
210 pub fn command_runner_available(&self) -> bool {
212 self.command_runner.read().is_available()
213 }
214
215 pub fn file_version_store(&self) -> FileVersionStore {
217 self.file_versions.clone()
218 }
219
220 pub fn todo_store(&self) -> TodoStore {
222 self.todo_store.clone()
223 }
224
225 pub fn todos(&self) -> Vec<TodoItem> {
227 self.todo_store.list()
228 }
229
230 pub fn web_fetch_extractor_slot(&self) -> Arc<RwLock<Option<Arc<dyn LLMProvider>>>> {
232 Arc::clone(&self.web_fetch_extractor)
233 }
234
235 pub fn set_web_fetch_extractor(&self, extractor: Option<Arc<dyn LLMProvider>>) {
237 *self.web_fetch_extractor.write() = extractor;
238 }
239
240 pub fn web_search_provider_slot(&self) -> WebSearchProviderSlot {
242 Arc::clone(&self.web_search_provider)
243 }
244
245 pub fn set_web_search_provider(&self, provider: Arc<dyn WebSearchProvider>) {
247 *self.web_search_provider.write() = provider;
248 }
249
250 pub fn web_search_available(&self) -> bool {
252 self.web_search_provider.read().is_available()
253 }
254
255 pub fn resolve(&self, id_or_alias: &str) -> Option<ResolvedTool> {
257 let tool_index = self.tool_index.read();
258 let requested_name = id_or_alias.to_string();
259 let normalized = Self::normalize_key(id_or_alias);
260
261 if let Some((canonical_id, tool_ref)) =
262 tool_index.get_key_value(id_or_alias).or_else(|| {
263 tool_index
264 .iter()
265 .find(|(id, _)| Self::normalize_key(id) == normalized)
266 })
267 {
268 return self.resolved_tool_from_ref(&requested_name, canonical_id, tool_ref);
269 }
270
271 if let Some(tool_id) = self.display_name_index.read().get(&normalized).cloned()
272 && let Some(tool_ref) = tool_index.get(&tool_id)
273 {
274 return self.resolved_tool_from_ref(&requested_name, &tool_id, tool_ref);
275 }
276
277 let alias_index = self.alias_index.read();
278 if let Some(tool_id) = alias_index.get(&normalized).cloned().or_else(|| {
279 alias_index.iter().find_map(|(alias_key, tool_id)| {
280 alias_key
281 .ends_with(&format!(":{}", normalized))
282 .then(|| tool_id.clone())
283 })
284 }) && let Some(tool_ref) = tool_index.get(&tool_id)
285 {
286 return self.resolved_tool_from_ref(&requested_name, &tool_id, tool_ref);
287 }
288
289 None
290 }
291
292 fn resolved_tool_from_ref(
293 &self,
294 requested_name: &str,
295 canonical_id: &str,
296 tool_ref: &ToolRef,
297 ) -> Option<ResolvedTool> {
298 let tool = self.resolve_tool_ref(tool_ref)?;
299 let provider_id = match tool_ref {
300 ToolRef::Builtin(_) => None,
301 ToolRef::Provider { provider_id, .. } => Some(provider_id.clone()),
302 };
303 Some(ResolvedTool {
304 identity: ToolIdentity {
305 requested_name: requested_name.to_string(),
306 canonical_id: canonical_id.to_string(),
307 display_name: tool.name().to_string(),
308 provider_id,
309 },
310 tool,
311 })
312 }
313
314 fn resolve_tool_ref(&self, tool_ref: &ToolRef) -> Option<Arc<dyn Tool>> {
315 match tool_ref {
316 ToolRef::Builtin(tool) => Some(tool.clone()),
317 ToolRef::Provider { tool, .. } => Some(tool.clone()),
318 }
319 }
320
321 pub fn list_ids(&self) -> Vec<String> {
322 self.tool_index.read().keys().cloned().collect()
323 }
324
325 pub fn list_infos(&self) -> Vec<ToolInfo> {
326 let tool_index = self.tool_index.read();
327 let mut infos = Vec::with_capacity(tool_index.len());
328
329 for tool_ref in tool_index.values() {
330 if let Some(tool) = self.resolve_tool_ref(tool_ref) {
331 infos.push(tool.info());
332 }
333 }
334
335 infos
336 }
337
338 pub fn len(&self) -> usize {
339 self.tool_index.read().len()
340 }
341
342 pub fn is_empty(&self) -> bool {
343 self.tool_index.read().is_empty()
344 }
345
346 pub fn map_tools<F>(&self, mut f: F) -> ToolRegistry
347 where
348 F: FnMut(Arc<dyn Tool>) -> Arc<dyn Tool>,
349 {
350 let mut mapped = ToolRegistry::new();
351 mapped.question_handler = Arc::clone(&self.question_handler);
352 mapped.diagnostics_provider = Arc::clone(&self.diagnostics_provider);
353 mapped.command_runner = Arc::clone(&self.command_runner);
354 mapped.todo_store = self.todo_store.clone();
355 mapped.file_versions = self.file_versions.clone();
356 mapped.web_fetch_extractor = Arc::clone(&self.web_fetch_extractor);
357 mapped.web_search_provider = Arc::clone(&self.web_search_provider);
358
359 {
360 let providers = self.providers.read();
361 let mut mapped_providers = mapped.providers.write();
362 for (id, provider) in providers.iter() {
363 mapped_providers.insert(id.clone(), provider.clone());
364 }
365 }
366
367 {
368 let aliases = self.alias_index.read();
369 let mut mapped_aliases = mapped.alias_index.write();
370 for (alias, tool_id) in aliases.iter() {
371 mapped_aliases.insert(alias.clone(), tool_id.clone());
372 }
373 }
374
375 {
376 let display_names = self.display_name_index.read();
377 let mut mapped_display_names = mapped.display_name_index.write();
378 for (name, tool_id) in display_names.iter() {
379 mapped_display_names.insert(name.clone(), tool_id.clone());
380 }
381 }
382
383 {
384 let builtin_aliases = self.builtin_aliases.read();
385 let mut mapped_builtin_aliases = mapped.builtin_aliases.write();
386 for (id, aliases) in builtin_aliases.iter() {
387 mapped_builtin_aliases.insert(id.clone(), aliases.clone());
388 }
389 }
390
391 let tool_index = self.tool_index.read();
392 let mut mapped_tool_index = mapped.tool_index.write();
393 let mut mapped_builtin_tools = mapped.builtin_tools.write();
394 for (id, tool_ref) in tool_index.iter() {
395 match tool_ref {
396 ToolRef::Builtin(tool) => {
397 let wrapped = f(tool.clone());
398 mapped_tool_index.insert(id.clone(), ToolRef::Builtin(wrapped.clone()));
399 mapped_builtin_tools.insert(id.clone(), wrapped);
400 }
401 ToolRef::Provider { provider_id, tool } => {
402 let wrapped = f(tool.clone());
403 mapped_tool_index.insert(
404 id.clone(),
405 ToolRef::Provider {
406 provider_id: provider_id.clone(),
407 tool: wrapped,
408 },
409 );
410 }
411 }
412 }
413
414 drop(mapped_builtin_tools);
415 drop(mapped_tool_index);
416 mapped
417 .registry_version
418 .store(self.version(), Ordering::SeqCst);
419 mapped
420 }
421
422 pub async fn register_provider(
423 &self,
424 provider: Arc<dyn ToolProvider>,
425 ) -> Result<(), ToolError> {
426 let provider_id = provider.id().to_string();
427
428 {
429 let providers = self.providers.read();
430 if providers.contains_key(&provider_id) {
431 return Err(ToolError::Duplicate(format!("Provider: {}", provider_id)));
432 }
433 }
434
435 let tools = provider.list_tools().await;
436
437 let mut resolved_tools = Vec::with_capacity(tools.len());
439 for descriptor in &tools {
440 resolved_tools.push((descriptor, provider.get_tool(&descriptor.id).await));
441 }
442
443 {
444 let mut tool_index = self.tool_index.write();
445 let mut alias_index = self.alias_index.write();
446 let mut display_name_index = self.display_name_index.write();
447
448 for (descriptor, tool) in resolved_tools {
449 if tool_index.contains_key(&descriptor.id) {
450 return Err(ToolError::Duplicate(descriptor.id.clone()));
451 }
452
453 if let Some(tool) = tool {
454 Self::insert_unique_index(
455 &mut display_name_index,
456 Self::normalize_key(&descriptor.name),
457 &descriptor.id,
458 );
459 tool_index.insert(
460 descriptor.id.clone(),
461 ToolRef::Provider {
462 provider_id: provider_id.clone(),
463 tool,
464 },
465 );
466
467 if let Some(ref aliases) = descriptor.aliases {
468 for (lang, name) in &aliases.names {
469 let key = format!("{}:{}", lang, Self::normalize_key(name));
470 Self::insert_unique_index(&mut alias_index, key, &descriptor.id);
471 Self::insert_unique_index(
472 &mut alias_index,
473 Self::normalize_key(name),
474 &descriptor.id,
475 );
476 }
477 }
478 }
479 }
480 }
481
482 self.providers.write().insert(provider_id, provider);
483 self.bump_version();
484
485 Ok(())
486 }
487
488 pub fn unregister_provider(&self, provider_id: &str) -> bool {
489 let removed = self.providers.write().remove(provider_id);
490
491 if removed.is_some() {
492 let mut tool_index = self.tool_index.write();
493 let mut alias_index = self.alias_index.write();
494 let mut display_name_index = self.display_name_index.write();
495
496 let tools_to_remove: Vec<String> = tool_index
497 .iter()
498 .filter_map(|(id, tool_ref)| {
499 if let ToolRef::Provider {
500 provider_id: pid, ..
501 } = tool_ref
502 && pid == provider_id
503 {
504 return Some(id.clone());
505 }
506 None
507 })
508 .collect();
509
510 for tool_id in &tools_to_remove {
511 tool_index.remove(tool_id);
512 }
513
514 alias_index.retain(|_, tool_id| !tools_to_remove.contains(tool_id));
515 display_name_index.retain(|_, tool_id| !tools_to_remove.contains(tool_id));
516 self.bump_version();
517
518 true
519 } else {
520 false
521 }
522 }
523
524 pub fn set_tool_aliases(&self, tool_id: &str, aliases: ToolAliases) {
525 if !self.tool_index.read().contains_key(tool_id) {
526 return;
527 }
528
529 {
530 let mut alias_index = self.alias_index.write();
531 for (lang, name) in &aliases.names {
532 let key = format!("{}:{}", lang, Self::normalize_key(name));
533 Self::insert_unique_index(&mut alias_index, key, tool_id);
534 Self::insert_unique_index(&mut alias_index, Self::normalize_key(name), tool_id);
535 }
536 }
537
538 self.builtin_aliases
539 .write()
540 .insert(tool_id.to_string(), aliases);
541 self.bump_version();
542 }
543
544 pub fn get_by_alias(&self, alias: &str, lang: &str) -> Option<Arc<dyn Tool>> {
545 let key = format!("{}:{}", lang, Self::normalize_key(alias));
546 let alias_index = self.alias_index.read();
547
548 if let Some(tool_id) = alias_index.get(&key) {
549 return self.get(tool_id);
550 }
551
552 None
553 }
554
555 pub fn list_providers(&self) -> Vec<String> {
556 self.providers.read().keys().cloned().collect()
557 }
558
559 pub async fn provider_health(&self, provider_id: &str) -> Option<ProviderHealth> {
560 let provider = self.providers.read().get(provider_id).cloned();
561 if let Some(provider) = provider {
562 Some(provider.health_check().await)
563 } else {
564 None
565 }
566 }
567
568 pub async fn refresh_provider(&self, provider_id: &str) -> Result<(), ToolProviderError> {
569 let provider = {
570 let providers = self.providers.read();
571 providers.get(provider_id).cloned()
572 };
573
574 if let Some(provider) = provider {
575 if provider.supports_refresh() {
576 provider.refresh().await?;
577
578 let tools = provider.list_tools().await;
579
580 let mut resolved_tools = Vec::with_capacity(tools.len());
582 for descriptor in &tools {
583 if let Some(tool) = provider.get_tool(&descriptor.id).await {
584 resolved_tools.push((descriptor, tool));
585 }
586 }
587
588 {
589 let mut tool_index = self.tool_index.write();
590 let mut alias_index = self.alias_index.write();
591 let mut display_name_index = self.display_name_index.write();
592
593 let old_tools: Vec<String> = tool_index
594 .iter()
595 .filter_map(|(id, tool_ref)| {
596 if let ToolRef::Provider {
597 provider_id: pid, ..
598 } = tool_ref
599 && pid == provider_id
600 {
601 return Some(id.clone());
602 }
603 None
604 })
605 .collect();
606
607 for tool_id in &old_tools {
608 tool_index.remove(tool_id);
609 }
610 alias_index.retain(|_, tool_id| !old_tools.contains(tool_id));
611 display_name_index.retain(|_, tool_id| !old_tools.contains(tool_id));
612
613 for (descriptor, tool) in resolved_tools {
614 Self::insert_unique_index(
615 &mut display_name_index,
616 Self::normalize_key(&descriptor.name),
617 &descriptor.id,
618 );
619 tool_index.insert(
620 descriptor.id.clone(),
621 ToolRef::Provider {
622 provider_id: provider_id.to_string(),
623 tool,
624 },
625 );
626
627 if let Some(ref aliases) = descriptor.aliases {
628 for (lang, name) in &aliases.names {
629 let key = format!("{}:{}", lang, Self::normalize_key(name));
630 Self::insert_unique_index(&mut alias_index, key, &descriptor.id);
631 Self::insert_unique_index(
632 &mut alias_index,
633 Self::normalize_key(name),
634 &descriptor.id,
635 );
636 }
637 }
638 }
639 self.bump_version();
640 }
641 }
642 Ok(())
643 } else {
644 Err(ToolProviderError::ToolNotFound(format!(
645 "Provider not found: {}",
646 provider_id
647 )))
648 }
649 }
650
651 pub fn generate_tools_prompt(&self) -> String {
652 self.generate_tools_prompt_with_lang(None, false)
653 }
654
655 pub fn generate_tools_prompt_with_parallel(&self, parallel: bool) -> String {
656 self.generate_tools_prompt_with_lang(None, parallel)
657 }
658
659 pub fn generate_tools_prompt_with_lang(
660 &self,
661 language: Option<&str>,
662 parallel: bool,
663 ) -> String {
664 let tool_index = self.tool_index.read();
665 if tool_index.is_empty() {
666 return String::new();
667 }
668
669 let builtin_aliases = self.builtin_aliases.read();
670 let mut prompt = String::from("Available tools:\n");
671
672 for (id, tool_ref) in tool_index.iter() {
673 if let Some(tool) = self.resolve_tool_ref(tool_ref) {
674 let (name, description) = if let Some(lang) = language {
675 if let Some(aliases) = builtin_aliases.get(id) {
676 let name = aliases
677 .names
678 .get(lang)
679 .map(|s| s.as_str())
680 .unwrap_or_else(|| tool.name());
681 let desc = aliases
682 .descriptions
683 .get(lang)
684 .map(|s| s.as_str())
685 .unwrap_or_else(|| tool.description());
686 (name, desc)
687 } else {
688 (tool.name(), tool.description())
689 }
690 } else {
691 (tool.name(), tool.description())
692 };
693
694 let schema = tool.input_schema();
695 let args_desc = if let Some(props) = schema.get("properties") {
696 serde_json::to_string(props).unwrap_or_default()
697 } else {
698 "{}".to_string()
699 };
700
701 prompt.push_str(&format!(
702 "- {}: {}. Arguments: {}\n",
703 name, description, args_desc
704 ));
705 }
706 }
707
708 Self::append_tool_format_instructions(&mut prompt, parallel);
709
710 prompt
711 }
712
713 pub fn generate_filtered_prompt(&self, tool_ids: &[String]) -> String {
714 self.generate_filtered_prompt_with_lang(tool_ids, None, false)
715 }
716
717 pub fn generate_filtered_prompt_with_parallel(
718 &self,
719 tool_ids: &[String],
720 parallel: bool,
721 ) -> String {
722 self.generate_filtered_prompt_with_lang(tool_ids, None, parallel)
723 }
724
725 pub fn generate_scoped_prompt_with_parallel(
727 &self,
728 tool_ids: &[String],
729 parallel: bool,
730 ) -> String {
731 self.generate_scoped_prompt_with_lang(tool_ids, None, parallel)
732 }
733
734 pub fn generate_scoped_prompt_with_lang(
735 &self,
736 tool_ids: &[String],
737 language: Option<&str>,
738 parallel: bool,
739 ) -> String {
740 if tool_ids.is_empty() {
741 return String::new();
742 }
743 self.generate_filtered_prompt_inner(tool_ids, language, parallel)
744 }
745
746 pub fn generate_filtered_prompt_with_lang(
747 &self,
748 tool_ids: &[String],
749 language: Option<&str>,
750 parallel: bool,
751 ) -> String {
752 if tool_ids.is_empty() {
753 return self.generate_tools_prompt_with_lang(language, parallel);
754 }
755
756 self.generate_filtered_prompt_inner(tool_ids, language, parallel)
757 }
758
759 fn generate_filtered_prompt_inner(
760 &self,
761 tool_ids: &[String],
762 language: Option<&str>,
763 parallel: bool,
764 ) -> String {
765 let tool_index = self.tool_index.read();
766 let builtin_aliases = self.builtin_aliases.read();
767 let mut prompt = String::from("Available tools:\n");
768 let mut found_any = false;
769
770 for id in tool_ids {
771 if let Some(tool_ref) = tool_index.get(id)
772 && let Some(tool) = self.resolve_tool_ref(tool_ref)
773 {
774 found_any = true;
775
776 let (name, description) = if let Some(lang) = language {
777 if let Some(aliases) = builtin_aliases.get(id) {
778 let name = aliases
779 .names
780 .get(lang)
781 .map(|s| s.as_str())
782 .unwrap_or_else(|| tool.name());
783 let desc = aliases
784 .descriptions
785 .get(lang)
786 .map(|s| s.as_str())
787 .unwrap_or_else(|| tool.description());
788 (name, desc)
789 } else {
790 (tool.name(), tool.description())
791 }
792 } else {
793 (tool.name(), tool.description())
794 };
795
796 let schema = tool.input_schema();
797 let args_desc = if let Some(props) = schema.get("properties") {
798 serde_json::to_string(props).unwrap_or_default()
799 } else {
800 "{}".to_string()
801 };
802
803 prompt.push_str(&format!(
804 "- {}: {}. Arguments: {}\n",
805 name, description, args_desc
806 ));
807 }
808 }
809
810 if !found_any {
811 return String::new();
812 }
813
814 Self::append_tool_format_instructions(&mut prompt, parallel);
815
816 prompt
817 }
818
819 pub fn generate_scoped_prompt_with_mode(
821 &self,
822 tool_ids: &[impl AsRef<str>],
823 language: Option<&str>,
824 parallel: bool,
825 mode: ToolSchemaPromptMode,
826 ) -> String {
827 if tool_ids.is_empty() {
828 return String::new();
829 }
830 match mode {
831 ToolSchemaPromptMode::Full => self.generate_scoped_prompt_with_lang(
832 &tool_ids
833 .iter()
834 .map(|s| s.as_ref().to_string())
835 .collect::<Vec<_>>(),
836 language,
837 parallel,
838 ),
839 ToolSchemaPromptMode::Compact => {
840 self.generate_compact_prompt_inner(tool_ids, language, parallel)
841 }
842 }
843 }
844
845 fn generate_compact_prompt_inner(
847 &self,
848 tool_ids: &[impl AsRef<str>],
849 language: Option<&str>,
850 parallel: bool,
851 ) -> String {
852 let tool_index = self.tool_index.read();
853 let builtin_aliases = self.builtin_aliases.read();
854 let mut prompt = String::from("Available tools:\n");
855 let mut found_any = false;
856
857 for id in tool_ids {
858 let id = id.as_ref();
859 if let Some(tool_ref) = tool_index.get(id)
860 && let Some(tool) = self.resolve_tool_ref(tool_ref)
861 {
862 found_any = true;
863
864 let (name, description) = if let Some(lang) = language {
865 if let Some(aliases) = builtin_aliases.get(id) {
866 let n = aliases
867 .names
868 .get(lang)
869 .map(|s| s.as_str())
870 .unwrap_or_else(|| tool.name());
871 let d = aliases
872 .descriptions
873 .get(lang)
874 .map(|s| s.as_str())
875 .unwrap_or_else(|| tool.description());
876 (n, d)
877 } else {
878 (tool.name(), tool.description())
879 }
880 } else {
881 (tool.name(), tool.description())
882 };
883
884 let schema = tool.input_schema();
885 let compact = compact_schema_descriptor(&schema);
886 prompt.push_str(&format!("- {}: {}. {}\n", name, description, compact));
887 }
888 }
889
890 if !found_any {
891 return String::new();
892 }
893
894 Self::append_tool_format_instructions(&mut prompt, parallel);
895 prompt
896 }
897
898 fn append_tool_format_instructions(prompt: &mut String, parallel: bool) {
902 prompt.push_str(
903 "\nWhen you need to use a tool, respond ONLY with valid JSON in this exact format:\n",
904 );
905 prompt.push_str("{\"tool\": \"tool_name\", \"arguments\": {...}}\n");
906 prompt.push_str("The \"tool\" value MUST be one of the exact tool names listed above. Do not invent tool names.\n");
907 if parallel {
908 prompt.push_str(
909 "\nWhen you need to call multiple tools at once, respond with a JSON array:\n",
910 );
911 prompt.push_str(
912 "[{\"tool\": \"tool_name1\", \"arguments\": {...}}, {\"tool\": \"tool_name2\", \"arguments\": {...}}]\n",
913 );
914 }
915 prompt.push_str("\nWhen you receive a tool result, summarize it naturally for the user.\n");
916 prompt.push_str("If no tool is needed, respond normally.");
917 }
918}
919
920fn compact_schema_descriptor(schema: &serde_json::Value) -> String {
923 let props = schema.get("properties").and_then(|p| p.as_object());
924 let required = schema.get("required").and_then(|r| r.as_array());
925 let mut parts = Vec::new();
926
927 if let Some(req) = required {
928 let req_fields: Vec<String> = req
929 .iter()
930 .filter_map(|v| v.as_str().map(str::to_string))
931 .collect();
932 if !req_fields.is_empty() {
933 parts.push(format!("required: [{}]", req_fields.join(", ")));
934 }
935 }
936
937 if let Some(props) = props {
938 let mut prop_parts = Vec::new();
939 for (key, value) in props.iter() {
940 let prop_type = value.get("type").and_then(|t| t.as_str()).unwrap_or("?");
941 let prop_desc = value
942 .get("description")
943 .and_then(|d| d.as_str())
944 .unwrap_or("");
945 let short_desc = if prop_desc.len() > 40 {
946 format!("{}...", &prop_desc[..40])
947 } else if !prop_desc.is_empty() {
948 prop_desc.to_string()
949 } else {
950 String::new()
951 };
952 if short_desc.is_empty() {
953 prop_parts.push(format!("{}({})", key, prop_type));
954 } else {
955 prop_parts.push(format!("{}({}): {}", key, prop_type, short_desc));
956 }
957 }
958 if !prop_parts.is_empty() {
959 parts.push(format!("args: {}", prop_parts.join("; ")));
960 }
961 }
962
963 if parts.is_empty() {
964 "Args: none".to_string()
965 } else {
966 parts.join(". ")
967 }
968}
969
970impl Default for ToolRegistry {
971 fn default() -> Self {
972 Self::new()
973 }
974}
975
976#[cfg(test)]
977mod tests {
978 use super::*;
979 use crate::ToolResult;
980 use async_trait::async_trait;
981 use serde_json::Value;
982
983 struct TestTool {
984 id: String,
985 }
986
987 #[async_trait]
988 impl Tool for TestTool {
989 fn id(&self) -> &str {
990 &self.id
991 }
992 fn name(&self) -> &str {
993 "Test"
994 }
995 fn description(&self) -> &str {
996 "A test tool"
997 }
998 fn input_schema(&self) -> Value {
999 serde_json::json!({"type": "object"})
1000 }
1001 async fn execute(
1002 &self,
1003 _args: Value,
1004 _ctx: ai_agents_core::ToolExecutionContext,
1005 ) -> ToolResult {
1006 ToolResult::ok("test")
1007 }
1008 }
1009
1010 #[test]
1011 fn test_register_and_get() {
1012 let mut registry = ToolRegistry::new();
1013 let tool = Arc::new(TestTool {
1014 id: "test".to_string(),
1015 });
1016
1017 registry.register(tool).unwrap();
1018 assert!(registry.get("test").is_some());
1019 assert_eq!(registry.len(), 1);
1020 }
1021
1022 #[test]
1023 fn test_duplicate_registration() {
1024 let mut registry = ToolRegistry::new();
1025 let tool1 = Arc::new(TestTool {
1026 id: "test".to_string(),
1027 });
1028 let tool2 = Arc::new(TestTool {
1029 id: "test".to_string(),
1030 });
1031
1032 registry.register(tool1).unwrap();
1033 assert!(registry.register(tool2).is_err());
1034 }
1035
1036 #[test]
1037 fn test_list_ids() {
1038 let mut registry = ToolRegistry::new();
1039 registry
1040 .register(Arc::new(TestTool {
1041 id: "a".to_string(),
1042 }))
1043 .unwrap();
1044 registry
1045 .register(Arc::new(TestTool {
1046 id: "b".to_string(),
1047 }))
1048 .unwrap();
1049
1050 let ids = registry.list_ids();
1051 assert_eq!(ids.len(), 2);
1052 assert!(ids.contains(&"a".to_string()));
1053 assert!(ids.contains(&"b".to_string()));
1054 }
1055
1056 #[test]
1057 fn test_generate_tools_prompt() {
1058 let empty_registry = ToolRegistry::new();
1059 let empty_prompt = empty_registry.generate_tools_prompt();
1060 assert!(empty_prompt.is_empty());
1061
1062 let mut registry = ToolRegistry::new();
1063 registry
1064 .register(Arc::new(TestTool {
1065 id: "test".to_string(),
1066 }))
1067 .unwrap();
1068
1069 let prompt = registry.generate_tools_prompt();
1070 assert!(prompt.contains("Available tools:"));
1071 assert!(prompt.contains("Test:"));
1072 assert!(prompt.contains("A test tool"));
1073 assert!(prompt.contains("tool_name"));
1074 }
1075
1076 #[test]
1077 fn test_generate_filtered_prompt_with_filter() {
1078 let mut registry = ToolRegistry::new();
1079 registry
1080 .register(Arc::new(TestTool {
1081 id: "tool_a".to_string(),
1082 }))
1083 .unwrap();
1084 registry
1085 .register(Arc::new(TestTool {
1086 id: "tool_b".to_string(),
1087 }))
1088 .unwrap();
1089 registry
1090 .register(Arc::new(TestTool {
1091 id: "tool_c".to_string(),
1092 }))
1093 .unwrap();
1094
1095 let prompt =
1096 registry.generate_filtered_prompt(&["tool_a".to_string(), "tool_c".to_string()]);
1097
1098 assert!(prompt.contains("tool_a") || prompt.contains("Test"));
1099 assert!(!prompt.contains("tool_b"));
1100 }
1101
1102 #[test]
1103 fn test_generate_filtered_prompt_empty_filter() {
1104 let mut registry = ToolRegistry::new();
1105 registry
1106 .register(Arc::new(TestTool {
1107 id: "tool_a".to_string(),
1108 }))
1109 .unwrap();
1110 registry
1111 .register(Arc::new(TestTool {
1112 id: "tool_b".to_string(),
1113 }))
1114 .unwrap();
1115
1116 let prompt = registry.generate_filtered_prompt(&[]);
1117 assert!(prompt.contains("Test"));
1118 }
1119
1120 #[test]
1121 fn test_generate_filtered_prompt_nonexistent_tools() {
1122 let mut registry = ToolRegistry::new();
1123 registry
1124 .register(Arc::new(TestTool {
1125 id: "tool_a".to_string(),
1126 }))
1127 .unwrap();
1128
1129 let prompt = registry.generate_filtered_prompt(&["nonexistent".to_string()]);
1130 assert!(prompt.is_empty());
1131
1132 let prompt2 =
1133 registry.generate_filtered_prompt(&["tool_a".to_string(), "nonexistent".to_string()]);
1134 assert!(prompt2.contains("Test"));
1135 }
1136
1137 #[test]
1138 fn test_set_tool_aliases() {
1139 let mut registry = ToolRegistry::new();
1140 registry
1141 .register(Arc::new(TestTool {
1142 id: "calculator".to_string(),
1143 }))
1144 .unwrap();
1145
1146 let aliases = ToolAliases::new()
1147 .with_name("ko", "계산기")
1148 .with_name("ja", "計算機")
1149 .with_description("ko", "수학 계산을 합니다");
1150
1151 registry.set_tool_aliases("calculator", aliases);
1152
1153 assert!(registry.get_by_alias("계산기", "ko").is_some());
1154 assert!(registry.get_by_alias("計算機", "ja").is_some());
1155 assert!(registry.get("calculator").is_some());
1156 }
1157
1158 #[test]
1159 fn test_get_by_alias_case_insensitive() {
1160 let mut registry = ToolRegistry::new();
1161 registry
1162 .register(Arc::new(TestTool {
1163 id: "search".to_string(),
1164 }))
1165 .unwrap();
1166
1167 let aliases = ToolAliases::new().with_name("ko", "검색");
1168 registry.set_tool_aliases("search", aliases);
1169
1170 assert!(registry.get_by_alias("검색", "ko").is_some());
1171 }
1172
1173 #[test]
1174 fn test_generate_prompt_with_language() {
1175 let mut registry = ToolRegistry::new();
1176 registry
1177 .register(Arc::new(TestTool {
1178 id: "calculator".to_string(),
1179 }))
1180 .unwrap();
1181
1182 let aliases = ToolAliases::new()
1183 .with_name("ko", "계산기")
1184 .with_description("ko", "수학 계산");
1185
1186 registry.set_tool_aliases("calculator", aliases);
1187
1188 let prompt_en = registry.generate_tools_prompt_with_lang(None, false);
1189 assert!(prompt_en.contains("Test"));
1190
1191 let prompt_ko = registry.generate_tools_prompt_with_lang(Some("ko"), false);
1192 assert!(prompt_ko.contains("계산기"));
1193 assert!(prompt_ko.contains("수학 계산"));
1194 }
1195
1196 #[test]
1197 fn test_generate_tools_prompt_parallel() {
1198 let mut registry = ToolRegistry::new();
1199 registry
1200 .register(Arc::new(TestTool {
1201 id: "tool_a".to_string(),
1202 }))
1203 .unwrap();
1204 registry
1205 .register(Arc::new(TestTool {
1206 id: "tool_b".to_string(),
1207 }))
1208 .unwrap();
1209
1210 let prompt_seq = registry.generate_tools_prompt();
1212 assert!(prompt_seq.contains("\"tool\": \"tool_name\""));
1213 assert!(!prompt_seq.contains("JSON array"));
1214 assert!(!prompt_seq.contains("tool_name1"));
1215
1216 let prompt_par = registry.generate_tools_prompt_with_parallel(true);
1218 assert!(prompt_par.contains("\"tool\": \"tool_name\""));
1219 assert!(prompt_par.contains("JSON array"));
1220 assert!(prompt_par.contains("tool_name1"));
1221 assert!(prompt_par.contains("tool_name2"));
1222 }
1223
1224 #[test]
1225 fn test_canonical_resolution_and_scoped_empty_prompt() {
1226 let mut registry = ToolRegistry::new();
1227 registry
1228 .register(Arc::new(TestTool {
1229 id: "calculator".to_string(),
1230 }))
1231 .unwrap();
1232 let aliases = ToolAliases::new().with_name("ko", "계산기");
1233 registry.set_tool_aliases("calculator", aliases);
1234
1235 let by_id = registry.resolve("calculator").unwrap();
1236 assert_eq!(by_id.identity.canonical_id, "calculator");
1237
1238 let by_alias = registry.resolve("계산기").unwrap();
1239 assert_eq!(by_alias.identity.canonical_id, "calculator");
1240
1241 let scoped = registry.generate_scoped_prompt_with_parallel(&[], false);
1242 assert!(scoped.is_empty());
1243 }
1244
1245 #[test]
1246 fn test_generate_filtered_prompt_parallel() {
1247 let mut registry = ToolRegistry::new();
1248 registry
1249 .register(Arc::new(TestTool {
1250 id: "tool_a".to_string(),
1251 }))
1252 .unwrap();
1253 registry
1254 .register(Arc::new(TestTool {
1255 id: "tool_b".to_string(),
1256 }))
1257 .unwrap();
1258
1259 let prompt_seq =
1261 registry.generate_filtered_prompt(&["tool_a".to_string(), "tool_b".to_string()]);
1262 assert!(!prompt_seq.contains("JSON array"));
1263
1264 let prompt_par = registry.generate_filtered_prompt_with_parallel(
1266 &["tool_a".to_string(), "tool_b".to_string()],
1267 true,
1268 );
1269 assert!(prompt_par.contains("JSON array"));
1270 assert!(prompt_par.contains("tool_name1"));
1271 }
1272}