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