mermaid_cli/providers/tool/
web.rs1use std::sync::Arc;
9
10use async_trait::async_trait;
11
12use crate::domain::{ToolDefinition, ToolMetadata, ToolOutcome, ToolRunMetadata};
13
14use super::super::ctx::{ExecContext, ProgressEvent};
15use super::ToolExecutor;
16use super::web_client::{WebFetchResult, WebSearchClient};
17
18pub struct WebSearchTool {
22 client: Arc<WebSearchClient>,
23}
24
25impl WebSearchTool {
26 pub fn new(api_key: String) -> Self {
27 Self {
28 client: Arc::new(WebSearchClient::new(api_key)),
29 }
30 }
31}
32
33#[async_trait]
34impl ToolExecutor for WebSearchTool {
35 fn name(&self) -> &'static str {
36 "web_search"
37 }
38
39 fn schema(&self) -> ToolDefinition {
40 ToolDefinition {
41 name: "web_search".to_string(),
42 description:
43 "Search the web via Ollama Cloud's search API. Takes either a single `query` + `max_results`, or an array of `queries` for parallel fan-out."
44 .to_string(),
45 input_schema: serde_json::json!({
46 "type": "object",
47 "properties": {
48 "query": { "type": "string" },
49 "max_results": { "type": "integer", "minimum": 1, "maximum": 10, "default": 5 },
50 "queries": {
51 "type": "array",
52 "items": {
53 "type": "object",
54 "properties": {
55 "query": { "type": "string" },
56 "max_results": { "type": "integer", "minimum": 1, "maximum": 10 }
57 },
58 "required": ["query"]
59 }
60 }
61 }
62 }),
63 }
64 }
65
66 async fn execute(&self, args: serde_json::Value, ctx: ExecContext) -> ToolOutcome {
67 let queries = match parse_queries(&args) {
68 Ok(q) => q,
69 Err(e) => return ToolOutcome::error(e, 0.0),
70 };
71 if queries.is_empty() {
72 return ToolOutcome::error("web_search requires at least one query", 0.0);
73 }
74 if let Some(blocked) = super::policy_gate::gate_external(
75 &ctx,
76 "web_search",
77 crate::runtime::ToolCategory::Web,
78 format!("web_search ({} queries)", queries.len()),
79 &args,
80 )
81 .await
82 {
83 return blocked;
84 }
85
86 let start = std::time::Instant::now();
87 let mut combined = String::new();
88 let mut result_count = 0usize;
89 let mut sources = Vec::new();
90 for (idx, (query, count)) in queries.iter().enumerate() {
91 let _ = ctx
92 .progress
93 .send(ProgressEvent::Status(format!(
94 "searching {}/{}: {}",
95 idx + 1,
96 queries.len(),
97 query
98 )))
99 .await;
100
101 let search = self.client.search_query(query, *count);
102 tokio::select! {
103 biased;
104 _ = ctx.token.cancelled() => return ToolOutcome::cancelled(),
105 result = search => {
106 match result {
107 Ok(results) => {
108 result_count += results.len();
109 sources.extend(results.iter().map(|result| result.url.clone()));
110 let formatted = self.client.format_results(&results);
111 if queries.len() > 1 {
112 combined.push_str(&format!("=== query: {} ===\n{}\n\n", query, formatted));
113 } else {
114 combined = formatted;
115 }
116 },
117 Err(e) => {
118 return ToolOutcome::error(
119 format!("web_search({}): {}", query, e),
120 start.elapsed().as_secs_f64(),
121 );
122 },
123 }
124 }
125 }
126 }
127
128 let duration_secs = start.elapsed().as_secs_f64();
129 let requested_count = queries.iter().map(|(_, count)| *count).sum();
130 let query_texts = queries.iter().map(|(query, _)| query.clone()).collect();
131 ToolOutcome::success(
132 combined,
133 format!(
134 "{} {} returned",
135 result_count,
136 if result_count == 1 {
137 "result"
138 } else {
139 "results"
140 }
141 ),
142 duration_secs,
143 )
144 .with_metadata(ToolRunMetadata {
145 detail: ToolMetadata::WebSearch {
146 queries: query_texts,
147 requested_count,
148 result_count,
149 sources,
150 },
151 result_count: Some(result_count),
152 ..ToolRunMetadata::default()
153 })
154 }
155}
156
157pub struct WebFetchTool {
160 client: Arc<WebSearchClient>,
161}
162
163impl WebFetchTool {
164 pub fn new(api_key: String) -> Self {
165 Self {
166 client: Arc::new(WebSearchClient::new(api_key)),
167 }
168 }
169}
170
171#[async_trait]
172impl ToolExecutor for WebFetchTool {
173 fn name(&self) -> &'static str {
174 "web_fetch"
175 }
176
177 fn schema(&self) -> ToolDefinition {
178 ToolDefinition {
179 name: "web_fetch".to_string(),
180 description: "Retrieve a single URL's main content as text (Ollama Cloud fetch API)."
181 .to_string(),
182 input_schema: serde_json::json!({
183 "type": "object",
184 "properties": { "url": { "type": "string" } },
185 "required": ["url"]
186 }),
187 }
188 }
189
190 async fn execute(&self, args: serde_json::Value, ctx: ExecContext) -> ToolOutcome {
191 let Some(url) = args.get("url").and_then(|v| v.as_str()) else {
192 return ToolOutcome::error("web_fetch requires 'url' (string)", 0.0);
193 };
194 if let Some(blocked) = super::policy_gate::gate_external(
195 &ctx,
196 "web_fetch",
197 crate::runtime::ToolCategory::Web,
198 format!("web_fetch {}", url),
199 &args,
200 )
201 .await
202 {
203 return blocked;
204 }
205 let start = std::time::Instant::now();
206 let fetch = self.client.fetch_url(url);
207
208 tokio::select! {
209 biased;
210 _ = ctx.token.cancelled() => ToolOutcome::cancelled(),
211 result = fetch => match result {
212 Ok(page) => {
213 let output = format_fetch(url, &page);
214 let duration_secs = start.elapsed().as_secs_f64();
215 let line_count = output.lines().count();
216 let byte_count = output.len();
217 let title = if page.title.is_empty() {
218 None
219 } else {
220 Some(page.title)
221 };
222 ToolOutcome::success(
223 output,
224 format!("{} {} fetched", line_count, if line_count == 1 { "line" } else { "lines" }),
225 duration_secs,
226 )
227 .with_metadata(ToolRunMetadata {
228 detail: ToolMetadata::WebFetch {
229 url: url.to_string(),
230 title,
231 line_count,
232 byte_count,
233 },
234 line_count: Some(line_count),
235 byte_count: Some(byte_count),
236 ..ToolRunMetadata::default()
237 })
238 },
239 Err(e) => ToolOutcome::error(
240 format!("web_fetch({}): {}", url, e),
241 start.elapsed().as_secs_f64(),
242 ),
243 },
244 }
245 }
246}
247
248fn format_fetch(url: &str, page: &WebFetchResult) -> String {
249 let title = if page.title.is_empty() {
250 "(no title)"
251 } else {
252 page.title.as_str()
253 };
254 format!("# {}\n\nURL: {}\n\n{}", title, url, page.content)
255}
256
257fn parse_queries(args: &serde_json::Value) -> Result<Vec<(String, usize)>, String> {
258 if let Some(arr) = args.get("queries").and_then(|v| v.as_array()) {
259 let mut out = Vec::with_capacity(arr.len());
260 for v in arr {
261 let Some(obj) = v.as_object() else {
262 return Err(
263 "web_search: 'queries' must be an array of {query, max_results}".to_string(),
264 );
265 };
266 let Some(query) = obj.get("query").and_then(|x| x.as_str()) else {
267 return Err("web_search: each query entry needs 'query' (string)".to_string());
268 };
269 let count = obj
270 .get("max_results")
271 .or_else(|| obj.get("result_count"))
272 .and_then(|x| x.as_u64())
273 .unwrap_or(5)
274 .clamp(1, 10) as usize;
275 out.push((query.to_string(), count));
276 }
277 return Ok(out);
278 }
279 if let Some(query) = args.get("query").and_then(|v| v.as_str()) {
280 let count = args
281 .get("max_results")
282 .or_else(|| args.get("result_count"))
283 .and_then(|v| v.as_u64())
284 .unwrap_or(5)
285 .clamp(1, 10) as usize;
286 return Ok(vec![(query.to_string(), count)]);
287 }
288 Err("web_search requires 'query' (string) or 'queries' (array)".to_string())
289}
290
291#[cfg(test)]
292mod tests {
293 use super::*;
294
295 #[test]
296 fn parse_queries_single_form() {
297 let args = serde_json::json!({"query": "rust async", "max_results": 3});
298 let q = parse_queries(&args).unwrap();
299 assert_eq!(q.len(), 1);
300 assert_eq!(q[0].0, "rust async");
301 assert_eq!(q[0].1, 3);
302 }
303
304 #[test]
305 fn parse_queries_array_form() {
306 let args = serde_json::json!({"queries": [
307 {"query": "a", "max_results": 2},
308 {"query": "b", "result_count": 5},
309 ]});
310 let q = parse_queries(&args).unwrap();
311 assert_eq!(q.len(), 2);
312 assert_eq!(q[1].1, 5);
313 }
314
315 #[test]
316 fn parse_queries_missing_errors() {
317 let args = serde_json::json!({});
318 assert!(parse_queries(&args).is_err());
319 }
320
321 #[test]
322 fn parse_queries_clamps_count() {
323 let args = serde_json::json!({"query": "q", "max_results": 999});
324 let q = parse_queries(&args).unwrap();
325 assert_eq!(q[0].1, 10);
326 let args = serde_json::json!({"query": "q", "max_results": 0});
327 let q = parse_queries(&args).unwrap();
328 assert_eq!(q[0].1, 1);
329 }
330}