connectivity improvements

This commit is contained in:
Graham McIntire 2026-01-23 10:33:13 -06:00
parent ee7ce7fe00
commit 9cb924fd77
No known key found for this signature in database
2 changed files with 66 additions and 9 deletions

View file

@ -1,4 +1,5 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc; use std::sync::Arc;
use std::time::Instant; use std::time::Instant;
use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::io::{AsyncReadExt, AsyncWriteExt};
@ -11,10 +12,18 @@ pub struct HealthStatus {
pub status: String, pub status: String,
pub version: String, pub version: String,
pub uptime_seconds: u64, pub uptime_seconds: u64,
pub connected: bool,
} }
/// Start a simple health check HTTP server using raw TCP /// Start a simple health check HTTP server using raw TCP.
pub async fn start_health_server(port: u16) -> Result<()> { ///
/// The health check returns:
/// - 200 OK with status "healthy" when WebSocket is connected
/// - 503 Service Unavailable with status "unhealthy" when WebSocket is disconnected
///
/// This allows Kubernetes readiness probes to remove the pod from service
/// when the WebSocket connection is broken.
pub async fn start_health_server(port: u16, connected: Arc<AtomicBool>) -> Result<()> {
let start_time = Arc::new(Instant::now()); let start_time = Arc::new(Instant::now());
let addr = format!("0.0.0.0:{}", port); let addr = format!("0.0.0.0:{}", port);
@ -25,6 +34,7 @@ pub async fn start_health_server(port: u16) -> Result<()> {
match listener.accept().await { match listener.accept().await {
Ok((mut socket, _)) => { Ok((mut socket, _)) => {
let start_time = Arc::clone(&start_time); let start_time = Arc::clone(&start_time);
let connected = Arc::clone(&connected);
tokio::spawn(async move { tokio::spawn(async move {
let mut buffer = [0u8; 1024]; let mut buffer = [0u8; 1024];
@ -36,11 +46,18 @@ pub async fn start_health_server(port: u16) -> Result<()> {
// Check if this is a GET request to /health // Check if this is a GET request to /health
if request.starts_with("GET /health") { if request.starts_with("GET /health") {
let is_connected = connected.load(Ordering::Relaxed);
let uptime = start_time.elapsed().as_secs(); let uptime = start_time.elapsed().as_secs();
let status = HealthStatus { let status = HealthStatus {
status: "healthy".to_string(), status: if is_connected {
"healthy".to_string()
} else {
"unhealthy".to_string()
},
version: env!("CARGO_PKG_VERSION").to_string(), version: env!("CARGO_PKG_VERSION").to_string(),
uptime_seconds: uptime, uptime_seconds: uptime,
connected: is_connected,
}; };
let json = serde_json::to_string(&status).unwrap_or_else(|_| { let json = serde_json::to_string(&status).unwrap_or_else(|_| {
@ -48,8 +65,16 @@ pub async fn start_health_server(port: u16) -> Result<()> {
.to_string() .to_string()
}); });
// Return 200 if connected, 503 if not
let http_status = if is_connected {
"200 OK"
} else {
"503 Service Unavailable"
};
let response = format!( let response = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{}", "HTTP/1.1 {}\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{}",
http_status,
json.len(), json.len(),
json json
); );
@ -79,26 +104,43 @@ mod tests {
use super::*; use super::*;
#[test] #[test]
fn test_health_status_serialization() { fn test_health_status_serialization_healthy() {
let status = HealthStatus { let status = HealthStatus {
status: "healthy".to_string(), status: "healthy".to_string(),
version: "0.1.0".to_string(), version: "0.1.0".to_string(),
uptime_seconds: 42, uptime_seconds: 42,
connected: true,
}; };
let json = serde_json::to_string(&status).unwrap(); let json = serde_json::to_string(&status).unwrap();
assert!(json.contains(r#""status":"healthy""#)); assert!(json.contains(r#""status":"healthy""#));
assert!(json.contains(r#""version":"0.1.0""#)); assert!(json.contains(r#""version":"0.1.0""#));
assert!(json.contains(r#""uptime_seconds":42"#)); assert!(json.contains(r#""uptime_seconds":42"#));
assert!(json.contains(r#""connected":true"#));
}
#[test]
fn test_health_status_serialization_unhealthy() {
let status = HealthStatus {
status: "unhealthy".to_string(),
version: "0.1.0".to_string(),
uptime_seconds: 100,
connected: false,
};
let json = serde_json::to_string(&status).unwrap();
assert!(json.contains(r#""status":"unhealthy""#));
assert!(json.contains(r#""connected":false"#));
} }
#[test] #[test]
fn test_health_status_deserialization() { fn test_health_status_deserialization() {
let json = r#"{"status":"healthy","version":"0.1.0","uptime_seconds":42}"#; let json = r#"{"status":"healthy","version":"0.1.0","uptime_seconds":42,"connected":true}"#;
let status: HealthStatus = serde_json::from_str(json).unwrap(); let status: HealthStatus = serde_json::from_str(json).unwrap();
assert_eq!(status.status, "healthy"); assert_eq!(status.status, "healthy");
assert_eq!(status.version, "0.1.0"); assert_eq!(status.version, "0.1.0");
assert_eq!(status.uptime_seconds, 42); assert_eq!(status.uptime_seconds, 42);
assert!(status.connected);
} }
#[test] #[test]
@ -107,11 +149,13 @@ mod tests {
status: "healthy".to_string(), status: "healthy".to_string(),
version: "0.1.0".to_string(), version: "0.1.0".to_string(),
uptime_seconds: 42, uptime_seconds: 42,
connected: true,
}; };
let cloned = status.clone(); let cloned = status.clone();
assert_eq!(status.status, cloned.status); assert_eq!(status.status, cloned.status);
assert_eq!(status.version, cloned.version); assert_eq!(status.version, cloned.version);
assert_eq!(status.uptime_seconds, cloned.uptime_seconds); assert_eq!(status.uptime_seconds, cloned.uptime_seconds);
assert_eq!(status.connected, cloned.connected);
} }
// Note: start_health_server is tested manually/via integration tests // Note: start_health_server is tested manually/via integration tests

View file

@ -7,6 +7,8 @@ mod websocket_client;
use clap::Parser; use clap::Parser;
use std::env; use std::env;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH}; use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tokio::time::sleep; use tokio::time::sleep;
use websocket_client::AgentClient; use websocket_client::AgentClient;
@ -194,9 +196,14 @@ async fn main() {
log_info!("WebSocket URL: {}", ws_url); log_info!("WebSocket URL: {}", ws_url);
// Start simple health endpoint (no storage needed for WebSocket mode) // Shared connection state for health check
tokio::spawn(async { // Starts as false (not connected), updated when WebSocket connects/disconnects
if let Err(e) = health::start_health_server(8080).await { let connected = Arc::new(AtomicBool::new(false));
let connected_for_health = Arc::clone(&connected);
// Start simple health endpoint with connection state
tokio::spawn(async move {
if let Err(e) = health::start_health_server(8080, connected_for_health).await {
log_warn!("Health server error: {}", e); log_warn!("Health server error: {}", e);
} }
}); });
@ -225,6 +232,8 @@ async fn main() {
let mut client = match AgentClient::connect(&ws_url, &args.token).await { let mut client = match AgentClient::connect(&ws_url, &args.token).await {
Ok(client) => { Ok(client) => {
log_info!("Successfully connected to server"); log_info!("Successfully connected to server");
// Mark as connected for health check
connected.store(true, Ordering::Relaxed);
// Reset retry delay on successful connection // Reset retry delay on successful connection
retry_delay = Duration::from_secs(1); retry_delay = Duration::from_secs(1);
attempt = 0; attempt = 0;
@ -232,6 +241,8 @@ async fn main() {
} }
Err(e) => { Err(e) => {
log_error!("Failed to connect to server: {}", e); log_error!("Failed to connect to server: {}", e);
// Mark as disconnected for health check
connected.store(false, Ordering::Relaxed);
continue; continue;
} }
}; };
@ -239,6 +250,8 @@ async fn main() {
// Run the agent event loop // Run the agent event loop
if let Err(e) = client.run().await { if let Err(e) = client.run().await {
log_error!("Agent disconnected: {}", e); log_error!("Agent disconnected: {}", e);
// Mark as disconnected for health check
connected.store(false, Ordering::Relaxed);
// Loop will retry with backoff // Loop will retry with backoff
} }
} }