Skip to main content

pe_tools/
selector.rs

1//! Tool selection — mechanism for filtering available tools before an LLM call.
2//!
3//! The library provides `ToolSelector` as the extension point and
4//! `AllToolsSelector` as the default (pass-through) implementation.
5//! Users implement their own selectors for context-aware tool filtering.
6
7use async_trait::async_trait;
8use pe_core::{Message, ToolSchema};
9
10/// Async trait for selecting which tools to present to the LLM.
11///
12/// Implementations receive the full set of available tools and the
13/// conversation history, then return the names of tools to include.
14///
15/// # Examples
16///
17/// ```
18/// use pe_tools::selector::{ToolSelector, AllToolsSelector};
19/// use pe_core::{ToolSchema, Message};
20///
21/// # tokio::runtime::Runtime::new().unwrap().block_on(async {
22/// let selector = AllToolsSelector;
23/// let tools = vec![
24///     ToolSchema {
25///         name: "search".into(),
26///         description: "Search the web".into(),
27///         parameters: serde_json::json!({}),
28///         strict: false,
29///     },
30/// ];
31/// let selected = selector.select(&tools, &[]).await;
32/// assert_eq!(selected, vec!["search".to_string()]);
33/// # });
34/// ```
35#[async_trait]
36pub trait ToolSelector: Send + Sync + 'static {
37    /// Select tool names from the available set.
38    ///
39    /// Returns a `Vec<String>` of tool names that should be included
40    /// in the next LLM call. Names must match `ToolSchema::name` values.
41    async fn select(&self, available: &[ToolSchema], messages: &[Message]) -> Vec<String>;
42}
43
44/// Default selector that returns all available tools.
45///
46/// This is the pass-through implementation: every registered tool
47/// is presented to the LLM. Users who need filtering implement
48/// their own `ToolSelector`.
49///
50/// # Examples
51///
52/// ```
53/// use pe_tools::selector::{ToolSelector, AllToolsSelector};
54/// use pe_core::{ToolSchema, Message};
55///
56/// # tokio::runtime::Runtime::new().unwrap().block_on(async {
57/// let selector = AllToolsSelector;
58/// let names = selector.select(&[], &[]).await;
59/// assert!(names.is_empty());
60/// # });
61/// ```
62#[derive(Debug, Clone, Default)]
63pub struct AllToolsSelector;
64
65#[async_trait]
66impl ToolSelector for AllToolsSelector {
67    async fn select(&self, available: &[ToolSchema], _messages: &[Message]) -> Vec<String> {
68        available.iter().map(|t| t.name.clone()).collect()
69    }
70}
71
72#[cfg(test)]
73mod tests {
74    use super::*;
75    use pe_core::Message;
76    use serde_json::json;
77
78    fn make_tool(name: &str) -> ToolSchema {
79        ToolSchema {
80            name: name.into(),
81            description: format!("{name} tool"),
82            parameters: json!({}),
83            strict: false,
84        }
85    }
86
87    #[tokio::test]
88    async fn all_tools_selector_returns_every_tool_name() {
89        let selector = AllToolsSelector;
90        let tools = vec![
91            make_tool("search"),
92            make_tool("calculator"),
93            make_tool("email"),
94        ];
95
96        let selected = selector.select(&tools, &[]).await;
97
98        assert_eq!(selected, vec!["search", "calculator", "email"]);
99    }
100
101    #[tokio::test]
102    async fn all_tools_selector_empty_input_returns_empty() {
103        let selector = AllToolsSelector;
104        let selected = selector.select(&[], &[]).await;
105        assert!(selected.is_empty());
106    }
107
108    #[tokio::test]
109    async fn all_tools_selector_ignores_messages() {
110        let selector = AllToolsSelector;
111        let tools = vec![make_tool("read_file")];
112        let messages = vec![
113            Message::human("Please read my file"),
114            Message::ai("I'll use the read_file tool"),
115        ];
116
117        let selected = selector.select(&tools, &messages).await;
118
119        assert_eq!(selected, vec!["read_file"]);
120    }
121
122    #[tokio::test]
123    async fn custom_selector_filters_by_conversation() {
124        /// A test selector that only includes tools mentioned in messages.
125        struct MentionedToolsSelector;
126
127        #[async_trait]
128        impl ToolSelector for MentionedToolsSelector {
129            async fn select(&self, available: &[ToolSchema], messages: &[Message]) -> Vec<String> {
130                let text: String = messages
131                    .iter()
132                    .filter_map(|m| match m {
133                        Message::Human(h) => h.content.as_text().map(|s| s.to_owned()),
134                        Message::Ai(a) => a.content.as_text().map(|s| s.to_owned()),
135                        Message::System(s) => Some(s.content.clone()),
136                        Message::Tool(t) => Some(t.content.clone()),
137                        _ => None,
138                    })
139                    .collect::<Vec<_>>()
140                    .join(" ");
141
142                available
143                    .iter()
144                    .filter(|t| text.contains(&t.name))
145                    .map(|t| t.name.clone())
146                    .collect()
147            }
148        }
149
150        let selector = MentionedToolsSelector;
151        let tools = vec![
152            make_tool("search"),
153            make_tool("calculator"),
154            make_tool("email"),
155        ];
156        let messages = vec![Message::human("I need to search for something")];
157
158        let selected = selector.select(&tools, &messages).await;
159
160        assert_eq!(selected, vec!["search"]);
161    }
162
163    #[tokio::test]
164    async fn custom_selector_can_return_empty() {
165        struct NoToolsSelector;
166
167        #[async_trait]
168        impl ToolSelector for NoToolsSelector {
169            async fn select(
170                &self,
171                _available: &[ToolSchema],
172                _messages: &[Message],
173            ) -> Vec<String> {
174                vec![]
175            }
176        }
177
178        let selector = NoToolsSelector;
179        let tools = vec![make_tool("search"), make_tool("calculator")];
180        let selected = selector.select(&tools, &[]).await;
181        assert!(selected.is_empty());
182    }
183}