Skip to main content

scv_tools/
web.rs

1//! Web tools: `web_fetch` and the configured-backend `web_search`.
2//!
3//! `web_fetch` refuses loopback, private, link-local, and other non-public
4//! addresses unless the user allows them. Host names are resolved once by a
5//! checking resolver whose addresses are the only ones the client connects
6//! to, so a name cannot be rebound to a local address between the check and
7//! the connection. IP-literal URLs, including redirect targets, are checked
8//! before any request.
9
10use std::{
11    net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr},
12    sync::Arc,
13    time::Duration,
14};
15
16use async_trait::async_trait;
17use futures_util::StreamExt;
18use reqwest::{
19    Url,
20    dns::{Addrs, Name, Resolve, Resolving},
21    redirect,
22};
23use scv_core::{Tool, ToolContext, ToolError, ToolOutput, ToolRegistry, ToolRisk, ToolSpec};
24use serde::Deserialize;
25use serde_json::{Value, json};
26
27use crate::{bounded, parse_args};
28
29const USER_AGENT: &str = concat!(
30    "scv/",
31    env!("CARGO_PKG_VERSION"),
32    " (+https://github.com/PeiyuanQi/scv)"
33);
34const MAX_URL_BYTES: usize = 4096;
35const MAX_QUERY_BYTES: usize = 512;
36const MAX_SEARCH_RESPONSE_BYTES: usize = 1024 * 1024;
37const HTML_WIDTH: usize = 120;
38
39/// Web tool settings resolved from the user's configuration.
40#[derive(Debug, Clone)]
41pub struct WebToolsConfig {
42    pub fetch_max_bytes: usize,
43    pub fetch_timeout: Duration,
44    pub max_redirects: usize,
45    /// HTTPS hosts fetched without approval. `*.example.com` matches
46    /// subdomains of `example.com` but not the domain itself.
47    pub auto_approve_domains: Vec<String>,
48    pub allow_private_addresses: bool,
49    pub search: Option<SearchBackend>,
50    pub max_search_results: usize,
51    pub output_limit: usize,
52}
53
54/// A search service that SCV queries itself. Provider-hosted search is
55/// configured on the provider instead and needs no SCV tool.
56#[derive(Clone)]
57pub enum SearchBackend {
58    Searxng { url: String },
59    Brave { url: String, api_key: String },
60}
61
62impl std::fmt::Debug for SearchBackend {
63    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
64        match self {
65            Self::Searxng { url } => formatter.debug_struct("Searxng").field("url", url).finish(),
66            Self::Brave { url, .. } => formatter
67                .debug_struct("Brave")
68                .field("url", url)
69                .field("api_key", &"[REDACTED]")
70                .finish(),
71        }
72    }
73}
74
75/// Registers `web_fetch`, and `web_search` when a search backend is configured.
76pub fn register(registry: &mut ToolRegistry, config: WebToolsConfig) -> Result<(), ToolError> {
77    let config = Arc::new(config);
78    let allow_private = config.allow_private_addresses;
79    registry.register(Arc::new(WebFetchTool {
80        config: Arc::clone(&config),
81        address_allowed: Arc::new(move |address: SocketAddr| {
82            allow_private || is_public(address.ip())
83        }),
84    }))?;
85    if let Some(backend) = config.search.clone() {
86        registry.register(Arc::new(WebSearchTool {
87            backend,
88            config: Arc::clone(&config),
89        }))?;
90    }
91    Ok(())
92}
93
94/// Decides whether the client may connect to an address. Production uses
95/// [`is_public`]; tests substitute a port-aware check for loopback servers.
96type AddressCheck = Arc<dyn Fn(SocketAddr) -> bool + Send + Sync>;
97
98struct WebFetchTool {
99    config: Arc<WebToolsConfig>,
100    address_allowed: AddressCheck,
101}
102
103#[derive(Deserialize)]
104#[serde(deny_unknown_fields)]
105struct FetchArgs {
106    url: String,
107    #[serde(default)]
108    offset: Option<usize>,
109}
110
111impl WebFetchTool {
112    fn parse_url(&self, value: &str) -> Result<Url, ToolError> {
113        if value.len() > MAX_URL_BYTES {
114            return Err(ToolError(format!("url exceeds {MAX_URL_BYTES} bytes")));
115        }
116        let url =
117            Url::parse(value.trim()).map_err(|error| ToolError(format!("invalid url: {error}")))?;
118        check_url_shape(&url)?;
119        Ok(url)
120    }
121
122    fn auto_approved(&self, url: &Url) -> bool {
123        url.scheme() == "https"
124            && url
125                .host_str()
126                .is_some_and(|host| domain_listed(&self.config.auto_approve_domains, host))
127    }
128}
129
130/// Only plain HTTP(S) URLs with a host and no embedded credentials.
131fn check_url_shape(url: &Url) -> Result<(), ToolError> {
132    if !matches!(url.scheme(), "http" | "https") {
133        return Err(ToolError(format!(
134            "web_fetch supports only http and https URLs, not {}",
135            url.scheme()
136        )));
137    }
138    if url.host_str().is_none_or(str::is_empty) {
139        return Err(ToolError("url has no host".into()));
140    }
141    if !url.username().is_empty() || url.password().is_some() {
142        return Err(ToolError(
143            "urls with embedded credentials are not allowed".into(),
144        ));
145    }
146    Ok(())
147}
148
149/// The literal IP of a URL host, if it is one.
150fn host_ip(url: &Url) -> Option<IpAddr> {
151    let host = url.host_str()?;
152    host.trim_start_matches('[')
153        .trim_end_matches(']')
154        .parse()
155        .ok()
156}
157
158fn check_literal(url: &Url, allowed: &AddressCheck) -> Result<(), String> {
159    if let Some(ip) = host_ip(url) {
160        let port = url.port_or_known_default().unwrap_or(0);
161        if !allowed(SocketAddr::new(ip, port)) {
162            return Err(format!(
163                "{ip} is a loopback, private, or otherwise non-public address"
164            ));
165        }
166    }
167    Ok(())
168}
169
170/// Case-insensitive host match against the allowlist.
171pub fn domain_listed(domains: &[String], host: &str) -> bool {
172    let host = host.trim_end_matches('.').to_ascii_lowercase();
173    domains.iter().any(|entry| {
174        let entry = entry.trim_end_matches('.').to_ascii_lowercase();
175        match entry.strip_prefix("*.") {
176            Some(parent) => host
177                .strip_suffix(parent)
178                .is_some_and(|prefix| prefix.len() > 1 && prefix.ends_with('.')),
179            None => host == entry,
180        }
181    })
182}
183
184/// Whether an address is a routable public one: not loopback, private,
185/// link-local, shared (CGNAT), multicast, documentation, reserved, or an IPv6
186/// form that embeds such an IPv4 address.
187pub fn is_public(ip: IpAddr) -> bool {
188    match ip {
189        IpAddr::V4(ip) => is_public_v4(ip),
190        IpAddr::V6(ip) => is_public_v6(ip),
191    }
192}
193
194fn is_public_v4(ip: Ipv4Addr) -> bool {
195    let [a, b, c, _] = ip.octets();
196    !(ip.is_unspecified()
197        || ip.is_loopback()
198        || ip.is_private()
199        || ip.is_link_local()
200        || ip.is_broadcast()
201        || ip.is_multicast()
202        || ip.is_documentation()
203        || a == 0
204        || (a == 100 && (64..128).contains(&b))
205        || (a == 192 && b == 0 && c == 0)
206        || (a == 198 && (b == 18 || b == 19))
207        || a >= 240)
208}
209
210fn is_public_v6(ip: Ipv6Addr) -> bool {
211    let segments = ip.segments();
212    if let Some(v4) = ip.to_ipv4_mapped() {
213        return is_public_v4(v4);
214    }
215    // IPv4-compatible (deprecated) and NAT64 addresses embed an IPv4 address
216    // in their low 32 bits.
217    let embedded = Ipv4Addr::from(ip.to_bits() as u32);
218    if segments[..6] == [0; 6] || segments[..6] == [0x64, 0xff9b, 0, 0, 0, 0] {
219        return !ip.is_unspecified() && !ip.is_loopback() && is_public_v4(embedded);
220    }
221    // 6to4 embeds its IPv4 address in bits 16..48.
222    if segments[0] == 0x2002 {
223        let v4 = Ipv4Addr::new(
224            (segments[1] >> 8) as u8,
225            segments[1] as u8,
226            (segments[2] >> 8) as u8,
227            segments[2] as u8,
228        );
229        return is_public_v4(v4);
230    }
231    !(ip.is_unspecified()
232        || ip.is_loopback()
233        || ip.is_multicast()
234        || (segments[0] & 0xfe00) == 0xfc00 // unique local
235        || (segments[0] & 0xffc0) == 0xfe80 // link-local
236        || (segments[0] & 0xffc0) == 0xfec0 // site-local
237        || (segments[0] == 0x2001 && segments[1] == 0x0db8) // documentation
238        || (segments[0] == 0x2001 && segments[1] == 0)) // Teredo
239}
240
241/// Resolves a name and fails if any of its addresses is refused, so the
242/// client only ever connects to addresses that passed the check.
243struct CheckedResolver {
244    allowed: AddressCheck,
245}
246
247impl Resolve for CheckedResolver {
248    fn resolve(&self, name: Name) -> Resolving {
249        let allowed = Arc::clone(&self.allowed);
250        let host = name.as_str().to_owned();
251        Box::pin(async move {
252            let addresses = resolve_checked(&host, &allowed).await?;
253            Ok(Box::new(addresses.into_iter()) as Addrs)
254        })
255    }
256}
257
258async fn resolve_checked(
259    host: &str,
260    allowed: &AddressCheck,
261) -> Result<Vec<SocketAddr>, Box<dyn std::error::Error + Send + Sync>> {
262    let addresses: Vec<SocketAddr> = tokio::net::lookup_host((host, 0)).await?.collect();
263    check_resolved(host, &addresses, allowed)?;
264    Ok(addresses)
265}
266
267/// Every resolved address must pass; one refused address refuses the name.
268fn check_resolved(
269    host: &str,
270    addresses: &[SocketAddr],
271    allowed: &AddressCheck,
272) -> Result<(), String> {
273    if addresses.is_empty() {
274        return Err(format!("{host} did not resolve to any address"));
275    }
276    if let Some(refused) = addresses.iter().find(|address| !allowed(**address)) {
277        return Err(format!(
278            "{host} resolves to {}, a loopback, private, or otherwise non-public address",
279            refused.ip()
280        ));
281    }
282    Ok(())
283}
284
285enum Body {
286    Html,
287    Text,
288}
289
290/// Classifies a response by its media type, sniffing only when none is given.
291fn classify(content_type: Option<&str>, bytes: &[u8]) -> Result<Body, String> {
292    let Some(media) = content_type.map(|value| {
293        value
294            .split(';')
295            .next()
296            .unwrap_or("")
297            .trim()
298            .to_ascii_lowercase()
299    }) else {
300        let head = String::from_utf8_lossy(&bytes[..bytes.len().min(1024)]).to_ascii_lowercase();
301        return if head.contains("<html") || head.contains("<!doctype html") {
302            Ok(Body::Html)
303        } else if std::str::from_utf8(bytes).is_ok()
304            || std::str::from_utf8(&bytes[..bytes.len().saturating_sub(4)]).is_ok()
305        {
306            Ok(Body::Text)
307        } else {
308            Err("the response has no content type and is not text".into())
309        };
310    };
311    if media == "text/html" || media == "application/xhtml+xml" {
312        return Ok(Body::Html);
313    }
314    let textual = media.starts_with("text/")
315        || media.ends_with("+json")
316        || media.ends_with("+xml")
317        || matches!(
318            media.as_str(),
319            "application/json"
320                | "application/xml"
321                | "application/javascript"
322                | "application/ecmascript"
323                | "application/x-javascript"
324                | "application/toml"
325                | "application/yaml"
326                | "application/x-yaml"
327                | "application/x-ndjson"
328                | "application/sql"
329                | "application/graphql"
330        );
331    if textual {
332        Ok(Body::Text)
333    } else {
334        Err(format!(
335            "web_fetch returns text only; {media} is not a text content type"
336        ))
337    }
338}
339
340/// One page of `text` starting at character `offset`, within `budget` bytes.
341/// Returns the page and the offset of the next page, if any.
342fn page(text: &str, offset: usize, budget: usize) -> (String, Option<usize>) {
343    let mut output = String::new();
344    for (taken, character) in text.chars().skip(offset).enumerate() {
345        if output.len() + character.len_utf8() > budget {
346            return (output, Some(offset + taken));
347        }
348        output.push(character);
349    }
350    (output, None)
351}
352
353fn error_chain(error: &reqwest::Error) -> String {
354    let mut message = error.to_string();
355    let mut source = std::error::Error::source(error);
356    while let Some(cause) = source {
357        let text = cause.to_string();
358        if !message.contains(&text) {
359            message.push_str(": ");
360            message.push_str(&text);
361        }
362        source = cause.source();
363    }
364    message
365}
366
367#[async_trait]
368impl Tool for WebFetchTool {
369    fn spec(&self) -> ToolSpec {
370        ToolSpec {
371            name: "web_fetch".into(),
372            description: format!(
373                "Fetch a public web page or HTTP API with a GET request and return it as readable text \
374                 (HTML is converted to text; JSON and plain text pass through). Use it to read \
375                 documentation, release notes, issues, or a URL the user gave, and to open results \
376                 from web search. Long pages are returned in parts: call again with the reported \
377                 offset. Only public addresses are reachable. HTTPS pages on {} are fetched without \
378                 approval; other hosts need approval because the URL is sent to that site.",
379                if self.config.auto_approve_domains.is_empty() {
380                    "no hosts".to_owned()
381                } else {
382                    self.config.auto_approve_domains.join(", ")
383                }
384            ),
385            parameters: json!({
386                "type":"object",
387                "properties":{
388                    "url":{"type":"string","description":"Absolute http or https URL"},
389                    "offset":{"type":"integer","minimum":0,"description":"Character offset of the part to return, from a previous call"}
390                },
391                "required":["url"],
392                "additionalProperties":false
393            }),
394        }
395    }
396
397    fn risk(&self, arguments: &Value) -> Result<ToolRisk, ToolError> {
398        let args: FetchArgs = parse_args(arguments)?;
399        let url = self.parse_url(&args.url)?;
400        Ok(if self.auto_approved(&url) {
401            ToolRisk::ReadOnly
402        } else {
403            ToolRisk::Network
404        })
405    }
406
407    fn approval_summary(&self, arguments: &Value) -> Result<String, ToolError> {
408        let args: FetchArgs = parse_args(arguments)?;
409        let url = self.parse_url(&args.url)?;
410        Ok(format!(
411            "Fetch {} with an HTTP GET (no cookies or credentials). The full URL is sent to {}.",
412            bounded(url.as_str(), 2000),
413            url.host_str().unwrap_or("the host")
414        ))
415    }
416
417    async fn execute(
418        &self,
419        arguments: Value,
420        context: ToolContext,
421    ) -> Result<ToolOutput, ToolError> {
422        let args: FetchArgs = parse_args(&arguments)?;
423        let url = self.parse_url(&args.url)?;
424        check_literal(&url, &self.address_allowed).map_err(ToolError)?;
425        let auto_approved = self.auto_approved(&url);
426        let max_redirects = self.config.max_redirects;
427        let domains = self.config.auto_approve_domains.clone();
428        let allowed = Arc::clone(&self.address_allowed);
429        let policy = redirect::Policy::custom(move |attempt| {
430            if attempt.previous().len() > max_redirects {
431                return attempt.error(format!("stopped after {max_redirects} redirects"));
432            }
433            let next = attempt.url().clone();
434            if let Err(error) = check_url_shape(&next) {
435                return attempt.error(error.0);
436            }
437            if let Err(error) = check_literal(&next, &allowed) {
438                return attempt.error(error);
439            }
440            if auto_approved
441                && !(next.scheme() == "https"
442                    && next
443                        .host_str()
444                        .is_some_and(|host| domain_listed(&domains, host)))
445            {
446                return attempt.error(format!(
447                    "redirected to {next}, outside the auto-approved hosts; call web_fetch with that URL to request approval"
448                ));
449            }
450            attempt.follow()
451        });
452        let client = reqwest::Client::builder()
453            .user_agent(USER_AGENT)
454            .timeout(self.config.fetch_timeout)
455            .connect_timeout(self.config.fetch_timeout.min(Duration::from_secs(10)))
456            .redirect(policy)
457            .referer(false)
458            .no_proxy()
459            .dns_resolver(Arc::new(CheckedResolver {
460                allowed: Arc::clone(&self.address_allowed),
461            }))
462            .build()
463            .map_err(|error| ToolError(format!("create HTTP client: {error}")))?;
464        let request = client.get(url.clone()).header(
465            reqwest::header::ACCEPT,
466            "text/html,application/xhtml+xml,text/plain;q=0.9,application/json;q=0.9,*/*;q=0.5",
467        );
468        let response = tokio::select! {
469            result = request.send() => result.map_err(|error| ToolError(format!("fetch {url}: {}", error_chain(&error))))?,
470            _ = context.cancellation.cancelled() => return Err(ToolError("web fetch cancelled".into())),
471        };
472        let status = response.status();
473        let final_url = response.url().clone();
474        let content_type = response
475            .headers()
476            .get(reqwest::header::CONTENT_TYPE)
477            .and_then(|value| value.to_str().ok())
478            .map(str::to_owned);
479        // Refuse a declared binary type before downloading it.
480        if content_type.is_some()
481            && let Err(message) = classify(content_type.as_deref(), &[])
482        {
483            return Ok(ToolOutput::failure(format!(
484                "URL: {final_url}\nStatus: {}\n{message}",
485                status.as_u16()
486            )));
487        }
488        let limit = self.config.fetch_max_bytes;
489        let mut stream = response.bytes_stream();
490        let mut bytes = Vec::with_capacity(limit.min(64 * 1024));
491        let mut download_truncated = false;
492        loop {
493            let chunk = tokio::select! {
494                chunk = stream.next() => chunk,
495                _ = context.cancellation.cancelled() => return Err(ToolError("web fetch cancelled".into())),
496            };
497            let Some(chunk) = chunk else { break };
498            let chunk = chunk
499                .map_err(|error| ToolError(format!("read {final_url}: {}", error_chain(&error))))?;
500            let remaining = limit - bytes.len();
501            if chunk.len() > remaining {
502                bytes.extend_from_slice(&chunk[..remaining]);
503                download_truncated = true;
504                break;
505            }
506            bytes.extend_from_slice(&chunk);
507        }
508        let kind = match classify(content_type.as_deref(), &bytes) {
509            Ok(kind) => kind,
510            Err(message) => {
511                return Ok(ToolOutput::failure(format!(
512                    "URL: {final_url}\nStatus: {}\n{message}",
513                    status.as_u16()
514                )));
515            }
516        };
517        let text = match kind {
518            Body::Text => String::from_utf8_lossy(&bytes).into_owned(),
519            Body::Html => tokio::task::spawn_blocking(move || {
520                html2text::from_read(bytes.as_slice(), HTML_WIDTH)
521                    .map_err(|error| ToolError(format!("convert HTML: {error}")))
522            })
523            .await
524            .map_err(|error| ToolError(format!("HTML conversion task failed: {error}")))??,
525        };
526        let offset = args.offset.unwrap_or(0);
527        let total = text.chars().count();
528        let header = format!(
529            "URL: {final_url}\nStatus: {}\nContent-Type: {}\nCharacters: {offset}-{{end}} of {total}{}\n\n",
530            status.as_u16(),
531            content_type.as_deref().unwrap_or("unknown"),
532            if download_truncated {
533                format!(" (download stopped at {limit} bytes)")
534            } else {
535                String::new()
536            }
537        );
538        let footer_reserve = 160;
539        let budget = self
540            .config
541            .output_limit
542            .saturating_sub(header.len() + footer_reserve)
543            .max(1);
544        let (body, next) = page(&text, offset, budget);
545        let end = offset + body.chars().count();
546        let mut content = header.replace("{end}", &end.to_string());
547        content.push_str(&body);
548        if let Some(next) = next {
549            content.push_str(&format!(
550                "\n\n[{} more characters; call web_fetch with offset={next} for the next part]",
551                total - next
552            ));
553        }
554        Ok(ToolOutput {
555            content,
556            is_error: status.is_client_error() || status.is_server_error(),
557            truncated: next.is_some() || download_truncated,
558        })
559    }
560}
561
562struct WebSearchTool {
563    backend: SearchBackend,
564    config: Arc<WebToolsConfig>,
565}
566
567#[derive(Deserialize)]
568#[serde(deny_unknown_fields)]
569struct SearchArgs {
570    query: String,
571    #[serde(default)]
572    count: Option<usize>,
573}
574
575impl WebSearchTool {
576    fn validate(&self, args: &SearchArgs) -> Result<(), ToolError> {
577        let query = args.query.trim();
578        if query.is_empty() {
579            return Err(ToolError("query must not be empty".into()));
580        }
581        if query.len() > MAX_QUERY_BYTES {
582            return Err(ToolError(format!("query exceeds {MAX_QUERY_BYTES} bytes")));
583        }
584        Ok(())
585    }
586
587    fn backend_name(&self) -> &'static str {
588        match self.backend {
589            SearchBackend::Searxng { .. } => "SearXNG",
590            SearchBackend::Brave { .. } => "Brave Search",
591        }
592    }
593}
594
595#[derive(Debug, PartialEq)]
596struct SearchResult {
597    title: String,
598    url: String,
599    snippet: String,
600}
601
602fn parse_results(backend: &SearchBackend, body: &Value) -> Result<Vec<SearchResult>, String> {
603    let (items, snippet_field) = match backend {
604        SearchBackend::Searxng { .. } => (body.get("results"), "content"),
605        SearchBackend::Brave { .. } => (
606            body.get("web").and_then(|web| web.get("results")),
607            "description",
608        ),
609    };
610    let Some(items) = items else {
611        return Ok(Vec::new());
612    };
613    let items = items
614        .as_array()
615        .ok_or_else(|| "search results are not a list".to_owned())?;
616    Ok(items
617        .iter()
618        .filter_map(|item| {
619            let url = item.get("url")?.as_str()?.to_owned();
620            let text = |field: &str| {
621                item.get(field)
622                    .and_then(Value::as_str)
623                    .map(strip_tags)
624                    .unwrap_or_default()
625            };
626            Some(SearchResult {
627                title: text("title"),
628                url,
629                snippet: text(snippet_field),
630            })
631        })
632        .collect())
633}
634
635/// Removes markup such as Brave's `<strong>` highlights from a snippet.
636fn strip_tags(value: &str) -> String {
637    let mut output = String::with_capacity(value.len());
638    let mut in_tag = false;
639    for character in value.chars() {
640        match character {
641            '<' => in_tag = true,
642            '>' if in_tag => in_tag = false,
643            _ if !in_tag => output.push(character),
644            _ => {}
645        }
646    }
647    output
648        .replace("&amp;", "&")
649        .replace("&lt;", "<")
650        .replace("&gt;", ">")
651        .replace("&quot;", "\"")
652        .replace("&#39;", "'")
653        .split_whitespace()
654        .collect::<Vec<_>>()
655        .join(" ")
656}
657
658fn format_results(query: &str, results: &[SearchResult], limit: usize) -> String {
659    if results.is_empty() {
660        return format!("No results for {query:?}.");
661    }
662    let mut output = format!("Results for {query:?}:\n");
663    for (index, result) in results.iter().enumerate() {
664        let entry = format!(
665            "\n{}. {}\n   {}\n   {}\n",
666            index + 1,
667            if result.title.is_empty() {
668                "(untitled)"
669            } else {
670                &result.title
671            },
672            result.url,
673            bounded(&result.snippet, 400)
674        );
675        if output.len() + entry.len() > limit {
676            break;
677        }
678        output.push_str(&entry);
679    }
680    output
681}
682
683#[async_trait]
684impl Tool for WebSearchTool {
685    fn spec(&self) -> ToolSpec {
686        ToolSpec {
687            name: "web_search".into(),
688            description: format!(
689                "Search the web with {} and return result titles, URLs, and snippets. Use it for \
690                 current facts, versions, documentation locations, and error messages, then open \
691                 the most relevant results with web_fetch.",
692                self.backend_name()
693            ),
694            parameters: json!({
695                "type":"object",
696                "properties":{
697                    "query":{"type":"string"},
698                    "count":{"type":"integer","minimum":1,"maximum":self.config.max_search_results,"description":"Number of results (default and maximum shown)"}
699                },
700                "required":["query"],
701                "additionalProperties":false
702            }),
703        }
704    }
705
706    // The query goes only to the search service the user configured, so it
707    // cannot carry data to a host the model chooses.
708    fn risk(&self, arguments: &Value) -> Result<ToolRisk, ToolError> {
709        let args: SearchArgs = parse_args(arguments)?;
710        self.validate(&args)?;
711        Ok(ToolRisk::ReadOnly)
712    }
713
714    fn approval_summary(&self, arguments: &Value) -> Result<String, ToolError> {
715        let args: SearchArgs = parse_args(arguments)?;
716        self.validate(&args)?;
717        Ok(format!(
718            "Search {} for {:?}",
719            self.backend_name(),
720            bounded(args.query.trim(), 500)
721        ))
722    }
723
724    async fn execute(
725        &self,
726        arguments: Value,
727        context: ToolContext,
728    ) -> Result<ToolOutput, ToolError> {
729        let args: SearchArgs = parse_args(&arguments)?;
730        self.validate(&args)?;
731        let query = args.query.trim().to_owned();
732        let count = args
733            .count
734            .unwrap_or(self.config.max_search_results)
735            .clamp(1, self.config.max_search_results);
736        let client = reqwest::Client::builder()
737            .user_agent(USER_AGENT)
738            .timeout(self.config.fetch_timeout)
739            .redirect(redirect::Policy::limited(3))
740            .build()
741            .map_err(|error| ToolError(format!("create HTTP client: {error}")))?;
742        let request = match &self.backend {
743            SearchBackend::Searxng { url } => {
744                let endpoint = format!("{}/search", url.trim_end_matches('/'));
745                client
746                    .get(endpoint)
747                    .query(&[("q", query.as_str()), ("format", "json")])
748            }
749            SearchBackend::Brave { url, api_key } => client
750                .get(url)
751                .query(&[("q", query.as_str()), ("count", &count.to_string())])
752                .header(reqwest::header::ACCEPT, "application/json")
753                .header("X-Subscription-Token", api_key),
754        };
755        let response = tokio::select! {
756            result = request.send() => result.map_err(|error| ToolError(format!("{} request failed: {}", self.backend_name(), error_chain(&error))))?,
757            _ = context.cancellation.cancelled() => return Err(ToolError("web search cancelled".into())),
758        };
759        let status = response.status();
760        let mut stream = response.bytes_stream();
761        let mut bytes = Vec::new();
762        loop {
763            let chunk = tokio::select! {
764                chunk = stream.next() => chunk,
765                _ = context.cancellation.cancelled() => return Err(ToolError("web search cancelled".into())),
766            };
767            let Some(chunk) = chunk else { break };
768            let chunk = chunk.map_err(|error| {
769                ToolError(format!(
770                    "{} response failed: {}",
771                    self.backend_name(),
772                    error_chain(&error)
773                ))
774            })?;
775            if bytes.len() + chunk.len() > MAX_SEARCH_RESPONSE_BYTES {
776                return Err(ToolError(format!(
777                    "{} response exceeded {MAX_SEARCH_RESPONSE_BYTES} bytes",
778                    self.backend_name()
779                )));
780            }
781            bytes.extend_from_slice(&chunk);
782        }
783        if !status.is_success() {
784            let body = String::from_utf8_lossy(&bytes[..bytes.len().min(300)]).into_owned();
785            let hint = match (&self.backend, status.as_u16()) {
786                (SearchBackend::Searxng { .. }, 403) => {
787                    " (enable the json format under search.formats in SearXNG's settings.yml)"
788                }
789                (SearchBackend::Brave { .. }, 401 | 403 | 422) => {
790                    " (check the Brave Search API key)"
791                }
792                _ => "",
793            };
794            return Ok(ToolOutput::failure(format!(
795                "{} returned HTTP {}{hint}: {}",
796                self.backend_name(),
797                status.as_u16(),
798                bounded(&body, 300)
799            )));
800        }
801        let body: Value = serde_json::from_slice(&bytes).map_err(|error| {
802            ToolError(format!(
803                "{} returned invalid JSON: {error}",
804                self.backend_name()
805            ))
806        })?;
807        let mut results = parse_results(&self.backend, &body).map_err(ToolError)?;
808        results.truncate(count);
809        Ok(ToolOutput::success(format_results(
810            &query,
811            &results,
812            self.config.output_limit,
813        )))
814    }
815}
816
817#[cfg(test)]
818mod tests {
819    use std::{
820        io::{Read, Write},
821        net::TcpListener,
822        thread,
823    };
824
825    use super::*;
826    use tokio_util::sync::CancellationToken;
827
828    fn config() -> WebToolsConfig {
829        WebToolsConfig {
830            fetch_max_bytes: 64 * 1024,
831            fetch_timeout: Duration::from_secs(5),
832            max_redirects: 3,
833            auto_approve_domains: vec!["docs.rs".into(), "*.example.org".into()],
834            allow_private_addresses: false,
835            search: None,
836            max_search_results: 5,
837            output_limit: 16 * 1024,
838        }
839    }
840
841    /// A fetch tool that may reach loopback only on the given ports, standing
842    /// in for "public" servers in tests.
843    fn fetch_tool(config: WebToolsConfig, ports: Vec<u16>) -> WebFetchTool {
844        WebFetchTool {
845            config: Arc::new(config),
846            address_allowed: Arc::new(move |address: SocketAddr| {
847                is_public(address.ip())
848                    || (address.ip().is_loopback() && ports.contains(&address.port()))
849            }),
850        }
851    }
852
853    fn context() -> ToolContext {
854        ToolContext {
855            workspace: std::env::temp_dir(),
856            cancellation: CancellationToken::new(),
857        }
858    }
859
860    /// Serves canned HTTP responses, one per connection, and returns the
861    /// request heads it received.
862    fn serve(responses: Vec<String>) -> (u16, thread::JoinHandle<Vec<String>>) {
863        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
864        let port = listener.local_addr().unwrap().port();
865        let handle = thread::spawn(move || {
866            let mut heads = Vec::new();
867            for response in responses {
868                let (mut stream, _) = listener.accept().unwrap();
869                let mut head = Vec::new();
870                let mut byte = [0u8; 1];
871                while !head.ends_with(b"\r\n\r\n") {
872                    if stream.read(&mut byte).unwrap() == 0 {
873                        break;
874                    }
875                    head.push(byte[0]);
876                }
877                heads.push(String::from_utf8_lossy(&head).into_owned());
878                let _ = stream.write_all(response.as_bytes());
879            }
880            heads
881        });
882        (port, handle)
883    }
884
885    fn http(status: &str, content_type: &str, body: &str, extra: &str) -> String {
886        format!(
887            "HTTP/1.1 {status}\r\nContent-Type: {content_type}\r\nContent-Length: {}\r\nConnection: close\r\n{extra}\r\n{body}",
888            body.len()
889        )
890    }
891
892    #[test]
893    fn public_address_classification() {
894        for private in [
895            "127.0.0.1",
896            "10.1.2.3",
897            "172.16.0.1",
898            "192.168.1.1",
899            "169.254.169.254",
900            "100.64.0.1",
901            "0.0.0.0",
902            "224.0.0.1",
903            "255.255.255.255",
904            "::1",
905            "::",
906            "fc00::1",
907            "fd12::1",
908            "fe80::1",
909            "::ffff:127.0.0.1",
910            "::ffff:10.0.0.1",
911            "64:ff9b::a00:1",
912            "2002:7f00:1::",
913            "2001:db8::1",
914        ] {
915            assert!(
916                !is_public(private.parse().unwrap()),
917                "{private} is not public"
918            );
919        }
920        for public in [
921            "1.1.1.1",
922            "140.82.112.3",
923            "2606:4700::1111",
924            "::ffff:8.8.8.8",
925            "64:ff9b::808:808",
926        ] {
927            assert!(is_public(public.parse().unwrap()), "{public} is public");
928        }
929    }
930
931    #[test]
932    fn allowlist_matches_hosts_and_wildcard_subdomains_only() {
933        let domains = vec!["docs.rs".to_owned(), "*.example.org".to_owned()];
934        assert!(domain_listed(&domains, "docs.rs"));
935        assert!(domain_listed(&domains, "DOCS.RS."));
936        assert!(!domain_listed(&domains, "evil-docs.rs"));
937        assert!(!domain_listed(&domains, "sub.docs.rs"));
938        assert!(domain_listed(&domains, "a.example.org"));
939        assert!(domain_listed(&domains, "a.b.example.org"));
940        assert!(!domain_listed(&domains, "example.org"));
941        assert!(!domain_listed(&domains, "badexample.org"));
942    }
943
944    #[test]
945    fn risk_is_read_only_only_for_https_allowlisted_hosts() {
946        let tool = fetch_tool(config(), Vec::new());
947        let risk = |url: &str| tool.risk(&json!({"url":url}));
948        assert_eq!(risk("https://docs.rs/serde").unwrap(), ToolRisk::ReadOnly);
949        assert_eq!(
950            risk("https://api.example.org/x").unwrap(),
951            ToolRisk::ReadOnly
952        );
953        assert_eq!(risk("http://docs.rs/serde").unwrap(), ToolRisk::Network);
954        assert_eq!(
955            risk("https://attacker.test/?q=secret").unwrap(),
956            ToolRisk::Network
957        );
958        for invalid in [
959            "ftp://docs.rs/x",
960            "file:///etc/passwd",
961            "https://user:pass@docs.rs/",
962            "not a url",
963        ] {
964            assert!(risk(invalid).is_err(), "{invalid} should be refused");
965        }
966        let summary = tool
967            .approval_summary(&json!({"url":"https://attacker.test/?q=1"}))
968            .unwrap();
969        assert!(summary.contains("attacker.test"), "{summary}");
970    }
971
972    #[test]
973    fn resolved_names_are_refused_when_any_address_is_not_public() {
974        let allowed: AddressCheck = Arc::new(|address: SocketAddr| is_public(address.ip()));
975        let public: SocketAddr = "1.1.1.1:0".parse().unwrap();
976        let private: SocketAddr = "10.0.0.1:0".parse().unwrap();
977        assert!(check_resolved("ok.test", &[public], &allowed).is_ok());
978        let error = check_resolved("mixed.test", &[public, private], &allowed).unwrap_err();
979        assert!(error.contains("10.0.0.1"), "{error}");
980        assert!(check_resolved("none.test", &[], &allowed).is_err());
981    }
982
983    #[tokio::test]
984    async fn html_is_converted_to_text_and_json_passes_through() {
985        let (port, server) = serve(vec![
986            http(
987                "200 OK",
988                "text/html; charset=utf-8",
989                "<html><head><title>T</title><script>var hidden=1;</script></head><body><h1>Serde</h1><p>Version <b>1.0.228</b> <a href=\"https://docs.rs/serde\">docs</a></p></body></html>",
990                "",
991            ),
992            http(
993                "200 OK",
994                "application/json",
995                r#"{"crate":{"max_version":"1.0.228"}}"#,
996                "",
997            ),
998        ]);
999        let tool = fetch_tool(config(), vec![port]);
1000        let html = tool
1001            .execute(
1002                json!({"url":format!("http://127.0.0.1:{port}/page")}),
1003                context(),
1004            )
1005            .await
1006            .unwrap();
1007        assert!(!html.is_error, "{}", html.content);
1008        assert!(html.content.contains("Serde"), "{}", html.content);
1009        assert!(html.content.contains("1.0.228"), "{}", html.content);
1010        assert!(
1011            html.content.contains("https://docs.rs/serde"),
1012            "{}",
1013            html.content
1014        );
1015        assert!(!html.content.contains("<b>"), "{}", html.content);
1016        let json = tool
1017            .execute(
1018                json!({"url":format!("http://127.0.0.1:{port}/api")}),
1019                context(),
1020            )
1021            .await
1022            .unwrap();
1023        assert!(
1024            json.content
1025                .contains(r#"{"crate":{"max_version":"1.0.228"}}"#),
1026            "{}",
1027            json.content
1028        );
1029        let heads = server.join().unwrap();
1030        assert!(
1031            heads[0].to_ascii_lowercase().contains("user-agent: scv/"),
1032            "{}",
1033            heads[0]
1034        );
1035        assert!(
1036            !heads[0].to_ascii_lowercase().contains("cookie"),
1037            "{}",
1038            heads[0]
1039        );
1040    }
1041
1042    #[tokio::test]
1043    async fn binary_content_is_refused_and_large_bodies_are_bounded_and_paged() {
1044        let long = "x".repeat(10_000);
1045        let (port, server) = serve(vec![
1046            http("200 OK", "image/png", "\u{89}PNG", ""),
1047            http("200 OK", "text/plain", &long, ""),
1048            http("200 OK", "text/plain", &long, ""),
1049        ]);
1050        let mut small = config();
1051        small.fetch_max_bytes = 4000;
1052        small.output_limit = 1500;
1053        let tool = fetch_tool(small, vec![port]);
1054        let binary = tool
1055            .execute(
1056                json!({"url":format!("http://127.0.0.1:{port}/a.png")}),
1057                context(),
1058            )
1059            .await
1060            .unwrap();
1061        assert!(binary.is_error);
1062        assert!(binary.content.contains("image/png"), "{}", binary.content);
1063        let first = tool
1064            .execute(
1065                json!({"url":format!("http://127.0.0.1:{port}/big")}),
1066                context(),
1067            )
1068            .await
1069            .unwrap();
1070        assert!(first.truncated);
1071        assert!(first.content.len() <= 1500, "{}", first.content.len());
1072        assert!(
1073            first.content.contains("download stopped at 4000 bytes"),
1074            "{}",
1075            first.content
1076        );
1077        let next: usize = first
1078            .content
1079            .rsplit("offset=")
1080            .next()
1081            .and_then(|rest| rest.split(' ').next())
1082            .unwrap()
1083            .parse()
1084            .unwrap();
1085        let second = tool
1086            .execute(
1087                json!({"url":format!("http://127.0.0.1:{port}/big"),"offset":next}),
1088                context(),
1089            )
1090            .await
1091            .unwrap();
1092        assert!(
1093            second.content.contains(&format!("Characters: {next}-")),
1094            "{}",
1095            second.content
1096        );
1097        server.join().unwrap();
1098    }
1099
1100    #[tokio::test]
1101    async fn loopback_and_private_targets_are_refused_directly_by_name_and_by_redirect() {
1102        let (blocked_port, _unused) = serve(Vec::new());
1103        let (port, server) = serve(vec![http(
1104            "302 Found",
1105            "text/plain",
1106            "",
1107            &format!("Location: http://127.0.0.1:{blocked_port}/admin\r\n"),
1108        )]);
1109        let tool = fetch_tool(config(), vec![port]);
1110        let direct = tool
1111            .execute(
1112                json!({"url":format!("http://127.0.0.1:{blocked_port}/")}),
1113                context(),
1114            )
1115            .await
1116            .unwrap_err();
1117        assert!(direct.0.contains("non-public"), "{}", direct.0);
1118        let metadata = tool
1119            .execute(
1120                json!({"url":"http://169.254.169.254/latest/meta-data/"}),
1121                context(),
1122            )
1123            .await
1124            .unwrap_err();
1125        assert!(metadata.0.contains("non-public"), "{}", metadata.0);
1126        let named = tool
1127            .execute(
1128                json!({"url":format!("http://localhost:{port}/")}),
1129                context(),
1130            )
1131            .await
1132            .unwrap_err();
1133        assert!(named.0.contains("non-public"), "{}", named.0);
1134        let redirected = tool
1135            .execute(
1136                json!({"url":format!("http://127.0.0.1:{port}/go")}),
1137                context(),
1138            )
1139            .await
1140            .unwrap_err();
1141        assert!(redirected.0.contains("non-public"), "{}", redirected.0);
1142        server.join().unwrap();
1143    }
1144
1145    #[tokio::test]
1146    async fn redirects_are_limited() {
1147        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1148        let port = listener.local_addr().unwrap().port();
1149        let looping = thread::spawn(move || {
1150            // The first request plus `max_redirects` (3) followed hops.
1151            for _ in 0..4 {
1152                let (mut stream, _) = listener.accept().unwrap();
1153                let mut buffer = [0u8; 1024];
1154                let _ = stream.read(&mut buffer);
1155                let _ = stream.write_all(
1156                    http(
1157                        "302 Found",
1158                        "text/plain",
1159                        "",
1160                        &format!("Location: http://127.0.0.1:{port}/again\r\n"),
1161                    )
1162                    .as_bytes(),
1163                );
1164            }
1165        });
1166        let tool = fetch_tool(config(), vec![port]);
1167        let error = tool
1168            .execute(
1169                json!({"url":format!("http://127.0.0.1:{port}/start")}),
1170                context(),
1171            )
1172            .await
1173            .unwrap_err();
1174        assert!(error.0.contains("redirects"), "{}", error.0);
1175        looping.join().unwrap();
1176    }
1177
1178    #[test]
1179    fn content_types_are_classified() {
1180        assert!(matches!(
1181            classify(Some("text/html; charset=utf-8"), b""),
1182            Ok(Body::Html)
1183        ));
1184        assert!(matches!(
1185            classify(Some("application/vnd.api+json"), b""),
1186            Ok(Body::Text)
1187        ));
1188        assert!(matches!(
1189            classify(Some("text/markdown"), b""),
1190            Ok(Body::Text)
1191        ));
1192        assert!(classify(Some("application/pdf"), b"").is_err());
1193        assert!(classify(Some("application/octet-stream"), b"").is_err());
1194        assert!(matches!(
1195            classify(None, b"<!DOCTYPE html><html>"),
1196            Ok(Body::Html)
1197        ));
1198        assert!(matches!(classify(None, b"plain words"), Ok(Body::Text)));
1199        assert!(classify(None, &[0xff, 0xfe, 0x00, 0x81, 0x90, 0xff, 0xfe, 0x00]).is_err());
1200    }
1201
1202    #[test]
1203    fn search_results_parse_for_each_backend() {
1204        let searxng = SearchBackend::Searxng {
1205            url: "http://s".into(),
1206        };
1207        let brave = SearchBackend::Brave {
1208            url: "http://b".into(),
1209            api_key: "brave-secret".into(),
1210        };
1211        assert_eq!(
1212            parse_results(
1213                &searxng,
1214                &json!({"results":[{"title":"Serde","url":"https://serde.rs","content":"A <b>framework</b>"},{"title":"no url"}]})
1215            )
1216            .unwrap(),
1217            vec![SearchResult {
1218                title: "Serde".into(),
1219                url: "https://serde.rs".into(),
1220                snippet: "A framework".into()
1221            }]
1222        );
1223        assert_eq!(
1224            parse_results(
1225                &brave,
1226                &json!({"web":{"results":[{"title":"Tokio &amp; async","url":"https://tokio.rs","description":"<strong>Tokio</strong> runtime"}]}})
1227            )
1228            .unwrap()[0],
1229            SearchResult {
1230                title: "Tokio & async".into(),
1231                url: "https://tokio.rs".into(),
1232                snippet: "Tokio runtime".into()
1233            }
1234        );
1235        assert!(parse_results(&brave, &json!({})).unwrap().is_empty());
1236        assert!(!format!("{brave:?}").contains("brave-secret"));
1237    }
1238
1239    #[tokio::test]
1240    async fn search_backends_are_queried_with_their_parameters() {
1241        let (port, server) = serve(vec![
1242            http(
1243                "200 OK",
1244                "application/json",
1245                r#"{"results":[{"title":"Serde","url":"https://serde.rs","content":"Serialization"}]}"#,
1246                "",
1247            ),
1248            http(
1249                "200 OK",
1250                "application/json",
1251                r#"{"web":{"results":[{"title":"Tokio","url":"https://tokio.rs","description":"Runtime"},{"title":"Two","url":"https://two.test","description":"x"}]}}"#,
1252                "",
1253            ),
1254            http("403 Forbidden", "text/plain", "forbidden", ""),
1255        ]);
1256        let mut searx = config();
1257        searx.search = Some(SearchBackend::Searxng {
1258            url: format!("http://127.0.0.1:{port}/"),
1259        });
1260        let tool = WebSearchTool {
1261            backend: searx.search.clone().unwrap(),
1262            config: Arc::new(searx),
1263        };
1264        assert_eq!(
1265            tool.risk(&json!({"query":"serde"})).unwrap(),
1266            ToolRisk::ReadOnly
1267        );
1268        assert!(tool.risk(&json!({"query":"  "})).is_err());
1269        let output = tool
1270            .execute(json!({"query":"serde json"}), context())
1271            .await
1272            .unwrap();
1273        assert!(
1274            output.content.contains("1. Serde\n   https://serde.rs"),
1275            "{}",
1276            output.content
1277        );
1278
1279        let mut brave = config();
1280        brave.search = Some(SearchBackend::Brave {
1281            url: format!("http://127.0.0.1:{port}/res/v1/web/search"),
1282            api_key: "brave-test-key".into(),
1283        });
1284        let tool = WebSearchTool {
1285            backend: brave.search.clone().unwrap(),
1286            config: Arc::new(brave),
1287        };
1288        let output = tool
1289            .execute(json!({"query":"tokio","count":1}), context())
1290            .await
1291            .unwrap();
1292        assert!(output.content.contains("Tokio"), "{}", output.content);
1293        assert!(!output.content.contains("two.test"), "{}", output.content);
1294        let failure = tool
1295            .execute(json!({"query":"tokio"}), context())
1296            .await
1297            .unwrap();
1298        assert!(failure.is_error);
1299        assert!(failure.content.contains("API key"), "{}", failure.content);
1300
1301        let heads = server.join().unwrap();
1302        assert!(
1303            heads[0].starts_with("GET /search?q=serde+json&format=json "),
1304            "{}",
1305            heads[0]
1306        );
1307        assert!(
1308            heads[1].starts_with("GET /res/v1/web/search?q=tokio&count=1 "),
1309            "{}",
1310            heads[1]
1311        );
1312        assert!(
1313            heads[1]
1314                .to_ascii_lowercase()
1315                .contains("x-subscription-token: brave-test-key"),
1316            "{}",
1317            heads[1]
1318        );
1319    }
1320
1321    #[test]
1322    fn registration_offers_search_only_with_a_backend() {
1323        let mut registry = ToolRegistry::default();
1324        register(&mut registry, config()).unwrap();
1325        let names: Vec<_> = registry.specs().into_iter().map(|spec| spec.name).collect();
1326        assert_eq!(names, vec!["web_fetch"]);
1327        let mut with_search = config();
1328        with_search.search = Some(SearchBackend::Searxng {
1329            url: "http://s".into(),
1330        });
1331        let mut registry = ToolRegistry::default();
1332        register(&mut registry, with_search).unwrap();
1333        let mut names: Vec<_> = registry.specs().into_iter().map(|spec| spec.name).collect();
1334        names.sort();
1335        assert_eq!(names, vec!["web_fetch", "web_search"]);
1336    }
1337}