1use std::io::{Read, Write};
14use std::net::{SocketAddr, TcpStream};
15use std::sync::Arc;
16use std::time::{Duration, Instant};
17
18use crate::net::{self, Net, Target};
19
20pub const MAX_BODY: usize = 256 * 1024;
22pub const MAX_REDIRECTS: usize = 5;
24
25#[derive(Clone)]
27pub struct HttpProbe {
28 pub url: String,
29 pub method: String,
31 pub headers: Vec<(String, String)>,
32 pub timeout: Duration,
33 pub follow_redirects: bool,
34 pub allow_private: bool,
35 pub connect_to: Option<SocketAddr>,
39 pub tls: Arc<rustls::ClientConfig>,
40}
41
42#[derive(Debug, Clone, Default)]
44pub struct HttpAnswer {
45 pub status: u16,
46 pub location: Option<String>,
47 pub body: Vec<u8>,
49 pub latency: Duration,
50 pub cert_expires: Option<u64>,
52 pub final_url: String,
54 pub access_refused: bool,
57}
58
59pub fn display_url(url: &str) -> String {
61 let u = url.split(['?', '#']).next().unwrap_or_default();
62 u.to_string()
63}
64
65pub fn http(p: &HttpProbe) -> Result<HttpAnswer, String> {
67 let deadline = Instant::now() + p.timeout;
68 let started = Instant::now();
69 let mut url = p.url.clone();
70 for hop in 0..=MAX_REDIRECTS {
71 let t = net::parse_url(&url).map_err(|e| format!("URL: {e}"))?;
72 let connect_to = if hop == 0 { p.connect_to } else { None };
74 let mut a = once(p, &t, connect_to, deadline)?;
75 a.final_url = display_url(&url);
76 let redirect = (300..400).contains(&a.status) && a.status != 304;
77 match (&a.location, redirect && p.follow_redirects) {
78 (Some(loc), true) if hop < MAX_REDIRECTS => url = join(&t, loc),
79 _ => {
80 a.latency = started.elapsed();
81 return Ok(a);
82 }
83 }
84 }
85 Err(format!("more than {MAX_REDIRECTS} redirects"))
86}
87
88pub fn join(base: &Target, loc: &str) -> String {
90 let l = loc.trim();
91 if l.starts_with("http://") || l.starts_with("https://") {
92 return l.to_string();
93 }
94 let scheme = if base.https { "https" } else { "http" };
95 let default = if base.https { 443 } else { 80 };
96 let host = if base.host.contains(':') {
97 format!("[{}]", base.host)
98 } else {
99 base.host.clone()
100 };
101 let authority = if base.port == default {
102 host
103 } else {
104 format!("{host}:{}", base.port)
105 };
106 if let Some(rest) = l.strip_prefix("//") {
107 return format!("{scheme}://{rest}");
108 }
109 if l.starts_with('/') {
110 return format!("{scheme}://{authority}{l}");
111 }
112 let dir = base.path.split('?').next().unwrap_or("/");
113 let dir = &dir[..dir.rfind('/').map(|i| i + 1).unwrap_or(1)];
114 format!("{scheme}://{authority}{dir}{l}")
115}
116
117fn left(deadline: Instant) -> Result<Duration, String> {
118 let d = deadline.saturating_duration_since(Instant::now());
119 if d.is_zero() {
120 Err("timed out".into())
121 } else {
122 Ok(d.max(Duration::from_millis(1)))
123 }
124}
125
126pub fn connect(
128 host: &str,
129 port: u16,
130 allow_private: bool,
131 deadline: Instant,
132) -> Result<TcpStream, String> {
133 let addrs = net::resolve(host, port, allow_private).map_err(|e| e.message)?;
134 let mut last = None;
135 for a in addrs {
136 match dial(a, deadline) {
137 Ok(s) => {
138 let peer = s.peer_addr().map_err(|e| format!("{host}: {e}"))?;
139 net::check_ip(peer.ip(), allow_private)
140 .map_err(|why| format!("refusing {host}: {why}"))?;
141 return Ok(s);
142 }
143 Err(e) => last = Some(e),
144 }
145 }
146 Err(format!(
147 "cannot connect to {host}:{port}: {}",
148 last.unwrap_or_default()
149 ))
150}
151
152pub fn dial(a: SocketAddr, deadline: Instant) -> Result<TcpStream, String> {
154 let s = TcpStream::connect_timeout(&a, left(deadline)?).map_err(|e| {
155 if e.kind() == std::io::ErrorKind::TimedOut || e.kind() == std::io::ErrorKind::WouldBlock {
156 "timed out connecting".to_string()
157 } else {
158 e.to_string()
159 }
160 })?;
161 s.set_nodelay(true).ok();
162 Ok(s)
163}
164
165pub fn tcp(
167 host: &str,
168 port: u16,
169 timeout: Duration,
170 allow_private: bool,
171) -> Result<Duration, String> {
172 let started = Instant::now();
173 let s = connect(host, port, allow_private, started + timeout)?;
174 drop(s);
175 Ok(started.elapsed())
176}
177
178fn host_header(t: &Target) -> String {
179 let h = if t.host.contains(':') {
180 format!("[{}]", t.host)
181 } else {
182 t.host.clone()
183 };
184 let default = if t.https { 443 } else { 80 };
185 if t.port == default {
186 h
187 } else {
188 format!("{h}:{}", t.port)
189 }
190}
191
192fn request_head(p: &HttpProbe, t: &Target) -> Result<String, String> {
193 let mut head = format!(
194 "{} {} HTTP/1.1\r\nHost: {}\r\nUser-Agent: isb-monitor/{}\r\nAccept: */*\r\nConnection: close\r\n",
195 p.method,
196 t.path,
197 host_header(t),
198 env!("CARGO_PKG_VERSION"),
199 );
200 for (k, v) in &p.headers {
201 if k.contains(['\r', '\n', ':']) || v.contains(['\r', '\n']) {
202 return Err(format!("a bad header {k:?}"));
203 }
204 head.push_str(&format!("{k}: {v}\r\n"));
205 }
206 head.push_str("\r\n");
207 Ok(head)
208}
209
210fn once(
212 p: &HttpProbe,
213 t: &Target,
214 connect_to: Option<SocketAddr>,
215 deadline: Instant,
216) -> Result<HttpAnswer, String> {
217 let head = request_head(p, t)?;
218 let tcp = match connect_to {
219 Some(a) => dial(a, deadline).map_err(|e| format!("cannot connect to {a}: {e}"))?,
220 None => connect(&t.host, t.port, p.allow_private, deadline)?,
221 };
222 let head_only = p.method == "HEAD";
223 if !t.https {
224 let sock = s_clone(&tcp)?;
225 let mut s = tcp;
226 let raw = exchange(&mut s, &sock, head.as_bytes(), deadline, &t.host)?;
227 return parse(&raw, head_only);
228 }
229 let net = Net {
230 allow_private: p.allow_private,
231 tls: p.tls.clone(),
232 };
233 let mut s = net::tls(&net, &t.host, tcp).map_err(|e| e.message)?;
234 let sock = s_clone(&s.sock)?;
235 while s.conn.is_handshaking() {
236 sock.set_read_timeout(Some(left(deadline)?)).ok();
237 s.conn
238 .complete_io(&mut s.sock)
239 .map_err(|e| tls_error(&t.host, &e))?;
240 }
241 let cert_expires = s
242 .conn
243 .peer_certificates()
244 .and_then(|c| c.first())
245 .and_then(|c| cert_not_after(c.as_ref()));
246 let raw = exchange(&mut s, &sock, head.as_bytes(), deadline, &t.host)?;
247 let mut a = parse(&raw, head_only)?;
248 a.cert_expires = cert_expires;
249 Ok(a)
250}
251
252fn s_clone(s: &TcpStream) -> Result<TcpStream, String> {
253 s.try_clone().map_err(|e| e.to_string())
254}
255
256fn tls_error(host: &str, e: &std::io::Error) -> String {
257 let m = e.to_string();
258 if e.kind() == std::io::ErrorKind::WouldBlock || e.kind() == std::io::ErrorKind::TimedOut {
259 format!("{host}: timed out in the TLS handshake")
260 } else {
261 format!("TLS to {host}: {m}")
262 }
263}
264
265fn exchange<S: Read + Write>(
268 s: &mut S,
269 sock: &TcpStream,
270 head: &[u8],
271 deadline: Instant,
272 host: &str,
273) -> Result<Vec<u8>, String> {
274 let io = |e: std::io::Error| {
275 if e.kind() == std::io::ErrorKind::WouldBlock || e.kind() == std::io::ErrorKind::TimedOut {
276 format!("{host}: timed out waiting for an answer")
277 } else {
278 format!("{host}: {e}")
279 }
280 };
281 sock.set_write_timeout(Some(left(deadline)?)).ok();
282 s.write_all(head).map_err(io)?;
283 s.flush().map_err(io)?;
284 let mut out = Vec::new();
285 let mut buf = [0u8; 16384];
286 loop {
287 sock.set_read_timeout(Some(
288 left(deadline).map_err(|_| format!("{host}: timed out waiting for an answer"))?,
289 ))
290 .ok();
291 match s.read(&mut buf) {
292 Ok(0) => break,
293 Ok(n) => {
294 out.extend_from_slice(&buf[..n]);
295 if out.len() >= MAX_BODY + 16384 || complete(&out) {
296 break;
297 }
298 }
299 Err(_) if header_end(&out).is_some() => break,
300 Err(e) => return Err(io(e)),
301 }
302 }
303 if out.is_empty() {
304 return Err(format!("{host}: closed the connection without answering"));
305 }
306 Ok(out)
307}
308
309fn header_end(b: &[u8]) -> Option<usize> {
310 b.windows(4).position(|w| w == b"\r\n\r\n").map(|p| p + 4)
311}
312
313fn complete(b: &[u8]) -> bool {
316 let Some(end) = header_end(b) else {
317 return false;
318 };
319 let mut headers = [httparse::EMPTY_HEADER; 64];
320 let mut r = httparse::Response::new(&mut headers);
321 if r.parse(b).is_err() {
322 return false;
323 }
324 let header = |name: &str| {
325 r.headers
326 .iter()
327 .find(|h| h.name.eq_ignore_ascii_case(name))
328 .and_then(|h| std::str::from_utf8(h.value).ok())
329 .map(|v| v.trim().to_ascii_lowercase())
330 };
331 if header("transfer-encoding").is_some_and(|v| v.contains("chunked")) {
332 return b[end..].ends_with(b"0\r\n\r\n");
333 }
334 match header("content-length").and_then(|v| v.parse::<usize>().ok()) {
335 Some(n) => b.len() >= end + n,
336 None => false,
337 }
338}
339
340pub fn parse(raw: &[u8], head_only: bool) -> Result<HttpAnswer, String> {
342 let mut headers = [httparse::EMPTY_HEADER; 64];
343 let mut r = httparse::Response::new(&mut headers);
344 let end = match r.parse(raw) {
345 Ok(httparse::Status::Complete(n)) => n,
346 Ok(httparse::Status::Partial) => return Err("an incomplete HTTP answer".into()),
347 Err(e) => return Err(format!("not an HTTP answer: {e}")),
348 };
349 let header = |name: &str| {
350 r.headers
351 .iter()
352 .find(|h| h.name.eq_ignore_ascii_case(name))
353 .and_then(|h| std::str::from_utf8(h.value).ok())
354 .map(|v| v.trim().to_string())
355 };
356 let chunked =
357 header("transfer-encoding").is_some_and(|v| v.to_ascii_lowercase().contains("chunked"));
358 let raw_body = if head_only { &[][..] } else { &raw[end..] };
359 let mut body = if chunked {
360 dechunk(raw_body)
361 } else {
362 raw_body.to_vec()
363 };
364 body.truncate(MAX_BODY);
365 let status = r.code.unwrap_or(0);
366 let access_refused = matches!(status, 401 | 403)
371 && (header("cf-access-domain").is_some()
372 || header("cf-access-aud").is_some()
373 || header("www-authenticate")
374 .is_some_and(|v| v.contains("cloudflare-access-protected-resource")));
375 Ok(HttpAnswer {
376 status,
377 location: header("location"),
378 body,
379 access_refused,
380 ..Default::default()
381 })
382}
383
384pub fn dechunk(mut b: &[u8]) -> Vec<u8> {
386 let mut out = Vec::new();
387 while let Some(i) = b.windows(2).position(|w| w == b"\r\n") {
388 let size = std::str::from_utf8(&b[..i])
389 .ok()
390 .and_then(|s| usize::from_str_radix(s.split(';').next()?.trim(), 16).ok());
391 let Some(n) = size else { break };
392 if n == 0 {
393 break;
394 }
395 let start = i + 2;
396 let stop = (start + n).min(b.len());
397 out.extend_from_slice(&b[start..stop]);
398 if start + n + 2 > b.len() {
399 break;
400 }
401 b = &b[start + n + 2..];
402 }
403 out
404}
405
406pub fn cert_not_after(der: &[u8]) -> Option<u64> {
411 let (_, cert, _) = tlv(der)?;
412 let (_, tbs, _) = tlv(cert)?;
413 let mut rest = tbs;
414 let (tag, _, r) = tlv(rest)?;
415 if tag == 0xa0 {
416 rest = r; }
418 for _ in 0..3 {
419 rest = tlv(rest)?.2; }
421 let (tag, validity, _) = tlv(rest)?;
422 if tag != 0x30 {
423 return None;
424 }
425 let (_, _, r) = tlv(validity)?; let (tag, t, _) = tlv(r)?;
427 let s = std::str::from_utf8(t).ok()?;
428 match tag {
429 0x17 => asn1_time(s, true),
430 0x18 => asn1_time(s, false),
431 _ => None,
432 }
433}
434
435fn tlv(b: &[u8]) -> Option<(u8, &[u8], &[u8])> {
437 let tag = *b.first()?;
438 let first = *b.get(1)?;
439 let (len, hdr) = if first < 0x80 {
440 (first as usize, 2)
441 } else {
442 let n = (first & 0x7f) as usize;
443 if n == 0 || n > 4 {
444 return None;
445 }
446 let mut len = 0usize;
447 for i in 0..n {
448 len = (len << 8) | *b.get(2 + i)? as usize;
449 }
450 (len, 2 + n)
451 };
452 let end = hdr.checked_add(len)?;
453 (end <= b.len()).then(|| (tag, &b[hdr..end], &b[end..]))
454}
455
456fn asn1_time(s: &str, utc: bool) -> Option<u64> {
458 let s = s.strip_suffix('Z')?;
459 let (year, rest) = if utc {
460 let y: i64 = s.get(..2)?.parse().ok()?;
461 (if y >= 50 { 1900 + y } else { 2000 + y }, s.get(2..)?)
462 } else {
463 (s.get(..4)?.parse().ok()?, s.get(4..)?)
464 };
465 let n = |i: usize| -> Option<i64> { rest.get(i..i + 2)?.parse().ok() };
466 let (mo, d, h, mi, se) = (n(0)?, n(2)?, n(4)?, n(6)?, n(8)?);
467 let days = days_from_civil(year, mo, d);
468 let t = days * 86400 + h * 3600 + mi * 60 + se;
469 u64::try_from(t).ok()
470}
471
472pub fn days_from_civil(y: i64, m: i64, d: i64) -> i64 {
474 let y = if m <= 2 { y - 1 } else { y };
475 let era = y.div_euclid(400);
476 let yoe = y - era * 400;
477 let mp = (m + 9) % 12;
478 let doy = (153 * mp + 2) / 5 + d - 1;
479 let doe = yoe * 365 + yoe / 4 - yoe / 100 + doy;
480 era * 146_097 + doe - 719_468
481}
482
483#[cfg(test)]
484mod tests {
485 use super::*;
486 use std::net::TcpListener;
487
488 fn probe(url: &str) -> HttpProbe {
489 HttpProbe {
490 url: url.into(),
491 method: "GET".into(),
492 headers: vec![("X-Token".into(), "t0k".into())],
493 timeout: Duration::from_secs(3),
494 follow_redirects: false,
495 allow_private: true,
496 connect_to: None,
497 tls: net::default_tls(),
498 }
499 }
500
501 fn server(answers: Vec<String>) -> (u16, std::thread::JoinHandle<Vec<String>>) {
504 let l = TcpListener::bind("127.0.0.1:0").unwrap();
505 let port = l.local_addr().unwrap().port();
506 let h = std::thread::spawn(move || {
507 let mut got = Vec::new();
508 for a in answers {
509 let (mut s, _) = l.accept().unwrap();
510 let mut buf = vec![0u8; 8192];
511 let n = s.read(&mut buf).unwrap();
512 got.push(String::from_utf8_lossy(&buf[..n]).to_string());
513 s.write_all(a.as_bytes()).unwrap();
514 }
515 got
516 });
517 (port, h)
518 }
519
520 #[test]
521 fn access_refusals_are_told_from_the_app() {
522 let refused = |raw: &str| parse(raw.as_bytes(), true).unwrap().access_refused;
523 assert!(refused(
525 "HTTP/1.1 403 Forbidden\r\ncf-access-domain: broker.example.com\r\n\r\n"
526 ));
527 assert!(refused(
528 "HTTP/1.1 401 Unauthorized\r\nwww-authenticate: Bearer realm=\"OAuth\", resource_metadata=\"https://a.example.com/.well-known/cloudflare-access-protected-resource/x\"\r\n\r\n"
529 ));
530 assert!(!refused("HTTP/1.1 403 Forbidden\r\n\r\n"));
532 assert!(!refused(
533 "HTTP/1.1 401 Unauthorized\r\nwww-authenticate: Basic realm=\"x\"\r\n\r\n"
534 ));
535 assert!(!refused(
536 "HTTP/1.1 200 OK\r\ncf-access-domain: broker.example.com\r\n\r\n"
537 ));
538 }
539
540 #[test]
541 fn get_with_headers_and_a_chunked_body() {
542 let (port, h) = server(vec![
543 "HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n6\r\n world\r\n0\r\n\r\n"
544 .into(),
545 ]);
546 let a = http(&probe(&format!("http://127.0.0.1:{port}/health?k=1"))).unwrap();
547 assert_eq!(a.status, 200);
548 assert_eq!(a.body, b"hello world");
549 assert_eq!(a.final_url, format!("http://127.0.0.1:{port}/health"));
550 let req = &h.join().unwrap()[0];
551 assert!(req.starts_with("GET /health?k=1 HTTP/1.1\r\n"), "{req}");
552 assert!(req.contains("X-Token: t0k\r\n"), "{req}");
553 }
554
555 #[test]
556 fn redirects_are_followed_only_when_asked() {
557 let (port, h) = server(vec![
558 "HTTP/1.1 302 Found\r\nLocation: /login\r\nContent-Length: 0\r\n\r\n".into(),
559 "HTTP/1.1 302 Found\r\nLocation: /login\r\nContent-Length: 0\r\n\r\n".into(),
560 "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok".into(),
561 ]);
562 let url = format!("http://127.0.0.1:{port}/");
563 let a = http(&probe(&url)).unwrap();
564 assert_eq!((a.status, a.location.as_deref()), (302, Some("/login")));
565 let mut p = probe(&url);
566 p.follow_redirects = true;
567 let a = http(&p).unwrap();
568 assert_eq!((a.status, a.body.as_slice()), (200, &b"ok"[..]));
569 assert!(a.final_url.ends_with("/login"));
570 h.join().unwrap();
571 }
572
573 #[test]
574 fn the_address_policy_holds_unless_dialled_by_reference() {
575 let (port, h) = server(vec!["HTTP/1.1 204 No Content\r\n\r\n".into()]);
576 let mut p = probe(&format!("http://127.0.0.1:{port}/"));
577 p.allow_private = false;
578 let e = http(&p).unwrap_err();
579 assert!(e.contains("refusing 127.0.0.1"), "{e}");
580 let mut p = probe("http://shop.example.com/");
582 p.allow_private = false;
583 p.connect_to = Some(format!("127.0.0.1:{port}").parse().unwrap());
584 assert_eq!(http(&p).unwrap().status, 204);
585 assert!(h.join().unwrap()[0].contains("Host: shop.example.com\r\n"));
586 assert!(tcp("127.0.0.1", port, Duration::from_secs(1), false).is_err());
588 }
589
590 #[test]
591 fn timeouts_and_refusals() {
592 let l = TcpListener::bind("127.0.0.1:0").unwrap();
594 let port = l.local_addr().unwrap().port();
595 let mut p = probe(&format!("http://127.0.0.1:{port}/"));
596 p.timeout = Duration::from_millis(300);
597 let started = Instant::now();
598 let e = http(&p).unwrap_err();
599 assert!(e.contains("timed out"), "{e}");
600 assert!(started.elapsed() < Duration::from_secs(2));
601 drop(l);
602 let e = http(&probe(&format!("http://127.0.0.1:{port}/"))).unwrap_err();
603 assert!(e.contains("cannot connect"), "{e}");
604 let e = tcp("127.0.0.1", port, Duration::from_secs(1), true).unwrap_err();
605 assert!(e.contains("cannot connect"), "{e}");
606 }
607
608 #[test]
609 fn redirect_targets() {
610 let t = net::parse_url("https://a.example.com:8443/x/y?q").unwrap();
611 assert_eq!(join(&t, "/z"), "https://a.example.com:8443/z");
612 assert_eq!(join(&t, "z"), "https://a.example.com:8443/x/z");
613 assert_eq!(join(&t, "http://b.example/"), "http://b.example/");
614 assert_eq!(join(&t, "//c.example/p"), "https://c.example/p");
615 assert_eq!(display_url("https://a/b?token=1#f"), "https://a/b");
616 }
617
618 #[test]
619 fn certificate_expiry() {
620 let mut params = rcgen::CertificateParams::new(vec!["shop.example.com".into()]).unwrap();
621 params.not_after = rcgen::date_time_ymd(2031, 7, 9);
622 let key = rcgen::KeyPair::generate().unwrap();
623 let cert = params.self_signed(&key).unwrap();
624 let want = days_from_civil(2031, 7, 9) as u64 * 86400;
625 assert_eq!(cert_not_after(cert.der().as_ref()), Some(want));
626 params.not_after = rcgen::date_time_ymd(2051, 1, 2);
628 let cert = params.self_signed(&key).unwrap();
629 assert_eq!(
630 cert_not_after(cert.der().as_ref()),
631 Some(days_from_civil(2051, 1, 2) as u64 * 86400)
632 );
633 assert_eq!(cert_not_after(b"\x30\x03\x02\x01\x01"), None);
634 assert_eq!(cert_not_after(&[]), None);
635 assert_eq!(days_from_civil(1970, 1, 1), 0);
636 assert_eq!(days_from_civil(2000, 3, 1), 11017);
637 }
638}