Skip to main content

ai_agents_tools/
registry.rs

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/// Schema rendering mode for tool prompt generation.
18#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
19#[serde(rename_all = "snake_case")]
20pub enum ToolSchemaPromptMode {
21    /// Include full JSON schema properties for every granted tool.
22    #[default]
23    Full,
24    /// Include only compact descriptors: name, description, required fields, and property types.
25    Compact,
26}
27
28/// Canonical identity produced by registry resolution.
29#[derive(Debug, Clone)]
30pub struct ToolIdentity {
31    /// Name, display name, or alias supplied by the caller.
32    pub requested_name: String,
33    /// Canonical tool ID used for policy and execution.
34    pub canonical_id: String,
35    /// Display name of the resolved tool.
36    pub display_name: String,
37    /// Provider ID for provider-backed tools.
38    pub provider_id: Option<String>,
39}
40
41/// Resolved tool handle plus canonical identity evidence.
42#[derive(Clone)]
43pub struct ResolvedTool {
44    /// Canonical identity returned by lookup.
45    pub identity: ToolIdentity,
46    /// Executable tool implementation.
47    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
59/// Registry for built-in, provider, alias, and localized tool lookup.
60pub 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    /// Creates an empty registry with versioned canonical indexes.
92    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    /// Returns the registry version used in tool execution evidence.
116    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    /// Registers a built-in or custom tool by canonical ID.
137    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    /// Resolves any accepted name to a canonical tool ID.
164    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    /// Resolves safety metadata for a registered tool.
170    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    /// Returns the shared question handler slot for host-bound tools.
176    pub fn question_handler_slot(&self) -> QuestionHandlerSlot {
177        Arc::clone(&self.question_handler)
178    }
179
180    /// Installs or clears the question handler used by `ask_user`.
181    pub fn set_question_handler(&self, handler: Option<Arc<dyn QuestionHandler>>) {
182        *self.question_handler.write() = handler;
183    }
184
185    /// Returns the shared diagnostics provider slot for host-bound tools.
186    pub fn diagnostics_provider_slot(&self) -> DiagnosticsProviderSlot {
187        Arc::clone(&self.diagnostics_provider)
188    }
189
190    /// Installs the diagnostics provider used by `diagnostics`.
191    pub fn set_diagnostics_provider(&self, provider: Arc<dyn DiagnosticsProvider>) {
192        *self.diagnostics_provider.write() = provider;
193    }
194
195    /// Returns whether the diagnostics provider can serve requests now.
196    pub fn diagnostics_available(&self) -> bool {
197        self.diagnostics_provider.read().is_available()
198    }
199
200    /// Returns the shared command runner slot for host-bound tools.
201    pub fn command_runner_slot(&self) -> CommandRunnerSlot {
202        Arc::clone(&self.command_runner)
203    }
204
205    /// Installs the command runner used by `command`.
206    pub fn set_command_runner(&self, runner: Arc<dyn CommandRunner>) {
207        *self.command_runner.write() = runner;
208    }
209
210    /// Returns whether the command runner can serve requests now.
211    pub fn command_runner_available(&self) -> bool {
212        self.command_runner.read().is_available()
213    }
214
215    /// Returns the session-local file version store shared with file tools.
216    pub fn file_version_store(&self) -> FileVersionStore {
217        self.file_versions.clone()
218    }
219
220    /// Returns the session-local todo store shared with `todo`.
221    pub fn todo_store(&self) -> TodoStore {
222        self.todo_store.clone()
223    }
224
225    /// Returns a snapshot of session-local todo items.
226    pub fn todos(&self) -> Vec<TodoItem> {
227        self.todo_store.list()
228    }
229
230    /// Returns the shared web-fetch extractor slot.
231    pub fn web_fetch_extractor_slot(&self) -> Arc<RwLock<Option<Arc<dyn LLMProvider>>>> {
232        Arc::clone(&self.web_fetch_extractor)
233    }
234
235    /// Installs or clears the LLM used for `web_fetch` prompt extraction.
236    pub fn set_web_fetch_extractor(&self, extractor: Option<Arc<dyn LLMProvider>>) {
237        *self.web_fetch_extractor.write() = extractor;
238    }
239
240    /// Returns the shared provider slot for `web_search`.
241    pub fn web_search_provider_slot(&self) -> WebSearchProviderSlot {
242        Arc::clone(&self.web_search_provider)
243    }
244
245    /// Installs the provider used by `web_search`.
246    pub fn set_web_search_provider(&self, provider: Arc<dyn WebSearchProvider>) {
247        *self.web_search_provider.write() = provider;
248    }
249
250    /// Returns whether the web search provider can serve requests now.
251    pub fn web_search_available(&self) -> bool {
252        self.web_search_provider.read().is_available()
253    }
254
255    /// Resolves IDs, display names, and aliases to one canonical tool.
256    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        // Resolve provider tools before taking index locks so provider I/O cannot block registry access.
438        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                // Resolve provider tools before taking index locks so refresh keeps the old snapshot available during provider I/O.
581                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    /// Generates a prompt for an explicit tool scope.
726    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    /// Generate a scoped prompt with a configurable schema rendering mode.
820    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    /// Generate a compact tool prompt with stable ordering and reduced schema.
846    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    /// Append tool call format instructions to a prompt.
899    /// When `parallel` is true, also instructs the LLM to use a JSON array
900    /// for multiple simultaneous tool calls.
901    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
920/// Build a compact schema descriptor from a full JSON schema.
921/// Includes required fields and property types only, with stable key ordering.
922fn 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        // Without parallel: no array instruction
1211        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        // With parallel: array instruction present
1217        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        // Filtered without parallel
1260        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        // Filtered with parallel
1265        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}