Skip to main content

ai_agents_tools/builtin/
web_search.rs

1use async_trait::async_trait;
2use parking_lot::RwLock;
3use schemars::JsonSchema;
4use serde::{Deserialize, Serialize};
5use serde_json::Value;
6use std::sync::Arc;
7
8use ai_agents_core::{
9    ResultLimitBinding, ResultLimitKind, Tool, ToolExecutionContext, ToolOperationKind,
10    ToolPolicyBindings, ToolResult, ToolSafetyMetadata, ToolSideEffectLevel,
11};
12
13use crate::generate_schema;
14use crate::types::{
15    UnavailableWebSearchProvider, WebSearchProviderSlot, WebSearchRequest, WebSearchResponse,
16    WebSearchSafeSearch,
17};
18
19const DEFAULT_MAX_RESULTS: usize = 5;
20const DEFAULT_MAX_OUTPUT_CHARS: usize = 12_000;
21
22/// Searches public web indexes through a host-provided search provider.
23pub struct WebSearchTool {
24    provider: WebSearchProviderSlot,
25}
26
27impl WebSearchTool {
28    /// Create a web search tool with an unavailable provider slot.
29    pub fn new() -> Self {
30        Self::with_provider_slot(Arc::new(RwLock::new(Arc::new(
31            UnavailableWebSearchProvider,
32        ))))
33    }
34
35    /// Create a web search tool backed by a shared provider slot.
36    pub fn with_provider_slot(provider: WebSearchProviderSlot) -> Self {
37        Self { provider }
38    }
39}
40
41impl Default for WebSearchTool {
42    fn default() -> Self {
43        Self::new()
44    }
45}
46
47#[derive(Debug, Deserialize, JsonSchema)]
48struct WebSearchInput {
49    /// Search query sent to the host provider.
50    query: String,
51    /// Maximum result count requested by the model. Defaults to 5.
52    #[serde(
53        default,
54        deserialize_with = "crate::deserialize_optional_positive_usize"
55    )]
56    #[schemars(range(min = 1))]
57    max_results: Option<usize>,
58    /// Optional result-domain filters such as docs.rs or rust-lang.org.
59    #[serde(default)]
60    include_domains: Vec<String>,
61    /// Optional language hint such as en or ja.
62    #[serde(default)]
63    language: Option<String>,
64    /// Optional region hint such as US or JP.
65    #[serde(default)]
66    region: Option<String>,
67    /// Optional safe-search preference.
68    #[serde(default)]
69    safe_search: Option<WebSearchSafeSearch>,
70}
71
72#[derive(Debug, Serialize)]
73struct WebSearchOutput {
74    available: bool,
75    query: String,
76    provider: Option<String>,
77    count: usize,
78    truncated: bool,
79    results: Vec<crate::types::WebSearchResultItem>,
80    message: Option<String>,
81}
82
83#[async_trait]
84impl Tool for WebSearchTool {
85    fn id(&self) -> &str {
86        "web_search"
87    }
88
89    fn name(&self) -> &str {
90        "Web Search"
91    }
92
93    fn description(&self) -> &str {
94        "Search public web indexes through a host-provided provider and return bounded cited results."
95    }
96
97    fn input_schema(&self) -> Value {
98        generate_schema::<WebSearchInput>()
99    }
100
101    fn safety_metadata(&self) -> ToolSafetyMetadata {
102        ToolSafetyMetadata {
103            read_only: true,
104            concurrency_safe: true,
105            operation: ToolOperationKind::Network,
106            side_effect_level: ToolSideEffectLevel::ExternalRead,
107            requires_network: true,
108            destructive: false,
109            open_world: true,
110            host_dependent: true,
111            requires_user_interaction: false,
112            supports_cancellation: true,
113            default_requires_approval: false,
114            should_defer_schema: false,
115            max_output_chars: Some(DEFAULT_MAX_OUTPUT_CHARS),
116            max_result_size_chars: Some(DEFAULT_MAX_OUTPUT_CHARS),
117        }
118    }
119
120    fn policy_bindings(&self) -> ToolPolicyBindings {
121        ToolPolicyBindings {
122            result_limit_fields: vec![ResultLimitBinding::new(
123                "max_results",
124                ResultLimitKind::MaxResults,
125            )],
126            ..Default::default()
127        }
128    }
129
130    async fn execute(&self, args: Value, ctx: ToolExecutionContext) -> ToolResult {
131        let input: WebSearchInput = match serde_json::from_value(args) {
132            Ok(input) => input,
133            Err(error) => return ToolResult::error(format!("Invalid input: {}", error)),
134        };
135        if let Err(error) = crate::validate_positive_max_results(ctx.limits.max_results) {
136            return ToolResult::error(format!("Invalid result limit: {error}"));
137        }
138        let max_results = input
139            .max_results
140            .unwrap_or(DEFAULT_MAX_RESULTS)
141            .min(ctx.limits.max_results.unwrap_or(DEFAULT_MAX_RESULTS));
142        let request = WebSearchRequest {
143            query: input.query.clone(),
144            max_results: Some(max_results),
145            include_domains: input.include_domains,
146            language: input.language,
147            region: input.region,
148            safe_search: input.safe_search,
149        };
150        let provider = self.provider.read().clone();
151        let response = provider.search(request).await;
152        result_from_response(input.query, response, max_results)
153    }
154}
155
156fn result_from_response(
157    query: String,
158    mut response: WebSearchResponse,
159    max_results: usize,
160) -> ToolResult {
161    let mut truncated = response.truncated;
162    if response.results.len() > max_results {
163        response.results.truncate(max_results);
164        truncated = true;
165    }
166    let output = WebSearchOutput {
167        available: response.available,
168        query,
169        provider: response.provider,
170        count: response.results.len(),
171        truncated,
172        results: response.results,
173        message: response.message,
174    };
175    let serialized = match serde_json::to_string(&output) {
176        Ok(value) => value,
177        Err(error) => return ToolResult::error(format!("Serialization error: {}", error)),
178    };
179    if output.available {
180        ToolResult::ok(serialized)
181    } else {
182        ToolResult {
183            success: false,
184            output: serialized,
185            metadata: None,
186        }
187    }
188}
189
190#[cfg(test)]
191mod tests {
192    use super::*;
193    use crate::types::{StaticWebSearchProvider, WebSearchProvider, WebSearchResultItem};
194    use parking_lot::Mutex;
195    use std::collections::HashMap;
196
197    struct RecordingProvider {
198        max_results: Mutex<Vec<usize>>,
199    }
200
201    #[async_trait::async_trait]
202    impl WebSearchProvider for RecordingProvider {
203        async fn search(&self, request: WebSearchRequest) -> WebSearchResponse {
204            self.max_results
205                .lock()
206                .push(request.max_results.unwrap_or_default());
207            WebSearchResponse {
208                available: true,
209                ..Default::default()
210            }
211        }
212    }
213
214    #[tokio::test]
215    async fn max_results_rejects_zero_and_preserves_positive_caps() {
216        let provider = Arc::new(RecordingProvider {
217            max_results: Mutex::new(Vec::new()),
218        });
219        let slot = Arc::new(RwLock::new(
220            Arc::clone(&provider) as Arc<dyn crate::types::WebSearchProvider>
221        ));
222        let tool = WebSearchTool::with_provider_slot(slot);
223
224        let zero_request = tool
225            .execute(
226                serde_json::json!({"query": "rust", "max_results": 0}),
227                ToolExecutionContext::test("web_search"),
228            )
229            .await;
230        assert!(!zero_request.success);
231        assert!(
232            zero_request
233                .output
234                .contains("max_results must be greater than 0")
235        );
236
237        let mut zero_context = ToolExecutionContext::test("web_search");
238        zero_context.limits.max_results = Some(0);
239        let invalid_context = tool
240            .execute(serde_json::json!({"query": "rust"}), zero_context)
241            .await;
242        assert!(!invalid_context.success);
243
244        let mut capped_context = ToolExecutionContext::test("web_search");
245        capped_context.limits.max_results = Some(2);
246        assert!(
247            tool.execute(
248                serde_json::json!({"query": "rust", "max_results": 4}),
249                capped_context.clone(),
250            )
251            .await
252            .success
253        );
254        assert!(
255            tool.execute(
256                serde_json::json!({"query": "rust", "max_results": 1}),
257                capped_context,
258            )
259            .await
260            .success
261        );
262
263        assert_eq!(*provider.max_results.lock(), vec![2, 1]);
264    }
265
266    #[tokio::test]
267    async fn static_provider_returns_bounded_results() {
268        let mut responses = HashMap::new();
269        responses.insert(
270            "rust async".to_string(),
271            WebSearchResponse {
272                available: true,
273                provider: Some("fixture".to_string()),
274                results: vec![
275                    WebSearchResultItem {
276                        title: "one".to_string(),
277                        url: "https://example.com/one".to_string(),
278                        snippet: "first".to_string(),
279                        source: None,
280                        published_at: None,
281                    },
282                    WebSearchResultItem {
283                        title: "two".to_string(),
284                        url: "https://example.com/two".to_string(),
285                        snippet: "second".to_string(),
286                        source: None,
287                        published_at: None,
288                    },
289                ],
290                truncated: false,
291                message: None,
292            },
293        );
294        let provider = Arc::new(RwLock::new(
295            Arc::new(StaticWebSearchProvider::new(responses))
296                as Arc<dyn crate::types::WebSearchProvider>,
297        ));
298        let tool = WebSearchTool::with_provider_slot(provider);
299        let result = tool
300            .execute(
301                serde_json::json!({"query":"rust async","max_results":1}),
302                ToolExecutionContext::test("web_search"),
303            )
304            .await;
305        assert!(result.success);
306        let output: serde_json::Value = serde_json::from_str(&result.output).unwrap();
307        assert_eq!(output["count"], 1);
308        assert_eq!(output["truncated"], true);
309    }
310}