Skip to main content

orinium_browser/platform/network/
core.rs

1//! ネットワークコア
2//! HTTP通信とレスポンス処理を担当する。
3
4use 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
21/// Per-thread driver for the shared network state.
22///
23/// The tokio runtime and its [`LocalSet`] are `!Send`, so each pool worker
24/// owns one of these; the expensive state they operate on lives in the
25/// [`SharedNetState`] behind an `Arc` and is reused across workers.
26///
27/// The [`SenderPool`] deliberately lives *here*, not in the shared state: a
28/// pooled connection's driver task is spawned onto the creating worker's
29/// local set and is only polled while that worker runs a fetch. Sharing
30/// senders across runtimes would let one worker check out a connection whose
31/// driver is parked on another (idle) worker, leaving the request awaiting
32/// frames nobody ever reads.
33pub(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    /// Runs a fetch to completion on this worker's local set.
58    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/// HTTP response
96#[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
105/// TLS, cache and config shared by every fetch worker.
106///
107/// All fields are thread-safe: the cache is internally synchronized, and the
108/// config is swapped atomically through an `Arc` so workers observe updates
109/// without blocking a fetch in progress. Connection pools are *not* shared:
110/// each [`AsyncNetworkCore`] owns its pool because a connection is only valid
111/// on the runtime that drives it.
112pub(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    /// Removes all cached responses.
133    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(&current.to_string())
162            {
163                log::info!("NetworkCache: hit for url={}", current);
164                return Ok(cached);
165            }
166
167            let resp = self
168                .send_request(&current, &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(&current, 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(&current.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    /// Keep-alive HTTP/1.1 server answering every request with
387    /// `ok:<request line>`, so pooled connections stay open across sequential
388    /// fetches. Each connection is served until the peer hangs up. Returns
389    /// the bound address.
390    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        // Worker A fetches and returns its keep-alive connection to its own
513        // pool; the connection stays open on A's local set.
514        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        // Worker B must open its own connection for the same host instead of
520        // checking out the sender parked on A's idle runtime, where nobody
521        // would ever poll it.
522        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        // A still reuses its own pooled connection.
532        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}