This repository has no description
16 kB
496 lines
1use std::future::Future;
2use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
3use std::pin::Pin;
4use std::sync::Arc;
5use std::time::Duration;
6
7use bytes::Bytes;
8use futures::TryStreamExt;
9use http::{HeaderMap, Method, StatusCode};
10use url::{Host, Url};
11
12#[derive(Debug, Clone, thiserror::Error)]
13pub enum NetworkError {
14 #[error("build: {0}")]
15 Build(String),
16 #[error("connect: {0}")]
17 Connect(String),
18 #[error("timeout: {0}")]
19 Timeout(String),
20 #[error("request: {0}")]
21 Request(String),
22 #[error("body: {0}")]
23 Body(String),
24 #[error("response exceeds {limit} bytes")]
25 TooLarge { limit: u64 },
26 #[error("refusing to reach non-public address {host}")]
27 Blocked { host: String },
28}
29
30pub fn is_blocked_ip(ip: IpAddr) -> bool {
31 match ip {
32 IpAddr::V4(v4) => is_blocked_v4(v4),
33 IpAddr::V6(v6) => match embedded_ipv4(v6) {
34 Some(embedded) => is_blocked_v4(embedded),
35 None => {
36 v6.is_loopback()
37 || v6.is_unspecified()
38 || v6.is_multicast()
39 || (v6.segments()[0] & 0xfe00) == 0xfc00
40 || (v6.segments()[0] & 0xffc0) == 0xfe80
41 }
42 },
43 }
44}
45
46fn is_blocked_v4(v4: Ipv4Addr) -> bool {
47 v4.is_loopback()
48 || v4.is_private()
49 || v4.is_link_local()
50 || v4.is_unspecified()
51 || v4.is_broadcast()
52 || v4.is_documentation()
53 || v4.is_multicast()
54 || v4.octets()[0] == 0
55 || v4.octets()[0] >= 240
56 || matches!(v4.octets(), [100, second, ..] if (64..=127).contains(&second))
57 || matches!(v4.octets(), [198, second, ..] if (18..=19).contains(&second))
58}
59
60fn embedded_ipv4(v6: Ipv6Addr) -> Option<Ipv4Addr> {
61 if let Some(mapped) = v6.to_ipv4() {
62 return Some(mapped);
63 }
64 let segments = v6.segments();
65 if segments[0] == 0x2002 {
66 return Some(Ipv4Addr::new(
67 (segments[1] >> 8) as u8,
68 segments[1] as u8,
69 (segments[2] >> 8) as u8,
70 segments[2] as u8,
71 ));
72 }
73 if segments[0] == 0x0064 && segments[1] == 0xff9b && segments[2..6] == [0, 0, 0, 0] {
74 return Some(Ipv4Addr::new(
75 (segments[6] >> 8) as u8,
76 segments[6] as u8,
77 (segments[7] >> 8) as u8,
78 segments[7] as u8,
79 ));
80 }
81 None
82}
83
84struct GuardedResolver;
85
86impl reqwest::dns::Resolve for GuardedResolver {
87 fn resolve(&self, name: reqwest::dns::Name) -> reqwest::dns::Resolving {
88 Box::pin(async move {
89 let host = name.as_str().to_owned();
90 let resolved = tokio::net::lookup_host((host.as_str(), 0)).await?;
91 let allowed: Vec<SocketAddr> =
92 resolved.filter(|addr| !is_blocked_ip(addr.ip())).collect();
93 if allowed.is_empty() {
94 return Err(Box::<dyn std::error::Error + Send + Sync>::from(format!(
95 "{host} resolves only to non-public addresses"
96 )));
97 }
98 Ok(Box::new(allowed.into_iter()) as reqwest::dns::Addrs)
99 })
100 }
101}
102
103#[derive(Debug, Clone, Copy)]
104pub struct HttpLimits {
105 pub connect_timeout: Duration,
106 pub read_timeout: Duration,
107 pub request_timeout: Duration,
108 pub max_response_bytes: u64,
109 pub block_private_addresses: bool,
110}
111
112impl Default for HttpLimits {
113 fn default() -> Self {
114 Self {
115 connect_timeout: Duration::from_secs(5),
116 read_timeout: Duration::from_secs(30),
117 request_timeout: Duration::from_secs(60),
118 max_response_bytes: 16 * 1024 * 1024,
119 block_private_addresses: true,
120 }
121 }
122}
123
124pub struct HttpRequest {
125 pub method: Method,
126 pub url: Url,
127 pub headers: HeaderMap,
128 pub body: Option<Bytes>,
129}
130
131impl HttpRequest {
132 pub fn get(url: Url) -> Self {
133 Self {
134 method: Method::GET,
135 url,
136 headers: HeaderMap::new(),
137 body: None,
138 }
139 }
140
141 pub fn post(url: Url, body: Bytes) -> Self {
142 Self {
143 method: Method::POST,
144 url,
145 headers: HeaderMap::new(),
146 body: Some(body),
147 }
148 }
149}
150
151#[derive(Debug)]
152pub struct HttpResponse {
153 pub status: StatusCode,
154 pub headers: HeaderMap,
155 pub body: Bytes,
156}
157
158pub type HttpFuture = Pin<Box<dyn Future<Output = Result<HttpResponse, NetworkError>> + Send>>;
159
160pub type ByteStream = Pin<Box<dyn futures::Stream<Item = Result<Bytes, NetworkError>> + Send>>;
161
162pub struct StreamedResponse {
163 pub status: StatusCode,
164 pub headers: HeaderMap,
165 pub body: ByteStream,
166}
167
168pub type StreamFuture =
169 Pin<Box<dyn Future<Output = Result<StreamedResponse, NetworkError>> + Send>>;
170
171pub trait HttpTransport: Send + Sync + 'static {
172 fn execute(&self, request: HttpRequest) -> HttpFuture;
173
174 fn execute_streamed(&self, request: HttpRequest) -> StreamFuture {
175 let response = self.execute(request);
176 Box::pin(async move {
177 let response = response.await?;
178 Ok(StreamedResponse {
179 status: response.status,
180 headers: response.headers,
181 body: Box::pin(futures::stream::once(std::future::ready(Ok(response.body)))),
182 })
183 })
184 }
185}
186
187pub struct ReqwestHttp {
188 client: reqwest::Client,
189 max_response_bytes: u64,
190 block_private_addresses: bool,
191}
192
193impl ReqwestHttp {
194 pub fn new(limits: HttpLimits) -> Result<Self, NetworkError> {
195 let mut builder = reqwest::Client::builder()
196 .connect_timeout(limits.connect_timeout)
197 .read_timeout(limits.read_timeout)
198 .timeout(limits.request_timeout)
199 .redirect(reqwest::redirect::Policy::none());
200 if limits.block_private_addresses {
201 builder = builder.dns_resolver(Arc::new(GuardedResolver));
202 }
203 if let Some(path) = std::env::var_os("KNOT_EXTRA_CA_FILE") {
204 let pem =
205 std::fs::read(&path).map_err(|error| NetworkError::Build(error.to_string()))?;
206 builder = reqwest::Certificate::from_pem_bundle(&pem)
207 .map_err(|error| NetworkError::Build(error.to_string()))?
208 .into_iter()
209 .fold(builder, reqwest::ClientBuilder::add_root_certificate);
210 }
211 let client = builder
212 .build()
213 .map_err(|error| NetworkError::Build(error.to_string()))?;
214 Ok(Self {
215 client,
216 max_response_bytes: limits.max_response_bytes,
217 block_private_addresses: limits.block_private_addresses,
218 })
219 }
220}
221
222impl HttpTransport for ReqwestHttp {
223 fn execute(&self, request: HttpRequest) -> HttpFuture {
224 let client = self.client.clone();
225 let limit = self.max_response_bytes;
226 let guard = self.block_private_addresses;
227 Box::pin(async move {
228 if let Some(host) = guard.then(|| blocked_literal(&request.url)).flatten() {
229 return Err(NetworkError::Blocked { host });
230 }
231 let mut builder = client
232 .request(request.method, request.url)
233 .headers(request.headers);
234 if let Some(body) = request.body {
235 builder = builder.body(body);
236 }
237 let response = builder.send().await.map_err(map_reqwest)?;
238 let status = response.status();
239 let headers = response.headers().clone();
240 if response.content_length().is_some_and(|len| len > limit) {
241 return Err(NetworkError::TooLarge { limit });
242 }
243 let body = bounded_body(response, limit).await?;
244 Ok(HttpResponse {
245 status,
246 headers,
247 body,
248 })
249 })
250 }
251
252 fn execute_streamed(&self, request: HttpRequest) -> StreamFuture {
253 let client = self.client.clone();
254 let guard = self.block_private_addresses;
255 Box::pin(async move {
256 if let Some(host) = guard.then(|| blocked_literal(&request.url)).flatten() {
257 return Err(NetworkError::Blocked { host });
258 }
259 let mut builder = client
260 .request(request.method, request.url)
261 .headers(request.headers);
262 if let Some(body) = request.body {
263 builder = builder.body(body);
264 }
265 let response = builder.send().await.map_err(map_reqwest)?;
266 let status = response.status();
267 let headers = response.headers().clone();
268 let body: ByteStream = Box::pin(response.bytes_stream().map_err(|error| {
269 if error.is_timeout() {
270 NetworkError::Timeout(error.to_string())
271 } else {
272 NetworkError::Body(error.to_string())
273 }
274 }));
275 Ok(StreamedResponse {
276 status,
277 headers,
278 body,
279 })
280 })
281 }
282}
283
284async fn bounded_body(response: reqwest::Response, limit: u64) -> Result<Bytes, NetworkError> {
285 response
286 .bytes_stream()
287 .map_err(|error| {
288 if error.is_timeout() {
289 NetworkError::Timeout(error.to_string())
290 } else {
291 NetworkError::Body(error.to_string())
292 }
293 })
294 .try_fold(Vec::new(), |mut buffer, chunk| async move {
295 if buffer.len() as u64 + chunk.len() as u64 > limit {
296 return Err(NetworkError::TooLarge { limit });
297 }
298 buffer.extend_from_slice(&chunk);
299 Ok(buffer)
300 })
301 .await
302 .map(Bytes::from)
303}
304
305fn blocked_literal(url: &Url) -> Option<String> {
306 match url.host()? {
307 Host::Ipv4(ip) if is_blocked_ip(IpAddr::V4(ip)) => Some(ip.to_string()),
308 Host::Ipv6(ip) if is_blocked_ip(IpAddr::V6(ip)) => Some(ip.to_string()),
309 _ => None,
310 }
311}
312
313fn map_reqwest(error: reqwest::Error) -> NetworkError {
314 if error.is_timeout() {
315 NetworkError::Timeout(error.to_string())
316 } else if error.is_connect() {
317 NetworkError::Connect(error.to_string())
318 } else {
319 NetworkError::Request(error.to_string())
320 }
321}
322
323pub struct FakeHttp<F> {
324 responder: F,
325}
326
327impl<F> FakeHttp<F>
328where
329 F: Fn(&HttpRequest) -> Result<HttpResponse, NetworkError> + Send + Sync + 'static,
330{
331 pub fn new(responder: F) -> Self {
332 Self { responder }
333 }
334}
335
336impl<F> HttpTransport for FakeHttp<F>
337where
338 F: Fn(&HttpRequest) -> Result<HttpResponse, NetworkError> + Send + Sync + 'static,
339{
340 fn execute(&self, request: HttpRequest) -> HttpFuture {
341 let result = (self.responder)(&request);
342 Box::pin(async move { result })
343 }
344}
345
346#[cfg(test)]
347mod tests {
348 use super::*;
349 use futures::StreamExt;
350 use std::io::{Read, Write};
351 use std::net::{SocketAddr, TcpListener};
352
353 #[test]
354 fn blocked_addresses_cover_the_internal_ranges() {
355 let blocked = [
356 "127.0.0.1",
357 "10.0.0.5",
358 "192.168.1.1",
359 "172.16.0.1",
360 "169.254.169.254",
361 "100.64.0.1",
362 "198.18.0.1",
363 "240.0.0.1",
364 "0.0.0.0",
365 "::1",
366 "::ffff:127.0.0.1",
367 "fd00::1",
368 "fe80::1",
369 "2002:7f00:1::",
370 "64:ff9b::7f00:1",
371 ];
372 for raw in blocked {
373 assert!(
374 is_blocked_ip(raw.parse().unwrap()),
375 "{raw} should be blocked"
376 );
377 }
378 let allowed = [
379 "1.1.1.1",
380 "8.8.8.8",
381 "93.184.216.34",
382 "2606:4700:4700::1111",
383 "2002:808:808::",
384 "64:ff9b::808:808",
385 ];
386 for raw in allowed {
387 assert!(
388 !is_blocked_ip(raw.parse().unwrap()),
389 "{raw} should be allowed"
390 );
391 }
392 }
393
394 fn tiny_limits(max_response_bytes: u64, request_timeout: Duration) -> HttpLimits {
395 HttpLimits {
396 connect_timeout: Duration::from_millis(200),
397 read_timeout: Duration::from_millis(200),
398 request_timeout,
399 max_response_bytes,
400 block_private_addresses: false,
401 }
402 }
403
404 fn serve_body(body: Vec<u8>, with_content_length: bool) -> SocketAddr {
405 let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
406 let addr = listener.local_addr().expect("addr");
407 std::thread::spawn(move || {
408 if let Ok((mut stream, _)) = listener.accept() {
409 let _ = stream.read(&mut [0u8; 1024]);
410 let header = if with_content_length {
411 format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len())
412 } else {
413 "HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_string()
414 };
415 let _ = stream.write_all(header.as_bytes());
416 let _ = stream.write_all(&body);
417 }
418 });
419 addr
420 }
421
422 fn serve_hang() -> SocketAddr {
423 let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
424 let addr = listener.local_addr().expect("addr");
425 std::thread::spawn(move || {
426 if let Ok((mut stream, _)) = listener.accept() {
427 let _ = stream.read(&mut [0u8; 1024]);
428 std::thread::sleep(Duration::from_secs(5));
429 drop(stream);
430 }
431 });
432 addr
433 }
434
435 async fn fetch(addr: SocketAddr, limits: HttpLimits) -> Result<HttpResponse, NetworkError> {
436 let transport = ReqwestHttp::new(limits).expect("client builds");
437 let url = Url::parse(&format!("http://{addr}/")).expect("url");
438 transport.execute(HttpRequest::get(url)).await
439 }
440
441 #[tokio::test]
442 async fn oversized_response_is_rejected_whether_declared_or_streamed() {
443 futures::stream::iter([true, false])
444 .for_each(|with_content_length| async move {
445 let addr = serve_body(vec![0u8; 4096], with_content_length);
446 let result = fetch(addr, tiny_limits(64, Duration::from_secs(2))).await;
447 assert!(matches!(result, Err(NetworkError::TooLarge { limit: 64 })));
448 })
449 .await;
450 }
451
452 #[tokio::test]
453 async fn small_response_within_limit_succeeds() {
454 let addr = serve_body(b"pong".to_vec(), true);
455 let response = fetch(addr, tiny_limits(64, Duration::from_secs(2)))
456 .await
457 .expect("response within limit");
458 assert_eq!(response.body.as_ref(), b"pong");
459 }
460
461 #[tokio::test]
462 async fn unresponsive_server_times_out() {
463 let addr = serve_hang();
464 let result = fetch(addr, tiny_limits(1024, Duration::from_millis(150))).await;
465 assert!(matches!(result, Err(NetworkError::Timeout(_))));
466 }
467
468 #[tokio::test]
469 async fn connect_failure_surfaces_typed_error() {
470 let listener = TcpListener::bind("127.0.0.1:0").expect("bind");
471 let addr = listener.local_addr().expect("addr");
472 drop(listener);
473 let result = fetch(addr, tiny_limits(1024, Duration::from_secs(2))).await;
474 assert!(matches!(
475 result,
476 Err(NetworkError::Connect(_) | NetworkError::Request(_) | NetworkError::Timeout(_))
477 ));
478 }
479
480 #[test]
481 fn fake_http_returns_canned_response() {
482 let transport = FakeHttp::new(|request: &HttpRequest| {
483 assert_eq!(request.method, Method::GET);
484 Ok(HttpResponse {
485 status: StatusCode::OK,
486 headers: HeaderMap::new(),
487 body: Bytes::from_static(b"pong"),
488 })
489 });
490 let request = HttpRequest::get(Url::parse("https://oyster.cafe/ping").unwrap());
491 let response =
492 futures::executor::block_on(transport.execute(request)).expect("fake response");
493 assert_eq!(response.status, StatusCode::OK);
494 assert_eq!(response.body.as_ref(), b"pong");
495 }
496}