1use std::sync::atomic::{AtomicU64, Ordering};
7use std::sync::{Arc, OnceLock};
8
9use super::{HttpTransportPolicy, TransportPolicyError};
10use std::time::{Duration, Instant};
11
12use futures::StreamExt as _;
13
14#[derive(Clone, Debug)]
15pub struct HttpRequest {
16 pub method: String,
18 pub url: String,
20 pub headers: Vec<(String, String)>,
22 pub body: Vec<u8>,
23 pub timeout_ms: u64,
24 pub max_response_size: u64,
26 pub decompress: bool,
28 pub transport_policy: Option<Arc<HttpTransportPolicy>>,
30 pub redirect_guard: Option<Arc<dyn RedirectGuard>>,
34 pub recorded_as: Option<RecordedRequest>,
39}
40
41#[derive(Clone, Debug, PartialEq, Eq)]
46pub struct RecordedRequest {
47 pub masked_url: String,
49 pub digest: String,
51}
52
53pub trait RedirectGuard: Send + Sync + std::fmt::Debug {
55 fn authorize(&self, hop: &RedirectHop<'_>) -> Result<(), RedirectDenied>;
56 fn audit_egress_denial(&self, _hop: &RedirectHop<'_>, _at: EgressAt) {}
57}
58
59#[derive(Debug, Clone, Copy, PartialEq, Eq)]
61pub enum EgressAt {
62 CurrentHop,
64 NewHop,
66}
67
68#[derive(Clone, Copy, Debug)]
70pub struct RedirectHop<'a> {
71 pub method: &'a str,
73 pub url: &'a url::Url,
74 pub method_rewritten: bool,
77 pub body_len: u64,
78}
79
80#[derive(Debug)]
83pub struct RedirectDenied(wasmtime::Error);
84
85impl RedirectDenied {
86 pub fn new(
88 caller: impl Into<String>,
89 capability: impl Into<String>,
90 reason: impl Into<String>,
91 ) -> Self {
92 Self(crate::runtime::host::permission_denied(
93 caller, capability, reason,
94 ))
95 }
96
97 pub(crate) fn from_error(error: wasmtime::Error) -> Self {
98 Self(error)
99 }
100
101 pub(crate) fn into_error(self) -> wasmtime::Error {
102 self.0
103 }
104}
105
106impl std::fmt::Display for RedirectDenied {
107 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
108 self.0.fmt(f)
109 }
110}
111
112impl std::error::Error for RedirectDenied {}
113
114#[derive(Clone, Debug)]
116pub struct HttpResponse {
117 pub status: u16,
118 pub status_text: String,
119 pub headers: Vec<(String, String)>,
121 pub body: Vec<u8>,
122 pub final_url: String,
123}
124
125#[derive(Clone, Debug)]
127pub struct DownloadMeta {
128 pub status: u16,
129 pub status_text: String,
130 pub headers: Vec<(String, String)>,
131 pub final_url: String,
132 pub bytes_written: u64,
133}
134
135#[derive(Debug)]
136pub enum HttpError {
137 Network(String),
138 EgressDenied(String),
139 Internal(String),
141 Policy(TransportPolicyError),
142 PermissionDenied(RedirectDenied),
144 Timeout,
145 TooLarge {
147 limit: u64,
148 },
149 UnsupportedMethod(String),
151 Other(String),
152}
153
154impl HttpError {
155 pub(super) fn kind(&self) -> &'static str {
157 match self {
158 HttpError::Network(_) => "network",
159 HttpError::EgressDenied(_) => "egress-denied",
160 HttpError::Internal(_) => "internal",
161 HttpError::Policy(_) => "policy",
162 HttpError::PermissionDenied(_) => "permission-denied",
163 HttpError::Timeout => "timeout",
164 HttpError::TooLarge { .. } => "too-large",
165 HttpError::UnsupportedMethod(_) => "unsupported-method",
166 HttpError::Other(_) => "other",
167 }
168 }
169}
170
171impl std::fmt::Display for HttpError {
172 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
173 match self {
174 HttpError::Internal(msg) => write!(f, "internal HTTP transport error: {msg}"),
175 HttpError::Policy(error) => error.fmt(f),
176 HttpError::PermissionDenied(denied) => denied.fmt(f),
177 HttpError::Network(msg) | HttpError::EgressDenied(msg) => {
178 write!(f, "network error: {msg}")
179 }
180 HttpError::Timeout => write!(f, "request timed out"),
181 HttpError::TooLarge { limit } => write!(
182 f,
183 "response too large (limit: {limit} bytes); consider http.download",
184 ),
185 HttpError::UnsupportedMethod(method) => {
186 write!(f, "unsupported HTTP method: {method}")
187 }
188 HttpError::Other(msg) => write!(f, "{msg}"),
189 }
190 }
191}
192
193impl std::error::Error for HttpError {}
194
195#[derive(Default)]
197pub struct DownloadProgress {
198 received: AtomicU64,
199 written: AtomicU64,
200}
201
202impl DownloadProgress {
203 pub fn received(&self, bytes: u64) {
205 let _ = self
206 .received
207 .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |total| {
208 Some(total.saturating_add(bytes))
209 });
210 }
211
212 pub fn bytes_received(&self) -> u64 {
213 self.received.load(Ordering::Relaxed)
214 }
215
216 pub(crate) fn written(&self, bytes: u64) {
217 let _ = self
218 .written
219 .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |total| {
220 Some(total.saturating_add(bytes))
221 });
222 }
223
224 pub(crate) fn bytes_written(&self) -> u64 {
225 self.written.load(Ordering::Relaxed)
226 }
227}
228
229struct ReceivedWriter<'a> {
230 writer: &'a mut (dyn std::io::Write + Send),
231 progress: &'a DownloadProgress,
232}
233
234impl std::io::Write for ReceivedWriter<'_> {
235 fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
236 let n = self.writer.write(bytes)?;
237 self.progress.received(n as u64);
238 Ok(n)
239 }
240
241 fn flush(&mut self) -> std::io::Result<()> {
242 self.writer.flush()
243 }
244}
245
246#[async_trait::async_trait]
255pub trait HttpClient: Send + Sync {
256 fn wants_recorded_request(&self) -> bool {
259 false
260 }
261
262 async fn send(&self, req: &HttpRequest) -> Result<HttpResponse, HttpError>;
263
264 async fn send_without_redirects_to(
274 &self,
275 _req: &HttpRequest,
276 _body: &mut (dyn std::io::Write + Send),
277 ) -> Result<HttpResponse, HttpError> {
278 Err(HttpError::Other(
279 "transport does not support requests without redirects".into(),
280 ))
281 }
282
283 async fn download(
288 &self,
289 req: &HttpRequest,
290 writer: &mut (dyn std::io::Write + Send),
291 ) -> Result<DownloadMeta, HttpError>;
292
293 async fn download_with_progress(
297 &self,
298 req: &HttpRequest,
299 writer: &mut (dyn std::io::Write + Send),
300 progress: &DownloadProgress,
301 ) -> Result<DownloadMeta, HttpError> {
302 self.download(req, &mut ReceivedWriter { writer, progress })
303 .await
304 }
305}
306
307pub struct ReqwestHttpClient {
315 client: OnceLock<Result<reqwest::Client, String>>,
316 policy: Arc<crate::stdlib::http::policy::NetworkPolicy>,
320}
321
322const MAX_REDIRECTS: usize = 10;
324
325const CROSS_ORIGIN_SENSITIVE_HEADERS: [&str; 5] = [
327 "authorization",
328 "cookie",
329 "cookie2",
330 "proxy-authorization",
331 "www-authenticate",
332];
333
334const BODY_HEADERS: [&str; 4] = [
336 "content-type",
337 "content-length",
338 "content-encoding",
339 "transfer-encoding",
340];
341
342impl Default for ReqwestHttpClient {
343 fn default() -> Self {
344 Self::new(Arc::new(
345 crate::stdlib::http::policy::NetworkPolicy::allow_all(),
346 ))
347 }
348}
349
350impl ReqwestHttpClient {
351 pub fn new(policy: Arc<crate::stdlib::http::policy::NetworkPolicy>) -> Self {
352 Self {
353 client: OnceLock::new(),
354 policy,
355 }
356 }
357
358 #[cfg(test)]
360 pub(super) fn with_client(
361 policy: Arc<crate::stdlib::http::policy::NetworkPolicy>,
362 customize: impl FnOnce(reqwest::ClientBuilder) -> reqwest::ClientBuilder,
363 ) -> Self {
364 let client = Self::new(policy);
365 let built = customize(client.client_builder())
366 .build()
367 .map_err(|error| error.to_string());
368 let _ = client.client.set(built);
369 client
370 }
371
372 fn client_builder(&self) -> reqwest::ClientBuilder {
373 self.policy
374 .client_builder()
375 .redirect(reqwest::redirect::Policy::none())
376 .referer(false)
377 }
378
379 fn client(&self) -> Result<&reqwest::Client, HttpError> {
380 self.client
381 .get_or_init(|| {
382 self.client_builder()
383 .build()
384 .map_err(|error| error.to_string())
385 })
386 .as_ref()
387 .map_err(|error| HttpError::Internal(error.clone()))
388 }
389
390 fn check_literal_ip(&self, url: &str) -> Result<(), HttpError> {
393 self.policy
394 .check_literal_host(url)
395 .map_err(HttpError::EgressDenied)
396 }
397
398 async fn follow_redirects(&self, req: &HttpRequest) -> Result<reqwest::Response, HttpError> {
402 let mut hop = self.initial_hop(req).inspect_err(|error| {
403 audit_egress(req, &req.url, &req.method, error, EgressAt::CurrentHop);
404 })?;
405 let initial = hop.url.clone();
406 let deadline = Deadline::new(req.timeout_ms);
407 let mut redirects = 0;
408 loop {
409 let resp = self.send_hop(&hop, &deadline).await.inspect_err(|error| {
410 audit_egress(
411 req,
412 hop.url.as_str(),
413 hop.method.as_str(),
414 error,
415 EgressAt::CurrentHop,
416 );
417 })?;
418 let Some(next) = redirect_location(&resp, &hop.url) else {
419 return Ok(resp);
420 };
421 if redirects >= MAX_REDIRECTS {
422 return Err(HttpError::Network(format!(
423 "too many redirects (limit: {MAX_REDIRECTS})"
424 )));
425 }
426 redirects += 1;
427 let status = resp.status();
428 drop(resp);
430 hop.redirect_to(status, next);
431 self.authorize_hop(req, &initial, &hop)?;
432 }
433 }
434
435 fn initial_hop<'a>(&self, req: &'a HttpRequest) -> Result<Hop<'a>, HttpError> {
437 let url = parse_url(&req.url)?;
438 if let Some(policy) = &req.transport_policy {
439 policy.check_destination(&url).map_err(HttpError::Policy)?;
440 }
441 self.check_literal_ip(&req.url)?;
442 Ok(Hop {
443 method: parse_method(&req.method)?,
444 url,
445 headers: req.headers.clone(),
446 body: &req.body,
447 method_rewritten: false,
448 })
449 }
450
451 fn authorize_hop(
454 &self,
455 req: &HttpRequest,
456 initial: &url::Url,
457 hop: &Hop<'_>,
458 ) -> Result<(), HttpError> {
459 if let Some(policy) = &req.transport_policy {
460 policy
461 .check_redirect(initial, &hop.url)
462 .map_err(HttpError::Policy)?;
463 }
464 if !matches!(hop.url.scheme(), "http" | "https") {
465 return Err(HttpError::Network(
466 "redirect to a URL that is not http or https".into(),
467 ));
468 }
469 self.check_literal_ip(hop.url.as_str())
470 .inspect_err(|error| {
471 audit_egress(
472 req,
473 hop.url.as_str(),
474 hop.method.as_str(),
475 error,
476 EgressAt::NewHop,
477 );
478 })?;
479 let Some(guard) = &req.redirect_guard else {
480 return Ok(());
481 };
482 guard
483 .authorize(&RedirectHop {
484 method: hop.method.as_str(),
485 url: &hop.url,
486 method_rewritten: hop.method_rewritten,
487 body_len: hop.body.len() as u64,
488 })
489 .map_err(HttpError::PermissionDenied)
490 }
491
492 async fn send_hop(
494 &self,
495 hop: &Hop<'_>,
496 deadline: &Deadline,
497 ) -> Result<reqwest::Response, HttpError> {
498 let mut rb = self
499 .client()?
500 .request(hop.method.clone(), hop.url.clone())
501 .timeout(deadline.remaining()?);
502 for (name, value) in &hop.headers {
503 rb = rb.header(name.as_str(), value.as_str());
504 }
505 if !hop.body.is_empty() {
506 rb = rb.body(hop.body.to_vec());
507 }
508 rb.send().await.map_err(map_reqwest_error)
509 }
510}
511
512struct Hop<'a> {
515 method: reqwest::Method,
516 url: url::Url,
518 headers: Vec<(String, String)>,
519 body: &'a [u8],
520 method_rewritten: bool,
522}
523
524impl Hop<'_> {
525 fn redirect_to(&mut self, status: reqwest::StatusCode, mut next: url::Url) {
530 let rewrites_method = match status.as_u16() {
531 301 | 302 => self.method == reqwest::Method::POST,
532 303 => ![reqwest::Method::GET, reqwest::Method::HEAD].contains(&self.method),
533 _ => false,
534 };
535 if rewrites_method {
536 self.method = reqwest::Method::GET;
537 self.method_rewritten = true;
538 }
539 if rewrites_method || status.as_u16() == 303 {
540 self.body = &[];
541 remove_headers(&mut self.headers, &BODY_HEADERS);
542 }
543 if self.url.origin() == next.origin() {
544 keep_userinfo(&self.url, &mut next);
545 } else {
546 remove_headers(&mut self.headers, &CROSS_ORIGIN_SENSITIVE_HEADERS);
547 }
548 self.url = next;
549 }
550}
551
552fn keep_userinfo(current: &url::Url, next: &mut url::Url) {
555 if !next.username().is_empty() || next.password().is_some() {
556 return;
557 }
558 let _ = next.set_username(current.username());
560 let _ = next.set_password(current.password());
561}
562
563fn remove_headers(headers: &mut Vec<(String, String)>, names: &[&str]) {
564 headers.retain(|(name, _)| {
565 !names
566 .iter()
567 .any(|removed| name.eq_ignore_ascii_case(removed))
568 });
569}
570
571fn redirect_location(resp: &reqwest::Response, current: &url::Url) -> Option<url::Url> {
575 if !matches!(resp.status().as_u16(), 301 | 302 | 303 | 307 | 308) {
576 return None;
577 }
578 let location = resp.headers().get(reqwest::header::LOCATION)?;
579 current
580 .join(std::str::from_utf8(location.as_bytes()).ok()?)
581 .ok()
582}
583
584fn parse_url(url: &str) -> Result<url::Url, HttpError> {
585 url::Url::parse(url).map_err(|_| HttpError::Network("invalid HTTP URL".into()))
586}
587
588fn parse_method(method: &str) -> Result<reqwest::Method, HttpError> {
589 reqwest::Method::from_bytes(method.to_ascii_uppercase().as_bytes())
590 .map_err(|_| HttpError::UnsupportedMethod(method.to_string()))
591}
592
593struct Deadline {
596 at: Option<Instant>,
598 timeout: Duration,
599}
600
601impl Deadline {
602 fn new(timeout_ms: u64) -> Self {
603 let timeout = Duration::from_millis(timeout_ms);
604 Self {
605 at: Instant::now().checked_add(timeout),
606 timeout,
607 }
608 }
609
610 fn remaining(&self) -> Result<Duration, HttpError> {
611 let Some(at) = self.at else {
612 return Ok(self.timeout);
613 };
614 let remaining = at.saturating_duration_since(Instant::now());
615 if remaining.is_zero() {
616 return Err(HttpError::Timeout);
617 }
618 Ok(remaining)
619 }
620}
621
622#[async_trait::async_trait]
623impl HttpClient for ReqwestHttpClient {
624 async fn send(&self, req: &HttpRequest) -> Result<HttpResponse, HttpError> {
625 let resp = self.follow_redirects(req).await?;
626 read_response(resp, req.max_response_size).await
627 }
628
629 async fn send_without_redirects_to(
630 &self,
631 req: &HttpRequest,
632 body: &mut (dyn std::io::Write + Send),
633 ) -> Result<HttpResponse, HttpError> {
634 let hop = self.initial_hop(req).inspect_err(|error| {
635 audit_egress(req, &req.url, &req.method, error, EgressAt::CurrentHop);
636 })?;
637 let resp = self
638 .send_hop(&hop, &Deadline::new(req.timeout_ms))
639 .await
640 .inspect_err(|error| {
641 audit_egress(
642 req,
643 hop.url.as_str(),
644 hop.method.as_str(),
645 error,
646 EgressAt::CurrentHop,
647 );
648 })?;
649 let status = resp.status().as_u16();
650 let status_text = resp.status().canonical_reason().unwrap_or("").to_string();
651 let final_url = resp.url().to_string();
652 let headers = collect_headers(resp.headers());
653 let limit = req.max_response_size;
654 let mut received: u64 = 0;
655 let mut stream = resp.bytes_stream();
656 while let Some(chunk) = stream.next().await {
657 let chunk = chunk.map_err(map_reqwest_error)?;
658 received = received.saturating_add(chunk.len() as u64);
659 if received > limit {
660 return Err(HttpError::TooLarge { limit });
661 }
662 body.write_all(&chunk)
663 .map_err(|error| HttpError::Other(error.to_string()))?;
664 }
665 Ok(HttpResponse {
666 status,
667 status_text,
668 headers,
669 body: Vec::new(),
670 final_url,
671 })
672 }
673
674 async fn download(
675 &self,
676 req: &HttpRequest,
677 writer: &mut (dyn std::io::Write + Send),
678 ) -> Result<DownloadMeta, HttpError> {
679 self.download_with_progress(req, writer, &DownloadProgress::default())
680 .await
681 }
682
683 async fn download_with_progress(
684 &self,
685 req: &HttpRequest,
686 writer: &mut (dyn std::io::Write + Send),
687 progress: &DownloadProgress,
688 ) -> Result<DownloadMeta, HttpError> {
689 if req.method != "GET" {
692 return Err(HttpError::Other(format!(
693 "http.download requires GET; got {}",
694 req.method,
695 )));
696 }
697 let resp = self.follow_redirects(req).await?;
698
699 let status = resp.status().as_u16();
700 let status_text = resp.status().canonical_reason().unwrap_or("").to_string();
701 let final_url = resp.url().to_string();
702 let headers = collect_headers(resp.headers());
703
704 let kind = detect_decompression(&headers, &req.url, req.decompress);
706 let limit = req.max_response_size;
707 let mut sink = DecodeSink::new(kind, LimitedWriter::new(writer, limit))?;
711 let mut wire: u64 = 0;
712 let mut stream = resp.bytes_stream();
713 while let Some(chunk) = stream.next().await {
714 let chunk = chunk.map_err(map_reqwest_error)?;
715 progress.received(chunk.len() as u64);
716 wire = wire.saturating_add(chunk.len() as u64);
717 if wire > limit {
718 return Err(HttpError::TooLarge { limit });
719 }
720 sink.write_all(&chunk)?;
721 }
722 let bytes_written = sink.finish()?;
723
724 Ok(DownloadMeta {
725 status,
726 status_text,
727 headers,
728 final_url,
729 bytes_written,
730 })
731 }
732}
733
734async fn read_response(resp: reqwest::Response, limit: u64) -> Result<HttpResponse, HttpError> {
737 let status = resp.status().as_u16();
738 let status_text = resp.status().canonical_reason().unwrap_or("").to_string();
739 let final_url = resp.url().to_string();
740 let headers = collect_headers(resp.headers());
741
742 let mut body_bytes: Vec<u8> = Vec::new();
743 let mut stream = resp.bytes_stream();
744 while let Some(chunk) = stream.next().await {
745 let chunk = chunk.map_err(map_reqwest_error)?;
746 if body_bytes.len() as u64 + chunk.len() as u64 > limit {
747 return Err(HttpError::TooLarge { limit });
748 }
749 body_bytes.extend_from_slice(&chunk);
750 }
751
752 Ok(HttpResponse {
753 status,
754 status_text,
755 headers,
756 body: body_bytes,
757 final_url,
758 })
759}
760
761fn collect_headers(map: &reqwest::header::HeaderMap) -> Vec<(String, String)> {
763 let mut headers = Vec::with_capacity(map.len());
764 for (name, value) in map {
765 if let Ok(v) = value.to_str() {
766 headers.push((name.as_str().to_ascii_lowercase(), v.to_string()));
767 }
768 }
769 headers
770}
771
772#[derive(Clone, Copy, Debug, PartialEq, Eq)]
773pub enum Decompression {
774 None,
775 Gzip,
776 Zstd,
777}
778
779pub fn detect_decompression(
782 headers: &[(String, String)],
783 url: &str,
784 decompress: bool,
785) -> Decompression {
786 if !decompress {
787 return Decompression::None;
788 }
789 let ce = headers
790 .iter()
791 .find(|(k, _)| k.eq_ignore_ascii_case("content-encoding"))
792 .map_or("", |(_, v)| v.as_str());
793 if ce.eq_ignore_ascii_case("gzip") || ce.eq_ignore_ascii_case("x-gzip") {
794 return Decompression::Gzip;
795 }
796 if ce.eq_ignore_ascii_case("zstd") {
797 return Decompression::Zstd;
798 }
799 let path = url.split(['?', '#']).next().unwrap_or(url);
801 let lower = path.to_ascii_lowercase();
802 if lower.ends_with(".gz") {
803 return Decompression::Gzip;
804 }
805 if lower.ends_with(".zst") {
806 return Decompression::Zstd;
807 }
808 Decompression::None
809}
810
811pub fn stream_to_writer(
816 mut reader: impl std::io::Read,
817 writer: &mut dyn std::io::Write,
818 kind: Decompression,
819 limit: u64,
820) -> Result<u64, HttpError> {
821 let mut sink = DecodeSink::new(kind, LimitedWriter::new(writer, limit))?;
822 let mut wire: u64 = 0;
823 let mut chunk = [0u8; 8 * 1024];
824 loop {
825 let n = reader
826 .read(&mut chunk)
827 .map_err(|e| HttpError::Network(format!("io: {e}")))?;
828 if n == 0 {
829 break;
830 }
831 wire = wire.saturating_add(n as u64);
832 if wire > limit {
833 return Err(HttpError::TooLarge { limit });
834 }
835 sink.write_all(&chunk[..n])?;
836 }
837 sink.finish()
838}
839
840struct LimitedWriter<W> {
843 inner: W,
844 written: u64,
845 limit: u64,
846 exceeded: bool,
847}
848
849impl<W: std::io::Write> LimitedWriter<W> {
850 fn new(inner: W, limit: u64) -> Self {
851 Self {
852 inner,
853 written: 0,
854 limit,
855 exceeded: false,
856 }
857 }
858}
859
860impl<W: std::io::Write> std::io::Write for LimitedWriter<W> {
861 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
862 if self.written.saturating_add(buf.len() as u64) > self.limit {
863 self.exceeded = true;
864 return Err(std::io::Error::other("download limit exceeded"));
865 }
866 let n = self.inner.write(buf)?;
867 self.written = self.written.saturating_add(n as u64);
868 Ok(n)
869 }
870
871 fn flush(&mut self) -> std::io::Result<()> {
872 self.inner.flush()
873 }
874}
875
876enum DecodeSink<W: std::io::Write> {
880 Plain(LimitedWriter<W>),
881 Gzip(flate2::write::MultiGzDecoder<LimitedWriter<W>>),
882 Zstd(zstd::stream::zio::Writer<LimitedWriter<W>, zstd::stream::raw::Decoder<'static>>),
883}
884
885impl<W: std::io::Write> DecodeSink<W> {
886 fn new(kind: Decompression, out: LimitedWriter<W>) -> Result<Self, HttpError> {
887 Ok(match kind {
888 Decompression::None => Self::Plain(out),
889 Decompression::Gzip => Self::Gzip(flate2::write::MultiGzDecoder::new(out)),
890 Decompression::Zstd => {
891 let decoder = zstd::stream::raw::Decoder::new()
892 .map_err(|e| HttpError::Network(format!("zstd: {e}")))?;
893 Self::Zstd(zstd::stream::zio::Writer::new(out, decoder))
894 }
895 })
896 }
897
898 fn write_all(&mut self, chunk: &[u8]) -> Result<(), HttpError> {
899 use std::io::Write;
900 let result = match self {
901 Self::Plain(out) => out.write_all(chunk),
902 Self::Gzip(decoder) => decoder.write_all(chunk),
903 Self::Zstd(decoder) => decoder.write_all(chunk),
904 };
905 result.map_err(|e| self.error(e))
906 }
907
908 fn finish(mut self) -> Result<u64, HttpError> {
910 use std::io::Write;
911 let result = match &mut self {
912 Self::Plain(out) => out.flush(),
913 Self::Gzip(decoder) => decoder.try_finish(),
914 Self::Zstd(decoder) => decoder.finish(),
915 };
916 result.map_err(|e| self.error(e))?;
917 Ok(self.out().written)
918 }
919
920 fn out(&self) -> &LimitedWriter<W> {
921 match self {
922 Self::Plain(out) => out,
923 Self::Gzip(decoder) => decoder.get_ref(),
924 Self::Zstd(decoder) => decoder.writer(),
925 }
926 }
927
928 fn error(&self, e: std::io::Error) -> HttpError {
929 let out = self.out();
930 if out.exceeded {
931 HttpError::TooLarge { limit: out.limit }
932 } else {
933 HttpError::Network(format!("io: {e}"))
934 }
935 }
936}
937
938fn audit_egress(req: &HttpRequest, url: &str, method: &str, error: &HttpError, at: EgressAt) {
939 if !matches!(error, HttpError::EgressDenied(_)) {
940 return;
941 }
942 if let (Some(guard), Ok(url)) = (&req.redirect_guard, url::Url::parse(url)) {
943 let hop = RedirectHop {
944 method,
945 url: &url,
946 method_rewritten: method != req.method,
947 body_len: 0,
948 };
949 guard.audit_egress_denial(&hop, at);
950 }
951}
952
953pub(super) fn http_failure_outcome(err: &HttpError) -> &'static str {
955 match err {
956 HttpError::Timeout => "timeout",
957 HttpError::TooLarge { .. } => "too_large",
958 HttpError::Network(_)
959 | HttpError::EgressDenied(_)
960 | HttpError::Internal(_)
961 | HttpError::Policy(_)
962 | HttpError::PermissionDenied(_)
963 | HttpError::UnsupportedMethod(_)
964 | HttpError::Other(_) => "error",
965 }
966}
967
968fn map_reqwest_error(err: reqwest::Error) -> HttpError {
975 if err.is_timeout() {
976 return HttpError::Timeout;
977 }
978 let mut source: Option<&(dyn std::error::Error + 'static)> = Some(&err);
979 while let Some(error) = source {
980 if error
981 .downcast_ref::<super::policy::EgressDenied>()
982 .is_some()
983 {
984 return HttpError::EgressDenied(describe_error_chain(&err.without_url()));
985 }
986 source = error.source();
987 }
988 HttpError::Network(describe_error_chain(&err.without_url()))
989}
990
991pub fn describe_error_chain(err: &dyn std::error::Error) -> String {
992 let top = err.to_string();
993 let mut deepest: Option<String> = None;
994 let mut current = err.source();
995 while let Some(cause) = current {
996 deepest = Some(cause.to_string());
997 current = cause.source();
998 }
999 match deepest {
1000 Some(reason) if !top.contains(&reason) => format!("{top}: {reason}"),
1001 _ => top,
1002 }
1003}
1004
1005pub fn is_policy_refusal(chain: &str) -> bool {
1008 chain.contains("blocked by network policy")
1009}
1010
1011pub fn default_http_client() -> Arc<dyn HttpClient> {
1012 Arc::new(ReqwestHttpClient::default())
1013}
1014
1015#[async_trait::async_trait]
1020pub trait AuthProxy: Send + Sync {
1021 async fn transform(
1022 &self,
1023 req: HttpRequest,
1024 caller: &str,
1025 ) -> Result<HttpRequest, AuthProxyError>;
1026}
1027
1028#[derive(Debug)]
1029pub enum AuthProxyError {
1030 UndeclaredSecret(String),
1032 MissingSecret(String),
1034 Other(String),
1035}
1036
1037impl std::fmt::Display for AuthProxyError {
1038 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1039 match self {
1040 AuthProxyError::UndeclaredSecret(name) => {
1041 write!(f, "undeclared secret '{name}'")
1042 }
1043 AuthProxyError::MissingSecret(name) => {
1044 write!(f, "missing secret '{name}'")
1045 }
1046 AuthProxyError::Other(msg) => write!(f, "{msg}"),
1047 }
1048 }
1049}
1050
1051impl std::error::Error for AuthProxyError {}
1052
1053pub struct NoopAuthProxy;
1054
1055#[async_trait::async_trait]
1056impl AuthProxy for NoopAuthProxy {
1057 async fn transform(
1058 &self,
1059 req: HttpRequest,
1060 _caller: &str,
1061 ) -> Result<HttpRequest, AuthProxyError> {
1062 Ok(req)
1063 }
1064}
1065
1066pub fn default_auth_proxy() -> Arc<dyn AuthProxy> {
1067 Arc::new(NoopAuthProxy)
1068}
1069
1070#[cfg(test)]
1071mod tests {
1072 use super::describe_error_chain;
1073
1074 #[derive(Debug)]
1075 struct Layer {
1076 message: &'static str,
1077 source: Option<Box<Layer>>,
1078 }
1079
1080 impl std::fmt::Display for Layer {
1081 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1082 f.write_str(self.message)
1083 }
1084 }
1085
1086 impl std::error::Error for Layer {
1087 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
1088 self.source
1089 .as_deref()
1090 .map(|layer| layer as &(dyn std::error::Error + 'static))
1091 }
1092 }
1093
1094 #[test]
1095 fn appends_the_root_cause_when_the_top_message_omits_it() {
1096 let err = Layer {
1097 message: "error sending request for url (http://x/)",
1098 source: Some(Box::new(Layer {
1099 message: "client error (Connect)",
1100 source: Some(Box::new(Layer {
1101 message: "blocked by network policy: x resolves only to private/loopback IP space",
1102 source: None,
1103 })),
1104 })),
1105 };
1106 assert_eq!(
1107 describe_error_chain(&err),
1108 "error sending request for url (http://x/): blocked by network policy: x resolves only to private/loopback IP space"
1109 );
1110 }
1111
1112 #[test]
1113 fn leaves_a_message_that_already_carries_its_cause_alone() {
1114 let err = Layer {
1115 message: "timeout: deadline elapsed",
1116 source: Some(Box::new(Layer {
1117 message: "deadline elapsed",
1118 source: None,
1119 })),
1120 };
1121 assert_eq!(describe_error_chain(&err), "timeout: deadline elapsed");
1122 let bare = Layer {
1123 message: "plain",
1124 source: None,
1125 };
1126 assert_eq!(describe_error_chain(&bare), "plain");
1127 }
1128}
1129
1130#[cfg(test)]
1131#[path = "transport_tls_tests.rs"]
1132mod tls_tests;
1133
1134#[cfg(test)]
1135mod decode_sink_tests {
1136 use std::io::Write as _;
1137
1138 use super::{DecodeSink, Decompression, HttpError, LimitedWriter};
1139
1140 fn gzip(data: &[u8]) -> Vec<u8> {
1141 let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::fast());
1142 encoder.write_all(data).unwrap();
1143 encoder.finish().unwrap()
1144 }
1145
1146 fn decode(kind: Decompression, wire: &[u8], limit: u64) -> (Result<u64, HttpError>, Vec<u8>) {
1148 let mut out: Vec<u8> = Vec::new();
1149 let result = (|| {
1150 let mut sink = DecodeSink::new(kind, LimitedWriter::new(&mut out, limit))?;
1151 for chunk in wire.chunks(7) {
1152 sink.write_all(chunk)?;
1153 }
1154 sink.finish()
1155 })();
1156 (result, out)
1157 }
1158
1159 #[test]
1160 fn plain_bytes_stop_at_the_limit() {
1161 let (ok, out) = decode(Decompression::None, b"hello", 5);
1162 assert_eq!(ok.unwrap(), 5);
1163 assert_eq!(out, b"hello");
1164 let (over, _) = decode(Decompression::None, b"hello!", 5);
1165 assert!(matches!(over, Err(HttpError::TooLarge { limit: 5 })));
1166 }
1167
1168 #[test]
1169 fn gzip_decodes_across_chunks_and_rejects_truncation() {
1170 let wire = gzip(b"hello gzipped download");
1171 let (ok, out) = decode(Decompression::Gzip, &wire, 1024);
1172 assert_eq!(ok.unwrap(), 22);
1173 assert_eq!(out, b"hello gzipped download");
1174
1175 let (truncated, _) = decode(Decompression::Gzip, &wire[..wire.len() - 4], 1024);
1176 assert!(
1177 matches!(truncated, Err(HttpError::Network(_))),
1178 "{truncated:?}"
1179 );
1180 }
1181
1182 #[test]
1183 fn a_small_gzip_body_that_inflates_past_the_limit_is_too_large() {
1184 let wire = gzip(&vec![b'x'; 100_000]);
1185 assert!(wire.len() < 1_000);
1186 let (result, out) = decode(Decompression::Gzip, &wire, 1_000);
1187 assert!(
1188 matches!(result, Err(HttpError::TooLarge { limit: 1_000 })),
1189 "{result:?}"
1190 );
1191 assert!(out.len() <= 1_000);
1192 }
1193
1194 #[test]
1195 fn zstd_decodes_and_rejects_an_incomplete_frame() {
1196 let wire = zstd::encode_all(&b"hello zstd download"[..], 1).unwrap();
1197 let (ok, out) = decode(Decompression::Zstd, &wire, 1024);
1198 assert_eq!(ok.unwrap(), 19);
1199 assert_eq!(out, b"hello zstd download");
1200
1201 let (truncated, _) = decode(Decompression::Zstd, &wire[..wire.len() - 3], 1024);
1202 assert!(
1203 matches!(truncated, Err(HttpError::Network(_))),
1204 "{truncated:?}"
1205 );
1206 }
1207}
1208
1209#[cfg(test)]
1210mod client_setup_tests {
1211 use super::*;
1212
1213 #[test]
1214 fn invalid_client_configuration_is_retained_as_an_internal_failure() {
1215 let client = ReqwestHttpClient::with_client(
1216 Arc::new(crate::stdlib::http::policy::NetworkPolicy::allow_all()),
1217 |builder| builder.user_agent("\n"),
1218 );
1219 for _ in 0..2 {
1220 assert!(matches!(client.client(), Err(HttpError::Internal(_))));
1221 }
1222 let healthy = ReqwestHttpClient::default();
1223 assert!(healthy.client().is_ok());
1224 assert!(healthy.client().is_ok());
1225 }
1226}