Skip to main content

mermaid_cli/providers/tool/
web.rs

1//! Web tools: `web_search` and `web_fetch`.
2//!
3//! Both delegate to `web_client::WebSearchClient` — a thin HTTP
4//! client for Ollama Cloud's web API (bearer-token path, via
5//! `OLLAMA_API_KEY`). The wrapper's job is cancellation plumbing +
6//! multi-query fan-out.
7
8use 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
18/// `web_search` — query Ollama Cloud's web-search endpoint. Accepts a
19/// single `{query, max_results}` OR a list of `{queries: [{query,
20/// max_results}]}` for parallel fan-out.
21pub 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
157/// `web_fetch` — retrieve a URL's readable content (Ollama Cloud's
158/// fetch endpoint). Single URL, single response.
159pub 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}