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