#![allow(dead_code)]
use std::net::IpAddr;
use std::time::{Duration, Instant};
use gregg_protocol::{StatusSnapshot, SCHEMA_VERSION_V1};
use crate::clock::Clock;
use crate::endpoint::Endpoint;
const MAX_RESPONSE_BYTES: usize = 64 * 1024;
#[derive(Debug)]
pub struct PollResult {
pub system_id: String,
pub endpoint: Endpoint,
pub outcome: PollOutcome,
pub latency: Duration,
}
#[derive(Debug, Clone, PartialEq)]
pub enum PollOutcome {
Online(Box<StatusSnapshot>),
Timeout,
ConnectionRefused,
DnsFailure,
NetworkError,
HttpStatus(u16),
BodyTooLarge,
DecodeError,
UnsupportedSchema,
InvalidSnapshot,
Cancelled,
}
#[derive(Debug)]
pub struct PollBatch {
pub generation: u64,
pub started_at: Instant,
pub completed_at: Instant,
pub results: Vec<PollResult>,
}
#[derive(Clone)]
pub struct HttpClient {
client: reqwest::Client,
}
impl HttpClient {
#[must_use]
pub fn new(timeout: Duration) -> Self {
let client = reqwest::Client::builder()
.timeout(timeout)
.redirect(reqwest::redirect::Policy::none())
.pool_max_idle_per_host(4)
.build()
.expect("reqwest client builder should not fail");
Self { client }
}
pub async fn poll(&self, endpoint: &Endpoint, clock: &impl Clock) -> PollResult {
let url = status_url(&endpoint.host, endpoint.port);
let start = clock.now();
let response = match self.client.get(&url).send().await {
Ok(r) => r,
Err(e) => {
let latency = start.elapsed();
let outcome = classify_reqwest_error(&e);
return PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome,
latency,
};
}
};
let status = response.status().as_u16();
if !response.status().is_success() {
let latency = start.elapsed();
return PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome: PollOutcome::HttpStatus(status),
latency,
};
}
let Ok(body) = response.bytes().await else {
let latency = start.elapsed();
return PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome: PollOutcome::NetworkError,
latency,
};
};
if body.len() > MAX_RESPONSE_BYTES {
let latency = start.elapsed();
return PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome: PollOutcome::BodyTooLarge,
latency,
};
}
let Ok(snapshot): Result<StatusSnapshot, _> = serde_json::from_slice(&body) else {
let latency = start.elapsed();
return PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome: PollOutcome::DecodeError,
latency,
};
};
if snapshot.schema_version != SCHEMA_VERSION_V1 {
let latency = start.elapsed();
return PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome: PollOutcome::UnsupportedSchema,
latency,
};
}
if snapshot.validate().is_err() {
let latency = start.elapsed();
return PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome: PollOutcome::InvalidSnapshot,
latency,
};
}
let latency = start.elapsed();
PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome: PollOutcome::Online(Box::new(snapshot)),
latency,
}
}
}
fn classify_reqwest_error(e: &reqwest::Error) -> PollOutcome {
if e.is_timeout() {
return PollOutcome::Timeout;
}
if is_connection_refused(&e) {
return PollOutcome::ConnectionRefused;
}
if is_dns_failure(&e) {
return PollOutcome::DnsFailure;
}
PollOutcome::NetworkError
}
fn is_connection_refused(e: &dyn std::error::Error) -> bool {
let msg = format!("{e}");
if msg.contains("connection refused") || msg.contains("Connection refused") {
return true;
}
let mut source: Option<&(dyn std::error::Error + 'static)> = e.source();
while let Some(err) = source {
let msg = format!("{err}");
if msg.contains("connection refused") || msg.contains("Connection refused") {
return true;
}
source = err.source();
}
false
}
fn is_dns_failure(e: &dyn std::error::Error) -> bool {
let msg = format!("{e}");
if msg.contains("dns") || msg.contains("resolve") || msg.contains("name") {
return true;
}
let mut source: Option<&(dyn std::error::Error + 'static)> = e.source();
while let Some(err) = source {
let msg = format!("{err}");
if msg.contains("dns") || msg.contains("resolve") || msg.contains("name") {
return true;
}
source = err.source();
}
false
}
#[must_use]
pub fn status_url(host: &str, port: u16) -> String {
if host.parse::<IpAddr>().is_ok() && host.contains(':') {
format!("http://[{host}]:{port}/v1/status")
} else {
format!("http://{host}:{port}/v1/status")
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::endpoint::Endpoint;
use gregg_protocol::test_support::LinuxSnapshotBuilder;
use std::io;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
async fn mock_server(body: Vec<u8>, status: &str) -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let status = status.to_string();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let header = format!(
"HTTP/1.1 {status}\r\nContent-Length: {}\r\n\r\n",
body.len()
);
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(&body).await.unwrap();
});
format!("http://127.0.0.1:{}", addr.port())
}
async fn mock_server_drop() -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
drop(stream);
});
format!("http://127.0.0.1:{}", addr.port())
}
async fn mock_server_closed_port() -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
drop(listener); format!("http://127.0.0.1:{}", addr.port())
}
fn endpoint_for(url: &str) -> Endpoint {
let stripped = url.strip_prefix("http://").unwrap();
let (host, port_str) = stripped.rsplit_once(':').unwrap();
let host = host
.strip_prefix('[')
.unwrap_or(host)
.strip_suffix(']')
.unwrap_or(host);
Endpoint {
id: "test-id".into(),
host: host.to_string(),
port: port_str.parse().unwrap(),
name: None,
}
}
fn valid_snapshot_json() -> String {
let snap = LinuxSnapshotBuilder::default().build();
serde_json::to_string(&snap).unwrap()
}
#[tokio::test]
async fn successful_poll_returns_online() {
let body = valid_snapshot_json();
let url = mock_server(body.into_bytes(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert_eq!(result.system_id, "test-id");
assert!(matches!(result.outcome, PollOutcome::Online(_)));
assert!(result.latency < Duration::from_secs(5));
}
#[tokio::test]
async fn successful_poll_with_macos_snapshot() {
let snap = gregg_protocol::test_support::MacosSnapshotBuilder::default().build();
let body = serde_json::to_string(&snap).unwrap();
let url = mock_server(body.into_bytes(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::Online(_)));
}
#[tokio::test]
async fn timeout_handling() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
tokio::time::sleep(Duration::from_secs(10)).await;
let header = "HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n{}";
let _ = stream.write_all(header.as_bytes()).await;
});
let ep = Endpoint {
id: "test-id".into(),
host: "127.0.0.1".into(),
port: addr.port(),
name: None,
};
let client = HttpClient::new(Duration::from_millis(50));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::Timeout));
}
#[tokio::test]
async fn connection_refused() {
let url = mock_server_closed_port().await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::ConnectionRefused));
}
#[tokio::test]
async fn non_2xx_status() {
let url = mock_server(b"not ready".to_vec(), "503 Service Unavailable").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::HttpStatus(503)));
}
#[tokio::test]
async fn oversized_body() {
let body = vec![b'x'; 65 * 1024];
let url = mock_server(body, "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::BodyTooLarge));
}
#[tokio::test]
async fn malformed_json() {
let url = mock_server(b"not json at all".to_vec(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::DecodeError));
}
#[tokio::test]
async fn unsupported_schema_version() {
let snap = LinuxSnapshotBuilder::default().build();
let mut json = serde_json::to_value(&snap).unwrap();
json["schema_version"] = serde_json::json!(99);
let body = serde_json::to_string(&json).unwrap();
let url = mock_server(body.into_bytes(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::UnsupportedSchema));
}
#[tokio::test]
async fn invalid_snapshot_validation_failure() {
let snap = LinuxSnapshotBuilder::default().build();
let mut json = serde_json::to_value(&snap).unwrap();
json["memory"]["used_bytes"] = serde_json::json!(999_999_999_999_i64);
json["memory"]["total_bytes"] = serde_json::json!(1);
let body = serde_json::to_string(&json).unwrap();
let url = mock_server(body.into_bytes(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::InvalidSnapshot));
}
#[tokio::test]
async fn url_construction_ipv4() {
let url = status_url("192.168.1.1", 11310);
assert_eq!(url, "http://192.168.1.1:11310/v1/status");
}
#[tokio::test]
async fn url_construction_ipv6() {
let url = status_url("::1", 8080);
assert_eq!(url, "http://[::1]:8080/v1/status");
}
#[tokio::test]
async fn url_construction_dns() {
let url = status_url("server.local", 11310);
assert_eq!(url, "http://server.local:11310/v1/status");
}
#[tokio::test]
async fn network_error_on_dropped_connection() {
let url = mock_server_drop().await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(
result.outcome,
PollOutcome::NetworkError | PollOutcome::DecodeError
),
"expected NetworkError or DecodeError, got {:?}",
result.outcome
);
}
#[test]
fn max_response_bytes_is_64k() {
assert_eq!(MAX_RESPONSE_BYTES, 64 * 1024);
}
#[test]
fn classify_timeout() {
let _ = classify_reqwest_error;
}
#[test]
fn is_connection_refused_returns_false_for_non_refused() {
let err = io::Error::other("some error");
assert!(!is_connection_refused(&err));
}
#[test]
fn is_connection_refused_returns_true_for_refused() {
let err = io::Error::new(io::ErrorKind::ConnectionRefused, "connection refused");
assert!(is_connection_refused(&err));
}
#[tokio::test]
async fn redirect_response_301() {
let url = mock_server(b"redirect".to_vec(), "301 Moved Permanently").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(matches!(result.outcome, PollOutcome::HttpStatus(301)));
}
#[tokio::test]
async fn partial_body_then_close() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let header = "HTTP/1.1 200 OK\r\nContent-Length: 1024\r\n\r\npartial";
let _ = stream.write_all(header.as_bytes()).await;
drop(stream);
});
let ep = Endpoint {
id: "test-id".into(),
host: "127.0.0.1".into(),
port: addr.port(),
name: None,
};
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(
result.outcome,
PollOutcome::NetworkError | PollOutcome::DecodeError
),
"expected NetworkError or DecodeError, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn empty_body_with_200() {
let url = mock_server(Vec::new(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(result.outcome, PollOutcome::DecodeError),
"expected DecodeError, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn wrong_content_type_with_valid_json() {
let body = valid_snapshot_json();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let header = format!(
"HTTP/1.1 200 OK\r\nContent-Type: text/html\r\nContent-Length: {}\r\n\r\n",
body.len()
);
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(body.as_bytes()).await.unwrap();
});
let ep = Endpoint {
id: "test-id".into(),
host: "127.0.0.1".into(),
port: addr.port(),
name: None,
};
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(result.outcome, PollOutcome::Online(_)),
"expected Online, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn large_valid_json_under_64k() {
let long_name = "x".repeat(60_000);
let snap = LinuxSnapshotBuilder::default().build();
let mut json = serde_json::to_value(&snap).unwrap();
json["system"]["name"] = serde_json::json!(long_name);
let body = serde_json::to_string(&json).unwrap();
assert!(body.len() < 64 * 1024);
let url = mock_server(body.into_bytes(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(result.outcome, PollOutcome::Online(_)),
"expected Online, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn unicode_in_system_name() {
let snap = LinuxSnapshotBuilder::default().build();
let mut json = serde_json::to_value(&snap).unwrap();
json["system"]["name"] = serde_json::json!("日本語サーバー");
let body = serde_json::to_string(&json).unwrap();
let url = mock_server(body.into_bytes(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(result.outcome, PollOutcome::Online(_)),
"expected Online, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn nested_invalid_json() {
let url = mock_server(b"{\"nested\": {\"invalid\": true}}".to_vec(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(result.outcome, PollOutcome::DecodeError),
"expected DecodeError, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn array_instead_of_object() {
let url = mock_server(b"[1, 2, 3]".to_vec(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(result.outcome, PollOutcome::DecodeError),
"expected DecodeError, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn null_json() {
let url = mock_server(b"null".to_vec(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(result.outcome, PollOutcome::DecodeError),
"expected DecodeError, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn stale_observation_timestamp_from_endpoint() {
let snap = LinuxSnapshotBuilder::default().build();
let mut json = serde_json::to_value(&snap).unwrap();
json["observation_timestamp_ms"] = serde_json::json!(0);
let body = serde_json::to_string(&json).unwrap();
let url = mock_server(body.into_bytes(), "200 OK").await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result = client.poll(&ep, &clock).await;
assert!(
matches!(result.outcome, PollOutcome::Online(_)),
"expected Online even with stale timestamp, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn config_change_between_polls() {
let body = valid_snapshot_json();
let url1 = mock_server(body.clone().into_bytes(), "200 OK").await;
let url2 = mock_server(body.into_bytes(), "200 OK").await;
let ep1 = endpoint_for(&url1);
let ep2 = endpoint_for(&url2);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result1 = client.poll(&ep1, &clock).await;
assert!(matches!(result1.outcome, PollOutcome::Online(_)));
assert_eq!(result1.system_id, "test-id");
let mut ep2 = ep2;
ep2.id = "new-system-id".into();
let result2 = client.poll(&ep2, &clock).await;
assert!(matches!(result2.outcome, PollOutcome::Online(_)));
assert_eq!(result2.system_id, "new-system-id");
}
#[tokio::test]
async fn cancel_during_poll() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
tokio::time::sleep(Duration::from_secs(10)).await;
let snap = LinuxSnapshotBuilder::default().build();
let body = serde_json::to_string(&snap).unwrap();
let header = format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len());
let _ = stream.write_all(header.as_bytes()).await;
let _ = stream.write_all(body.as_bytes()).await;
});
let ep = Endpoint {
id: "test-id".into(),
host: "127.0.0.1".into(),
port: addr.port(),
name: None,
};
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let result =
tokio::time::timeout(Duration::from_millis(100), client.poll(&ep, &clock)).await;
assert!(result.is_ok() || result.is_err());
}
#[tokio::test]
async fn multiple_rapid_polls_same_result() {
let body = valid_snapshot_json();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let status_line = "200 OK".to_string();
tokio::spawn(async move {
for _ in 0..10 {
let (mut stream, _) = listener.accept().await.unwrap();
let body = body.clone();
let status_line = status_line.clone();
tokio::spawn(async move {
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let header = format!(
"HTTP/1.1 {status_line}\r\nContent-Length: {}\r\n\r\n",
body.len()
);
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(body.as_bytes()).await.unwrap();
});
}
});
let ep = Endpoint {
id: "test-id".into(),
host: "127.0.0.1".into(),
port: addr.port(),
name: None,
};
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let mut outcomes = Vec::new();
for _ in 0..10 {
let result = client.poll(&ep, &clock).await;
outcomes.push(result.outcome.clone());
}
for outcome in &outcomes {
assert!(
matches!(outcome, PollOutcome::Online(_)),
"expected Online for all polls, got {outcome:?}",
);
}
}
}