1use std::cell::RefCell;
6use std::collections::hash_map::HashMap;
7use std::convert::TryFrom;
8use std::sync::{Arc, LazyLock};
9use std::time::Duration;
10use std::{fmt, io};
11
12use futures::task::{Context, Poll};
13use futures::{Future, TryFutureExt};
14use http::uri::{Authority, Uri as Destination};
15use http_body_util::combinators::BoxBody;
16use hyper::body::Bytes;
17use hyper::rt::Executor;
18use hyper_rustls::{HttpsConnector as HyperRustlsHttpsConnector, MaybeHttpsStream};
19use hyper_util::client::legacy::Client;
20use hyper_util::client::legacy::connect::proxy::Tunnel;
21use hyper_util::client::legacy::connect::{
22 Connected, Connection, HttpConnector as HyperHttpConnector,
23};
24use hyper_util::rt::{TokioIo, TokioTimer};
25use log::warn;
26use parking_lot::Mutex;
27use rustls::client::danger::ServerCertVerifier;
28use rustls::client::{ClientConnection, EchStatus};
29use rustls::crypto::{CryptoProvider, aws_lc_rs};
30use rustls::{CipherSuite, ClientConfig, NamedGroup, ProtocolVersion};
31use rustls_pki_types::{CertificateDer, ServerName, UnixTime};
32use servo_config::pref;
33use tokio::net::TcpStream;
34use tower::Service;
35
36use crate::async_runtime::spawn_task;
37use crate::hosts::replace_host;
38
39pub const BUF_SIZE: usize = 32768;
40
41pub const ALPN_H2: &str = "h2";
43
44#[derive(Clone)]
45pub struct ServoHttpConnector {
46 inner: HyperHttpConnector,
47}
48
49impl ServoHttpConnector {
50 fn new() -> ServoHttpConnector {
51 let mut inner = HyperHttpConnector::new();
52 inner.enforce_http(false);
53 inner.set_happy_eyeballs_timeout(None);
54 inner.set_connect_timeout(Some(Duration::from_secs(pref!(network_connection_timeout))));
55 ServoHttpConnector { inner }
56 }
57}
58
59impl Service<Destination> for ServoHttpConnector {
60 type Response = TokioIo<TcpStream>;
61 type Error = ConnectionError;
62 type Future =
63 std::pin::Pin<Box<dyn Future<Output = Result<TokioIo<TcpStream>, ConnectionError>> + Send>>;
64
65 fn call(&mut self, dest: Destination) -> Self::Future {
66 let mut new_dest = dest.clone();
68 let mut parts = dest.into_parts();
69
70 if let Some(auth) = parts.authority {
71 let host = auth.host();
72 let host = replace_host(host);
73
74 let authority = if let Some(port) = auth.port() {
75 format!("{}:{}", host, port.as_str())
76 } else {
77 (*host).to_string()
78 };
79
80 if let Ok(authority) = Authority::from_maybe_shared(authority) {
81 parts.authority = Some(authority);
82 if let Ok(dest) = Destination::from_parts(parts) {
83 new_dest = dest
84 }
85 }
86 }
87
88 Box::pin(
89 self.inner
90 .call(new_dest)
91 .map_err(|e| ConnectionError::HttpError(format!("{e}"))),
92 )
93 }
94
95 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
96 Ok(()).into()
97 }
98}
99
100type BoxError = Box<dyn std::error::Error + Send + Sync>;
101
102#[derive(Clone)]
103pub struct InstrumentedConnector<T> {
104 inner: HyperRustlsHttpsConnector<T>,
105}
106
107impl<T> InstrumentedConnector<T> {
108 fn new(inner: HyperRustlsHttpsConnector<T>) -> Self {
109 Self { inner }
110 }
111}
112
113impl<T> From<HyperRustlsHttpsConnector<T>> for InstrumentedConnector<T> {
114 fn from(inner: HyperRustlsHttpsConnector<T>) -> Self {
115 Self::new(inner)
116 }
117}
118
119pub struct InstrumentedStream<T> {
120 inner: MaybeHttpsStream<T>,
121 tls_info: RefCell<Option<TlsHandshakeInfo>>,
122}
123
124impl<T: Unpin> Unpin for InstrumentedStream<T> {}
125
126impl<T> fmt::Debug for InstrumentedStream<T> {
127 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
128 f.debug_struct("InstrumentedStream")
129 .field("tls_info", &self.tls_info)
130 .finish()
131 }
132}
133
134#[derive(Clone, Debug)]
135pub struct TlsHandshakeInfo {
136 pub protocol_version: Option<ProtocolVersion>,
137 pub cipher_suite: Option<CipherSuite>,
138 pub kea_group_name: Option<NamedGroup>,
139 pub signature_scheme_name: Option<String>,
140 pub alpn_protocol: Option<String>,
141 pub certificate_chain_der: Vec<Vec<u8>>,
142 pub used_ech: bool,
143}
144
145impl TlsHandshakeInfo {
146 fn from_connection(conn: &ClientConnection) -> Self {
147 let protocol_version = conn.protocol_version();
148 let cipher_suite = conn.negotiated_cipher_suite().map(|suite| suite.suite());
149 let kea_group_name = conn
150 .negotiated_key_exchange_group()
151 .map(|group| group.name());
152 let certificate_chain_der = conn
153 .peer_certificates()
154 .map(|certs| certs.iter().map(|cert| cert.as_ref().to_vec()).collect())
155 .unwrap_or_default();
156 let alpn_protocol = conn
157 .alpn_protocol()
158 .map(|proto| String::from_utf8_lossy(proto).into_owned());
159 let used_ech = matches!(conn.ech_status(), EchStatus::Accepted);
160
161 Self {
162 protocol_version,
163 cipher_suite,
164 kea_group_name,
165 signature_scheme_name: None,
166 alpn_protocol,
167 certificate_chain_der,
168 used_ech,
169 }
170 }
171}
172
173impl<T> InstrumentedStream<T>
174where
175 T: Connection + hyper::rt::Read + hyper::rt::Write + Unpin,
176{
177 fn from_maybe_https_stream(stream: MaybeHttpsStream<T>) -> Self {
178 match stream {
179 MaybeHttpsStream::Http(inner) => Self {
180 inner: MaybeHttpsStream::Http(inner),
181 tls_info: RefCell::new(None),
182 },
183 MaybeHttpsStream::Https(tls_stream) => {
184 let (_tcp, tls) = tls_stream.inner().get_ref();
185 let tls_info = TlsHandshakeInfo::from_connection(tls);
186
187 Self {
188 inner: MaybeHttpsStream::Https(tls_stream),
189 tls_info: RefCell::new(Some(tls_info)),
190 }
191 },
192 }
193 }
194}
195
196impl<T> Connection for InstrumentedStream<T>
197where
198 T: Connection + hyper::rt::Read + hyper::rt::Write + Unpin,
199{
200 fn connected(&self) -> Connected {
201 let connected = match &self.inner {
202 MaybeHttpsStream::Http(stream) => stream.connected(),
203 MaybeHttpsStream::Https(stream) => {
204 let (tcp, tls) = stream.inner().get_ref();
205 if tls.alpn_protocol() == Some(ALPN_H2.as_bytes()) {
206 tcp.inner().connected().negotiated_h2()
207 } else {
208 tcp.inner().connected()
209 }
210 },
211 };
212 if let Some(info) = self.tls_info.borrow_mut().take() {
213 connected.extra(info)
214 } else {
215 connected
216 }
217 }
218}
219
220impl<T> hyper::rt::Read for InstrumentedStream<T>
221where
222 T: Connection + hyper::rt::Read + hyper::rt::Write + Unpin,
223{
224 fn poll_read(
225 self: std::pin::Pin<&mut Self>,
226 cx: &mut Context<'_>,
227 buf: hyper::rt::ReadBufCursor<'_>,
228 ) -> Poll<Result<(), io::Error>> {
229 std::pin::Pin::new(&mut self.get_mut().inner).poll_read(cx, buf)
230 }
231}
232
233impl<T> hyper::rt::Write for InstrumentedStream<T>
234where
235 T: Connection + hyper::rt::Read + hyper::rt::Write + Unpin,
236{
237 fn poll_write(
238 self: std::pin::Pin<&mut Self>,
239 cx: &mut Context<'_>,
240 buf: &[u8],
241 ) -> Poll<Result<usize, io::Error>> {
242 std::pin::Pin::new(&mut self.get_mut().inner).poll_write(cx, buf)
243 }
244
245 fn poll_flush(
246 self: std::pin::Pin<&mut Self>,
247 cx: &mut Context<'_>,
248 ) -> Poll<Result<(), io::Error>> {
249 std::pin::Pin::new(&mut self.get_mut().inner).poll_flush(cx)
250 }
251
252 fn poll_shutdown(
253 self: std::pin::Pin<&mut Self>,
254 cx: &mut Context<'_>,
255 ) -> Poll<Result<(), io::Error>> {
256 std::pin::Pin::new(&mut self.get_mut().inner).poll_shutdown(cx)
257 }
258
259 fn is_write_vectored(&self) -> bool {
260 self.inner.is_write_vectored()
261 }
262
263 fn poll_write_vectored(
264 self: std::pin::Pin<&mut Self>,
265 cx: &mut Context<'_>,
266 bufs: &[io::IoSlice<'_>],
267 ) -> Poll<Result<usize, io::Error>> {
268 std::pin::Pin::new(&mut self.get_mut().inner).poll_write_vectored(cx, bufs)
269 }
270}
271
272impl<T> Service<Destination> for InstrumentedConnector<T>
273where
274 T: Service<Destination>,
275 T::Response: Connection + hyper::rt::Read + hyper::rt::Write + Send + Unpin + 'static,
276 T::Future: Send + 'static,
277 T::Error: Into<BoxError>,
278{
279 type Response = InstrumentedStream<T::Response>;
280 type Error = BoxError;
281 type Future = std::pin::Pin<
282 Box<dyn Future<Output = Result<InstrumentedStream<T::Response>, BoxError>> + Send>,
283 >;
284
285 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
286 self.inner.poll_ready(cx).map_err(Into::into)
287 }
288
289 fn call(&mut self, dst: Destination) -> Self::Future {
290 let future = self.inner.call(dst);
291 Box::pin(async move {
292 let stream = future.await.map_err(|error| -> BoxError { error })?;
293 Ok(InstrumentedStream::from_maybe_https_stream(stream))
294 })
295 }
296}
297
298pub type Connector = InstrumentedConnector<ServoHttpConnector>;
299pub type TlsConfig = ClientConfig;
300
301#[derive(Clone, Debug, Default)]
302struct CertificateErrorOverrideManagerInternal {
303 certificates_failing_to_verify: HashMap<ServerName<'static>, CertificateDer<'static>>,
306 overrides: Vec<CertificateDer<'static>>,
309}
310
311#[derive(Clone, Debug, Default)]
316pub struct CertificateErrorOverrideManager(Arc<Mutex<CertificateErrorOverrideManagerInternal>>);
317
318impl CertificateErrorOverrideManager {
319 pub fn new() -> Self {
320 Self(Default::default())
321 }
322
323 pub fn add_override(&self, certificate: &CertificateDer<'static>) {
326 self.0.lock().overrides.push(certificate.clone());
327 }
328
329 pub(crate) fn remove_certificate_failing_verification(
333 &self,
334 host: &str,
335 ) -> Option<CertificateDer<'static>> {
336 let server_name = match ServerName::try_from(host) {
337 Ok(name) => name.to_owned(),
338 Err(error) => {
339 warn!("Could not convert host string into RustTLS ServerName: {error:?}");
340 return None;
341 },
342 };
343 self.0
344 .lock()
345 .certificates_failing_to_verify
346 .remove(&server_name)
347 }
348}
349
350#[derive(Clone, Debug, Default)]
351pub enum CACertificates<'de> {
352 #[default]
353 Default,
354 Override(Vec<CertificateDer<'de>>),
355}
356
357#[servo_tracing::instrument(skip_all)]
364pub fn create_tls_config(
365 ca_certificates: CACertificates<'static>,
366 ignore_certificate_errors: bool,
367 override_manager: CertificateErrorOverrideManager,
368) -> TlsConfig {
369 let verifier = CertificateVerificationOverrideVerifier::new(
370 ca_certificates,
371 ignore_certificate_errors,
372 override_manager,
373 );
374 rustls::ClientConfig::builder()
377 .dangerous()
378 .with_custom_certificate_verifier(Arc::new(verifier))
379 .with_no_client_auth()
380}
381
382#[derive(Clone)]
383struct TokioExecutor {}
384
385impl<F> Executor<F> for TokioExecutor
386where
387 F: Future<Output = ()> + 'static + std::marker::Send,
388{
389 fn execute(&self, fut: F) {
390 spawn_task(fut);
391 }
392}
393
394static CRYPTO_PROVIDER_CACHE: LazyLock<Arc<CryptoProvider>> = LazyLock::new(|| {
395 CryptoProvider::get_default()
396 .cloned()
397 .unwrap_or_else(|| {
400 warn!("Default crypto provider not initialized before first access in connector.");
401 Arc::new(aws_lc_rs::default_provider())
402 })
403});
404
405static RUSTLS_PLATFORM_VERIFIER_CACHE: LazyLock<Arc<rustls_platform_verifier::Verifier>> =
410 LazyLock::new(|| {
411 Arc::new(
412 rustls_platform_verifier::Verifier::new(CRYPTO_PROVIDER_CACHE.clone())
413 .expect("Could not initialize platform certificate verifier"),
414 )
415 });
416
417#[inline]
424pub fn prewarm_tls() {
425 #[servo_tracing::instrument]
426 fn prewarm_tls_impl() {
427 let mut sink = [0u8; 32];
428 let _ = CRYPTO_PROVIDER_CACHE.secure_random.fill(&mut sink);
430 }
433
434 if let Err(error) = std::thread::Builder::new()
435 .name("Net-TLS-prewarm".into())
436 .spawn(prewarm_tls_impl)
437 {
438 warn!("Failed to spawn thread to prewarm TLS: {error:?}");
439 }
440}
441
442#[derive(Debug)]
443struct CertificateVerificationOverrideVerifier {
444 main_verifier: Arc<dyn ServerCertVerifier>,
445 ignore_certificate_errors: bool,
446 override_manager: CertificateErrorOverrideManager,
447}
448
449impl CertificateVerificationOverrideVerifier {
450 fn new(
451 ca_certficates: CACertificates<'static>,
452 ignore_certificate_errors: bool,
453 override_manager: CertificateErrorOverrideManager,
454 ) -> Self {
455 let use_webpki_roots = cfg!(target_os = "android") || pref!(network_use_webpki_roots);
464 let main_verifier = if !use_webpki_roots {
465 let verifier = match ca_certficates {
466 CACertificates::Default => RUSTLS_PLATFORM_VERIFIER_CACHE.clone(),
467 CACertificates::Override(_certificates) => {
470 #[cfg(target_os = "android")]
471 unreachable!("Android should always use the WebPKI verifier.");
472 #[cfg(not(target_os = "android"))]
473 {
474 let verifier = rustls_platform_verifier::Verifier::new_with_extra_roots(
475 _certificates,
476 CRYPTO_PROVIDER_CACHE.clone(),
477 )
478 .expect("Could not initialize platform certificate verifier");
479 Arc::new(verifier)
480 }
481 },
482 };
483 verifier as Arc<dyn ServerCertVerifier>
484 } else {
485 let mut root_store =
486 rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
487 match ca_certficates {
488 CACertificates::Default => {},
489 CACertificates::Override(certificates) => {
490 for certificate in certificates {
491 if root_store.add(certificate).is_err() {
492 log::error!("Could not add an override certificate.");
493 }
494 }
495 },
496 }
497 rustls::client::WebPkiServerVerifier::builder(root_store.into())
498 .build()
499 .expect("Could not initialize platform certificate verifier.")
500 as Arc<dyn ServerCertVerifier>
501 };
502
503 Self {
504 main_verifier,
505 ignore_certificate_errors,
506 override_manager,
507 }
508 }
509}
510
511impl rustls::client::danger::ServerCertVerifier for CertificateVerificationOverrideVerifier {
512 fn verify_tls12_signature(
513 &self,
514 message: &[u8],
515 cert: &CertificateDer<'_>,
516 dss: &rustls::DigitallySignedStruct,
517 ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
518 self.main_verifier
519 .verify_tls12_signature(message, cert, dss)
520 }
521
522 fn verify_tls13_signature(
523 &self,
524 message: &[u8],
525 cert: &CertificateDer<'_>,
526 dss: &rustls::DigitallySignedStruct,
527 ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
528 self.main_verifier
529 .verify_tls13_signature(message, cert, dss)
530 }
531
532 fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
533 self.main_verifier.supported_verify_schemes()
534 }
535
536 fn verify_server_cert(
537 &self,
538 end_entity: &CertificateDer<'_>,
539 intermediates: &[CertificateDer<'_>],
540 server_name: &ServerName<'_>,
541 ocsp_response: &[u8],
542 now: UnixTime,
543 ) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
544 let error = match self.main_verifier.verify_server_cert(
545 end_entity,
546 intermediates,
547 server_name,
548 ocsp_response,
549 now,
550 ) {
551 Ok(result) => return Ok(result),
552 Err(error) => error,
553 };
554
555 if self.ignore_certificate_errors {
556 warn!("Ignoring certficate error: {error:?}");
557 return Ok(rustls::client::danger::ServerCertVerified::assertion());
558 }
559
560 for cert_with_exception in &*self.override_manager.0.lock().overrides {
562 if *end_entity == *cert_with_exception {
563 return Ok(rustls::client::danger::ServerCertVerified::assertion());
564 }
565 }
566 self.override_manager
567 .0
568 .lock()
569 .certificates_failing_to_verify
570 .insert(server_name.to_owned(), end_entity.clone().into_owned());
571 Err(error)
572 }
573}
574
575pub type BoxedBody = BoxBody<Bytes, hyper::Error>;
576
577#[derive(Debug)]
578pub enum ConnectionError {
580 HttpError(String),
581 ProxyError(String),
583}
584
585impl std::fmt::Display for ConnectionError {
586 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
587 write!(f, "{self:?}")
588 }
589}
590
591impl std::error::Error for ConnectionError {}
592
593#[derive(Clone)]
594pub struct ProxyConnector {
597 client: ServoHttpConnector,
599 matcher: std::sync::Arc<hyper_util::client::proxy::matcher::Matcher>,
601}
602
603impl ProxyConnector {
604 fn new() -> Self {
605 let matcher_builder = hyper_util::client::proxy::matcher::Matcher::builder()
606 .http(servo_config::pref!(network_http_proxy_uri))
607 .https(servo_config::pref!(network_https_proxy_uri))
608 .no(servo_config::pref!(network_http_no_proxy));
609 ProxyConnector {
610 client: ServoHttpConnector::new(),
611 matcher: std::sync::Arc::new(matcher_builder.build()),
612 }
613 }
614}
615
616impl Service<Destination> for ProxyConnector {
618 type Response = TokioIo<TcpStream>;
619 type Error = ConnectionError;
620 type Future =
621 std::pin::Pin<Box<dyn Future<Output = Result<TokioIo<TcpStream>, ConnectionError>> + Send>>;
622
623 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
624 self.client
625 .poll_ready(cx)
626 .map_err(|e| ConnectionError::ProxyError(format!("{e}")))
627 }
628
629 fn call(&mut self, req: Destination) -> Self::Future {
630 match self.matcher.intercept(&req) {
631 Some(intercept) => {
632 let mut tunnel = Tunnel::new(intercept.uri().clone(), self.client.clone());
633 let final_tunnel = if let Some(auth) = intercept.basic_auth() {
634 tunnel.with_auth(auth.clone())
635 } else {
636 tunnel
637 }
638 .call(req)
639 .map_err(|e| ConnectionError::ProxyError(format!("{e}")));
640 Box::pin(final_tunnel)
641 },
642 None => Box::pin(
643 self.client
644 .call(req)
645 .map_err(|e| ConnectionError::ProxyError(format!("{e}"))),
646 ),
647 }
648 }
649}
650
651pub type ServoClient = Client<InstrumentedConnector<ProxyConnector>, BoxedBody>;
652
653pub fn create_http_client(tls_config: TlsConfig) -> ServoClient {
654 let connector = hyper_rustls::HttpsConnectorBuilder::new()
655 .with_tls_config(tls_config)
656 .https_or_http()
657 .enable_http1()
658 .enable_http2()
659 .wrap_connector(ProxyConnector::new());
660
661 Client::builder(TokioExecutor {})
662 .http1_title_case_headers(true)
663 .timer(TokioTimer::new())
665 .pool_timer(TokioTimer::new())
667 .pool_max_idle_per_host(6)
672 .build(InstrumentedConnector::from(connector))
673}