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
22pub struct WebSearchTool {
24 provider: WebSearchProviderSlot,
25}
26
27impl WebSearchTool {
28 pub fn new() -> Self {
30 Self::with_provider_slot(Arc::new(RwLock::new(Arc::new(
31 UnavailableWebSearchProvider,
32 ))))
33 }
34
35 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 query: String,
51 #[serde(
53 default,
54 deserialize_with = "crate::deserialize_optional_positive_usize"
55 )]
56 #[schemars(range(min = 1))]
57 max_results: Option<usize>,
58 #[serde(default)]
60 include_domains: Vec<String>,
61 #[serde(default)]
63 language: Option<String>,
64 #[serde(default)]
66 region: Option<String>,
67 #[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}