use std::net::IpAddr;
#[cfg(test)]
use std::sync::atomic::{AtomicUsize, Ordering};
#[cfg(test)]
use std::sync::Arc;
use std::time::{Duration, Instant};
use futures_util::StreamExt;
use gregg_protocol::v2::{StatusPayloadV2, SCHEMA_VERSION_V2};
use gregg_protocol::{StatusSnapshot, SCHEMA_VERSION_V1};
use crate::clock::Clock;
use crate::endpoint::Endpoint;
#[derive(Clone, Copy)]
enum ExpectedSchema {
V1,
V2,
}
#[cfg(test)]
#[derive(Clone)]
pub struct PollActivityObserver {
active: Arc<AtomicUsize>,
peak: Arc<AtomicUsize>,
}
#[cfg(test)]
impl PollActivityObserver {
#[must_use]
pub fn new() -> Self {
Self {
active: Arc::new(AtomicUsize::new(0)),
peak: Arc::new(AtomicUsize::new(0)),
}
}
#[must_use]
pub fn guard(&self) -> PollActivityGuard<'_> {
let prev = self.active.fetch_add(1, Ordering::SeqCst);
let new_active = prev + 1;
self.peak.fetch_max(new_active, Ordering::SeqCst);
PollActivityGuard { observer: self }
}
#[must_use]
pub fn peak(&self) -> usize {
self.peak.load(Ordering::Relaxed)
}
}
#[cfg(test)]
pub struct PollActivityGuard<'a> {
observer: &'a PollActivityObserver,
}
#[cfg(test)]
impl Drop for PollActivityGuard<'_> {
fn drop(&mut self) {
self.observer.active.fetch_sub(1, Ordering::SeqCst);
}
}
const MAX_RESPONSE_BYTES: usize = 64 * 1024;
#[derive(Debug)]
pub struct PollResult {
pub system_id: String,
#[allow(dead_code)] pub endpoint: Endpoint,
pub outcome: PollOutcome,
pub latency: Duration,
}
#[derive(Debug, Clone, PartialEq)]
pub enum PollOutcome {
Online(Box<StatusSnapshot>),
OnlineV2(Box<StatusPayloadV2>),
Timeout,
ConnectionRefused,
DnsFailure,
NetworkError,
HttpStatus(u16),
BodyTooLarge,
DecodeError,
UnsupportedSchema,
InvalidSnapshot,
Cancelled,
}
#[derive(Debug)]
pub struct PollBatch {
pub generation: u64,
#[allow(dead_code)] pub started_at: Instant,
pub completed_at: Instant,
pub results: Vec<PollResult>,
}
#[derive(Clone)]
pub struct HttpClient {
client: reqwest::Client,
#[cfg(test)]
observer: Option<PollActivityObserver>,
}
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,
#[cfg(test)]
observer: None,
}
}
#[cfg(test)]
#[must_use]
pub fn new_with_observer(timeout: Duration, observer: PollActivityObserver) -> 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,
observer: Some(observer),
}
}
pub async fn poll(&self, endpoint: &Endpoint, clock: &impl Clock) -> PollResult {
#[cfg(test)]
let _guard = self.observer.as_ref().map(|o| o.guard());
let start = clock.now();
let v2_url = v2_status_url(&endpoint.host, endpoint.port);
let v2_result = self
.poll_single_url(&v2_url, endpoint, ExpectedSchema::V2, start)
.await;
if matches!(&v2_result.outcome, PollOutcome::HttpStatus(404)) {
let v1_url = status_url(&endpoint.host, endpoint.port);
return self
.poll_single_url(&v1_url, endpoint, ExpectedSchema::V1, start)
.await;
}
v2_result
}
fn make_result(
endpoint: &Endpoint,
outcome: PollOutcome,
start: std::time::Instant,
) -> PollResult {
PollResult {
system_id: endpoint.id.clone(),
endpoint: endpoint.clone(),
outcome,
latency: start.elapsed(),
}
}
async fn poll_single_url(
&self,
url: &str,
endpoint: &Endpoint,
expected_schema: ExpectedSchema,
start: std::time::Instant,
) -> PollResult {
let response = match self.client.get(url).send().await {
Ok(r) => r,
Err(e) => {
return Self::make_result(endpoint, classify_reqwest_error(&e), start);
}
};
let status = response.status().as_u16();
if !response.status().is_success() {
return Self::make_result(endpoint, PollOutcome::HttpStatus(status), start);
}
if let Some(content_length) = response.content_length() {
if content_length > MAX_RESPONSE_BYTES as u64 {
return Self::make_result(endpoint, PollOutcome::BodyTooLarge, start);
}
}
let body = match Self::read_body(response).await {
Ok(body) => body,
Err(outcome) => return Self::make_result(endpoint, outcome, start),
};
Self::parse_response(&body, endpoint, expected_schema, start)
}
async fn read_body(response: reqwest::Response) -> Result<Vec<u8>, PollOutcome> {
let mut stream = response.bytes_stream();
let mut body = Vec::new();
while let Some(chunk_result) = stream.next().await {
let c = chunk_result.map_err(|_| PollOutcome::NetworkError)?;
if body.len() + c.len() > MAX_RESPONSE_BYTES {
return Err(PollOutcome::BodyTooLarge);
}
body.extend_from_slice(&c);
}
Ok(body)
}
fn parse_response(
body: &[u8],
endpoint: &Endpoint,
expected_schema: ExpectedSchema,
start: std::time::Instant,
) -> PollResult {
if matches!(expected_schema, ExpectedSchema::V2) {
let Ok(payload) = serde_json::from_slice::<StatusPayloadV2>(body) else {
return Self::make_result(endpoint, PollOutcome::DecodeError, start);
};
if payload.snapshot.schema_version != SCHEMA_VERSION_V2 {
return Self::make_result(endpoint, PollOutcome::UnsupportedSchema, start);
}
if payload.validate().is_err() {
return Self::make_result(endpoint, PollOutcome::InvalidSnapshot, start);
}
return Self::make_result(endpoint, PollOutcome::OnlineV2(Box::new(payload)), start);
}
let Ok(snapshot): Result<StatusSnapshot, _> = serde_json::from_slice(body) else {
return Self::make_result(endpoint, PollOutcome::DecodeError, start);
};
if snapshot.schema_version != SCHEMA_VERSION_V1 {
return Self::make_result(endpoint, PollOutcome::UnsupportedSchema, start);
}
if snapshot.validate().is_err() {
return Self::make_result(endpoint, PollOutcome::InvalidSnapshot, start);
}
Self::make_result(endpoint, PollOutcome::Online(Box::new(snapshot)), start)
}
}
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")
}
}
#[must_use]
pub fn v2_status_url(host: &str, port: u16) -> String {
if host.parse::<IpAddr>().is_ok() && host.contains(':') {
format!("http://[{host}]:{port}/v2/status")
} else {
format!("http://{host}:{port}/v2/status")
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::endpoint::Endpoint;
use gregg_protocol::test_support::LinuxSnapshotBuilder;
use gregg_protocol::test_support::LinuxSnapshotV2Builder;
use gregg_protocol::v2::{DriveMetrics, MAX_DRIVE_ENTRIES, MAX_DRIVE_NAME_BYTES};
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 {
loop {
let Ok((mut stream, _)) = listener.accept().await else {
break;
};
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 request = String::from_utf8_lossy(&buf[..total]);
let response_status = if request
.lines()
.next()
.is_some_and(|line| line.contains("/v2/"))
&& !body.windows(9).any(|window| window == b"snapshot\":")
{
"404 Not Found"
} else {
&status
};
let header = format!(
"HTTP/1.1 {response_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 | PollOutcome::NetworkError
),
"expected ConnectionRefused or NetworkError, got {:?}",
result.outcome
);
}
#[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 oversized_body_chunked_delivery() {
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 first = vec![b'x'; 60 * 1024];
let second = vec![b'x'; 10 * 1024];
let header = "HTTP/1.1 200 OK\r\nContent-Length: 71680\r\n\r\n";
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(&first).await.unwrap();
stream.write_all(&second).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::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::DecodeError),
"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_v1_v2(
Some((b"[1, 2, 3]".to_vec(), "200 OK".to_string())),
(b"should not reach".to_vec(), "200 OK".to_string()),
)
.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_v1_v2(None, (body.clone().into_bytes(), "200 OK".to_string())).await;
let url2 = mock_server_v1_v2(None, (body.into_bytes(), "200 OK".to_string())).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 client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let mut outcomes = Vec::new();
for _ in 0..10 {
let url = mock_server(body.clone().into_bytes(), "200 OK").await;
let ep = endpoint_for(&url);
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:?}",
);
}
}
async fn mock_server_v1_v2(
v2_response: Option<(Vec<u8>, String)>,
v1_response: (Vec<u8>, String),
) -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let Ok((mut stream, _)) = listener.accept().await else {
break;
};
let v2_resp = v2_response.clone();
let v1_resp = v1_response.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 request = String::from_utf8_lossy(&buf[..total]);
let is_v2 = request
.lines()
.next()
.is_some_and(|line| line.contains("/v2/"));
let (body, status) = if is_v2 {
match &v2_resp {
Some((body, status)) => (body.clone(), status.clone()),
None => (b"not found".to_vec(), "404 Not Found".to_string()),
}
} else {
(v1_resp.0.clone(), v1_resp.1.clone())
};
let header = format!(
"HTTP/1.1 {status}\r\nContent-Length: {}\r\n\r\n",
body.len()
);
let _ = stream.write_all(header.as_bytes()).await;
let _ = stream.write_all(&body).await;
});
}
});
format!("http://127.0.0.1:{}", addr.port())
}
#[tokio::test]
async fn v2_404_falls_back_to_v1() {
let v1_snap = LinuxSnapshotBuilder::default().build();
let v1_body = serde_json::to_string(&v1_snap).unwrap();
let url = mock_server_v1_v2(None, (v1_body.into_bytes(), "200 OK".to_string())).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(_)),
"v2 404 should fall back to v1 Online, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn v2_malformed_does_not_fall_back() {
let url = mock_server_v1_v2(
Some((b"not json".to_vec(), "200 OK".to_string())),
(b"should not reach".to_vec(), "200 OK".to_string()),
)
.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),
"malformed v2 should not fall back, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn invalid_v2_drives_do_not_fall_back_to_v1() {
let valid = LinuxSnapshotV2Builder::default()
.drives(Some(vec![DriveMetrics {
name: "/broken".into(),
used_bytes: 1,
total_bytes: 10,
available_bytes: None,
}]))
.build_payload();
let mut v2_json = serde_json::to_value(valid).unwrap();
v2_json["drives"][0]["used_bytes"] = serde_json::json!(11);
let v2_body = serde_json::to_vec(&v2_json).unwrap();
let v1_body = serde_json::to_vec(&LinuxSnapshotBuilder::default().build()).unwrap();
let url = mock_server_v1_v2(
Some((v2_body, "200 OK".to_string())),
(v1_body, "200 OK".to_string()),
)
.await;
let ep = endpoint_for(&url);
let client = HttpClient::new(Duration::from_secs(5));
let result = client.poll(&ep, &crate::clock::RealClock).await;
assert!(matches!(result.outcome, PollOutcome::InvalidSnapshot));
}
#[test]
fn maximum_valid_v2_drive_payload_fits_response_cap_with_margin() {
let drives = (0..MAX_DRIVE_ENTRIES)
.map(|index| DriveMetrics {
name: format!(
"/{index}{}",
"x".repeat(MAX_DRIVE_NAME_BYTES - index.to_string().len() - 1)
),
used_bytes: u64::MAX / 2,
total_bytes: u64::MAX,
available_bytes: None,
})
.collect();
let payload = LinuxSnapshotV2Builder::default()
.drives(Some(drives))
.build_payload();
let serialized = serde_json::to_vec(&payload).unwrap();
assert!(
serialized.len() <= MAX_RESPONSE_BYTES - 1024,
"maximum valid v2 payload is {} bytes, cap is {}",
serialized.len(),
MAX_RESPONSE_BYTES
);
}
#[tokio::test]
async fn v2_503_does_not_fall_back() {
let url = mock_server_v1_v2(
Some((b"not ready".to_vec(), "503 Service Unavailable".to_string())),
(b"should not reach".to_vec(), "200 OK".to_string()),
)
.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)),
"v2 503 should not fall back, got {:?}",
result.outcome
);
}
#[tokio::test]
async fn v2_success_does_not_call_v1() {
let v2_snap = gregg_protocol::test_support::LinuxSnapshotV2Builder::default().build();
let v2_body = serde_json::to_string(&v2_snap).unwrap();
let url = mock_server_v1_v2(
Some((v2_body.into_bytes(), "200 OK".to_string())),
(b"should not reach".to_vec(), "200 OK".to_string()),
)
.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::OnlineV2(_)),
"v2 success should return OnlineV2, got {:?}",
result.outcome
);
}
}