1use async_trait::async_trait;
2use futures::StreamExt;
3use parking_lot::RwLock;
4use regex::Regex;
5use schemars::JsonSchema;
6use serde::{Deserialize, Serialize};
7use serde_json::Value;
8use std::collections::HashMap;
9use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
10use std::sync::Arc;
11use std::time::{Duration, Instant};
12
13use ai_agents_core::{
14 ChatMessage, DomainPolicyBinding, LLMConfig, LLMProvider, ResultLimitBinding, ResultLimitKind,
15 Tool, ToolApprovalRecord, ToolApprovalStatus, ToolCallClassification, ToolExecutionContext,
16 ToolOperationKind, ToolPolicyBindings, ToolResult, ToolSafetyMetadata, ToolSideEffectLevel,
17};
18
19use crate::generate_schema;
20
21const DEFAULT_MAX_RESPONSE_BYTES: usize = 1_048_576;
22const DEFAULT_MAX_OUTPUT_CHARS: usize = 20_000;
23const DEFAULT_CACHE_TTL_SECONDS: u64 = 900;
24const DEFAULT_MAX_REDIRECTS: usize = 5;
25const DEFAULT_TIMEOUT_MS: u64 = 15_000;
26const MAX_CACHE_ENTRIES: usize = 128;
27
28#[derive(Debug, Clone, PartialEq, Eq)]
30pub struct WebFetchTransportRequest {
31 pub url: String,
32 pub max_response_bytes: usize,
33}
34
35#[derive(Debug, Clone, PartialEq, Eq)]
37pub struct WebFetchTransportResponse {
38 pub status: u16,
39 pub content_type: Option<String>,
40 pub location: Option<String>,
41 pub body: Vec<u8>,
42}
43
44#[async_trait]
48pub trait WebFetchTransport: Send + Sync {
49 async fn send(
50 &self,
51 request: WebFetchTransportRequest,
52 ) -> Result<WebFetchTransportResponse, String>;
53
54 async fn send_validated(
57 &self,
58 request: WebFetchTransportRequest,
59 _addresses: &[SocketAddr],
60 ) -> Result<WebFetchTransportResponse, String> {
61 self.send(request).await
62 }
63}
64
65#[async_trait]
67pub trait WebFetchResolver: Send + Sync {
68 async fn resolve(&self, host: &str, port: u16) -> Result<Vec<IpAddr>, String>;
69}
70
71struct ReqwestWebFetchTransport;
72
73#[async_trait]
74impl WebFetchTransport for ReqwestWebFetchTransport {
75 async fn send(
76 &self,
77 _request: WebFetchTransportRequest,
78 ) -> Result<WebFetchTransportResponse, String> {
79 Err(
80 "Validated network addresses are required for the default web fetch transport"
81 .to_string(),
82 )
83 }
84
85 async fn send_validated(
86 &self,
87 request: WebFetchTransportRequest,
88 addresses: &[SocketAddr],
89 ) -> Result<WebFetchTransportResponse, String> {
90 let url = reqwest::Url::parse(&request.url)
91 .map_err(|error| format!("Request URL is invalid: {}", error))?;
92 let host = url
93 .host_str()
94 .ok_or_else(|| "Request URL host is required".to_string())?;
95 let port = url.port_or_known_default().unwrap_or(443);
96 if addresses.is_empty() || addresses.iter().any(|address| address.port() != port) {
97 return Err("Validated network addresses do not match the request port".to_string());
98 }
99 if let Ok(ip) = host.parse::<IpAddr>()
100 && addresses.iter().any(|address| address.ip() != ip)
101 {
102 return Err("Validated network addresses do not match the request host".to_string());
103 }
104
105 let mut client = reqwest::Client::builder()
108 .redirect(reqwest::redirect::Policy::none())
109 .timeout(Duration::from_millis(DEFAULT_TIMEOUT_MS))
110 .no_proxy();
111 if host.parse::<IpAddr>().is_err() {
112 client = client.resolve_to_addrs(host, addresses);
113 }
114 let client = client
115 .build()
116 .map_err(|error| format!("Request client construction failed: {}", error))?;
117 let response = client
118 .get(url)
119 .send()
120 .await
121 .map_err(|error| format!("Request failed: {}", error))?;
122 let status = response.status().as_u16();
123 let content_type = response
124 .headers()
125 .get(reqwest::header::CONTENT_TYPE)
126 .and_then(|value| value.to_str().ok())
127 .map(str::to_string);
128 let location = response
129 .headers()
130 .get(reqwest::header::LOCATION)
131 .map(|value| {
132 value
133 .to_str()
134 .map(str::to_string)
135 .map_err(|_| "Redirect Location header is not valid UTF-8".to_string())
136 })
137 .transpose()?;
138 let mut body = Vec::new();
139 let mut stream = response.bytes_stream();
140 while let Some(chunk) = stream.next().await {
141 let chunk = chunk.map_err(|error| format!("Response stream failed: {}", error))?;
142 if body.len().saturating_add(chunk.len()) > request.max_response_bytes {
143 return Err(format!(
144 "Response exceeded max_response_bytes {}",
145 request.max_response_bytes
146 ));
147 }
148 body.extend_from_slice(&chunk);
149 }
150 Ok(WebFetchTransportResponse {
151 status,
152 content_type,
153 location,
154 body,
155 })
156 }
157}
158
159struct TokioWebFetchResolver;
160
161#[async_trait]
162impl WebFetchResolver for TokioWebFetchResolver {
163 async fn resolve(&self, host: &str, port: u16) -> Result<Vec<IpAddr>, String> {
164 tokio::net::lookup_host((host, port))
165 .await
166 .map(|addresses| addresses.map(|address| address.ip()).collect())
167 .map_err(|error| format!("DNS lookup failed: {}", error))
168 }
169}
170
171pub struct WebFetchTool {
173 transport: Arc<dyn WebFetchTransport>,
174 resolver: Arc<dyn WebFetchResolver>,
175 cache: Arc<RwLock<WebFetchCache>>,
176 extractor: Arc<RwLock<Option<Arc<dyn LLMProvider>>>>,
177}
178
179impl WebFetchTool {
180 pub fn new() -> Self {
182 Self::with_extractor_slot(Arc::new(RwLock::new(None)))
183 }
184
185 pub fn with_extractor_slot(extractor: Arc<RwLock<Option<Arc<dyn LLMProvider>>>>) -> Self {
187 Self::with_extractor_slot_and_transport(
188 extractor,
189 Arc::new(ReqwestWebFetchTransport),
190 Arc::new(TokioWebFetchResolver),
191 )
192 }
193
194 pub fn with_transport_and_resolver(
196 transport: Arc<dyn WebFetchTransport>,
197 resolver: Arc<dyn WebFetchResolver>,
198 ) -> Self {
199 Self::with_extractor_slot_and_transport(Arc::new(RwLock::new(None)), transport, resolver)
200 }
201
202 pub fn with_extractor_slot_and_transport(
204 extractor: Arc<RwLock<Option<Arc<dyn LLMProvider>>>>,
205 transport: Arc<dyn WebFetchTransport>,
206 resolver: Arc<dyn WebFetchResolver>,
207 ) -> Self {
208 Self {
209 transport,
210 resolver,
211 cache: Arc::new(RwLock::new(WebFetchCache::new(MAX_CACHE_ENTRIES))),
212 extractor,
213 }
214 }
215}
216
217impl Default for WebFetchTool {
218 fn default() -> Self {
219 Self::new()
220 }
221}
222
223#[derive(Debug, Deserialize, JsonSchema)]
224struct WebFetchInput {
225 url: String,
227 #[serde(default)]
229 prompt: Option<String>,
230 #[serde(default)]
232 max_chars: Option<usize>,
233 #[serde(default)]
235 cache_ttl_seconds: Option<u64>,
236 #[serde(default)]
238 max_response_bytes: Option<usize>,
239 #[serde(default)]
241 max_redirects: Option<usize>,
242}
243
244#[derive(Debug, Deserialize, Clone, Default)]
245struct WebFetchPolicyInput {
246 #[serde(default)]
247 allowed_domains: Vec<String>,
248 #[serde(default)]
249 blocked_domains: Vec<String>,
250 #[serde(default)]
251 domain_allow: Vec<String>,
252 #[serde(default)]
253 domain_deny: Vec<String>,
254 #[serde(default)]
255 domain_requires_approval: Vec<String>,
256 #[serde(default)]
257 domain_unavailable: Vec<String>,
258 #[serde(default)]
259 allowed_schemes: Vec<String>,
260 #[serde(default)]
261 allowed_ports: Vec<u16>,
262 #[serde(default = "default_true")]
263 blocked_private_networks: bool,
264 #[serde(default)]
265 max_redirects: Option<usize>,
266}
267
268impl WebFetchPolicyInput {
269 fn from_context(value: &Value) -> Self {
271 let mut policy = serde_json::from_value::<Self>(value.clone()).unwrap_or_default();
272 policy.domain_allow.extend(
273 value
274 .get("domains")
275 .and_then(|domains| domains.get("allow"))
276 .and_then(Value::as_array)
277 .into_iter()
278 .flatten()
279 .filter_map(Value::as_str)
280 .map(str::to_string),
281 );
282 policy.domain_deny.extend(
283 value
284 .get("domains")
285 .and_then(|domains| domains.get("deny"))
286 .and_then(Value::as_array)
287 .into_iter()
288 .flatten()
289 .filter_map(Value::as_str)
290 .map(str::to_string),
291 );
292 policy.domain_requires_approval.extend(
293 value
294 .get("domains")
295 .and_then(|domains| domains.get("requires_approval"))
296 .and_then(Value::as_array)
297 .into_iter()
298 .flatten()
299 .filter_map(Value::as_str)
300 .map(str::to_string),
301 );
302 policy.domain_unavailable.extend(
303 value
304 .get("domains")
305 .and_then(|domains| domains.get("unavailable"))
306 .and_then(Value::as_array)
307 .into_iter()
308 .flatten()
309 .filter_map(Value::as_str)
310 .map(str::to_string),
311 );
312 policy
313 }
314}
315
316#[derive(Debug, Serialize, Clone)]
317struct WebFetchOutput {
318 url: String,
319 final_url: String,
320 status: u16,
321 content_type: Option<String>,
322 content: String,
323 truncated: bool,
324 from_cache: bool,
325 redirects: Vec<String>,
326 extraction_prompt_used: bool,
327 extraction_available: bool,
328}
329
330#[derive(Debug, Clone)]
331struct CacheEntry {
332 output: WebFetchOutput,
333 expires_at: Instant,
334 sequence: u64,
335}
336
337struct WebFetchCache {
339 entries: HashMap<String, CacheEntry>,
340 capacity: usize,
341 next_sequence: u64,
342}
343
344impl WebFetchCache {
345 fn new(capacity: usize) -> Self {
346 Self {
347 entries: HashMap::new(),
348 capacity,
349 next_sequence: 0,
350 }
351 }
352
353 fn get(&mut self, key: &str, now: Instant) -> Option<WebFetchOutput> {
354 self.remove_expired(now);
355 self.entries.get(key).map(|entry| entry.output.clone())
356 }
357
358 fn remove(&mut self, key: &str) {
359 self.entries.remove(key);
360 }
361
362 fn insert(
363 &mut self,
364 key: String,
365 output: WebFetchOutput,
366 lifetime: Duration,
367 stored_at: Instant,
368 ) -> Result<(), String> {
369 let expires_at = stored_at
370 .checked_add(lifetime)
371 .ok_or_else(cache_ttl_unrepresentable_error)?;
372 self.remove_expired(stored_at);
373 if self.capacity == 0 {
374 return Ok(());
375 }
376
377 if !self.entries.contains_key(&key)
378 && self.entries.len() >= self.capacity
379 && let Some(oldest_key) = self
380 .entries
381 .iter()
382 .min_by(|(left_key, left), (right_key, right)| {
383 left.sequence
384 .cmp(&right.sequence)
385 .then_with(|| left_key.cmp(right_key))
386 })
387 .map(|(key, _)| key.clone())
388 {
389 self.entries.remove(&oldest_key);
390 }
391
392 let sequence = self.next_sequence;
393 self.next_sequence = self.next_sequence.wrapping_add(1);
394 self.entries.insert(
395 key,
396 CacheEntry {
397 output,
398 expires_at,
399 sequence,
400 },
401 );
402 Ok(())
403 }
404
405 fn remove_expired(&mut self, now: Instant) {
406 self.entries.retain(|_, entry| now < entry.expires_at);
407 }
408}
409
410#[derive(Debug, Clone, Copy, PartialEq, Eq)]
414enum UrlPolicyTarget {
415 Initial,
416 Redirect,
417}
418
419#[async_trait]
420impl Tool for WebFetchTool {
421 fn id(&self) -> &str {
422 "web_fetch"
423 }
424
425 fn name(&self) -> &str {
426 "Web Fetch"
427 }
428
429 fn description(&self) -> &str {
430 "Fetch public web content with URL, redirect, DNS/IP, byte, and output safety checks."
431 }
432
433 fn input_schema(&self) -> Value {
434 generate_schema::<WebFetchInput>()
435 }
436
437 fn safety_metadata(&self) -> ToolSafetyMetadata {
438 ToolSafetyMetadata {
439 read_only: true,
440 concurrency_safe: true,
441 operation: ToolOperationKind::Network,
442 side_effect_level: ToolSideEffectLevel::ExternalRead,
443 requires_network: true,
444 destructive: false,
445 open_world: true,
446 host_dependent: false,
447 requires_user_interaction: false,
448 supports_cancellation: true,
449 default_requires_approval: false,
450 should_defer_schema: false,
451 max_output_chars: Some(DEFAULT_MAX_OUTPUT_CHARS),
452 max_result_size_chars: Some(DEFAULT_MAX_OUTPUT_CHARS),
453 }
454 }
455
456 fn classify_call(&self, _args: &Value) -> ToolCallClassification {
457 ToolCallClassification::from_metadata(&self.safety_metadata())
458 }
459
460 fn policy_bindings(&self) -> ToolPolicyBindings {
461 ToolPolicyBindings {
462 domain_fields: vec![DomainPolicyBinding::url("url")],
463 result_limit_fields: vec![
464 ResultLimitBinding::new("max_chars", ResultLimitKind::MaxOutputChars),
465 ResultLimitBinding::new("max_response_bytes", ResultLimitKind::MaxResponseBytes),
466 ResultLimitBinding::new("max_redirects", ResultLimitKind::MaxRedirects),
467 ],
468 ..Default::default()
469 }
470 }
471
472 async fn execute(&self, args: Value, ctx: ToolExecutionContext) -> ToolResult {
473 let input: WebFetchInput = match serde_json::from_value(args) {
474 Ok(input) => input,
475 Err(error) => return ToolResult::error(format!("Invalid input: {}", error)),
476 };
477 let policy = WebFetchPolicyInput::from_context(&ctx.policy_snapshot);
478 let max_output_chars = input.max_chars.unwrap_or(DEFAULT_MAX_OUTPUT_CHARS).min(
479 ctx.limits
480 .max_output_chars
481 .unwrap_or(DEFAULT_MAX_OUTPUT_CHARS),
482 );
483 let max_response_bytes = input
484 .max_response_bytes
485 .unwrap_or(DEFAULT_MAX_RESPONSE_BYTES)
486 .min(DEFAULT_MAX_RESPONSE_BYTES * 8)
487 .min(
488 ctx.limits
489 .max_response_bytes
490 .unwrap_or(DEFAULT_MAX_RESPONSE_BYTES * 8),
491 );
492 let max_redirects = effective_max_redirects(
493 input.max_redirects,
494 policy.max_redirects,
495 ctx.limits.max_redirects,
496 );
497 let cache_ttl = input.cache_ttl_seconds.unwrap_or(DEFAULT_CACHE_TTL_SECONDS);
498 let cache_lifetime = match checked_cache_lifetime(cache_ttl, Instant::now()) {
499 Ok(lifetime) => lifetime,
500 Err(error) => return ToolResult::error(error),
501 };
502
503 let original_url = match reqwest::Url::parse(&input.url) {
504 Ok(url) => url,
505 Err(error) => return ToolResult::error(format!("Invalid URL: {}", error)),
506 };
507 if let Err(result) =
508 require_initial_domain_approval(&original_url, &policy, ctx.approval.as_ref())
509 {
510 return result;
511 }
512 if let Err(result) = validate_url_with_policy(
513 &original_url,
514 Some(&policy),
515 UrlPolicyTarget::Initial,
516 self.resolver.as_ref(),
517 )
518 .await
519 {
520 return result;
521 }
522 let cache_key = format!(
523 "{}|{}|{}|{}|{}|{}",
524 input.url,
525 input.prompt.as_deref().unwrap_or(""),
526 max_output_chars,
527 max_response_bytes,
528 max_redirects,
529 policy_cache_fingerprint(&policy)
530 );
531 if cache_lifetime.is_some() {
532 let cached_output = { self.cache.write().get(&cache_key, Instant::now()) };
533 if let Some(mut output) = cached_output {
534 match validate_cached_output_with_policy(
535 &output,
536 &policy,
537 max_redirects,
538 self.resolver.as_ref(),
539 )
540 .await
541 {
542 Ok(true) => {
543 output.from_cache = true;
544 return web_result(&output, output.truncated, max_output_chars, true, None);
545 }
546 Ok(false) => {
547 self.cache.write().remove(&cache_key);
548 }
549 Err(result) => return result,
550 }
551 }
552 }
553
554 let mut current_url = original_url.clone();
555 let mut redirects = Vec::new();
556
557 let response = loop {
558 if ctx.cancellation.is_cancelled() {
559 return ToolResult::error(
560 ctx.cancellation
561 .reason()
562 .unwrap_or("Tool execution cancelled"),
563 );
564 }
565 let policy_target = if redirects.is_empty() {
566 UrlPolicyTarget::Initial
567 } else {
568 UrlPolicyTarget::Redirect
569 };
570 let validated_target = match validate_url_with_policy(
571 ¤t_url,
572 Some(&policy),
573 policy_target,
574 self.resolver.as_ref(),
575 )
576 .await
577 {
578 Ok(target) => target,
579 Err(result) => return result,
580 };
581 let response = match self
582 .transport
583 .send_validated(
584 WebFetchTransportRequest {
585 url: current_url.to_string(),
586 max_response_bytes,
587 },
588 &validated_target.addresses,
589 )
590 .await
591 {
592 Ok(response) => response,
593 Err(error) => return ToolResult::error(error),
594 };
595 if response.body.len() > max_response_bytes {
596 return ToolResult::error(format!(
597 "Response exceeded max_response_bytes {}",
598 max_response_bytes
599 ));
600 }
601 if (300..400).contains(&response.status) {
602 if redirects.len() >= max_redirects {
603 return ToolResult::error("Redirect limit exceeded");
604 }
605 let Some(location) = response.location.as_deref() else {
606 return ToolResult::error("Redirect response missing Location header");
607 };
608 let next_url = match current_url.join(location) {
609 Ok(url) => url,
610 Err(error) => {
611 return ToolResult::error(format!("Invalid redirect URL: {}", error));
612 }
613 };
614 redirects.push(next_url.to_string());
615 current_url = next_url;
616 continue;
617 }
618 break response;
619 };
620
621 let status = response.status;
622 let content_type = response.content_type;
623 if ctx.cancellation.is_cancelled() {
624 return ToolResult::error(
625 ctx.cancellation
626 .reason()
627 .unwrap_or("Tool execution cancelled"),
628 );
629 }
630 let raw_text = match String::from_utf8(response.body) {
631 Ok(text) => text,
632 Err(_) => return ToolResult::error("Response body is not UTF-8 text"),
633 };
634 let converted = if content_type
635 .as_deref()
636 .is_some_and(|value| value.to_ascii_lowercase().contains("html"))
637 || raw_text.trim_start().starts_with('<')
638 {
639 html_to_text(&raw_text)
640 } else {
641 raw_text
642 };
643 let (content, truncated) = truncate_chars(converted, max_output_chars);
644 let mut extraction_available = false;
645 let mut final_content = content;
646 let mut extraction_error = None;
647 if let Some(prompt) = input.prompt.as_deref()
648 && let Some(extractor) = { self.extractor.read().clone() }
649 {
650 match extract_with_llm(extractor, prompt, &final_content).await {
651 Ok(answer) => {
652 final_content = answer;
653 extraction_available = true;
654 }
655 Err(error) => {
656 extraction_error = Some(error);
657 }
658 }
659 }
660 let (final_content, extraction_truncated) = truncate_chars(final_content, max_output_chars);
661 let output = WebFetchOutput {
662 url: original_url.to_string(),
663 final_url: current_url.to_string(),
664 status,
665 content_type,
666 content: final_content,
667 truncated: truncated || extraction_truncated,
668 from_cache: false,
669 redirects,
670 extraction_prompt_used: input.prompt.is_some(),
671 extraction_available,
672 };
673 if let Some(cache_lifetime) = cache_lifetime
674 && let Err(error) =
675 self.cache
676 .write()
677 .insert(cache_key, output.clone(), cache_lifetime, Instant::now())
678 {
679 return ToolResult::error(error);
680 }
681 web_result(
682 &output,
683 output.truncated,
684 max_output_chars,
685 false,
686 extraction_error,
687 )
688 }
689}
690
691struct ValidatedWebTarget {
692 addresses: Vec<SocketAddr>,
693}
694
695async fn validate_url_with_policy(
696 url: &reqwest::Url,
697 policy: Option<&WebFetchPolicyInput>,
698 policy_target: UrlPolicyTarget,
699 resolver: &dyn WebFetchResolver,
700) -> Result<ValidatedWebTarget, ToolResult> {
701 if let Some(policy) = policy {
702 check_configured_url_policy(url, policy, policy_target)?;
703 }
704 validate_url(url, resolver).await
705}
706
707fn require_initial_domain_approval(
709 url: &reqwest::Url,
710 policy: &WebFetchPolicyInput,
711 approval: Option<&ToolApprovalRecord>,
712) -> Result<(), ToolResult> {
713 let Some(host) = url.host_str().map(normalize_host) else {
714 return Err(ToolResult::error("URL host is required"));
715 };
716 let Some(pattern) = policy
717 .domain_requires_approval
718 .iter()
719 .find(|pattern| host_matches(pattern, &host))
720 else {
721 return Ok(());
722 };
723 if approval.is_some_and(|record| {
724 matches!(
725 record.status,
726 ToolApprovalStatus::Approved | ToolApprovalStatus::Modified
727 )
728 }) {
729 return Ok(());
730 }
731 Err(ToolResult::error(format!(
732 "Domain '{}' requires an approved tool execution context",
733 pattern
734 )))
735}
736
737fn checked_cache_lifetime(seconds: u64, now: Instant) -> Result<Option<Duration>, String> {
738 if seconds == 0 {
739 return Ok(None);
740 }
741 let lifetime = Duration::from_secs(seconds);
742 now.checked_add(lifetime)
743 .ok_or_else(cache_ttl_unrepresentable_error)?;
744 Ok(Some(lifetime))
745}
746
747fn cache_ttl_unrepresentable_error() -> String {
748 "cache_ttl_seconds is too large to represent on this platform".to_string()
749}
750
751async fn validate_cached_output_with_policy(
753 output: &WebFetchOutput,
754 policy: &WebFetchPolicyInput,
755 max_redirects: usize,
756 resolver: &dyn WebFetchResolver,
757) -> Result<bool, ToolResult> {
758 if output.redirects.len() > max_redirects {
759 return Ok(false);
760 }
761 let final_url = reqwest::Url::parse(&output.final_url)
762 .map_err(|error| ToolResult::error(format!("Cached final URL is invalid: {}", error)))?;
763 let final_target = if output.redirects.is_empty() {
764 UrlPolicyTarget::Initial
765 } else {
766 UrlPolicyTarget::Redirect
767 };
768 let _ = validate_url_with_policy(&final_url, Some(policy), final_target, resolver).await?;
769 for redirect in &output.redirects {
770 let url = reqwest::Url::parse(redirect).map_err(|error| {
771 ToolResult::error(format!("Cached redirect URL is invalid: {}", error))
772 })?;
773 let _ = validate_url_with_policy(&url, Some(policy), UrlPolicyTarget::Redirect, resolver)
774 .await?;
775 }
776 Ok(true)
777}
778
779fn effective_max_redirects(
780 request_limit: Option<usize>,
781 policy_limit: Option<usize>,
782 context_limit: Option<usize>,
783) -> usize {
784 [request_limit, policy_limit, context_limit]
785 .into_iter()
786 .flatten()
787 .min()
788 .unwrap_or(DEFAULT_MAX_REDIRECTS)
789 .min(DEFAULT_MAX_REDIRECTS * 4)
790}
791
792fn policy_cache_fingerprint(policy: &WebFetchPolicyInput) -> String {
793 serde_json::json!({
794 "allowed_domains": policy.allowed_domains,
795 "blocked_domains": policy.blocked_domains,
796 "domain_allow": policy.domain_allow,
797 "domain_deny": policy.domain_deny,
798 "domain_requires_approval": policy.domain_requires_approval,
799 "domain_unavailable": policy.domain_unavailable,
800 "allowed_schemes": policy.allowed_schemes,
801 "allowed_ports": policy.allowed_ports,
802 "blocked_private_networks": policy.blocked_private_networks,
803 "max_redirects": policy.max_redirects,
804 })
805 .to_string()
806}
807
808async fn validate_url(
809 url: &reqwest::Url,
810 resolver: &dyn WebFetchResolver,
811) -> Result<ValidatedWebTarget, ToolResult> {
812 match url.scheme() {
813 "http" | "https" => {}
814 other => {
815 return Err(ToolResult::error(format!(
816 "URL scheme '{}' is not allowed",
817 other
818 )));
819 }
820 }
821 if url.username() != "" || url.password().is_some() {
822 return Err(ToolResult::error(
823 "URLs with embedded credentials are not allowed",
824 ));
825 }
826 let Some(host) = url.host_str() else {
827 return Err(ToolResult::error("URL host is required"));
828 };
829 if is_metadata_host(host) || is_localhost_name(host) {
830 return Err(ToolResult::error(
831 "Localhost and metadata-service hosts are blocked",
832 ));
833 }
834 let port = url.port_or_known_default().unwrap_or(443);
835 if let Ok(ip) = host.parse::<IpAddr>() {
836 validate_ip(ip)?;
837 return Ok(ValidatedWebTarget {
838 addresses: vec![SocketAddr::new(ip, port)],
839 });
840 }
841 let addresses = resolver
842 .resolve(host, port)
843 .await
844 .map_err(ToolResult::error)?;
845 if addresses.is_empty() {
846 return Err(ToolResult::error("DNS lookup returned no addresses"));
847 }
848 for address in &addresses {
849 validate_ip(*address)?;
850 }
851 Ok(ValidatedWebTarget {
852 addresses: addresses
853 .into_iter()
854 .map(|address| SocketAddr::new(address, port))
855 .collect(),
856 })
857}
858
859fn check_configured_url_policy(
860 url: &reqwest::Url,
861 policy: &WebFetchPolicyInput,
862 policy_target: UrlPolicyTarget,
863) -> Result<(), ToolResult> {
864 let scheme = url.scheme();
865 if !policy.allowed_schemes.is_empty()
866 && !policy
867 .allowed_schemes
868 .iter()
869 .any(|allowed| allowed.eq_ignore_ascii_case(scheme))
870 {
871 return Err(ToolResult::error(format!(
872 "URL scheme '{}' is not allowed by policy",
873 scheme
874 )));
875 }
876
877 if !policy.allowed_ports.is_empty() {
878 let port = url.port_or_known_default().unwrap_or(0);
879 if !policy.allowed_ports.contains(&port) {
880 return Err(ToolResult::error(format!(
881 "URL port '{}' is not allowed by policy",
882 port
883 )));
884 }
885 }
886
887 let Some(host) = url.host_str().map(normalize_host) else {
888 return Err(ToolResult::error("URL host is required"));
889 };
890 if policy.blocked_private_networks && (is_metadata_host(&host) || is_localhost_name(&host)) {
891 return Err(ToolResult::error(
892 "Private, localhost, link-local, or metadata host is blocked by policy",
893 ));
894 }
895
896 for pattern in policy
897 .blocked_domains
898 .iter()
899 .chain(policy.domain_deny.iter())
900 {
901 if host_matches(pattern, &host) {
902 return Err(ToolResult::error(format!(
903 "Domain '{}' is blocked by policy",
904 pattern
905 )));
906 }
907 }
908
909 for pattern in &policy.domain_unavailable {
910 if host_matches(pattern, &host) {
911 return Err(ToolResult::error(format!(
912 "Domain '{}' is unavailable by policy",
913 pattern
914 )));
915 }
916 }
917
918 if policy_target == UrlPolicyTarget::Redirect {
919 for pattern in &policy.domain_requires_approval {
920 if host_matches(pattern, &host) {
921 return Err(ToolResult::error(format!(
922 "Domain '{}' requires approval and cannot be reached by redirect",
923 pattern
924 )));
925 }
926 }
927 }
928
929 let allowed: Vec<&String> = policy
930 .allowed_domains
931 .iter()
932 .chain(policy.domain_allow.iter())
933 .collect();
934 if !allowed.is_empty() && !allowed.iter().any(|pattern| host_matches(pattern, &host)) {
935 return Err(ToolResult::error(
936 "URL domain is not in the configured allowlist",
937 ));
938 }
939
940 Ok(())
941}
942
943fn validate_ip(ip: IpAddr) -> Result<(), ToolResult> {
944 if is_blocked_ip(ip) {
945 Err(ToolResult::error(
946 "Non-public and special-use IP addresses are blocked",
947 ))
948 } else {
949 Ok(())
950 }
951}
952
953fn is_blocked_ip(ip: IpAddr) -> bool {
957 match ip {
958 IpAddr::V4(ip) => [
959 (Ipv4Addr::new(0, 0, 0, 0), 8),
960 (Ipv4Addr::new(10, 0, 0, 0), 8),
961 (Ipv4Addr::new(100, 64, 0, 0), 10),
962 (Ipv4Addr::new(127, 0, 0, 0), 8),
963 (Ipv4Addr::new(169, 254, 0, 0), 16),
964 (Ipv4Addr::new(172, 16, 0, 0), 12),
965 (Ipv4Addr::new(192, 0, 0, 0), 24),
966 (Ipv4Addr::new(192, 0, 2, 0), 24),
967 (Ipv4Addr::new(192, 88, 99, 0), 24),
968 (Ipv4Addr::new(192, 168, 0, 0), 16),
969 (Ipv4Addr::new(198, 18, 0, 0), 15),
970 (Ipv4Addr::new(198, 51, 100, 0), 24),
971 (Ipv4Addr::new(203, 0, 113, 0), 24),
972 (Ipv4Addr::new(224, 0, 0, 0), 4),
973 (Ipv4Addr::new(240, 0, 0, 0), 4),
974 ]
975 .into_iter()
976 .any(|(network, prefix)| ipv4_in_prefix(ip, network, prefix)),
977 IpAddr::V6(ip) => [
978 (Ipv6Addr::UNSPECIFIED, 96),
979 (Ipv6Addr::new(0, 0, 0, 0, 0, 0xffff, 0, 0), 96),
980 (Ipv6Addr::new(0x64, 0xff9b, 0, 0, 0, 0, 0, 0), 96),
981 (Ipv6Addr::new(0x64, 0xff9b, 1, 0, 0, 0, 0, 0), 48),
982 (Ipv6Addr::new(0x100, 0, 0, 0, 0, 0, 0, 0), 64),
983 (Ipv6Addr::new(0x100, 0, 0, 1, 0, 0, 0, 0), 64),
984 (Ipv6Addr::new(0x2001, 0, 0, 0, 0, 0, 0, 0), 23),
985 (Ipv6Addr::new(0x2001, 0x0db8, 0, 0, 0, 0, 0, 0), 32),
986 (Ipv6Addr::new(0x2002, 0, 0, 0, 0, 0, 0, 0), 16),
987 (Ipv6Addr::new(0x3ffe, 0, 0, 0, 0, 0, 0, 0), 16),
988 (Ipv6Addr::new(0x3fff, 0, 0, 0, 0, 0, 0, 0), 20),
989 (Ipv6Addr::new(0x5f00, 0, 0, 0, 0, 0, 0, 0), 16),
990 (Ipv6Addr::new(0xfc00, 0, 0, 0, 0, 0, 0, 0), 7),
991 (Ipv6Addr::new(0xfe80, 0, 0, 0, 0, 0, 0, 0), 10),
992 (Ipv6Addr::new(0xfec0, 0, 0, 0, 0, 0, 0, 0), 10),
993 (Ipv6Addr::new(0xff00, 0, 0, 0, 0, 0, 0, 0), 8),
994 ]
995 .into_iter()
996 .any(|(network, prefix)| ipv6_in_prefix(ip, network, prefix)),
997 }
998}
999
1000fn ipv4_in_prefix(ip: Ipv4Addr, network: Ipv4Addr, prefix: u8) -> bool {
1001 let shift = 32 - u32::from(prefix);
1002 u32::from(ip) >> shift == u32::from(network) >> shift
1003}
1004
1005fn ipv6_in_prefix(ip: Ipv6Addr, network: Ipv6Addr, prefix: u8) -> bool {
1006 let shift = 128 - u32::from(prefix);
1007 u128::from_be_bytes(ip.octets()) >> shift == u128::from_be_bytes(network.octets()) >> shift
1008}
1009
1010fn normalize_host(host: &str) -> String {
1011 host.trim_end_matches('.').to_ascii_lowercase()
1012}
1013
1014fn host_matches(pattern: &str, host: &str) -> bool {
1015 let pattern = normalize_host(pattern.trim_start_matches("*."));
1016 host == pattern || host.ends_with(&format!(".{}", pattern))
1017}
1018
1019fn is_localhost_name(host: &str) -> bool {
1020 let host = normalize_host(host);
1021 host == "localhost" || host.ends_with(".localhost")
1022}
1023
1024fn is_metadata_host(host: &str) -> bool {
1025 let host = normalize_host(host);
1026 matches!(
1027 host.as_str(),
1028 "metadata.google.internal" | "metadata" | "169.254.169.254" | "100.100.100.200"
1029 )
1030}
1031
1032fn html_to_text(html: &str) -> String {
1033 let without_scripts = Regex::new(r"(?is)<(script|style)[^>]*>.*?</(script|style)>")
1034 .map(|regex| regex.replace_all(html, " ").to_string())
1035 .unwrap_or_else(|_| html.to_string());
1036 let with_breaks = Regex::new(r"(?i)<\s*(br|/p|/div|/h[1-6]|/li)\s*/?>")
1037 .map(|regex| regex.replace_all(&without_scripts, "\n").to_string())
1038 .unwrap_or(without_scripts);
1039 let without_tags = Regex::new(r"(?is)<[^>]+>")
1040 .map(|regex| regex.replace_all(&with_breaks, " ").to_string())
1041 .unwrap_or(with_breaks);
1042 decode_html_entities(&without_tags)
1043 .lines()
1044 .map(str::trim)
1045 .filter(|line| !line.is_empty())
1046 .collect::<Vec<_>>()
1047 .join("\n")
1048}
1049
1050fn decode_html_entities(text: &str) -> String {
1051 text.replace(" ", " ")
1052 .replace("&", "&")
1053 .replace("<", "<")
1054 .replace(">", ">")
1055 .replace(""", "\"")
1056 .replace("'", "'")
1057}
1058
1059fn default_true() -> bool {
1060 true
1061}
1062
1063fn truncate_chars(text: String, max_chars: usize) -> (String, bool) {
1064 let mut chars = text.chars();
1065 let truncated: String = chars.by_ref().take(max_chars).collect();
1066 if chars.next().is_some() {
1067 (truncated, true)
1068 } else {
1069 (text, false)
1070 }
1071}
1072
1073async fn extract_with_llm(
1074 extractor: Arc<dyn LLMProvider>,
1075 prompt: &str,
1076 content: &str,
1077) -> Result<String, String> {
1078 let messages = vec![
1079 ChatMessage::system(
1080 "Extract only the information requested by the user from the provided web content. Return a concise answer.",
1081 ),
1082 ChatMessage::user(format!(
1083 "Extraction request:\n{}\n\nWeb content:\n{}",
1084 prompt, content
1085 )),
1086 ];
1087 let config = LLMConfig {
1088 max_tokens: Some(800),
1089 temperature: Some(0.0),
1090 ..LLMConfig::default()
1091 };
1092 extractor
1093 .complete(&messages, Some(&config))
1094 .await
1095 .map(|response| response.content)
1096 .map_err(|error| error.to_string())
1097}
1098
1099fn web_result(
1100 output: &WebFetchOutput,
1101 truncated: bool,
1102 max_output_chars: usize,
1103 from_cache: bool,
1104 extraction_error: Option<String>,
1105) -> ToolResult {
1106 let json = match serde_json::to_string(output) {
1107 Ok(json) => json,
1108 Err(error) => return ToolResult::error(format!("Serialization error: {}", error)),
1109 };
1110 let mut metadata = HashMap::new();
1111 metadata.insert("truncated".to_string(), Value::Bool(truncated));
1112 metadata.insert(
1113 "max_output_chars".to_string(),
1114 Value::from(max_output_chars),
1115 );
1116 metadata.insert("from_cache".to_string(), Value::Bool(from_cache));
1117 if output.extraction_prompt_used {
1118 let status = if output.extraction_available {
1119 "executed".to_string()
1120 } else if extraction_error.is_some() {
1121 "failed".to_string()
1122 } else {
1123 "unavailable".to_string()
1124 };
1125 metadata.insert("nested_llm_extraction".to_string(), Value::String(status));
1126 if let Some(error) = extraction_error {
1127 metadata.insert("nested_llm_error".to_string(), Value::String(error));
1128 }
1129 }
1130 ToolResult::ok_with_metadata(json, metadata)
1131}
1132
1133#[cfg(test)]
1134mod tests {
1135 use super::*;
1136 use tokio::io::{AsyncReadExt, AsyncWriteExt};
1137
1138 #[derive(Default)]
1139 struct FixtureTransport {
1140 routes: RwLock<HashMap<String, WebFetchTransportResponse>>,
1141 requests: RwLock<Vec<String>>,
1142 validated_addresses: RwLock<Vec<Vec<SocketAddr>>>,
1143 }
1144
1145 impl FixtureTransport {
1146 fn route(&self, url: &str, response: WebFetchTransportResponse) {
1147 self.routes.write().insert(url.to_string(), response);
1148 }
1149
1150 fn requests(&self) -> Vec<String> {
1151 self.requests.read().clone()
1152 }
1153
1154 fn validated_addresses(&self) -> Vec<Vec<SocketAddr>> {
1155 self.validated_addresses.read().clone()
1156 }
1157 }
1158
1159 #[async_trait]
1160 impl WebFetchTransport for FixtureTransport {
1161 async fn send(
1162 &self,
1163 request: WebFetchTransportRequest,
1164 ) -> Result<WebFetchTransportResponse, String> {
1165 self.requests.write().push(request.url.clone());
1166 self.routes
1167 .read()
1168 .get(&request.url)
1169 .cloned()
1170 .ok_or_else(|| format!("Unconfigured web fetch route: {}", request.url))
1171 }
1172
1173 async fn send_validated(
1174 &self,
1175 request: WebFetchTransportRequest,
1176 addresses: &[SocketAddr],
1177 ) -> Result<WebFetchTransportResponse, String> {
1178 self.validated_addresses.write().push(addresses.to_vec());
1179 self.send(request).await
1180 }
1181 }
1182
1183 struct FixtureResolver {
1184 addresses: HashMap<String, Vec<IpAddr>>,
1185 default: Vec<IpAddr>,
1186 requests: RwLock<Vec<String>>,
1187 }
1188
1189 impl FixtureResolver {
1190 fn requests(&self) -> Vec<String> {
1191 self.requests.read().clone()
1192 }
1193 }
1194
1195 impl Default for FixtureResolver {
1196 fn default() -> Self {
1197 Self {
1198 addresses: HashMap::new(),
1199 default: vec![IpAddr::from([93, 184, 216, 34])],
1200 requests: RwLock::new(Vec::new()),
1201 }
1202 }
1203 }
1204
1205 #[async_trait]
1206 impl WebFetchResolver for FixtureResolver {
1207 async fn resolve(&self, host: &str, _port: u16) -> Result<Vec<IpAddr>, String> {
1208 self.requests.write().push(host.to_string());
1209 Ok(self
1210 .addresses
1211 .get(host)
1212 .cloned()
1213 .unwrap_or_else(|| self.default.clone()))
1214 }
1215 }
1216
1217 fn response(
1218 status: u16,
1219 content_type: Option<&str>,
1220 location: Option<&str>,
1221 body: &str,
1222 ) -> WebFetchTransportResponse {
1223 WebFetchTransportResponse {
1224 status,
1225 content_type: content_type.map(str::to_string),
1226 location: location.map(str::to_string),
1227 body: body.as_bytes().to_vec(),
1228 }
1229 }
1230
1231 fn fixture_tool(transport: Arc<FixtureTransport>) -> WebFetchTool {
1232 WebFetchTool::with_transport_and_resolver(transport, Arc::new(FixtureResolver::default()))
1233 }
1234
1235 fn approval_context(status: Option<ToolApprovalStatus>) -> ToolExecutionContext {
1236 let mut context = ToolExecutionContext::test("web_fetch");
1237 context.policy_snapshot = serde_json::json!({
1238 "domains": {
1239 "requires_approval": ["approval.test"]
1240 }
1241 });
1242 context.approval = status.map(|status| ToolApprovalRecord {
1243 status,
1244 reason: None,
1245 modified_arguments: None,
1246 });
1247 context
1248 }
1249
1250 fn output(result: &ToolResult) -> Value {
1251 serde_json::from_str(&result.output).expect("web fetch output should be JSON")
1252 }
1253
1254 fn cached_output(url: &str) -> WebFetchOutput {
1255 WebFetchOutput {
1256 url: url.to_string(),
1257 final_url: url.to_string(),
1258 status: 200,
1259 content_type: Some("text/plain".to_string()),
1260 content: url.to_string(),
1261 truncated: false,
1262 from_cache: false,
1263 redirects: Vec::new(),
1264 extraction_prompt_used: false,
1265 extraction_available: false,
1266 }
1267 }
1268
1269 #[tokio::test]
1270 async fn blocks_non_public_ip_ranges_without_requesting() {
1271 for url in [
1272 "http://0.0.0.1/",
1273 "http://10.0.0.1/",
1274 "http://100.64.0.1/",
1275 "http://127.0.0.1/",
1276 "http://169.254.1.1/",
1277 "http://169.254.169.254/latest",
1278 "http://192.0.0.1/",
1279 "http://192.0.2.1/",
1280 "http://192.88.99.2/",
1281 "http://198.18.0.1/",
1282 "http://240.0.0.1/",
1283 "http://255.255.255.255/",
1284 "http://[::1]/",
1285 "http://[64:ff9b::1]/",
1286 "http://[64:ff9b:1::1]/",
1287 "http://[100::1]/",
1288 "http://[100:0:0:1::1]/",
1289 "http://[2001:2::1]/",
1290 "http://[2001:db8::1]/",
1291 "http://[2002::1]/",
1292 "http://[3fff::1]/",
1293 "http://[5f00::1]/",
1294 "http://[fe80::1]/",
1295 "http://[fec0::1]/",
1296 "http://[::ffff:10.0.0.1]/",
1297 ] {
1298 let result = WebFetchTool::new()
1299 .execute(
1300 serde_json::json!({"url": url}),
1301 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1302 )
1303 .await;
1304 assert!(!result.success, "URL should be blocked: {}", url);
1305 }
1306 }
1307
1308 #[tokio::test]
1309 async fn default_transport_connects_only_to_supplied_address() {
1310 let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0))
1311 .await
1312 .unwrap();
1313 let address = listener.local_addr().unwrap();
1314 let server = tokio::spawn(async move {
1315 let (mut stream, _) = listener.accept().await.unwrap();
1316 let mut request = vec![0u8; 4096];
1317 let read = stream.read(&mut request).await.unwrap();
1318 stream
1319 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 5\r\nConnection: close\r\n\r\nbound")
1320 .await
1321 .unwrap();
1322 String::from_utf8_lossy(&request[..read]).into_owned()
1323 });
1324 let url = format!("http://binding.invalid:{}/bound", address.port());
1325
1326 let response = ReqwestWebFetchTransport
1327 .send_validated(
1328 WebFetchTransportRequest {
1329 url,
1330 max_response_bytes: 64,
1331 },
1332 &[address],
1333 )
1334 .await
1335 .unwrap();
1336 let request = server.await.unwrap().to_ascii_lowercase();
1337
1338 assert_eq!(response.body, b"bound");
1339 assert!(request.contains(&format!("host: binding.invalid:{}", address.port())));
1340 }
1341
1342 #[tokio::test]
1343 async fn validated_addresses_are_passed_to_transport() {
1344 let transport = Arc::new(FixtureTransport::default());
1345 transport.route(
1346 "https://public.test/page",
1347 response(200, Some("text/plain"), None, "bound"),
1348 );
1349 let mut resolver = FixtureResolver::default();
1350 resolver
1351 .addresses
1352 .insert("public.test".to_string(), vec![IpAddr::from([1, 1, 1, 1])]);
1353 let tool = WebFetchTool::with_transport_and_resolver(
1354 Arc::clone(&transport) as Arc<dyn WebFetchTransport>,
1355 Arc::new(resolver),
1356 );
1357
1358 let result = tool
1359 .execute(
1360 serde_json::json!({"url": "https://public.test/page"}),
1361 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1362 )
1363 .await;
1364
1365 assert!(result.success);
1366 assert_eq!(
1367 transport.validated_addresses(),
1368 vec![vec![SocketAddr::from(([1, 1, 1, 1], 443))]]
1369 );
1370 }
1371
1372 #[tokio::test]
1373 async fn blocks_embedded_credentials_without_requesting() {
1374 let transport = Arc::new(FixtureTransport::default());
1375 let result = fixture_tool(Arc::clone(&transport))
1376 .execute(
1377 serde_json::json!({"url": "https://user:password@public.test/"}),
1378 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1379 )
1380 .await;
1381
1382 assert!(!result.success);
1383 assert!(result.output.contains("embedded credentials"));
1384 assert!(transport.requests().is_empty());
1385 assert!(transport.validated_addresses().is_empty());
1386 }
1387
1388 #[tokio::test]
1389 async fn dns_alias_to_metadata_address_is_blocked_before_transport() {
1390 let transport = Arc::new(FixtureTransport::default());
1391 let mut resolver = FixtureResolver::default();
1392 resolver.addresses.insert(
1393 "metadata-alias.test".to_string(),
1394 vec![IpAddr::from([100, 100, 100, 200])],
1395 );
1396 let tool = WebFetchTool::with_transport_and_resolver(
1397 Arc::clone(&transport) as Arc<dyn WebFetchTransport>,
1398 Arc::new(resolver),
1399 );
1400
1401 let result = tool
1402 .execute(
1403 serde_json::json!({"url": "http://metadata-alias.test/latest"}),
1404 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1405 )
1406 .await;
1407
1408 assert!(!result.success);
1409 assert!(result.output.contains("Non-public and special-use"));
1410 assert!(transport.requests().is_empty());
1411 assert!(transport.validated_addresses().is_empty());
1412 }
1413
1414 #[test]
1415 fn public_address_examples_remain_allowed() {
1416 for address in ["1.1.1.1", "8.8.8.8", "2606:4700:4700::1111"] {
1417 let address = address.parse::<IpAddr>().unwrap();
1418 assert!(
1419 !is_blocked_ip(address),
1420 "address should remain allowed: {address}"
1421 );
1422 }
1423 }
1424
1425 #[tokio::test]
1426 async fn in_memory_transport_fetches_html_and_text() {
1427 let transport = Arc::new(FixtureTransport::default());
1428 transport.route(
1429 "https://public.test/page",
1430 response(
1431 200,
1432 Some("text/html"),
1433 None,
1434 "<h1>Title</h1><p>Hello & bye</p>",
1435 ),
1436 );
1437 transport.route(
1438 "https://public.test/plain",
1439 response(200, Some("text/plain"), None, "plain response"),
1440 );
1441 let tool = fixture_tool(transport);
1442
1443 let html = tool
1444 .execute(
1445 serde_json::json!({"url": "https://public.test/page"}),
1446 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1447 )
1448 .await;
1449 let text = tool
1450 .execute(
1451 serde_json::json!({"url": "https://public.test/plain"}),
1452 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1453 )
1454 .await;
1455
1456 assert!(html.success);
1457 assert_eq!(output(&html)["content"], "Title\nHello & bye");
1458 assert!(text.success);
1459 assert_eq!(output(&text)["content"], "plain response");
1460 }
1461
1462 #[tokio::test]
1463 async fn in_memory_transport_follows_exact_route_redirects() {
1464 let transport = Arc::new(FixtureTransport::default());
1465 transport.route(
1466 "https://public.test/start",
1467 response(302, None, Some("/final"), ""),
1468 );
1469 transport.route(
1470 "https://public.test/final",
1471 response(200, Some("text/plain"), None, "done"),
1472 );
1473 let tool = fixture_tool(Arc::clone(&transport));
1474
1475 let result = tool
1476 .execute(
1477 serde_json::json!({"url": "https://public.test/start"}),
1478 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1479 )
1480 .await;
1481
1482 assert!(result.success);
1483 assert_eq!(output(&result)["final_url"], "https://public.test/final");
1484 assert_eq!(
1485 transport.requests(),
1486 vec![
1487 "https://public.test/start".to_string(),
1488 "https://public.test/final".to_string()
1489 ]
1490 );
1491 assert_eq!(transport.validated_addresses().len(), 2);
1492 }
1493
1494 #[tokio::test]
1495 async fn approved_and_modified_initial_domains_reach_transport() {
1496 let transport = Arc::new(FixtureTransport::default());
1497 transport.route(
1498 "https://approval.test/page",
1499 response(200, Some("text/plain"), None, "approved"),
1500 );
1501 let tool = fixture_tool(Arc::clone(&transport));
1502
1503 for status in [ToolApprovalStatus::Approved, ToolApprovalStatus::Modified] {
1504 let result = tool
1505 .execute(
1506 serde_json::json!({
1507 "url": "https://approval.test/page",
1508 "cache_ttl_seconds": 0
1509 }),
1510 approval_context(Some(status)),
1511 )
1512 .await;
1513 assert!(result.success);
1514 }
1515
1516 assert_eq!(
1517 transport.requests(),
1518 vec![
1519 "https://approval.test/page".to_string(),
1520 "https://approval.test/page".to_string()
1521 ]
1522 );
1523 }
1524
1525 #[tokio::test]
1526 async fn approval_required_initial_domain_fails_closed_before_dns_and_transport() {
1527 let transport = Arc::new(FixtureTransport::default());
1528 let resolver = Arc::new(FixtureResolver::default());
1529 let tool = WebFetchTool::with_transport_and_resolver(
1530 Arc::clone(&transport) as Arc<dyn WebFetchTransport>,
1531 Arc::clone(&resolver) as Arc<dyn WebFetchResolver>,
1532 );
1533 let statuses = [
1534 None,
1535 Some(ToolApprovalStatus::NotRequired),
1536 Some(ToolApprovalStatus::Rejected),
1537 Some(ToolApprovalStatus::Timeout),
1538 Some(ToolApprovalStatus::Unavailable),
1539 ];
1540
1541 for status in statuses {
1542 let result = tool
1543 .execute(
1544 serde_json::json!({"url": "https://approval.test/page"}),
1545 approval_context(status),
1546 )
1547 .await;
1548 assert!(!result.success);
1549 assert!(
1550 result
1551 .output
1552 .contains("requires an approved tool execution context")
1553 );
1554 }
1555
1556 assert!(resolver.requests().is_empty());
1557 assert!(transport.requests().is_empty());
1558 assert!(transport.validated_addresses().is_empty());
1559 }
1560
1561 #[tokio::test]
1562 async fn redirect_to_approval_required_domain_stops_before_second_dns_and_request() {
1563 let transport = Arc::new(FixtureTransport::default());
1564 transport.route(
1565 "https://public.test/start",
1566 response(302, None, Some("https://approval.test/page"), ""),
1567 );
1568 transport.route(
1569 "https://approval.test/page",
1570 response(200, Some("text/plain"), None, "not reached"),
1571 );
1572 let resolver = Arc::new(FixtureResolver::default());
1573 let tool = WebFetchTool::with_transport_and_resolver(
1574 Arc::clone(&transport) as Arc<dyn WebFetchTransport>,
1575 Arc::clone(&resolver) as Arc<dyn WebFetchResolver>,
1576 );
1577
1578 let result = tool
1579 .execute(
1580 serde_json::json!({"url": "https://public.test/start"}),
1581 approval_context(Some(ToolApprovalStatus::Approved)),
1582 )
1583 .await;
1584
1585 assert!(!result.success);
1586 assert!(result.output.contains("cannot be reached by redirect"));
1587 assert_eq!(transport.requests(), vec!["https://public.test/start"]);
1588 assert!(resolver.requests().iter().all(|host| host == "public.test"));
1589 }
1590
1591 #[tokio::test]
1592 async fn blocks_private_redirect_before_second_request() {
1593 let transport = Arc::new(FixtureTransport::default());
1594 transport.route(
1595 "https://public.test/start",
1596 response(302, None, Some("http://private.test/secret"), ""),
1597 );
1598 let mut resolver = FixtureResolver::default();
1599 resolver.addresses.insert(
1600 "private.test".to_string(),
1601 vec![IpAddr::from([10, 0, 0, 1])],
1602 );
1603 let tool = WebFetchTool::with_transport_and_resolver(
1604 Arc::clone(&transport) as Arc<dyn WebFetchTransport>,
1605 Arc::new(resolver),
1606 );
1607
1608 let result = tool
1609 .execute(
1610 serde_json::json!({"url": "https://public.test/start"}),
1611 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1612 )
1613 .await;
1614
1615 assert!(!result.success);
1616 assert_eq!(transport.requests(), vec!["https://public.test/start"]);
1617 assert!(result.output.contains("Non-public and special-use"));
1618 }
1619
1620 #[tokio::test]
1621 async fn enforces_byte_limit_on_injected_responses() {
1622 let transport = Arc::new(FixtureTransport::default());
1623 transport.route(
1624 "https://public.test/large",
1625 response(200, Some("text/plain"), None, "123456"),
1626 );
1627 let result = fixture_tool(transport)
1628 .execute(
1629 serde_json::json!({
1630 "url": "https://public.test/large",
1631 "max_response_bytes": 5
1632 }),
1633 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1634 )
1635 .await;
1636
1637 assert!(!result.success);
1638 assert!(result.output.contains("max_response_bytes 5"));
1639 }
1640
1641 #[tokio::test]
1642 async fn response_byte_limit_applies_independently_to_each_redirect_response() {
1643 let transport = Arc::new(FixtureTransport::default());
1644 transport.route(
1645 "https://public.test/start",
1646 response(302, None, Some("/final"), "12345"),
1647 );
1648 transport.route(
1649 "https://public.test/final",
1650 response(200, Some("text/plain"), None, "abcde"),
1651 );
1652
1653 let result = fixture_tool(Arc::clone(&transport))
1654 .execute(
1655 serde_json::json!({
1656 "url": "https://public.test/start",
1657 "max_response_bytes": 5
1658 }),
1659 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1660 )
1661 .await;
1662
1663 assert!(result.success);
1664 assert_eq!(output(&result)["content"], "abcde");
1665 assert_eq!(transport.requests().len(), 2);
1666 }
1667
1668 #[test]
1669 fn cache_evicts_oldest_entry_at_capacity() {
1670 let stored_at = Instant::now();
1671 let mut cache = WebFetchCache::new(2);
1672 let lifetime = Duration::from_secs(60);
1673 cache
1674 .insert(
1675 "first".to_string(),
1676 cached_output("https://public.test/first"),
1677 lifetime,
1678 stored_at,
1679 )
1680 .unwrap();
1681 cache
1682 .insert(
1683 "second".to_string(),
1684 cached_output("https://public.test/second"),
1685 lifetime,
1686 stored_at + Duration::from_secs(1),
1687 )
1688 .unwrap();
1689 cache
1690 .insert(
1691 "third".to_string(),
1692 cached_output("https://public.test/third"),
1693 lifetime,
1694 stored_at + Duration::from_secs(2),
1695 )
1696 .unwrap();
1697
1698 let checked_at = stored_at + Duration::from_secs(3);
1699 assert_eq!(cache.entries.len(), 2);
1700 assert!(cache.get("first", checked_at).is_none());
1701 assert!(cache.get("second", checked_at).is_some());
1702 assert!(cache.get("third", checked_at).is_some());
1703 }
1704
1705 #[test]
1706 fn cache_lazily_removes_expired_entries_without_refreshing_hits() {
1707 let stored_at = Instant::now();
1708 let mut cache = WebFetchCache::new(3);
1709 cache
1710 .insert(
1711 "short".to_string(),
1712 cached_output("https://public.test/short"),
1713 Duration::from_secs(1),
1714 stored_at,
1715 )
1716 .unwrap();
1717 cache
1718 .insert(
1719 "long".to_string(),
1720 cached_output("https://public.test/long"),
1721 Duration::from_secs(10),
1722 stored_at,
1723 )
1724 .unwrap();
1725
1726 assert!(
1727 cache
1728 .get("short", stored_at + Duration::from_millis(500))
1729 .is_some()
1730 );
1731 assert!(
1732 cache
1733 .get("long", stored_at + Duration::from_secs(2))
1734 .is_some()
1735 );
1736 assert_eq!(cache.entries.len(), 1);
1737 assert!(
1738 cache
1739 .get("short", stored_at + Duration::from_secs(2))
1740 .is_none()
1741 );
1742 }
1743
1744 #[tokio::test]
1745 async fn unrepresentable_cache_ttl_fails_before_dns_and_transport() {
1746 let transport = Arc::new(FixtureTransport::default());
1747 let resolver = Arc::new(FixtureResolver::default());
1748 let tool = WebFetchTool::with_transport_and_resolver(
1749 Arc::clone(&transport) as Arc<dyn WebFetchTransport>,
1750 Arc::clone(&resolver) as Arc<dyn WebFetchResolver>,
1751 );
1752
1753 let result = tool
1754 .execute(
1755 serde_json::json!({
1756 "url": "https://public.test/page",
1757 "cache_ttl_seconds": u64::MAX
1758 }),
1759 ToolExecutionContext::test("web_fetch"),
1760 )
1761 .await;
1762
1763 assert!(!result.success);
1764 assert!(result.output.contains("cache_ttl_seconds is too large"));
1765 assert!(resolver.requests().is_empty());
1766 assert!(transport.requests().is_empty());
1767 assert!(transport.validated_addresses().is_empty());
1768 }
1769
1770 #[tokio::test]
1771 async fn caches_injected_transport_responses() {
1772 let transport = Arc::new(FixtureTransport::default());
1773 transport.route(
1774 "https://public.test/cached",
1775 response(200, Some("text/plain"), None, "cached"),
1776 );
1777 let tool = fixture_tool(Arc::clone(&transport));
1778 let args = serde_json::json!({
1779 "url": "https://public.test/cached",
1780 "cache_ttl_seconds": 60
1781 });
1782
1783 let first = tool
1784 .execute(
1785 args.clone(),
1786 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1787 )
1788 .await;
1789 let second = tool
1790 .execute(
1791 args,
1792 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1793 )
1794 .await;
1795
1796 assert!(first.success && second.success);
1797 assert_eq!(transport.requests().len(), 1);
1798 assert_eq!(output(&first)["from_cache"], false);
1799 assert_eq!(output(&second)["from_cache"], true);
1800 }
1801
1802 #[test]
1803 fn effective_redirect_limit_uses_the_strictest_supplied_value() {
1804 assert_eq!(effective_max_redirects(None, None, None), 5);
1805 assert_eq!(effective_max_redirects(Some(10), Some(3), Some(4)), 3);
1806 assert_eq!(effective_max_redirects(Some(0), Some(5), None), 0);
1807 assert_eq!(effective_max_redirects(Some(99), None, None), 20);
1808 }
1809
1810 #[tokio::test]
1811 async fn stricter_request_redirect_limit_uses_a_fresh_request_and_fails() {
1812 let transport = Arc::new(FixtureTransport::default());
1813 transport.route(
1814 "https://public.test/start",
1815 response(302, None, Some("/middle"), ""),
1816 );
1817 transport.route(
1818 "https://public.test/middle",
1819 response(302, None, Some("/final"), ""),
1820 );
1821 transport.route(
1822 "https://public.test/final",
1823 response(200, Some("text/plain"), None, "done"),
1824 );
1825 let tool = fixture_tool(Arc::clone(&transport));
1826
1827 let first = tool
1828 .execute(
1829 serde_json::json!({
1830 "url": "https://public.test/start",
1831 "cache_ttl_seconds": 60,
1832 "max_redirects": 5
1833 }),
1834 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1835 )
1836 .await;
1837 let second = tool
1838 .execute(
1839 serde_json::json!({
1840 "url": "https://public.test/start",
1841 "cache_ttl_seconds": 60,
1842 "max_redirects": 1
1843 }),
1844 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1845 )
1846 .await;
1847
1848 assert!(first.success);
1849 assert!(!second.success);
1850 assert!(second.output.contains("Redirect limit exceeded"));
1851 assert_eq!(
1852 transport.requests(),
1853 vec![
1854 "https://public.test/start".to_string(),
1855 "https://public.test/middle".to_string(),
1856 "https://public.test/final".to_string(),
1857 "https://public.test/start".to_string(),
1858 "https://public.test/middle".to_string(),
1859 ]
1860 );
1861 }
1862
1863 #[tokio::test]
1864 async fn zero_redirect_limit_reuses_direct_response_but_rejects_redirects() {
1865 let transport = Arc::new(FixtureTransport::default());
1866 transport.route(
1867 "https://public.test/direct",
1868 response(200, Some("text/plain"), None, "direct"),
1869 );
1870 transport.route(
1871 "https://public.test/start",
1872 response(302, None, Some("/final"), ""),
1873 );
1874 let tool = fixture_tool(Arc::clone(&transport));
1875 let direct_args = serde_json::json!({
1876 "url": "https://public.test/direct",
1877 "cache_ttl_seconds": 60,
1878 "max_redirects": 0
1879 });
1880
1881 let direct_first = tool
1882 .execute(
1883 direct_args.clone(),
1884 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1885 )
1886 .await;
1887 let direct_second = tool
1888 .execute(
1889 direct_args,
1890 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1891 )
1892 .await;
1893 let redirected = tool
1894 .execute(
1895 serde_json::json!({
1896 "url": "https://public.test/start",
1897 "cache_ttl_seconds": 60,
1898 "max_redirects": 0
1899 }),
1900 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1901 )
1902 .await;
1903
1904 assert!(direct_first.success && direct_second.success);
1905 assert_eq!(output(&direct_second)["from_cache"], true);
1906 assert!(!redirected.success);
1907 assert!(redirected.output.contains("Redirect limit exceeded"));
1908 assert_eq!(transport.requests().len(), 2);
1909 }
1910
1911 #[tokio::test]
1912 async fn policy_and_context_redirect_limits_participate_in_cache_separation() {
1913 let transport = Arc::new(FixtureTransport::default());
1914 transport.route(
1915 "https://public.test/start",
1916 response(302, None, Some("/final"), ""),
1917 );
1918 transport.route(
1919 "https://public.test/final",
1920 response(200, Some("text/plain"), None, "done"),
1921 );
1922 let tool = fixture_tool(Arc::clone(&transport));
1923 let args = serde_json::json!({
1924 "url": "https://public.test/start",
1925 "cache_ttl_seconds": 60
1926 });
1927 let mut permissive = ai_agents_core::ToolExecutionContext::test("web_fetch");
1928 permissive.policy_snapshot = serde_json::json!({"max_redirects": 5});
1929 permissive.limits.max_redirects = Some(5);
1930 let mut strict_policy = ai_agents_core::ToolExecutionContext::test("web_fetch");
1931 strict_policy.policy_snapshot = serde_json::json!({"max_redirects": 0});
1932 let mut strict_context = ai_agents_core::ToolExecutionContext::test("web_fetch");
1933 strict_context.limits.max_redirects = Some(0);
1934
1935 let first = tool.execute(args.clone(), permissive).await;
1936 let policy_result = tool.execute(args.clone(), strict_policy).await;
1937 let context_result = tool.execute(args, strict_context).await;
1938
1939 assert!(first.success);
1940 assert!(!policy_result.success);
1941 assert!(!context_result.success);
1942 assert!(policy_result.output.contains("Redirect limit exceeded"));
1943 assert!(context_result.output.contains("Redirect limit exceeded"));
1944 assert_eq!(transport.requests().len(), 4);
1945 }
1946
1947 #[tokio::test]
1948 async fn compatible_cache_hit_revalidates_dns_without_transport_request() {
1949 let transport = Arc::new(FixtureTransport::default());
1950 transport.route(
1951 "https://public.test/cached",
1952 response(200, Some("text/plain"), None, "cached"),
1953 );
1954 let resolver = Arc::new(FixtureResolver::default());
1955 let tool = WebFetchTool::with_transport_and_resolver(
1956 Arc::clone(&transport) as Arc<dyn WebFetchTransport>,
1957 Arc::clone(&resolver) as Arc<dyn WebFetchResolver>,
1958 );
1959 let args = serde_json::json!({
1960 "url": "https://public.test/cached",
1961 "cache_ttl_seconds": 60,
1962 "max_redirects": 0
1963 });
1964
1965 let first = tool
1966 .execute(
1967 args.clone(),
1968 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1969 )
1970 .await;
1971 let second = tool
1972 .execute(
1973 args,
1974 ai_agents_core::ToolExecutionContext::test("web_fetch"),
1975 )
1976 .await;
1977
1978 assert!(first.success && second.success);
1979 assert_eq!(transport.requests().len(), 1);
1980 assert_eq!(resolver.requests().len(), 4);
1981 assert_eq!(output(&second)["from_cache"], true);
1982 }
1983
1984 #[tokio::test]
1985 async fn cache_separates_output_and_response_byte_caps() {
1986 let transport = Arc::new(FixtureTransport::default());
1987 transport.route(
1988 "https://public.test/capped",
1989 response(200, Some("text/plain"), None, "abcdef"),
1990 );
1991 let tool = fixture_tool(Arc::clone(&transport));
1992
1993 let first = tool
1994 .execute(
1995 serde_json::json!({
1996 "url": "https://public.test/capped",
1997 "cache_ttl_seconds": 60,
1998 "max_chars": 6,
1999 "max_response_bytes": 6
2000 }),
2001 ai_agents_core::ToolExecutionContext::test("web_fetch"),
2002 )
2003 .await;
2004 let narrower_output = tool
2005 .execute(
2006 serde_json::json!({
2007 "url": "https://public.test/capped",
2008 "cache_ttl_seconds": 60,
2009 "max_chars": 3,
2010 "max_response_bytes": 6
2011 }),
2012 ai_agents_core::ToolExecutionContext::test("web_fetch"),
2013 )
2014 .await;
2015 let narrower_response = tool
2016 .execute(
2017 serde_json::json!({
2018 "url": "https://public.test/capped",
2019 "cache_ttl_seconds": 60,
2020 "max_chars": 6,
2021 "max_response_bytes": 5
2022 }),
2023 ai_agents_core::ToolExecutionContext::test("web_fetch"),
2024 )
2025 .await;
2026
2027 assert!(first.success && narrower_output.success);
2028 assert_eq!(output(&narrower_output)["content"], "abc");
2029 assert_eq!(output(&narrower_output)["truncated"], true);
2030 assert!(!narrower_response.success);
2031 assert!(narrower_response.output.contains("max_response_bytes 5"));
2032 assert_eq!(transport.requests().len(), 3);
2033 }
2034
2035 #[tokio::test]
2036 async fn oversized_redirect_body_is_rejected_before_following_location() {
2037 let transport = Arc::new(FixtureTransport::default());
2038 transport.route(
2039 "https://public.test/start",
2040 response(302, None, Some("/final"), "123456"),
2041 );
2042 transport.route(
2043 "https://public.test/final",
2044 response(200, Some("text/plain"), None, "done"),
2045 );
2046
2047 let result = fixture_tool(Arc::clone(&transport))
2048 .execute(
2049 serde_json::json!({
2050 "url": "https://public.test/start",
2051 "max_response_bytes": 5
2052 }),
2053 ai_agents_core::ToolExecutionContext::test("web_fetch"),
2054 )
2055 .await;
2056
2057 assert!(!result.success);
2058 assert!(result.output.contains("max_response_bytes 5"));
2059 assert_eq!(transport.requests(), vec!["https://public.test/start"]);
2060 }
2061
2062 #[tokio::test]
2063 async fn cached_redirect_count_above_limit_is_a_safe_miss() {
2064 let resolver = FixtureResolver::default();
2065 let cached = WebFetchOutput {
2066 url: "https://public.test/start".to_string(),
2067 final_url: "https://public.test/final".to_string(),
2068 status: 200,
2069 content_type: Some("text/plain".to_string()),
2070 content: "done".to_string(),
2071 truncated: false,
2072 from_cache: false,
2073 redirects: vec!["https://public.test/final".to_string()],
2074 extraction_prompt_used: false,
2075 extraction_available: false,
2076 };
2077
2078 let compatible = validate_cached_output_with_policy(
2079 &cached,
2080 &WebFetchPolicyInput::default(),
2081 1,
2082 &resolver,
2083 )
2084 .await
2085 .unwrap();
2086 let incompatible = validate_cached_output_with_policy(
2087 &cached,
2088 &WebFetchPolicyInput::default(),
2089 0,
2090 &resolver,
2091 )
2092 .await
2093 .unwrap();
2094
2095 assert!(compatible);
2096 assert!(!incompatible);
2097 }
2098
2099 #[tokio::test]
2100 async fn applies_policy_before_injected_transport() {
2101 let transport = Arc::new(FixtureTransport::default());
2102 transport.route(
2103 "https://blocked.test/page",
2104 response(200, Some("text/plain"), None, "not reached"),
2105 );
2106 let tool = fixture_tool(Arc::clone(&transport));
2107 let mut context = ai_agents_core::ToolExecutionContext::test("web_fetch");
2108 context.policy_snapshot = serde_json::json!({
2109 "blocked_domains": ["blocked.test"]
2110 });
2111
2112 let result = tool
2113 .execute(
2114 serde_json::json!({"url": "https://blocked.test/page"}),
2115 context,
2116 )
2117 .await;
2118
2119 assert!(!result.success);
2120 assert!(transport.requests().is_empty());
2121 assert!(result.output.contains("blocked by policy"));
2122 }
2123
2124 #[tokio::test]
2125 async fn reports_unconfigured_exact_routes() {
2126 let transport = Arc::new(FixtureTransport::default());
2127 let result = fixture_tool(Arc::clone(&transport))
2128 .execute(
2129 serde_json::json!({"url": "https://public.test/missing"}),
2130 ai_agents_core::ToolExecutionContext::test("web_fetch"),
2131 )
2132 .await;
2133
2134 assert!(!result.success);
2135 assert!(result.output.contains("Unconfigured web fetch route"));
2136 assert_eq!(transport.requests(), vec!["https://public.test/missing"]);
2137 }
2138
2139 #[test]
2140 fn configured_policy_blocks_redirect_targets_outside_allowlist() {
2141 let policy = WebFetchPolicyInput {
2142 domain_allow: vec!["docs.rs".to_string()],
2143 allowed_schemes: vec!["https".to_string()],
2144 allowed_ports: vec![443],
2145 ..WebFetchPolicyInput::default()
2146 };
2147 let allowed = reqwest::Url::parse("https://docs.rs/serde/latest/serde/").unwrap();
2148 let denied = reqwest::Url::parse("https://example.com/").unwrap();
2149
2150 assert!(check_configured_url_policy(&allowed, &policy, UrlPolicyTarget::Redirect).is_ok());
2151 assert!(check_configured_url_policy(&denied, &policy, UrlPolicyTarget::Redirect).is_err());
2152 }
2153}