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::new(std::env::temp_dir(), CancellationToken::new())
855    }
856
857    /// Serves canned HTTP responses, one per connection, and returns the
858    /// request heads it received.
859    fn serve(responses: Vec<String>) -> (u16, thread::JoinHandle<Vec<String>>) {
860        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
861        let port = listener.local_addr().unwrap().port();
862        let handle = thread::spawn(move || {
863            let mut heads = Vec::new();
864            for response in responses {
865                let (mut stream, _) = listener.accept().unwrap();
866                let mut head = Vec::new();
867                let mut byte = [0u8; 1];
868                while !head.ends_with(b"\r\n\r\n") {
869                    if stream.read(&mut byte).unwrap() == 0 {
870                        break;
871                    }
872                    head.push(byte[0]);
873                }
874                heads.push(String::from_utf8_lossy(&head).into_owned());
875                let _ = stream.write_all(response.as_bytes());
876            }
877            heads
878        });
879        (port, handle)
880    }
881
882    fn http(status: &str, content_type: &str, body: &str, extra: &str) -> String {
883        format!(
884            "HTTP/1.1 {status}\r\nContent-Type: {content_type}\r\nContent-Length: {}\r\nConnection: close\r\n{extra}\r\n{body}",
885            body.len()
886        )
887    }
888
889    #[test]
890    fn public_address_classification() {
891        for private in [
892            "127.0.0.1",
893            "10.1.2.3",
894            "172.16.0.1",
895            "192.168.1.1",
896            "169.254.169.254",
897            "100.64.0.1",
898            "0.0.0.0",
899            "224.0.0.1",
900            "255.255.255.255",
901            "::1",
902            "::",
903            "fc00::1",
904            "fd12::1",
905            "fe80::1",
906            "::ffff:127.0.0.1",
907            "::ffff:10.0.0.1",
908            "64:ff9b::a00:1",
909            "2002:7f00:1::",
910            "2001:db8::1",
911        ] {
912            assert!(
913                !is_public(private.parse().unwrap()),
914                "{private} is not public"
915            );
916        }
917        for public in [
918            "1.1.1.1",
919            "140.82.112.3",
920            "2606:4700::1111",
921            "::ffff:8.8.8.8",
922            "64:ff9b::808:808",
923        ] {
924            assert!(is_public(public.parse().unwrap()), "{public} is public");
925        }
926    }
927
928    #[test]
929    fn allowlist_matches_hosts_and_wildcard_subdomains_only() {
930        let domains = vec!["docs.rs".to_owned(), "*.example.org".to_owned()];
931        assert!(domain_listed(&domains, "docs.rs"));
932        assert!(domain_listed(&domains, "DOCS.RS."));
933        assert!(!domain_listed(&domains, "evil-docs.rs"));
934        assert!(!domain_listed(&domains, "sub.docs.rs"));
935        assert!(domain_listed(&domains, "a.example.org"));
936        assert!(domain_listed(&domains, "a.b.example.org"));
937        assert!(!domain_listed(&domains, "example.org"));
938        assert!(!domain_listed(&domains, "badexample.org"));
939    }
940
941    #[test]
942    fn risk_is_read_only_only_for_https_allowlisted_hosts() {
943        let tool = fetch_tool(config(), Vec::new());
944        let risk = |url: &str| tool.risk(&json!({"url":url}));
945        assert_eq!(risk("https://docs.rs/serde").unwrap(), ToolRisk::ReadOnly);
946        assert_eq!(
947            risk("https://api.example.org/x").unwrap(),
948            ToolRisk::ReadOnly
949        );
950        assert_eq!(risk("http://docs.rs/serde").unwrap(), ToolRisk::Network);
951        assert_eq!(
952            risk("https://attacker.test/?q=secret").unwrap(),
953            ToolRisk::Network
954        );
955        for invalid in [
956            "ftp://docs.rs/x",
957            "file:///etc/passwd",
958            "https://user:pass@docs.rs/",
959            "not a url",
960        ] {
961            assert!(risk(invalid).is_err(), "{invalid} should be refused");
962        }
963        let summary = tool
964            .approval_summary(&json!({"url":"https://attacker.test/?q=1"}))
965            .unwrap();
966        assert!(summary.contains("attacker.test"), "{summary}");
967    }
968
969    #[test]
970    fn resolved_names_are_refused_when_any_address_is_not_public() {
971        let allowed: AddressCheck = Arc::new(|address: SocketAddr| is_public(address.ip()));
972        let public: SocketAddr = "1.1.1.1:0".parse().unwrap();
973        let private: SocketAddr = "10.0.0.1:0".parse().unwrap();
974        assert!(check_resolved("ok.test", &[public], &allowed).is_ok());
975        let error = check_resolved("mixed.test", &[public, private], &allowed).unwrap_err();
976        assert!(error.contains("10.0.0.1"), "{error}");
977        assert!(check_resolved("none.test", &[], &allowed).is_err());
978    }
979
980    #[tokio::test]
981    async fn html_is_converted_to_text_and_json_passes_through() {
982        let (port, server) = serve(vec![
983            http(
984                "200 OK",
985                "text/html; charset=utf-8",
986                "<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>",
987                "",
988            ),
989            http(
990                "200 OK",
991                "application/json",
992                r#"{"crate":{"max_version":"1.0.228"}}"#,
993                "",
994            ),
995        ]);
996        let tool = fetch_tool(config(), vec![port]);
997        let html = tool
998            .execute(
999                json!({"url":format!("http://127.0.0.1:{port}/page")}),
1000                context(),
1001            )
1002            .await
1003            .unwrap();
1004        assert!(!html.is_error, "{}", html.content);
1005        assert!(html.content.contains("Serde"), "{}", html.content);
1006        assert!(html.content.contains("1.0.228"), "{}", html.content);
1007        assert!(
1008            html.content.contains("https://docs.rs/serde"),
1009            "{}",
1010            html.content
1011        );
1012        assert!(!html.content.contains("<b>"), "{}", html.content);
1013        let json = tool
1014            .execute(
1015                json!({"url":format!("http://127.0.0.1:{port}/api")}),
1016                context(),
1017            )
1018            .await
1019            .unwrap();
1020        assert!(
1021            json.content
1022                .contains(r#"{"crate":{"max_version":"1.0.228"}}"#),
1023            "{}",
1024            json.content
1025        );
1026        let heads = server.join().unwrap();
1027        assert!(
1028            heads[0].to_ascii_lowercase().contains("user-agent: scv/"),
1029            "{}",
1030            heads[0]
1031        );
1032        assert!(
1033            !heads[0].to_ascii_lowercase().contains("cookie"),
1034            "{}",
1035            heads[0]
1036        );
1037    }
1038
1039    #[tokio::test]
1040    async fn binary_content_is_refused_and_large_bodies_are_bounded_and_paged() {
1041        let long = "x".repeat(10_000);
1042        let (port, server) = serve(vec![
1043            http("200 OK", "image/png", "\u{89}PNG", ""),
1044            http("200 OK", "text/plain", &long, ""),
1045            http("200 OK", "text/plain", &long, ""),
1046        ]);
1047        let mut small = config();
1048        small.fetch_max_bytes = 4000;
1049        small.output_limit = 1500;
1050        let tool = fetch_tool(small, vec![port]);
1051        let binary = tool
1052            .execute(
1053                json!({"url":format!("http://127.0.0.1:{port}/a.png")}),
1054                context(),
1055            )
1056            .await
1057            .unwrap();
1058        assert!(binary.is_error);
1059        assert!(binary.content.contains("image/png"), "{}", binary.content);
1060        let first = tool
1061            .execute(
1062                json!({"url":format!("http://127.0.0.1:{port}/big")}),
1063                context(),
1064            )
1065            .await
1066            .unwrap();
1067        assert!(first.truncated);
1068        assert!(first.content.len() <= 1500, "{}", first.content.len());
1069        assert!(
1070            first.content.contains("download stopped at 4000 bytes"),
1071            "{}",
1072            first.content
1073        );
1074        let next: usize = first
1075            .content
1076            .rsplit("offset=")
1077            .next()
1078            .and_then(|rest| rest.split(' ').next())
1079            .unwrap()
1080            .parse()
1081            .unwrap();
1082        let second = tool
1083            .execute(
1084                json!({"url":format!("http://127.0.0.1:{port}/big"),"offset":next}),
1085                context(),
1086            )
1087            .await
1088            .unwrap();
1089        assert!(
1090            second.content.contains(&format!("Characters: {next}-")),
1091            "{}",
1092            second.content
1093        );
1094        server.join().unwrap();
1095    }
1096
1097    #[tokio::test]
1098    async fn loopback_and_private_targets_are_refused_directly_by_name_and_by_redirect() {
1099        let (blocked_port, _unused) = serve(Vec::new());
1100        let (port, server) = serve(vec![http(
1101            "302 Found",
1102            "text/plain",
1103            "",
1104            &format!("Location: http://127.0.0.1:{blocked_port}/admin\r\n"),
1105        )]);
1106        let tool = fetch_tool(config(), vec![port]);
1107        let direct = tool
1108            .execute(
1109                json!({"url":format!("http://127.0.0.1:{blocked_port}/")}),
1110                context(),
1111            )
1112            .await
1113            .unwrap_err();
1114        assert!(direct.0.contains("non-public"), "{}", direct.0);
1115        let metadata = tool
1116            .execute(
1117                json!({"url":"http://169.254.169.254/latest/meta-data/"}),
1118                context(),
1119            )
1120            .await
1121            .unwrap_err();
1122        assert!(metadata.0.contains("non-public"), "{}", metadata.0);
1123        let named = tool
1124            .execute(
1125                json!({"url":format!("http://localhost:{port}/")}),
1126                context(),
1127            )
1128            .await
1129            .unwrap_err();
1130        assert!(named.0.contains("non-public"), "{}", named.0);
1131        let redirected = tool
1132            .execute(
1133                json!({"url":format!("http://127.0.0.1:{port}/go")}),
1134                context(),
1135            )
1136            .await
1137            .unwrap_err();
1138        assert!(redirected.0.contains("non-public"), "{}", redirected.0);
1139        server.join().unwrap();
1140    }
1141
1142    #[tokio::test]
1143    async fn redirects_are_limited() {
1144        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1145        let port = listener.local_addr().unwrap().port();
1146        let looping = thread::spawn(move || {
1147            // The first request plus `max_redirects` (3) followed hops.
1148            for _ in 0..4 {
1149                let (mut stream, _) = listener.accept().unwrap();
1150                let mut buffer = [0u8; 1024];
1151                let _ = stream.read(&mut buffer);
1152                let _ = stream.write_all(
1153                    http(
1154                        "302 Found",
1155                        "text/plain",
1156                        "",
1157                        &format!("Location: http://127.0.0.1:{port}/again\r\n"),
1158                    )
1159                    .as_bytes(),
1160                );
1161            }
1162        });
1163        let tool = fetch_tool(config(), vec![port]);
1164        let error = tool
1165            .execute(
1166                json!({"url":format!("http://127.0.0.1:{port}/start")}),
1167                context(),
1168            )
1169            .await
1170            .unwrap_err();
1171        assert!(error.0.contains("redirects"), "{}", error.0);
1172        looping.join().unwrap();
1173    }
1174
1175    #[test]
1176    fn content_types_are_classified() {
1177        assert!(matches!(
1178            classify(Some("text/html; charset=utf-8"), b""),
1179            Ok(Body::Html)
1180        ));
1181        assert!(matches!(
1182            classify(Some("application/vnd.api+json"), b""),
1183            Ok(Body::Text)
1184        ));
1185        assert!(matches!(
1186            classify(Some("text/markdown"), b""),
1187            Ok(Body::Text)
1188        ));
1189        assert!(classify(Some("application/pdf"), b"").is_err());
1190        assert!(classify(Some("application/octet-stream"), b"").is_err());
1191        assert!(matches!(
1192            classify(None, b"<!DOCTYPE html><html>"),
1193            Ok(Body::Html)
1194        ));
1195        assert!(matches!(classify(None, b"plain words"), Ok(Body::Text)));
1196        assert!(classify(None, &[0xff, 0xfe, 0x00, 0x81, 0x90, 0xff, 0xfe, 0x00]).is_err());
1197    }
1198
1199    #[test]
1200    fn search_results_parse_for_each_backend() {
1201        let searxng = SearchBackend::Searxng {
1202            url: "http://s".into(),
1203        };
1204        let brave = SearchBackend::Brave {
1205            url: "http://b".into(),
1206            api_key: "brave-secret".into(),
1207        };
1208        assert_eq!(
1209            parse_results(
1210                &searxng,
1211                &json!({"results":[{"title":"Serde","url":"https://serde.rs","content":"A <b>framework</b>"},{"title":"no url"}]})
1212            )
1213            .unwrap(),
1214            vec![SearchResult {
1215                title: "Serde".into(),
1216                url: "https://serde.rs".into(),
1217                snippet: "A framework".into()
1218            }]
1219        );
1220        assert_eq!(
1221            parse_results(
1222                &brave,
1223                &json!({"web":{"results":[{"title":"Tokio &amp; async","url":"https://tokio.rs","description":"<strong>Tokio</strong> runtime"}]}})
1224            )
1225            .unwrap()[0],
1226            SearchResult {
1227                title: "Tokio & async".into(),
1228                url: "https://tokio.rs".into(),
1229                snippet: "Tokio runtime".into()
1230            }
1231        );
1232        assert!(parse_results(&brave, &json!({})).unwrap().is_empty());
1233        assert!(!format!("{brave:?}").contains("brave-secret"));
1234    }
1235
1236    #[tokio::test]
1237    async fn search_backends_are_queried_with_their_parameters() {
1238        let (port, server) = serve(vec![
1239            http(
1240                "200 OK",
1241                "application/json",
1242                r#"{"results":[{"title":"Serde","url":"https://serde.rs","content":"Serialization"}]}"#,
1243                "",
1244            ),
1245            http(
1246                "200 OK",
1247                "application/json",
1248                r#"{"web":{"results":[{"title":"Tokio","url":"https://tokio.rs","description":"Runtime"},{"title":"Two","url":"https://two.test","description":"x"}]}}"#,
1249                "",
1250            ),
1251            http("403 Forbidden", "text/plain", "forbidden", ""),
1252        ]);
1253        let mut searx = config();
1254        searx.search = Some(SearchBackend::Searxng {
1255            url: format!("http://127.0.0.1:{port}/"),
1256        });
1257        let tool = WebSearchTool {
1258            backend: searx.search.clone().unwrap(),
1259            config: Arc::new(searx),
1260        };
1261        assert_eq!(
1262            tool.risk(&json!({"query":"serde"})).unwrap(),
1263            ToolRisk::ReadOnly
1264        );
1265        assert!(tool.risk(&json!({"query":"  "})).is_err());
1266        let output = tool
1267            .execute(json!({"query":"serde json"}), context())
1268            .await
1269            .unwrap();
1270        assert!(
1271            output.content.contains("1. Serde\n   https://serde.rs"),
1272            "{}",
1273            output.content
1274        );
1275
1276        let mut brave = config();
1277        brave.search = Some(SearchBackend::Brave {
1278            url: format!("http://127.0.0.1:{port}/res/v1/web/search"),
1279            api_key: "brave-test-key".into(),
1280        });
1281        let tool = WebSearchTool {
1282            backend: brave.search.clone().unwrap(),
1283            config: Arc::new(brave),
1284        };
1285        let output = tool
1286            .execute(json!({"query":"tokio","count":1}), context())
1287            .await
1288            .unwrap();
1289        assert!(output.content.contains("Tokio"), "{}", output.content);
1290        assert!(!output.content.contains("two.test"), "{}", output.content);
1291        let failure = tool
1292            .execute(json!({"query":"tokio"}), context())
1293            .await
1294            .unwrap();
1295        assert!(failure.is_error);
1296        assert!(failure.content.contains("API key"), "{}", failure.content);
1297
1298        let heads = server.join().unwrap();
1299        assert!(
1300            heads[0].starts_with("GET /search?q=serde+json&format=json "),
1301            "{}",
1302            heads[0]
1303        );
1304        assert!(
1305            heads[1].starts_with("GET /res/v1/web/search?q=tokio&count=1 "),
1306            "{}",
1307            heads[1]
1308        );
1309        assert!(
1310            heads[1]
1311                .to_ascii_lowercase()
1312                .contains("x-subscription-token: brave-test-key"),
1313            "{}",
1314            heads[1]
1315        );
1316    }
1317
1318    #[test]
1319    fn registration_offers_search_only_with_a_backend() {
1320        let mut registry = ToolRegistry::default();
1321        register(&mut registry, config()).unwrap();
1322        let names: Vec<_> = registry.specs().into_iter().map(|spec| spec.name).collect();
1323        assert_eq!(names, vec!["web_fetch"]);
1324        let mut with_search = config();
1325        with_search.search = Some(SearchBackend::Searxng {
1326            url: "http://s".into(),
1327        });
1328        let mut registry = ToolRegistry::default();
1329        register(&mut registry, with_search).unwrap();
1330        let mut names: Vec<_> = registry.specs().into_iter().map(|spec| spec.name).collect();
1331        names.sort();
1332        assert_eq!(names, vec!["web_fetch", "web_search"]);
1333    }
1334}