1use 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#[derive(Debug, Clone)]
41pub struct WebToolsConfig {
42 pub fetch_max_bytes: usize,
43 pub fetch_timeout: Duration,
44 pub max_redirects: usize,
45 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#[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
75pub 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
94type 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
130fn 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
149fn 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
170pub 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
184pub 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 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 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 || (segments[0] & 0xffc0) == 0xfe80 || (segments[0] & 0xffc0) == 0xfec0 || (segments[0] == 0x2001 && segments[1] == 0x0db8) || (segments[0] == 0x2001 && segments[1] == 0)) }
240
241struct 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
267fn 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
290fn 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
340fn 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 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
635fn 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("&", "&")
649 .replace("<", "<")
650 .replace(">", ">")
651 .replace(""", "\"")
652 .replace("'", "'")
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 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 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 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 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 & 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}