use super::*;
use tokio::time::{sleep, Duration};
pub struct FailoverManager {
config: FailoverConfig,
servers: Arc<RwLock<Vec<DhcpServer>>>,
current_server: Arc<RwLock<Option<usize>>>,
}
impl FailoverManager {
pub fn new(config: FailoverConfig, servers: Vec<DhcpServer>) -> Self {
Self {
config,
servers: Arc::new(RwLock::new(servers)),
current_server: Arc::new(RwLock::new(None)),
}
}
pub async fn add_server(&self, server: DhcpServer) {
let mut servers = self.servers.write().await;
if !servers.iter().any(|s| s.address == server.address) {
servers.push(server);
}
}
pub async fn get_best_server(&self) -> Option<DhcpServer> {
if !self.config.enabled {
return self.get_any_server().await;
}
let servers = self.servers.read().await;
if self.config.prefer_previous_server {
if let Some(index) = *self.current_server.read().await {
if index < servers.len() && servers[index].healthy {
return Some(servers[index].clone());
}
}
}
let best = servers.iter()
.filter(|s| s.healthy)
.min_by_key(|s| s.response_time.unwrap_or(Duration::MAX));
best.cloned()
}
async fn get_any_server(&self) -> Option<DhcpServer> {
let servers = self.servers.read().await;
servers.iter()
.find(|s| s.healthy)
.cloned()
}
pub async fn mark_success(&self, server_addr: &IpAddr, response_time: Duration) {
let mut servers = self.servers.write().await;
for (index, server) in servers.iter_mut().enumerate() {
if server.address == *server_addr {
server.mark_success(response_time);
*self.current_server.write().await = Some(index);
break;
}
}
}
pub async fn mark_failure(&self, server_addr: &IpAddr) {
let mut servers = self.servers.write().await;
for server in servers.iter_mut() {
if server.address == *server_addr {
server.mark_failure();
break;
}
}
}
pub async fn health_check(&self) {
if !self.config.enabled {
return;
}
debug!("Running DHCP server health check");
let servers = self.servers.read().await;
for server in servers.iter() {
debug!("Health check for server {}: {}", server.address, server.healthy);
}
}
pub fn start_health_checks(self: Arc<Self>) {
if !self.config.enabled {
return;
}
let interval = Duration::from_secs(self.config.health_check_interval);
tokio::spawn(async move {
loop {
sleep(interval).await;
self.health_check().await;
}
});
}
pub async fn get_servers(&self) -> Vec<DhcpServer> {
self.servers.read().await.clone()
}
pub async fn healthy_count(&self) -> usize {
self.servers.read().await.iter().filter(|s| s.healthy).count()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_failover_manager() {
let config = FailoverConfig::default();
let servers = vec![
DhcpServer::new(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))),
DhcpServer::new(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 2))),
];
let manager = FailoverManager::new(config, servers);
assert_eq!(manager.healthy_count().await, 2);
let best = manager.get_best_server().await;
assert!(best.is_some());
}
#[tokio::test]
async fn test_server_failure() {
let config = FailoverConfig::default();
let server_addr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
let servers = vec![DhcpServer::new(server_addr)];
let manager = FailoverManager::new(config, servers);
manager.mark_failure(&server_addr).await;
manager.mark_failure(&server_addr).await;
manager.mark_failure(&server_addr).await;
assert_eq!(manager.healthy_count().await, 0);
}
#[tokio::test]
async fn test_server_recovery() {
let config = FailoverConfig::default();
let server_addr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
let servers = vec![DhcpServer::new(server_addr)];
let manager = FailoverManager::new(config, servers);
manager.mark_failure(&server_addr).await;
manager.mark_failure(&server_addr).await;
manager.mark_failure(&server_addr).await;
manager.mark_success(&server_addr, Duration::from_millis(10)).await;
assert_eq!(manager.healthy_count().await, 1);
}
}