use std::time::Duration;
use tokio::{
io::{AsyncReadExt as _, AsyncWriteExt as _},
net::TcpStream,
};
use tracing::trace;
const H2_PREFACE: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n";
const H2_SETTINGS: &[u8] = &[0, 0, 0, 4, 0, 0, 0, 0, 0];
const H2_SETTINGS_ACK: &[u8] = &[0, 0, 0, 4, 1, 0, 0, 0, 0];
const H2_GOAWAY: &[u8] = &[0, 0, 8, 7, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0];
const H2_FRAME_TYPE_SETTINGS: u8 = 0x04;
const H2_FRAME_HEADER_LEN: usize = 9;
const H2_CLOSE_GRACE: Duration = Duration::from_millis(100);
pub async fn tcp_probe(addr: &str, timeout: Duration) -> bool {
match tokio::time::timeout(timeout, TcpStream::connect(addr)).await {
Ok(Ok(_stream)) => {
trace!(addr, "tcp health check succeeded");
true
},
Ok(Err(e)) => {
trace!(addr, error = %e, "tcp health check connect failed");
false
},
Err(_) => {
trace!(addr, "tcp health check timed out");
false
},
}
}
pub async fn http_probe(addr: &str, path: &str, expected_status: u16, timeout: Duration) -> bool {
http_probe_with_request(addr, &build_http_probe_request(path), expected_status, timeout).await
}
pub(crate) fn build_http_probe_request(path: &str) -> String {
format!("GET {path} HTTP/1.1\r\nHost: health-check\r\nConnection: close\r\n\r\n")
}
pub(crate) async fn http_probe_with_request(
addr: &str,
request: &str,
expected_status: u16,
timeout: Duration,
) -> bool {
let result = tokio::time::timeout(timeout, http_probe_inner(addr, request, expected_status)).await;
if let Ok(ok) = result {
ok
} else {
trace!(addr, "health check timed out");
false
}
}
pub(crate) fn parse_status_code(response: &str) -> Option<u16> {
let mut parts = response.lines().next()?.splitn(3, ' ');
parts.next().filter(|version| is_http_version(version))?;
parts
.next()
.filter(|status| status.len() == 3 && status.bytes().all(|byte| byte.is_ascii_digit()))?
.parse()
.ok()
}
fn is_http_version(version: &str) -> bool {
matches!(
version.strip_prefix("HTTP/").map(str::as_bytes),
Some([major, b'.', minor]) if major.is_ascii_digit() && minor.is_ascii_digit()
)
}
pub async fn h2_probe(addr: &str, timeout: Duration) -> bool {
let result = tokio::time::timeout(timeout, h2_probe_inner(addr)).await;
if let Ok(ok) = result {
ok
} else {
trace!(addr, "h2 health check timed out");
false
}
}
pub(crate) fn is_settings_frame(buf: &[u8]) -> bool {
buf.len() >= H2_FRAME_HEADER_LEN && buf.get(3) == Some(&H2_FRAME_TYPE_SETTINGS)
}
async fn http_probe_inner(addr: &str, request: &str, expected_status: u16) -> bool {
let mut stream = match TcpStream::connect(addr).await {
Ok(s) => s,
Err(e) => {
trace!(addr, error = %e, "health check connect failed");
return false;
},
};
if let Err(e) = stream.write_all(request.as_bytes()).await {
trace!(addr, error = %e, "health check write failed");
return false;
}
match read_status_line(&mut stream, addr).await {
Some(data) => parse_status_code(&data) == Some(expected_status),
None => false,
}
}
#[expect(clippy::indexing_slicing, reason = "bounded by filled counter")]
async fn read_status_line(stream: &mut TcpStream, addr: &str) -> Option<String> {
let mut buf = [0_u8; 256];
let mut filled = 0;
loop {
match stream.read(&mut buf[filled..]).await {
Ok(0) => break,
Ok(n) => {
filled += n;
if buf[..filled].windows(2).any(|w| w == b"\r\n") || filled >= buf.len() {
break;
}
},
Err(e) => {
trace!(addr, error = %e, "health check read failed");
return None;
},
}
}
if filled == 0 {
trace!(addr, "health check received empty response");
return None;
}
Some(String::from_utf8_lossy(&buf[..filled]).into_owned())
}
async fn h2_probe_inner(addr: &str) -> bool {
let mut stream = match TcpStream::connect(addr).await {
Ok(s) => s,
Err(e) => {
trace!(addr, error = %e, "h2 health check connect failed");
return false;
},
};
if !h2_send_preface(&mut stream, addr).await {
return false;
}
if !h2_read_settings(&mut stream, addr).await {
return false;
}
drop(tokio::time::timeout(H2_CLOSE_GRACE, h2_close_gracefully(&mut stream)).await);
true
}
async fn h2_send_preface(stream: &mut TcpStream, addr: &str) -> bool {
if let Err(e) = stream.write_all(H2_PREFACE).await {
trace!(addr, error = %e, "h2 health check preface write failed");
return false;
}
if let Err(e) = stream.write_all(H2_SETTINGS).await {
trace!(addr, error = %e, "h2 health check settings write failed");
return false;
}
true
}
async fn h2_read_settings(stream: &mut TcpStream, addr: &str) -> bool {
let mut header = [0_u8; H2_FRAME_HEADER_LEN];
if let Err(e) = stream.read_exact(&mut header).await {
trace!(addr, error = %e, "h2 health check read failed");
return false;
}
if !is_settings_frame(&header) {
trace!(addr, "h2 health check did not receive SETTINGS frame");
return false;
}
true
}
async fn h2_close_gracefully(stream: &mut TcpStream) {
drop(stream.write_all(H2_SETTINGS_ACK).await);
drop(stream.write_all(H2_GOAWAY).await);
let mut drain = [0_u8; 256];
while stream.read(&mut drain).await.unwrap_or(0) > 0 {}
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::indexing_slicing, reason = "tests")]
mod tests {
use super::*;
#[test]
fn parse_status_200() {
assert_eq!(
parse_status_code("HTTP/1.1 200 OK\r\nContent-Length: 0\r\n"),
Some(200),
"should parse 200 from status line"
);
}
#[test]
fn parse_status_503() {
assert_eq!(
parse_status_code("HTTP/1.1 503 Service Unavailable\r\n"),
Some(503),
"should parse 503 from status line"
);
}
#[test]
fn parse_status_204() {
assert_eq!(
parse_status_code("HTTP/1.1 204 No Content\r\n"),
Some(204),
"should parse 204 from status line"
);
}
#[test]
fn parse_status_garbage() {
assert_eq!(
parse_status_code("not a valid http response"),
None,
"should return None for garbage input"
);
}
#[test]
fn parse_status_empty() {
assert_eq!(parse_status_code(""), None, "should return None for empty input");
}
#[test]
fn parse_status_partial() {
assert_eq!(
parse_status_code("HTTP/1.1"),
None,
"should return None for incomplete status line"
);
}
#[test]
fn parse_status_rejects_non_http_status_lines() {
for line in ["SMTP 200 ready\r\n", "garbage 200 anything"] {
assert_eq!(
parse_status_code(line),
None,
"a non-HTTP greeting must not be read as status 200: {line:?}"
);
}
}
#[test]
fn parse_status_http10() {
assert_eq!(
parse_status_code("HTTP/1.0 301 Moved Permanently\r\n"),
Some(301),
"should parse HTTP/1.0 status lines"
);
}
#[test]
fn parse_status_rejects_malformed_http_version() {
for line in ["HTTP/nope 200 ready\r\n", "HTTP/ 200 OK\r\n", "HTTP/11 200 OK\r\n"] {
assert_eq!(
parse_status_code(line),
None,
"HTTP-version must be HTTP/ DIGIT . DIGIT (RFC 9112 §4): {line:?}"
);
}
}
#[test]
fn parse_status_rejects_non_three_digit_status() {
for line in ["HTTP/1.1 0200 weird\r\n", "HTTP/1.1 20 OK\r\n", "HTTP/1.1 2000 OK\r\n"] {
assert_eq!(
parse_status_code(line),
None,
"status-code must be exactly 3DIGIT (RFC 9112 §4): {line:?}"
);
}
}
#[tokio::test]
async fn tcp_probe_refuses_nonexistent() {
let result = tcp_probe("127.0.0.1:1", Duration::from_millis(100)).await;
assert!(!result, "should fail for non-listening port");
}
#[tokio::test]
async fn http_probe_refuses_nonexistent() {
let result = http_probe("127.0.0.1:1", "/", 200, Duration::from_millis(100)).await;
assert!(!result, "should fail for non-listening port");
}
#[tokio::test]
async fn tcp_probe_succeeds_on_listener() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe = tokio::spawn(async move { tcp_probe(&addr, Duration::from_secs(1)).await });
let (_socket, _peer) = listener.accept().await.unwrap();
let result = probe.await.unwrap();
assert!(result, "should succeed when endpoint is listening");
}
#[tokio::test]
async fn http_probe_succeeds_with_matching_status() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe = tokio::spawn(async move { http_probe(&probe_addr, "/health", 200, Duration::from_secs(1)).await });
let (mut socket, _peer) = listener.accept().await.unwrap();
let mut buf = [0_u8; 512];
let _ = socket.read(&mut buf).await.unwrap();
socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
.await
.unwrap();
socket.shutdown().await.unwrap();
let result = probe.await.unwrap();
assert!(result, "should succeed with matching 200 status");
}
#[test]
fn is_settings_frame_valid() {
assert!(
is_settings_frame(&[0, 0, 0, 4, 0, 0, 0, 0, 0]),
"valid SETTINGS frame should be recognized"
);
}
#[test]
fn is_settings_frame_with_ack() {
assert!(
is_settings_frame(&[0, 0, 0, 4, 1, 0, 0, 0, 0]),
"SETTINGS ACK frame should be recognized"
);
}
#[test]
fn is_settings_frame_wrong_type() {
assert!(
!is_settings_frame(&[0, 0, 0, 1, 0, 0, 0, 0, 0]),
"non-SETTINGS frame type should be rejected"
);
}
#[test]
fn is_settings_frame_too_short() {
assert!(
!is_settings_frame(&[0, 0, 0, 4]),
"buffer shorter than frame header should be rejected"
);
}
#[test]
fn is_settings_frame_empty() {
assert!(!is_settings_frame(&[]), "empty buffer should be rejected");
}
#[tokio::test]
async fn h2_probe_refuses_nonexistent() {
let result = h2_probe("127.0.0.1:1", Duration::from_millis(100)).await;
assert!(!result, "should fail for non-listening port");
}
#[tokio::test]
async fn h2_probe_succeeds_with_settings_response() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe = tokio::spawn(async move { h2_probe(&probe_addr, Duration::from_secs(2)).await });
let (mut socket, _peer) = listener.accept().await.unwrap();
let mut buf = [0_u8; 512];
let _ = socket.read(&mut buf).await.unwrap();
socket.write_all(H2_SETTINGS).await.unwrap();
socket.write_all(H2_SETTINGS_ACK).await.unwrap();
socket.shutdown().await.unwrap();
let result = probe.await.unwrap();
assert!(result, "should succeed when server responds with SETTINGS");
}
#[tokio::test]
async fn h2_probe_accepts_settings_header_split_across_reads() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe = tokio::spawn(async move { h2_probe(&probe_addr, Duration::from_secs(2)).await });
let (mut socket, _peer) = listener.accept().await.unwrap();
let mut buf = [0_u8; 512];
let _ = socket.read(&mut buf).await.unwrap();
socket.write_all(&H2_SETTINGS[..4]).await.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
socket.write_all(&H2_SETTINGS[4..]).await.unwrap();
socket.shutdown().await.unwrap();
let result = probe.await.unwrap();
assert!(
result,
"a SETTINGS header delivered in two segments is still a valid handshake"
);
}
#[tokio::test]
async fn h2_probe_succeeds_when_server_keeps_connection_open() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe = tokio::spawn(async move { h2_probe(&probe_addr, Duration::from_secs(2)).await });
let (mut socket, _peer) = listener.accept().await.unwrap();
let mut buf = [0_u8; 512];
let _ = socket.read(&mut buf).await.unwrap();
socket.write_all(H2_SETTINGS).await.unwrap();
let result = probe.await.unwrap();
assert!(
result,
"a server that answers SETTINGS but never closes is still healthy"
);
drop(socket);
}
#[tokio::test]
async fn h2_probe_fails_with_non_settings_response() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe = tokio::spawn(async move { h2_probe(&probe_addr, Duration::from_secs(2)).await });
let (mut socket, _peer) = listener.accept().await.unwrap();
let mut buf = [0_u8; 512];
let _ = socket.read(&mut buf).await.unwrap();
socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n")
.await
.unwrap();
socket.shutdown().await.unwrap();
let result = probe.await.unwrap();
assert!(
!result,
"should fail when server responds with HTTP/1.1 instead of SETTINGS"
);
}
#[tokio::test]
async fn h2_probe_times_out_on_no_response() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe = tokio::spawn(async move { h2_probe(&probe_addr, Duration::from_millis(100)).await });
let (_socket, _peer) = listener.accept().await.unwrap();
let result = probe.await.unwrap();
assert!(!result, "should fail when server does not respond within timeout");
}
#[tokio::test]
async fn http_probe_fails_with_wrong_status() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe = tokio::spawn(async move { http_probe(&probe_addr, "/", 200, Duration::from_secs(1)).await });
let (mut socket, _peer) = listener.accept().await.unwrap();
let mut buf = [0_u8; 512];
let _ = socket.read(&mut buf).await.unwrap();
socket
.write_all(b"HTTP/1.1 503 Service Unavailable\r\nConnection: close\r\n\r\n")
.await
.unwrap();
socket.shutdown().await.unwrap();
let result = probe.await.unwrap();
assert!(!result, "should fail when status code does not match");
}
#[tokio::test]
async fn http_probe_times_out_on_silent_backend() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe =
tokio::spawn(async move { http_probe(&probe_addr, "/healthz", 200, Duration::from_millis(200)).await });
let (_socket, _peer) = listener.accept().await.unwrap();
tokio::time::sleep(Duration::from_millis(400)).await;
let result = probe.await.unwrap();
assert!(!result, "a silent backend must fail the probe via timeout");
}
#[tokio::test]
async fn http_probe_fails_on_empty_response() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let probe_addr = addr.clone();
let probe = tokio::spawn(async move { http_probe(&probe_addr, "/healthz", 200, Duration::from_secs(2)).await });
let (mut socket, _peer) = listener.accept().await.unwrap();
let mut buf = [0_u8; 512];
let _bytes_read = socket.read(&mut buf).await.unwrap();
socket.shutdown().await.unwrap();
let result = probe.await.unwrap();
assert!(!result, "an empty response must fail the probe");
drop(socket);
}
}