返回 CodeWhale
fetch.rs
根目录 / crates / tui / src / tools / web / fetch.rs
1 //! Unified guarded fetch pipeline for `fetch_url` and `web.run`.
2
3 use std::collections::BTreeMap;
4 use std::sync::Arc;
5 use std::time::{Duration, Instant};
6 #[cfg(not(test))]
7 use std::time::{SystemTime, UNIX_EPOCH};
8
9 use futures_util::StreamExt;
10 use serde::Serialize;
11
12 use super::adapter::AdapterResult;
13 use super::cache::{self, CachedFetch};
14 use super::extract::is_js_shell_error;
15 use super::guard::{
16 DnsPin, guarded_reqwest_client_builder, validate_fetch_target, validate_network_policy,
17 };
18 use crate::features::Feature;
19 use crate::tools::spec::{ToolContext, ToolError};
20 use crate::worker_profile::ShellPolicy;
21
22 pub(crate) const DEFAULT_TIMEOUT: Duration = Duration::from_secs(15);
23 pub(crate) const HARD_MAX_TIMEOUT: Duration = Duration::from_secs(60);
24 pub(crate) const DEFAULT_MAX_BYTES: usize = 1_000_000;
25 pub(crate) const HARD_MAX_BYTES: usize = 10 * 1024 * 1024;
26 const MAX_REDIRECTS: usize = 5;
27 const USER_AGENT: &str = concat!(
28 "Mozilla/5.0 (compatible; codewhale/",
29 env!("CARGO_PKG_VERSION"),
30 "; +https://github.com/codewhale-hq/CodeWhale)"
31 );
32
33 #[derive(Debug, Clone)]
34 pub(crate) struct FetchOptions {
35 pub(crate) timeout: Duration,
36 pub(crate) max_bytes: usize,
37 pub(crate) accept: &'static str,
38 pub(crate) user_agent: &'static str,
39 }
40
41 impl FetchOptions {
42 pub(crate) fn new(timeout: Duration, max_bytes: usize, accept: &'static str) -> Self {
43 Self {
44 timeout: timeout.min(HARD_MAX_TIMEOUT),
45 max_bytes: max_bytes.clamp(1, HARD_MAX_BYTES),
46 accept,
47 user_agent: USER_AGENT,
48 }
49 }
50
51 /// Request with the shared browser user-agent instead of the Codewhale
52 /// one. [`fetch_readable`] uses this only as the one-shot fallback after
53 /// a site refused the default agent with 401/403.
54 #[must_use]
55 pub(crate) fn with_browser_user_agent(mut self) -> Self {
56 self.user_agent = super::scrape::BROWSER_USER_AGENT;
57 self
58 }
59
60 fn uses_browser_user_agent(&self) -> bool {
61 self.user_agent == super::scrape::BROWSER_USER_AGENT
62 }
63 }
64
65 #[derive(Debug, Clone)]
66 pub(crate) struct FetchedPayload {
67 pub(crate) url: String,
68 pub(crate) status: u16,
69 pub(crate) headers: BTreeMap<String, String>,
70 pub(crate) content_type: String,
71 pub(crate) bytes: Arc<Vec<u8>>,
72 pub(crate) truncated: bool,
73 pub(crate) cache_hit: bool,
74 pub(crate) retries: usize,
75 pub(crate) redirects: usize,
76 }
77
78 /// Whether one request may be answered from a cache, or must revalidate.
79 ///
80 /// `Revalidate` bypasses the session fetch cache *and* asks every intermediary
81 /// to revalidate. An edge cache can hold a prerendered variant while an origin
82 /// MISS serves the client-side shell, so the same URL alternates between
83 /// readable and unreadable depending on which variant answered (#5904).
84 #[derive(Debug, Clone, Copy, PartialEq, Eq)]
85 enum CacheMode {
86 Default,
87 Revalidate,
88 }
89
90 impl CacheMode {
91 const fn is_revalidate(self) -> bool {
92 matches!(self, Self::Revalidate)
93 }
94 }
95
96 /// Response headers that explain the cache state behind a 200.
97 ///
98 /// These are the four that distinguish "the edge served a prerendered page"
99 /// from "the origin served the JavaScript shell", so both the success and the
100 /// failure receipt carry whichever of them the response actually had.
101 const CACHE_STATE_HEADERS: [&str; 4] = [
102 "age",
103 "cf-cache-status",
104 "x-nextjs-prerender",
105 "x-vercel-cache",
106 ];
107
108 /// One request inside a readable-fetch sequence, as it appears on the receipt.
109 #[derive(Debug, Clone, Serialize)]
110 pub(crate) struct FetchAttempt {
111 /// 1-based position in the sequence.
112 pub(crate) attempt: usize,
113 pub(crate) status: u16,
114 /// Whether the session fetch cache answered this attempt.
115 pub(crate) cache_hit: bool,
116 /// Whether this attempt sent `Cache-Control: no-cache` / `Pragma: no-cache`
117 /// and skipped the session cache.
118 pub(crate) cache_busted: bool,
119 /// Whether this attempt is the one that yielded a readable document.
120 pub(crate) produced_content: bool,
121 /// Whether this attempt used the browser user-agent after the default
122 /// agent was refused with 401/403.
123 #[serde(skip_serializing_if = "std::ops::Not::not")]
124 pub(crate) browser_user_agent: bool,
125 /// `age`, `cf-cache-status`, `x-nextjs-prerender`, `x-vercel-cache` — only
126 /// those the response actually carried.
127 #[serde(skip_serializing_if = "BTreeMap::is_empty")]
128 pub(crate) cache_headers: BTreeMap<String, String>,
129 }
130
131 impl FetchAttempt {
132 fn record(
133 payload: &FetchedPayload,
134 attempt: usize,
135 mode: CacheMode,
136 browser_user_agent: bool,
137 ) -> Self {
138 Self {
139 attempt,
140 status: payload.status,
141 cache_hit: payload.cache_hit,
142 cache_busted: mode.is_revalidate(),
143 produced_content: false,
144 browser_user_agent,
145 cache_headers: cache_state_headers(&payload.headers),
146 }
147 }
148
149 fn summarize(&self) -> String {
150 let mut facts = vec![format!("HTTP {}", self.status)];
151 if self.cache_hit {
152 facts.push("session cache hit".to_string());
153 }
154 if self.browser_user_agent {
155 facts.push("browser user-agent".to_string());
156 }
157 for (name, value) in &self.cache_headers {
158 facts.push(format!("{name}={value}"));
159 }
160 let label = if self.cache_busted {
161 format!("attempt {} (Cache-Control: no-cache)", self.attempt)
162 } else {
163 format!("attempt {}", self.attempt)
164 };
165 format!("{label}: {}", facts.join(", "))
166 }
167 }
168
169 fn cache_state_headers(headers: &BTreeMap<String, String>) -> BTreeMap<String, String> {
170 CACHE_STATE_HEADERS
171 .iter()
172 .filter_map(|name| {
173 headers
174 .get(*name)
175 .map(|value| ((*name).to_string(), value.clone()))
176 })
177 .collect()
178 }
179
180 /// The extraction step [`fetch_readable`] may run against either attempt.
181 ///
182 /// Spelled as an explicit boxed future over an *owned* payload rather than as
183 /// an `AsyncFn` over a borrowed one: the callers live inside
184 /// `async fn execute(&self, .., &ToolContext)` futures that must stay `Send`,
185 /// and a higher-ranked borrow of the payload would force the returned future
186 /// to outlive the tool context it reads.
187 pub(crate) type ExtractFuture<'a, T> =
188 std::pin::Pin<Box<dyn Future<Output = AdapterResult<T>> + Send + 'a>>;
189
190 /// A fetch that produced a readable document, plus the attempts it took.
191 #[derive(Debug)]
192 pub(crate) struct ReadableFetch<T> {
193 pub(crate) payload: FetchedPayload,
194 pub(crate) document: T,
195 pub(crate) attempts: Vec<FetchAttempt>,
196 }
197
198 /// Fetch `url` and extract it, re-fetching once past every cache when a 2xx
199 /// response yields no readable content.
200 ///
201 /// This is the single place that turns the JS-shell case into either a second
202 /// chance or an error the model can act on. `extract` runs against the fetched
203 /// payload; only [`is_js_shell_error`] failures earn the second request, so
204 /// transport failures keep the existing single-retry behavior of one request.
205 pub(crate) async fn fetch_readable<'e, T, F>(
206 url: &str,
207 options: &FetchOptions,
208 context: &ToolContext,
209 tool_label: &str,
210 extract: F,
211 ) -> Result<ReadableFetch<T>, ToolError>
212 where
213 F: Fn(FetchedPayload) -> ExtractFuture<'e, T>,
214 {
215 fetch_readable_inner(url, options, context, tool_label, None, extract).await
216 }
217
218 #[cfg(test)]
219 pub(crate) async fn fetch_readable_with_initial_pin<'e, T, F>(
220 url: &str,
221 options: &FetchOptions,
222 context: &ToolContext,
223 tool_label: &str,
224 initial_pin: DnsPin,
225 extract: F,
226 ) -> Result<ReadableFetch<T>, ToolError>
227 where
228 F: Fn(FetchedPayload) -> ExtractFuture<'e, T>,
229 {
230 fetch_readable_inner(
231 url,
232 options,
233 context,
234 tool_label,
235 Some(initial_pin),
236 extract,
237 )
238 .await
239 }
240
241 async fn fetch_readable_inner<'e, T, F>(
242 url: &str,
243 options: &FetchOptions,
244 context: &ToolContext,
245 tool_label: &str,
246 test_initial_pin: Option<DnsPin>,
247 extract: F,
248 ) -> Result<ReadableFetch<T>, ToolError>
249 where
250 F: Fn(FetchedPayload) -> ExtractFuture<'e, T>,
251 {
252 let mut attempts: Vec<FetchAttempt> = Vec::with_capacity(3);
253 let mut options = options.clone();
254 for mode in [CacheMode::Default, CacheMode::Revalidate] {
255 let mut payload = fetch_inner(
256 url,
257 &options,
258 context,
259 tool_label,
260 test_initial_pin.clone(),
261 mode,
262 )
263 .await?;
264 // Many sites refuse non-browser agents outright. One retry as a
265 // browser is the fallback; a second refusal is final.
266 if matches!(payload.status, 401 | 403) && !options.uses_browser_user_agent() {
267 attempts.push(FetchAttempt::record(
268 &payload,
269 attempts.len() + 1,
270 mode,
271 false,
272 ));
273 options = options.with_browser_user_agent();
274 payload = fetch_inner(
275 url,
276 &options,
277 context,
278 tool_label,
279 test_initial_pin.clone(),
280 mode,
281 )
282 .await?;
283 }
284 let mut record = FetchAttempt::record(
285 &payload,
286 attempts.len() + 1,
287 mode,
288 options.uses_browser_user_agent(),
289 );
290 let final_url = payload.url.clone();
291 match extract(payload.clone()).await {
292 Ok(document) => {
293 record.produced_content = true;
294 attempts.push(record);
295 return Ok(ReadableFetch {
296 payload,
297 document,
298 attempts,
299 });
300 }
301 // A 2xx whose body held no readable content is the one failure a
302 // second request can fix: the first response may have been a
303 // cached client-side shell.
304 Err(error)
305 if error.content()
306 && is_js_shell_error(&error.error)
307 && (200..300).contains(&payload.status)
308 && mode == CacheMode::Default =>
309 {
310 attempts.push(record);
311 }
312 Err(error) => {
313 attempts.push(record);
314 return Err(if error.content() && is_js_shell_error(&error.error) {
315 js_shell_failure(&final_url, &attempts, context)
316 } else {
317 error.error
318 });
319 }
320 }
321 }
322 unreachable!("the revalidate pass either returns a document or an error");
323 }
324
325 /// The terminal JS-shell error, carrying the failure receipt and the recovery
326 /// the *calling role* actually owns.
327 fn js_shell_failure(url: &str, attempts: &[FetchAttempt], context: &ToolContext) -> ToolError {
328 let receipt = attempts
329 .iter()
330 .map(FetchAttempt::summarize)
331 .collect::<Vec<_>>()
332 .join("; ");
333 ToolError::execution_failed(format!(
334 "{marker} {url} after {count} attempts, the second past every cache ({receipt}). The response parsed but held no readable body, which usually means the page renders its content with JavaScript. Recovery: {recovery}",
335 marker = super::extract::JS_SHELL_MARKER,
336 count = attempts.len(),
337 recovery = js_shell_recovery(url, context),
338 ))
339 }
340
341 /// Whether the model-facing `Web` tool is reachable from this context.
342 ///
343 /// Both facts already exist: the web family is feature-gated, and a
344 /// network-denied Fleet worker carries `network_access: Some(false)` on the
345 /// authority envelope that also removes the web tools from its registry
346 /// (`fleet::role::NETWORK_TOOL_DENYLIST`). Nothing new is registered here.
347 fn web_tool_available(context: &ToolContext) -> bool {
348 context.features.enabled(Feature::WebSearch) && network_authorized(context)
349 }
350
351 /// Whether this role could shell out to `curl` as a last resort. Read-only and
352 /// shell-less roles cannot: the read-only grammar rejects a network fetch.
353 fn shell_fallback_available(context: &ToolContext) -> bool {
354 context.shell_policy == ShellPolicy::Full && network_authorized(context)
355 }
356
357 fn network_authorized(context: &ToolContext) -> bool {
358 context
359 .tool_authority
360 .as_deref()
361 .is_none_or(|authority| authority.network_access != Some(false))
362 }
363
364 /// The raw-HTML fetch a JS-shell recovery suggests, as a `Web` tool input.
365 fn raw_fetch_call(url: &str) -> serde_json::Value {
366 serde_json::json!({"action": "fetch", "url": url, "format": "raw"})
367 }
368
369 /// Honest next steps after a JS shell. `web.run` shares this fetch, its
370 /// browser-agent retry and this extractor, and no Codewhale web tool runs
371 /// JavaScript, so re-opening the page is never suggested.
372 fn js_shell_recovery(url: &str, context: &ToolContext) -> String {
373 let same_result = "no Codewhale web tool runs JavaScript (`web.run` shares this fetch and extractor), so opening the URL again returns the same shell.";
374 if web_tool_available(context) {
375 return format!(
376 "{same_result} Many JavaScript pages embed their content as JSON in a script tag; to look for it, fetch the raw HTML with `Web {call}`. Otherwise search for another source of the same content.",
377 call = raw_fetch_call(url),
378 );
379 }
380 if shell_fallback_available(context) {
381 format!(
382 "{same_result} Inspect the raw HTML with a shell fetch (`curl -sSL`) for embedded data, or use a rendering tool."
383 )
384 } else {
385 format!(
386 "{same_result} This role is read-only and cannot fall back to a shell fetch, so report this URL as unreadable rather than substituting another source."
387 )
388 }
389 }
390
391 #[cfg(test)]
392 pub(crate) async fn fetch_with_initial_pin(
393 url: &str,
394 options: &FetchOptions,
395 context: &ToolContext,
396 tool_label: &str,
397 initial_pin: DnsPin,
398 ) -> Result<FetchedPayload, ToolError> {
399 fetch_inner(
400 url,
401 options,
402 context,
403 tool_label,
404 Some(initial_pin),
405 CacheMode::Default,
406 )
407 .await
408 }
409
410 async fn fetch_inner(
411 url: &str,
412 options: &FetchOptions,
413 context: &ToolContext,
414 tool_label: &str,
415 test_initial_pin: Option<DnsPin>,
416 cache_mode: CacheMode,
417 ) -> Result<FetchedPayload, ToolError> {
418 let initial_url = reqwest::Url::parse(url)
419 .map_err(|err| ToolError::invalid_input(format!("invalid URL: {err}")))?;
420 if !matches!(initial_url.scheme(), "http" | "https") {
421 return Err(ToolError::invalid_input(
422 "only http:// and https:// URLs are supported",
423 ));
424 }
425
426 // Validation precedes cache lookup so a policy tightened during the
427 // session cannot be bypassed by a previously cached response.
428 let validated_initial_pin = match test_initial_pin {
429 Some(pin) => pin,
430 None => validate_fetch_target(&initial_url, context, tool_label).await?,
431 };
432
433 if let Some(cached) = (!cache_mode.is_revalidate())
434 .then(|| {
435 cache::get(
436 &context.state_namespace,
437 &initial_url,
438 options.accept,
439 options.max_bytes,
440 )
441 })
442 .flatten()
443 {
444 let cached_url = reqwest::Url::parse(&cached.url).map_err(|err| {
445 ToolError::execution_failed(format!("cached response URL was invalid: {err}"))
446 })?;
447 let cached_host = cached_url.host_str().ok_or_else(|| {
448 ToolError::execution_failed("cached response URL did not include a host")
449 })?;
450 // No network request occurs on a cache hit, so DNS/SSRF validation is
451 // unnecessary. The final redirect destination still needs a policy
452 // check in case the session policy was tightened after insertion.
453 validate_network_policy(cached_host, context, tool_label)?;
454 return Ok(from_cached(cached, true, 0));
455 }
456
457 let deadline = Instant::now() + options.timeout;
458 let mut last_transient = None;
459 for attempt in 0..=1 {
460 let remaining = deadline.saturating_duration_since(Instant::now());
461 if remaining.is_zero() {
462 break;
463 }
464 match fetch_attempt(
465 initial_url.clone(),
466 options,
467 context,
468 tool_label,
469 remaining,
470 validated_initial_pin.clone(),
471 cache_mode,
472 )
473 .await
474 {
475 Ok(payload) if is_transient_status(payload.status) && attempt == 0 => {
476 last_transient = Some(format!("HTTP {}", payload.status));
477 }
478 Ok(payload) => {
479 let fetched = from_cached(payload.clone(), false, attempt);
480 if (200..300).contains(&payload.status) {
481 cache::insert(
482 &context.state_namespace,
483 &initial_url,
484 options.accept,
485 payload,
486 );
487 }
488 return Ok(fetched);
489 }
490 Err(AttemptError::Fatal(error)) => return Err(error),
491 Err(AttemptError::Transient(message)) if attempt == 0 => {
492 last_transient = Some(message);
493 }
494 Err(AttemptError::Transient(message)) => {
495 return Err(ToolError::execution_failed(format!(
496 "request failed after one retry: {message}"
497 )));
498 }
499 }
500
501 let delay = retry_delay();
502 if deadline.saturating_duration_since(Instant::now()) <= delay {
503 break;
504 }
505 tokio::time::sleep(delay).await;
506 }
507
508 Err(ToolError::execution_failed(format!(
509 "request timed out before retry completed{}",
510 last_transient
511 .map(|message| format!(" (last failure: {message})"))
512 .unwrap_or_default()
513 )))
514 }
515
516 #[derive(Debug)]
517 enum AttemptError {
518 Fatal(ToolError),
519 Transient(String),
520 }
521
522 async fn fetch_attempt(
523 initial_url: reqwest::Url,
524 options: &FetchOptions,
525 context: &ToolContext,
526 tool_label: &str,
527 timeout: Duration,
528 initial_pin: DnsPin,
529 cache_mode: CacheMode,
530 ) -> Result<CachedFetch, AttemptError> {
531 let mut current_url = initial_url;
532 let mut redirects = 0usize;
533 let mut initial_pin = initial_pin;
534 let deadline = Instant::now() + timeout;
535
536 let response = loop {
537 let dns_pin = if redirects == 0 {
538 match initial_pin.take() {
539 Some(pin) => Some(pin),
540 None => validate_fetch_target(&current_url, context, tool_label)
541 .await
542 .map_err(AttemptError::Fatal)?,
543 }
544 } else {
545 validate_fetch_target(&current_url, context, tool_label)
546 .await
547 .map_err(AttemptError::Fatal)?
548 };
549
550 let remaining = deadline.saturating_duration_since(Instant::now());
551 if remaining.is_zero() {
552 return Err(AttemptError::Transient(
553 "request timed out while following redirects".to_string(),
554 ));
555 }
556 let mut builder = guarded_reqwest_client_builder()
557 .timeout(remaining)
558 .user_agent(options.user_agent)
559 .redirect(reqwest::redirect::Policy::none());
560 if let Some((hostname, validated_ip)) = dns_pin {
561 builder = builder.resolve(&hostname, std::net::SocketAddr::new(validated_ip, 0));
562 }
563 let client = builder.build().map_err(|err| {
564 AttemptError::Fatal(ToolError::execution_failed(format!(
565 "failed to build HTTP client: {err}"
566 )))
567 })?;
568 let mut request = client
569 .get(current_url.clone())
570 .header("Accept", options.accept)
571 .header("Accept-Language", "en-US,en;q=0.5");
572 if cache_mode.is_revalidate() {
573 // `no-cache` (revalidate), not `no-store`: the shared caches still
574 // get to serve a validated copy, which is what recovers a page
575 // whose prerendered variant exists but was not the one served.
576 // `Pragma` is the HTTP/1.0 spelling some CDNs still honor.
577 request = request
578 .header("Cache-Control", "no-cache")
579 .header("Pragma", "no-cache");
580 }
581 let response = request
582 .send()
583 .await
584 .map_err(|err| AttemptError::Transient(err.to_string()))?;
585
586 if !response.status().is_redirection() {
587 break response;
588 }
589 if redirects >= MAX_REDIRECTS {
590 return Err(AttemptError::Fatal(ToolError::execution_failed(
591 "request exceeded the five-redirect limit",
592 )));
593 }
594 let Some(location) = response
595 .headers()
596 .get(reqwest::header::LOCATION)
597 .and_then(|value| value.to_str().ok())
598 else {
599 break response;
600 };
601 current_url = response.url().join(location).map_err(|err| {
602 AttemptError::Fatal(ToolError::execution_failed(format!(
603 "invalid redirect location: {err}"
604 )))
605 })?;
606 redirects += 1;
607 };
608
609 let final_url = response.url().to_string();
610 let status = response.status().as_u16();
611 let content_type = response
612 .headers()
613 .get(reqwest::header::CONTENT_TYPE)
614 .and_then(|value| value.to_str().ok())
615 .unwrap_or("application/octet-stream")
616 .to_string();
617 let headers = response_headers(response.headers());
618 let mut stream = response.bytes_stream();
619 let mut bytes = Vec::with_capacity(options.max_bytes.min(64 * 1024));
620 let mut truncated = false;
621 while let Some(chunk) = stream.next().await {
622 let chunk = chunk.map_err(|err| AttemptError::Transient(err.to_string()))?;
623 let remaining = options.max_bytes.saturating_sub(bytes.len());
624 if chunk.len() > remaining {
625 bytes.extend_from_slice(&chunk[..remaining]);
626 truncated = true;
627 break;
628 }
629 bytes.extend_from_slice(&chunk);
630 if bytes.len() == options.max_bytes {
631 // A response exactly at the cap may be complete. Ask for one more
632 // chunk to distinguish exact length from actual truncation.
633 if let Some(next) = stream.next().await {
634 let next = next.map_err(|err| AttemptError::Transient(err.to_string()))?;
635 truncated = !next.is_empty();
636 }
637 break;
638 }
639 }
640
641 Ok(CachedFetch {
642 url: final_url,
643 status,
644 headers,
645 content_type,
646 bytes: Arc::new(bytes),
647 truncated,
648 redirects,
649 })
650 }
651
652 fn from_cached(payload: CachedFetch, cache_hit: bool, retries: usize) -> FetchedPayload {
653 FetchedPayload {
654 url: payload.url,
655 status: payload.status,
656 headers: payload.headers,
657 content_type: payload.content_type,
658 bytes: payload.bytes,
659 truncated: payload.truncated,
660 cache_hit,
661 retries,
662 redirects: payload.redirects,
663 }
664 }
665
666 fn response_headers(headers: &reqwest::header::HeaderMap) -> BTreeMap<String, String> {
667 headers
668 .iter()
669 .filter(|(name, _)| {
670 !matches!(
671 name.as_str(),
672 "authorization"
673 | "proxy-authorization"
674 | "cookie"
675 | "set-cookie"
676 | "set-cookie2"
677 | "x-api-key"
678 | "api-key"
679 )
680 })
681 .filter_map(|(name, value)| {
682 value
683 .to_str()
684 .ok()
685 .map(|value| (name.as_str().to_ascii_lowercase(), value.to_string()))
686 })
687 .collect()
688 }
689
690 fn is_transient_status(status: u16) -> bool {
691 (500..600).contains(&status)
692 }
693
694 fn retry_delay() -> Duration {
695 #[cfg(test)]
696 return Duration::ZERO;
697
698 #[cfg(not(test))]
699 {
700 let jitter_ms = SystemTime::now()
701 .duration_since(UNIX_EPOCH)
702 .map(|duration| u64::from(duration.subsec_nanos()) % 41)
703 .unwrap_or(0);
704 Duration::from_millis(30 + jitter_ms)
705 }
706 }
707
708 #[cfg(test)]
709 mod tests {
710 use std::collections::BTreeMap;
711 use std::net::{IpAddr, Ipv4Addr};
712 use std::sync::Arc;
713 use std::sync::atomic::{AtomicUsize, Ordering};
714
715 use serde_json::json;
716 use wiremock::matchers::{method, path};
717 use wiremock::{Mock, MockServer, Request, Respond, ResponseTemplate};
718
719 use super::*;
720
721 fn context(namespace: &str) -> ToolContext {
722 ToolContext::new(".").with_state_namespace(namespace)
723 }
724
725 fn pin() -> DnsPin {
726 Some((
727 "public.example".to_string(),
728 IpAddr::V4(Ipv4Addr::LOCALHOST),
729 ))
730 }
731
732 #[derive(Clone)]
733 struct FailOnce {
734 calls: Arc<AtomicUsize>,
735 }
736
737 impl Respond for FailOnce {
738 fn respond(&self, _request: &Request) -> ResponseTemplate {
739 if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
740 ResponseTemplate::new(503).set_body_json(json!({"error": "retry"}))
741 } else {
742 ResponseTemplate::new(200)
743 .insert_header("content-type", "text/plain")
744 .set_body_string("recovered response")
745 }
746 }
747 }
748
749 #[tokio::test]
750 async fn transient_server_error_retries_once_then_caches() {
751 let server = MockServer::start().await;
752 let calls = Arc::new(AtomicUsize::new(0));
753 Mock::given(method("GET"))
754 .and(path("/retry"))
755 .respond_with(FailOnce {
756 calls: Arc::clone(&calls),
757 })
758 .mount(&server)
759 .await;
760 let url = format!("http://public.example:{}/retry", server.address().port());
761 let options = FetchOptions::new(Duration::from_secs(5), 1_024, "text/plain");
762 let context = context("fetch-retry-cache");
763
764 let first = fetch_with_initial_pin(&url, &options, &context, "test", pin())
765 .await
766 .expect("retry succeeds");
767 assert_eq!(first.status, 200);
768 assert_eq!(first.retries, 1);
769 assert!(!first.cache_hit);
770 assert_eq!(&*first.bytes, b"recovered response");
771
772 let second = fetch_with_initial_pin(&url, &options, &context, "test", pin())
773 .await
774 .expect("cache hit");
775 assert!(second.cache_hit);
776 assert_eq!(calls.load(Ordering::SeqCst), 2);
777 }
778
779 #[tokio::test]
780 async fn truncated_cache_refetches_when_larger_body_is_requested() {
781 let server = MockServer::start().await;
782 Mock::given(method("GET"))
783 .and(path("/large"))
784 .respond_with(
785 ResponseTemplate::new(200)
786 .insert_header("content-type", "text/plain")
787 .set_body_string("0123456789"),
788 )
789 .mount(&server)
790 .await;
791 let url = format!("http://public.example:{}/large", server.address().port());
792 let context = context("fetch-truncated-refetch");
793
794 let small = fetch_with_initial_pin(
795 &url,
796 &FetchOptions::new(Duration::from_secs(5), 4, "text/plain"),
797 &context,
798 "test",
799 pin(),
800 )
801 .await
802 .expect("small fetch");
803 assert_eq!(&*small.bytes, b"0123");
804 assert!(small.truncated);
805
806 let large = fetch_with_initial_pin(
807 &url,
808 &FetchOptions::new(Duration::from_secs(5), 16, "text/plain"),
809 &context,
810 "test",
811 pin(),
812 )
813 .await
814 .expect("larger refetch");
815 assert_eq!(&*large.bytes, b"0123456789");
816 assert!(!large.truncated);
817 assert!(!large.cache_hit);
818 }
819
820 /// A Vercel-style edge that serves the client-side shell to an ordinary
821 /// request and the prerendered page to a revalidating one (#5904).
822 #[derive(Clone)]
823 struct ShellUntilRevalidated {
824 calls: Arc<AtomicUsize>,
825 always_shell: bool,
826 }
827
828 const JS_SHELL_BODY: &str = "<html><head><title>Pricing</title></head><body><div id='root'></div><script>boot()</script></body></html>";
829 const PRERENDERED_BODY: &str = "<html><head><title>Pricing</title></head><body><main><h1>Pricing</h1><p>The prerendered variant carries the full pricing table for every plan.</p></main></body></html>";
830
831 impl Respond for ShellUntilRevalidated {
832 fn respond(&self, request: &Request) -> ResponseTemplate {
833 self.calls.fetch_add(1, Ordering::SeqCst);
834 let revalidating = request
835 .headers
836 .get("cache-control")
837 .and_then(|value| value.to_str().ok())
838 .is_some_and(|value| value.contains("no-cache"));
839 if revalidating && !self.always_shell {
840 ResponseTemplate::new(200)
841 .insert_header("content-type", "text/html")
842 .insert_header("x-vercel-cache", "HIT")
843 .insert_header("x-nextjs-prerender", "1")
844 .insert_header("age", "12")
845 .set_body_string(PRERENDERED_BODY)
846 } else {
847 ResponseTemplate::new(200)
848 .insert_header("content-type", "text/html")
849 .insert_header("x-vercel-cache", "MISS")
850 .set_body_string(JS_SHELL_BODY)
851 }
852 }
853 }
854
855 async fn js_shell_server(always_shell: bool) -> (MockServer, Arc<AtomicUsize>) {
856 let server = MockServer::start().await;
857 let calls = Arc::new(AtomicUsize::new(0));
858 Mock::given(method("GET"))
859 .and(path("/pricing"))
860 .respond_with(ShellUntilRevalidated {
861 calls: Arc::clone(&calls),
862 always_shell,
863 })
864 .mount(&server)
865 .await;
866 (server, calls)
867 }
868
869 fn extract_html_document(
870 payload: FetchedPayload,
871 ) -> ExtractFuture<'static, super::super::extract::ExtractedDocument> {
872 Box::pin(async move {
873 super::super::extract::extract_document(
874 &payload.url,
875 Some(&payload.content_type),
876 &payload.bytes,
877 None,
878 )
879 .await
880 })
881 }
882
883 #[tokio::test]
884 async fn js_shell_is_refetched_past_every_cache_before_it_becomes_an_error() {
885 let (server, calls) = js_shell_server(false).await;
886 let url = format!("http://public.example:{}/pricing", server.address().port());
887 let context = context("js-shell-recovers");
888
889 let readable = fetch_readable_with_initial_pin(
890 &url,
891 &FetchOptions::new(Duration::from_secs(5), 65_536, "text/html"),
892 &context,
893 "fetch_url",
894 pin(),
895 extract_html_document,
896 )
897 .await
898 .expect("the revalidated response carries the prerendered page");
899
900 assert!(
901 readable.document.markdown.contains("full pricing table"),
902 "the second attempt's content must be what the caller receives: {}",
903 readable.document.markdown
904 );
905 assert_eq!(calls.load(Ordering::SeqCst), 2, "exactly one extra request");
906 assert_eq!(readable.attempts.len(), 2);
907 assert!(!readable.attempts[0].cache_busted);
908 assert!(!readable.attempts[0].produced_content);
909 assert_eq!(
910 readable.attempts[0].cache_headers.get("x-vercel-cache"),
911 Some(&"MISS".to_string()),
912 "the failing attempt keeps the header that explains its cache state"
913 );
914 assert!(readable.attempts[1].cache_busted);
915 assert!(readable.attempts[1].produced_content);
916 assert_eq!(
917 readable.attempts[1].cache_headers.get("x-vercel-cache"),
918 Some(&"HIT".to_string())
919 );
920 assert_eq!(
921 readable.attempts[1].cache_headers.get("x-nextjs-prerender"),
922 Some(&"1".to_string())
923 );
924 assert_eq!(
925 readable.attempts[1].cache_headers.get("age"),
926 Some(&"12".to_string())
927 );
928 }
929
930 #[tokio::test]
931 async fn two_shells_fail_with_the_escalation_the_calling_role_owns() {
932 let (server, calls) = js_shell_server(true).await;
933 let url = format!("http://public.example:{}/pricing", server.address().port());
934 let options = FetchOptions::new(Duration::from_secs(5), 65_536, "text/html");
935
936 let error = fetch_readable_with_initial_pin(
937 &url,
938 &options,
939 &context("js-shell-browser-role"),
940 "fetch_url",
941 pin(),
942 extract_html_document,
943 )
944 .await
945 .expect_err("two shells must fail");
946 let message = error.to_string();
947 assert_eq!(calls.load(Ordering::SeqCst), 2, "no third request");
948 assert!(
949 message.contains("no Codewhale web tool runs JavaScript"),
950 "the recovery must not send the model to a re-open that fails the same way: {message}"
951 );
952 // The suggested call must be one the `Web` tool schema accepts.
953 let call = message
954 .split_once("`Web ")
955 .and_then(|(_, tail)| tail.split_once('`'))
956 .map(|(call, _)| call)
957 .expect("the recovery names a Web call");
958 let call: serde_json::Value = serde_json::from_str(call).expect("suggested call is JSON");
959 assert_eq!(call, raw_fetch_call(&url));
960 let schema = crate::tools::spec::ToolSpec::input_schema(
961 &crate::tools::web_tool::WebTool::new("Web"),
962 );
963 let validator = jsonschema::options()
964 .with_draft(jsonschema::Draft::Draft202012)
965 .build(&schema)
966 .expect("Web schema compiles");
967 assert!(
968 validator.is_valid(&call),
969 "suggested call must satisfy the Web schema: {call}"
970 );
971 assert!(
972 message.contains("attempt 1")
973 && message.contains("attempt 2 (Cache-Control: no-cache)"),
974 "the failure receipt names both attempts: {message}"
975 );
976 assert!(
977 message.contains("x-vercel-cache=MISS"),
978 "the failure receipt carries the cache-state headers: {message}"
979 );
980
981 // A read-only worker whose envelope denies network keeps `Web{fetch}`
982 // but loses `web.run` and any shell fallback, so the error must say so
983 // instead of naming a surface the role cannot call.
984 let mut denied = context("js-shell-read-only-role");
985 denied.shell_policy = ShellPolicy::ReadOnly;
986 denied.execution.tool_authority =
987 Some(Arc::new(crate::tools::spec::ToolAuthorityEnvelope {
988 schema_version: 1,
989 owner: "scout".to_string(),
990 authority: crate::tools::spec::ToolMutationAuthority::ReadOnly,
991 network_access: Some(false),
992 shell: crate::tools::spec::ToolShellAuthority::ReadOnly,
993 verification: crate::tools::spec::ToolVerificationAuthority::None,
994 writable_roots: Vec::new(),
995 writable_files: Vec::new(),
996 coordination_contracts: Vec::new(),
997 }));
998 let error = fetch_readable_with_initial_pin(
999 &url,
1000 &options,
1001 &denied,
1002 "fetch_url",
1003 pin(),
1004 extract_html_document,
1005 )
1006 .await
1007 .expect_err("two shells must fail");
1008 let message = error.to_string();
1009 assert!(
1010 !message.contains("`Web "),
1011 "a role without the web tools must not be sent to one: {message}"
1012 );
1013 assert!(
1014 message.contains("cannot fall back to a shell fetch"),
1015 "read-only roles must not be sent to curl: {message}"
1016 );
1017 }
1018
1019 #[derive(Clone)]
1020 struct RefuseBots {
1021 calls: Arc<AtomicUsize>,
1022 }
1023
1024 impl Respond for RefuseBots {
1025 fn respond(&self, request: &Request) -> ResponseTemplate {
1026 self.calls.fetch_add(1, Ordering::SeqCst);
1027 let agent = request
1028 .headers
1029 .get("user-agent")
1030 .and_then(|value| value.to_str().ok())
1031 .unwrap_or_default();
1032 if agent.contains("codewhale") {
1033 ResponseTemplate::new(403).set_body_string("bots not welcome")
1034 } else {
1035 ResponseTemplate::new(200)
1036 .insert_header("content-type", "text/plain")
1037 .set_body_string("browser body")
1038 }
1039 }
1040 }
1041
1042 #[tokio::test]
1043 async fn forbidden_default_agent_retries_once_as_a_browser() {
1044 let server = MockServer::start().await;
1045 let calls = Arc::new(AtomicUsize::new(0));
1046 Mock::given(method("GET"))
1047 .and(path("/guarded"))
1048 .respond_with(RefuseBots {
1049 calls: Arc::clone(&calls),
1050 })
1051 .mount(&server)
1052 .await;
1053 let url = format!("http://public.example:{}/guarded", server.address().port());
1054
1055 let readable = fetch_readable_with_initial_pin(
1056 &url,
1057 &FetchOptions::new(Duration::from_secs(5), 1_024, "text/plain"),
1058 &context("fetch-403-browser-retry"),
1059 "fetch_url",
1060 pin(),
1061 |payload: FetchedPayload| {
1062 Box::pin(async move { Ok(String::from_utf8_lossy(&payload.bytes).into_owned()) })
1063 },
1064 )
1065 .await
1066 .expect("the browser-agent retry reads the page");
1067
1068 assert_eq!(readable.document, "browser body");
1069 assert_eq!(readable.payload.status, 200);
1070 assert_eq!(calls.load(Ordering::SeqCst), 2, "exactly one browser retry");
1071 assert_eq!(readable.attempts.len(), 2);
1072 assert_eq!(readable.attempts[0].status, 403);
1073 assert!(!readable.attempts[0].browser_user_agent);
1074 assert!(readable.attempts[1].browser_user_agent);
1075 assert!(readable.attempts[1].produced_content);
1076 }
1077
1078 #[tokio::test]
1079 async fn a_browser_refusal_is_final() {
1080 let server = MockServer::start().await;
1081 let calls = Arc::new(AtomicUsize::new(0));
1082 Mock::given(method("GET"))
1083 .and(path("/closed"))
1084 .respond_with({
1085 let calls = Arc::clone(&calls);
1086 move |_: &Request| {
1087 calls.fetch_add(1, Ordering::SeqCst);
1088 ResponseTemplate::new(403)
1089 }
1090 })
1091 .mount(&server)
1092 .await;
1093 let url = format!("http://public.example:{}/closed", server.address().port());
1094
1095 let readable = fetch_readable_with_initial_pin(
1096 &url,
1097 &FetchOptions::new(Duration::from_secs(5), 1_024, "text/plain"),
1098 &context("fetch-403-final"),
1099 "fetch_url",
1100 pin(),
1101 |payload: FetchedPayload| Box::pin(async move { Ok(payload.status) }),
1102 )
1103 .await
1104 .expect("the 403 is handed to the caller, which owns non-2xx rendering");
1105
1106 assert_eq!(readable.document, 403);
1107 assert_eq!(calls.load(Ordering::SeqCst), 2, "no third request");
1108 }
1109
1110 #[tokio::test]
1111 async fn transport_failures_do_not_earn_a_cache_busting_refetch() {
1112 let server = MockServer::start().await;
1113 let calls = Arc::new(AtomicUsize::new(0));
1114 Mock::given(method("GET"))
1115 .and(path("/flaky"))
1116 .respond_with(FailOnce {
1117 calls: Arc::clone(&calls),
1118 })
1119 .mount(&server)
1120 .await;
1121 let url = format!("http://public.example:{}/flaky", server.address().port());
1122 let context = context("js-shell-transport-retry");
1123
1124 let readable = fetch_readable_with_initial_pin(
1125 &url,
1126 &FetchOptions::new(Duration::from_secs(5), 1_024, "text/plain"),
1127 &context,
1128 "fetch_url",
1129 pin(),
1130 |payload: FetchedPayload| {
1131 Box::pin(async move { Ok(String::from_utf8_lossy(&payload.bytes).into_owned()) })
1132 },
1133 )
1134 .await
1135 .expect("the existing transport retry still recovers");
1136
1137 assert_eq!(readable.document, "recovered response");
1138 assert_eq!(
1139 calls.load(Ordering::SeqCst),
1140 2,
1141 "the 503 costs the existing single transport retry and nothing more"
1142 );
1143 assert_eq!(
1144 readable.attempts.len(),
1145 1,
1146 "a transport retry is not a readable-fetch attempt"
1147 );
1148 assert_eq!(readable.payload.retries, 1);
1149 assert!(!readable.attempts[0].cache_busted);
1150 assert!(readable.attempts[0].produced_content);
1151 }
1152
1153 #[test]
1154 fn fetch_user_agent_tracks_the_crate_version() {
1155 assert!(
1156 USER_AGENT.contains(concat!("codewhale/", env!("CARGO_PKG_VERSION"))),
1157 "guarded-fetch UA must never pin a stale release: {USER_AGENT}"
1158 );
1159 }
1160
1161 #[test]
1162 fn response_headers_drop_set_cookie_values() {
1163 let mut headers = reqwest::header::HeaderMap::new();
1164 headers.insert("content-type", "text/plain".parse().unwrap());
1165 headers.insert("set-cookie", "session=secret".parse().unwrap());
1166 headers.insert("x-api-key", "secret".parse().unwrap());
1167
1168 let filtered = response_headers(&headers);
1169
1170 assert_eq!(
1171 filtered.get("content-type").map(String::as_str),
1172 Some("text/plain")
1173 );
1174 assert!(!filtered.contains_key("set-cookie"));
1175 assert!(!filtered.contains_key("x-api-key"));
1176 }
1177
1178 #[tokio::test]
1179 async fn tightened_network_policy_blocks_an_existing_cache_entry() {
1180 use crate::network_policy::{Decision, NetworkPolicy, NetworkPolicyDecider};
1181
1182 let url = reqwest::Url::parse("https://example.com/cached").unwrap();
1183 cache::insert(
1184 "policy-cache",
1185 &url,
1186 "text/plain",
1187 CachedFetch {
1188 url: url.to_string(),
1189 status: 200,
1190 headers: BTreeMap::new(),
1191 content_type: "text/plain".to_string(),
1192 bytes: Arc::new(b"cached".to_vec()),
1193 truncated: false,
1194 redirects: 0,
1195 },
1196 );
1197 let policy = NetworkPolicy {
1198 default: Decision::Deny.into(),
1199 allow: Vec::new(),
1200 deny: Vec::new(),
1201 proxy: Vec::new(),
1202 proxy_fake_ip_cidrs: Vec::new(),
1203 audit: false,
1204 };
1205 let context =
1206 context("policy-cache").with_network_policy(NetworkPolicyDecider::new(policy, None));
1207
1208 let error = fetch_inner(
1209 url.as_str(),
1210 &FetchOptions::new(Duration::from_secs(1), 100, "text/plain"),
1211 &context,
1212 "fetch_url",
1213 None,
1214 CacheMode::Default,
1215 )
1216 .await
1217 .expect_err("policy must win over cache");
1218 assert!(error.to_string().contains("blocked by network policy"));
1219 }
1220
1221 #[tokio::test]
1222 async fn tightened_network_policy_checks_cached_redirect_destination() {
1223 use crate::network_policy::{Decision, NetworkPolicy, NetworkPolicyDecider};
1224
1225 let initial_url = reqwest::Url::parse("https://8.8.8.8/cached").unwrap();
1226 cache::insert(
1227 "redirect-policy-cache",
1228 &initial_url,
1229 "text/plain",
1230 CachedFetch {
1231 url: "https://1.1.1.1/redirected".to_string(),
1232 status: 200,
1233 headers: BTreeMap::new(),
1234 content_type: "text/plain".to_string(),
1235 bytes: Arc::new(b"cached".to_vec()),
1236 truncated: false,
1237 redirects: 1,
1238 },
1239 );
1240 let policy = NetworkPolicy {
1241 default: Decision::Allow.into(),
1242 allow: Vec::new(),
1243 deny: vec!["1.1.1.1".to_string()],
1244 proxy: Vec::new(),
1245 proxy_fake_ip_cidrs: Vec::new(),
1246 audit: false,
1247 };
1248 let context = context("redirect-policy-cache")
1249 .with_network_policy(NetworkPolicyDecider::new(policy, None));
1250
1251 let error = fetch_inner(
1252 initial_url.as_str(),
1253 &FetchOptions::new(Duration::from_secs(1), 100, "text/plain"),
1254 &context,
1255 "fetch_url",
1256 None,
1257 CacheMode::Default,
1258 )
1259 .await
1260 .expect_err("final redirect policy must win over cache");
1261 assert!(error.to_string().contains("1.1.1.1"));
1262 assert!(error.to_string().contains("blocked by network policy"));
1263 }
1264 }
1265
1265 lines RUST