This repository has no description
0

Configure Feed

Select the types of activity you want to include in your feed.

core / knot2 / crates / knot-runtime / src / http.rs
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}