orinium_browser/platform/network/
core.rs1use super::{HostKey, HttpSender, NetworkConfig, NetworkError, NetworkRequest, SenderPool};
5
6use http_body_util::{BodyExt, Full};
7use hyper::{
8 Method, Request, Uri,
9 body::{Bytes, Incoming},
10 client::conn,
11 http::uri::Scheme,
12};
13use hyper_util::rt::TokioIo;
14use rustls::{ClientConfig, RootCertStore};
15use rustls_native_certs::load_native_certs;
16use serde::{Deserialize, Serialize};
17use std::sync::{Arc, RwLock};
18use tokio::{net::TcpStream, runtime::Runtime, task::LocalSet};
19use tokio_rustls::TlsConnector;
20
21pub(super) struct AsyncNetworkCore {
34 local: LocalSet,
35 rt: Runtime,
36 inner: Arc<SharedNetState>,
37 sender_pool: Arc<std::sync::RwLock<SenderPool>>,
38}
39
40impl AsyncNetworkCore {
41 pub fn new(inner: Arc<SharedNetState>) -> Self {
42 let rt = tokio::runtime::Builder::new_current_thread()
43 .enable_all()
44 .build()
45 .expect("failed to build tokio runtime");
46
47 let local = LocalSet::new();
48
49 Self {
50 rt,
51 local,
52 inner,
53 sender_pool: Arc::new(std::sync::RwLock::new(SenderPool::new())),
54 }
55 }
56
57 pub fn fetch_request_blocking(
59 &self,
60 request: &NetworkRequest,
61 ) -> Result<Response, NetworkError> {
62 self.local
63 .block_on(&self.rt, async { self.fetch_request(request).await })
64 }
65}
66
67#[derive(Debug, Deserialize, Serialize, Clone)]
68pub struct StatusCode(u16);
69
70impl From<hyper::StatusCode> for StatusCode {
71 fn from(value: hyper::StatusCode) -> Self {
72 Self(value.as_u16())
73 }
74}
75
76impl StatusCode {
77 pub fn as_u16(&self) -> u16 {
78 self.0
79 }
80
81 pub fn is_success(&self) -> bool {
82 (200..300).contains(&self.0)
83 }
84
85 pub fn is_redirection(&self) -> bool {
86 (300..400).contains(&self.0)
87 }
88
89 pub fn canonical_reason(&self) -> Option<&'static str> {
90 let hyper_code: hyper::StatusCode = self.as_u16().try_into().ok()?;
91 hyper_code.canonical_reason()
92 }
93}
94
95#[derive(Debug, Deserialize, Serialize, Clone)]
97pub struct Response {
98 pub url: String,
99 pub status: StatusCode,
100 pub reason_phrase: String,
101 pub headers: Vec<(String, String)>,
102 pub body: Vec<u8>,
103}
104
105pub(super) struct SharedNetState {
113 tls_config: Arc<ClientConfig>,
114 network_config: RwLock<Arc<NetworkConfig>>,
115 cache: super::Cache,
116}
117
118impl SharedNetState {
119 pub fn new() -> Self {
120 Self {
121 tls_config: Arc::new(Self::build_tls_config()),
122 network_config: RwLock::new(Arc::new(NetworkConfig::default())),
123 cache: super::Cache::new(),
124 }
125 }
126
127 pub fn set_network_config(&self, config: NetworkConfig) {
128 self.cache.set_enabled(config.enable_cache);
129 *self.network_config.write().unwrap() = Arc::new(config);
130 }
131
132 pub fn clear_cache(&self) {
134 self.cache.clear();
135 }
136
137 fn build_tls_config() -> ClientConfig {
138 let mut roots = RootCertStore::empty();
139 let result = load_native_certs();
140
141 for cert in result.certs {
142 let _ = roots.add(cert);
143 }
144
145 ClientConfig::builder()
146 .with_root_certificates(roots)
147 .with_no_client_auth()
148 }
149}
150
151impl AsyncNetworkCore {
152 pub async fn fetch_request(&self, request: &NetworkRequest) -> Result<Response, NetworkError> {
153 let mut current: Uri = request.url.parse().map_err(|_| NetworkError::InvalidUri)?;
154 let mut method = Method::from_bytes(request.method.as_bytes())
155 .map_err(|_| NetworkError::HttpRequestFailed)?;
156 let mut body = request.body.clone();
157 let mut redirects = 0usize;
158
159 loop {
160 if method == Method::GET
161 && let Some(cached) = self.inner.cache.get(¤t.to_string())
162 {
163 log::info!("NetworkCache: hit for url={}", current);
164 return Ok(cached);
165 }
166
167 let resp = self
168 .send_request(¤t, &method, &request.headers, &body)
169 .await?;
170
171 if self.inner.network_config.read().unwrap().follow_redirects
172 && hyper::StatusCode::try_from(resp.status.0)
173 .map_err(|_| NetworkError::InvalidIpcStatusCode)?
174 .is_redirection()
175 {
176 if redirects >= 10 {
177 return Err(NetworkError::TooManyRedirects);
178 }
179
180 if let Some(loc) = resp
181 .headers
182 .iter()
183 .find(|(k, _)| k.eq_ignore_ascii_case("location"))
184 .map(|(_, v)| v)
185 {
186 current = resolve_redirect(¤t, loc)?;
187 if resp.status.as_u16() == 303
188 || ((resp.status.as_u16() == 301 || resp.status.as_u16() == 302)
189 && method == Method::POST)
190 {
191 method = Method::GET;
192 body.clear();
193 }
194 redirects += 1;
195 continue;
196 }
197 }
198
199 if method == Method::GET && resp.status.is_success() {
200 self.inner.cache.set(¤t.to_string(), &resp);
201 }
202
203 return Ok(resp);
204 }
205 }
206
207 async fn send_request(
208 &self,
209 uri: &Uri,
210 method: &Method,
211 headers: &[(String, String)],
212 body: &[u8],
213 ) -> Result<Response, NetworkError> {
214 let host = uri.host().ok_or(NetworkError::MissingHost)?;
215 let scheme = uri.scheme().unwrap_or(&Scheme::HTTP);
216 let port = uri
217 .port_u16()
218 .unwrap_or(if scheme == &Scheme::HTTPS { 443 } else { 80 });
219
220 let key = HostKey {
221 scheme: scheme.clone(),
222 host: host.to_string(),
223 port,
224 };
225
226 let mut sender = self.get_or_create_sender(&key).await?;
227
228 let user_agent = self.inner.network_config.read().unwrap().user_agent.clone();
229 let mut request = Request::builder()
230 .method(method.clone())
231 .uri(uri.path_and_query().map_or("/", |p| p.as_str()))
232 .header("Host", host)
233 .header("User-Agent", user_agent);
234 if !headers
235 .iter()
236 .any(|(name, _)| name.eq_ignore_ascii_case("accept-language"))
237 {
238 request = request.header(
239 "Accept-Language",
240 crate::platform::locale::accept_language_header(),
241 );
242 }
243 for (name, value) in headers {
244 request = request.header(name, value);
245 }
246 let req = request
247 .body(Full::new(Bytes::copy_from_slice(body)))
248 .map_err(|_| NetworkError::HttpRequestFailed)?;
249
250 let mut res = match &mut sender {
251 HttpSender::Http1(s) => s
252 .send_request(req)
253 .await
254 .map_err(|_| NetworkError::HttpRequestFailed)?,
255 _ => {
256 return Err(NetworkError::UnsupportedHttpVersion);
257 }
258 };
259
260 let response = Self::collect_response(uri.to_string(), &mut res).await?;
261
262 self.sender_pool
263 .write()
264 .unwrap()
265 .add_connection(key, sender);
266
267 Ok(response)
268 }
269
270 async fn collect_response(
271 url: String,
272 res: &mut hyper::Response<Incoming>,
273 ) -> Result<Response, NetworkError> {
274 let status = res.status();
275 let reason_phrase = status.canonical_reason().unwrap_or("").to_string();
276
277 let headers = res
278 .headers()
279 .iter()
280 .map(|(k, v)| (k.as_str().to_string(), v.to_str().unwrap_or("").to_string()))
281 .collect();
282
283 let mut body = Vec::new();
284 while let Some(frame) = res.frame().await {
285 let frame = frame.map_err(|_| NetworkError::HttpResponseFailed)?;
286 if let Some(chunk) = frame.data_ref() {
287 body.extend_from_slice(chunk);
288 }
289 }
290
291 Ok(Response {
292 url,
293 status: status.into(),
294 reason_phrase,
295 headers,
296 body,
297 })
298 }
299
300 async fn get_or_create_sender(&self, key: &HostKey) -> Result<HttpSender, NetworkError> {
301 if let Some(s) = self.sender_pool.write().unwrap().get_connection(key) {
302 return Ok(s);
303 }
304
305 self.create_connection(key).await
306 }
307
308 async fn create_connection(&self, key: &HostKey) -> Result<HttpSender, NetworkError> {
309 let addr = format!("{}:{}", key.host, key.port);
310 let stream = TcpStream::connect(addr)
311 .await
312 .map_err(|_| NetworkError::ConnectionFailed)?;
313
314 if key.scheme == Scheme::HTTPS {
315 let tls = TlsConnector::from(Arc::clone(&self.inner.tls_config));
316 let key = key.clone();
317 let domain = rustls::pki_types::ServerName::try_from(key.host.clone())
318 .map_err(|_| NetworkError::InvalidDnsName)?;
319
320 let stream = tls
321 .connect(domain, stream)
322 .await
323 .map_err(|_| NetworkError::TlsFailed)?;
324
325 let (sender, conn) = conn::http1::handshake(TokioIo::new(stream))
326 .await
327 .map_err(|_| NetworkError::HttpHandshakeFailed)?;
328
329 self.spawn_connection_task(conn, key);
330 Ok(HttpSender::Http1(sender))
331 } else {
332 let (sender, conn) = conn::http1::handshake(TokioIo::new(stream))
333 .await
334 .map_err(|_| NetworkError::HttpHandshakeFailed)?;
335
336 self.spawn_connection_task(conn, key.clone());
337 Ok(HttpSender::Http1(sender))
338 }
339 }
340
341 fn spawn_connection_task(
342 &self,
343 conn: conn::http1::Connection<
344 TokioIo<impl tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + 'static>,
345 Full<Bytes>,
346 >,
347 key: HostKey,
348 ) {
349 let pool = Arc::clone(&self.sender_pool);
350 tokio::task::spawn_local(async move {
351 let _ = conn.await;
352 pool.write().unwrap().remove_connection(&key);
353 });
354 }
355}
356
357fn resolve_redirect(base: &Uri, location: &str) -> Result<Uri, NetworkError> {
358 if location.starts_with("http://") || location.starts_with("https://") {
359 return location.parse().map_err(|_| NetworkError::InvalidUri);
360 }
361
362 let scheme = base.scheme_str().unwrap_or("https");
363 let authority = base.authority().ok_or(NetworkError::InvalidUri)?;
364
365 let next = if location.starts_with("//") {
366 format!("{scheme}:{location}")
367 } else if location.starts_with('/') {
368 format!("{scheme}://{}{location}", authority)
369 } else {
370 let base_path = base.path();
371 let prefix = base_path.rsplit_once('/').map_or("", |x| x.0);
372 format!("{scheme}://{}{prefix}/{location}", authority)
373 };
374
375 next.parse().map_err(|_| NetworkError::InvalidUri)
376}
377
378#[cfg(test)]
379mod tests {
380 use super::*;
381 use std::io::{Read, Write};
382 use std::net::TcpListener;
383 use std::thread;
384 use std::time::Duration;
385
386 fn spawn_keep_alive_server(listener: TcpListener) -> String {
391 let address = listener.local_addr().unwrap().to_string();
392 thread::spawn(move || {
393 while let Ok((stream, _)) = listener.accept() {
394 thread::spawn(move || serve_keep_alive_connection(stream));
395 }
396 });
397 address
398 }
399
400 fn serve_keep_alive_connection(mut stream: std::net::TcpStream) {
401 stream
402 .set_read_timeout(Some(Duration::from_secs(10)))
403 .unwrap();
404 let mut buffer = [0_u8; 4096];
405 loop {
406 let mut request = Vec::new();
407 loop {
408 let read = match stream.read(&mut buffer) {
409 Ok(0) | Err(_) => return,
410 Ok(read) => read,
411 };
412 request.extend_from_slice(&buffer[..read]);
413 if request.windows(4).any(|bytes| bytes == b"\r\n\r\n") {
414 break;
415 }
416 }
417 let request_line = String::from_utf8_lossy(&request)
418 .lines()
419 .next()
420 .unwrap_or("")
421 .to_string();
422 let body = format!("ok:{request_line}");
423 let response = format!(
424 "HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n{body}",
425 body.len()
426 );
427 if stream.write_all(response.as_bytes()).is_err() {
428 return;
429 }
430 }
431 }
432
433 #[test]
434 fn enable_cache_config_is_applied_to_cache() {
435 let state = SharedNetState::new();
436 assert!(state.cache.is_enabled());
437
438 state.set_network_config(NetworkConfig {
439 enable_cache: false,
440 ..NetworkConfig::default()
441 });
442 assert!(!state.cache.is_enabled());
443 }
444
445 #[test]
446 fn sends_request_method_headers_and_body() {
447 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
448 let address = listener.local_addr().unwrap();
449 let server = thread::spawn(move || {
450 let (mut stream, _) = listener.accept().unwrap();
451 stream
452 .set_read_timeout(Some(Duration::from_secs(5)))
453 .unwrap();
454 let mut request = Vec::new();
455 let mut buffer = [0_u8; 1024];
456 loop {
457 let read = stream.read(&mut buffer).unwrap();
458 if read == 0 {
459 break;
460 }
461 request.extend_from_slice(&buffer[..read]);
462 let Some(header_end) = request.windows(4).position(|bytes| bytes == b"\r\n\r\n")
463 else {
464 continue;
465 };
466 let headers = String::from_utf8_lossy(&request[..header_end]);
467 let content_length = headers
468 .lines()
469 .find_map(|line| {
470 let (name, value) = line.split_once(':')?;
471 name.eq_ignore_ascii_case("content-length")
472 .then(|| value.trim().parse::<usize>().unwrap())
473 })
474 .unwrap_or(0);
475 if request.len() >= header_end + 4 + content_length {
476 break;
477 }
478 }
479 stream
480 .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
481 .unwrap();
482 request
483 });
484
485 let core = AsyncNetworkCore::new(Arc::new(SharedNetState::new()));
486 let response = core
487 .fetch_request_blocking(&NetworkRequest {
488 url: format!("http://{address}/submit"),
489 method: "POST".to_string(),
490 headers: vec![("X-Orinium-Test".to_string(), "yes".to_string())],
491 body: b"hello".to_vec(),
492 })
493 .unwrap();
494 assert_eq!(response.body, b"ok");
495
496 let request = String::from_utf8(server.join().unwrap()).unwrap();
497 assert!(request.starts_with("POST /submit HTTP/1.1\r\n"));
498 assert!(request.to_ascii_lowercase().contains("x-orinium-test: yes"));
499 assert!(request.to_ascii_lowercase().contains("accept-language: "));
500 assert!(request.ends_with("\r\n\r\nhello"));
501 }
502
503 #[test]
504 fn pooled_connections_are_not_shared_between_worker_runtimes() {
505 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
506 let address = spawn_keep_alive_server(listener);
507
508 let shared = Arc::new(SharedNetState::new());
509 let core_a = AsyncNetworkCore::new(Arc::clone(&shared));
510 let core_b = AsyncNetworkCore::new(shared);
511
512 let first = core_a
515 .fetch_request_blocking(&NetworkRequest::get(format!("http://{address}/first")))
516 .unwrap();
517 assert_eq!(first.body, b"ok:GET /first HTTP/1.1");
518
519 assert!(
523 core_b.sender_pool.read().unwrap().is_empty(),
524 "A's pooled connection must not be visible to another worker"
525 );
526 let second = core_b
527 .fetch_request_blocking(&NetworkRequest::get(format!("http://{address}/second")))
528 .unwrap();
529 assert_eq!(second.body, b"ok:GET /second HTTP/1.1");
530
531 let third = core_a
533 .fetch_request_blocking(&NetworkRequest::get(format!("http://{address}/third")))
534 .unwrap();
535 assert_eq!(third.body, b"ok:GET /third HTTP/1.1");
536 }
537}