diff --git a/src/algorithms/least_connection.rs b/src/algorithms/least_connection.rs index 5b9e53b..d0716c9 100644 --- a/src/algorithms/least_connection.rs +++ b/src/algorithms/least_connection.rs @@ -1,10 +1,12 @@ -use crate::{error::Error, middleware::Server}; +use std::collections::HashMap; -pub async fn least_connection(available_servers: &[Server]) -> Result { - let server = available_servers - .iter() - .min_by_key(|server| server.load()) +use crate::error::Error; + +pub async fn least_connection(server_loads: HashMap) -> Result { + let (url, _) = server_loads + .into_iter() + .min_by_key(|(_, load)| *load) .ok_or_else(|| Error::NoServerAvailable)?; - Ok(server.clone()) + Ok(url) } diff --git a/src/algorithms/mod.rs b/src/algorithms/mod.rs index e0d7ae0..dee1207 100644 --- a/src/algorithms/mod.rs +++ b/src/algorithms/mod.rs @@ -1,4 +1,6 @@ -use crate::{error::Error, middleware::Server}; +use reqwest::Url; + +use crate::{db::RedisClient, error::Error, middleware::ServerClient}; mod least_connection; mod resource_based; @@ -27,18 +29,34 @@ impl From for Algorithm { } impl Algorithm { - pub async fn select_server(&self, available_servers: &[Server]) -> Result { - match self { - Algorithm::LeastConnection => { - least_connection::least_connection(available_servers).await - } - Algorithm::ResourceBased => resource_based::resource_based(available_servers).await, + pub async fn select_server( + &self, + mut redis_client: RedisClient, + ) -> Result { + let server_loads = redis_client.get_all_server_load().await?; + let weights = redis_client.get_all_server_weights().await?; + let url = match self { + Algorithm::LeastConnection => least_connection::least_connection(server_loads).await, + Algorithm::ResourceBased => unimplemented!(), Algorithm::WeightedLeastConnection => { - weighted_least_connection::weighted_least_connection(available_servers).await + weighted_least_connection::weighted_least_connection(server_loads, weights).await } Algorithm::WeightedResponseTime => { - weighted_response_time::weighted_response_time(available_servers).await + weighted_response_time::weighted_response_time( + redis_client.get_all_server_mean_latency().await?, + weights, + ) + .await } - } + }?; + + redis_client.update_server_load(&url, 1).await?; + + let url = url.parse::().map_err(|e| Error::Other(e.into()))?; + + Ok(ServerClient { + url, + client: reqwest::Client::new(), + }) } } diff --git a/src/algorithms/resource_based.rs b/src/algorithms/resource_based.rs index 64e47ac..c58b722 100644 --- a/src/algorithms/resource_based.rs +++ b/src/algorithms/resource_based.rs @@ -1,6 +1,8 @@ -use crate::{error::Error, middleware::Server}; +use crate::{error::Error, middleware::StaticServerData}; // TODO: Implement resource-based load balancing algorithm -pub async fn resource_based(available_servers: &[Server]) -> Result { +pub async fn _resource_based( + available_servers: &[StaticServerData], +) -> Result { Ok(available_servers[0].clone()) } diff --git a/src/algorithms/weighted_least_connection.rs b/src/algorithms/weighted_least_connection.rs index 84f2f87..43416d2 100644 --- a/src/algorithms/weighted_least_connection.rs +++ b/src/algorithms/weighted_least_connection.rs @@ -1,10 +1,18 @@ -use crate::{error::Error, middleware::Server}; +use std::collections::HashMap; -pub async fn weighted_least_connection(available_servers: &[Server]) -> Result { - let server = available_servers - .iter() - .min_by_key(|server| server.load() / server.weight()) +use crate::error::Error; + +pub async fn weighted_least_connection( + server_loads: HashMap, + weights: HashMap, +) -> Result { + let (url, _) = server_loads + .into_iter() + .min_by_key(|(key, load)| { + let weight = weights.get(key).unwrap_or(&1); + *load / weight + }) .ok_or_else(|| Error::NoServerAvailable)?; - Ok(server.clone()) + Ok(url) } diff --git a/src/algorithms/weighted_response_time.rs b/src/algorithms/weighted_response_time.rs index 0d289b2..84725a6 100644 --- a/src/algorithms/weighted_response_time.rs +++ b/src/algorithms/weighted_response_time.rs @@ -1,10 +1,18 @@ -use crate::{error::Error, middleware::Server}; +use std::collections::HashMap; -pub async fn weighted_response_time(available_servers: &[Server]) -> Result { - let server = available_servers - .iter() - .min_by_key(|server| server.mean_latency()) +use crate::error::Error; + +pub async fn weighted_response_time( + latencies: HashMap, + weights: HashMap, +) -> Result { + let (url, _) = latencies + .into_iter() + .min_by_key(|(key, latency)| { + let weight = weights.get(key).unwrap_or(&1); + *latency / weight + }) .ok_or_else(|| Error::NoServerAvailable)?; - Ok(server.clone()) + Ok(url) } diff --git a/src/app.rs b/src/app.rs new file mode 100644 index 0000000..6376d08 --- /dev/null +++ b/src/app.rs @@ -0,0 +1,75 @@ +use axum::{Router, routing::get}; +use tokio::{net::TcpListener, task::JoinHandle}; +use tower_http::{ + cors::{AllowHeaders, AllowMethods, AllowOrigin, CorsLayer}, + trace::TraceLayer, +}; + +use crate::config::State; +use crate::middleware::request_route; +use crate::route::health::status; +use crate::services::{latency_tracker_worker, server_status_worker}; + +type JoinHandleWrapper = JoinHandle>; + +/// Application struct to hold main and background worker tasks +pub struct App { + main: JoinHandleWrapper, + server_status_background_worker: JoinHandleWrapper, + latency_tracker_background_worker: JoinHandleWrapper, +} + +impl App { + pub async fn setup( + state: State, + listener: TcpListener, + ) -> Result> { + let server = Router::new() + .route("/status", get(status)) + .layer( + CorsLayer::new() + .allow_headers(AllowHeaders::any()) + .allow_origin(AllowOrigin::any()) + .allow_methods(AllowMethods::any()), + ) + .layer(TraceLayer::new_for_http()) + // add rate limitter middleware mechanism + .layer(axum::middleware::from_fn_with_state( + state.clone(), + request_route, + )) + .with_state(state.clone()); + + let main = tokio::spawn(async move { axum::serve(listener, server).await }); + + let redis_conn_1 = state.redis_conn.clone(); + let redis_conn_2 = state.redis_conn.clone(); + + let server_status_background_worker = tokio::spawn(async move { + let _: () = server_status_worker(redis_conn_1).await; + Ok(()) + }); + + let latency_tracker_background_worker = tokio::spawn(async move { + let _: () = latency_tracker_worker(redis_conn_2).await; + Ok(()) + }); + + Ok(Self { + main, + server_status_background_worker, + latency_tracker_background_worker, + }) + } + + pub async fn start(self) -> Result<(), Box> { + match tokio::try_join!( + self.main, + self.server_status_background_worker, + self.latency_tracker_background_worker + ) { + Ok(_) => Ok(()), + Err(err) => Err(err)?, + } + } +} diff --git a/src/config.rs b/src/config.rs index 8c7dbcf..23cca22 100644 --- a/src/config.rs +++ b/src/config.rs @@ -3,12 +3,12 @@ use serde::Deserialize; use crate::{ algorithms::Algorithm, db::{self, RedisClient}, - middleware::{Server, ServerClients}, + middleware::StaticServerData, }; #[derive(Deserialize)] pub struct SystemConfig { - pub available_servers: String, // TODO: This should be hosted in redis + pub available_servers: String, pub port: u16, pub redis_url: String, pub algorithm: String, @@ -20,13 +20,12 @@ impl SystemConfig { dotenvy::dotenv_override().ok(); envy::from_env::() - .map_err(|e| anyhow::anyhow!("Failed to load environment variables: {}", e)) + .map_err(|e| anyhow::anyhow!("Failed to load environment variable(s): {}", e)) } } #[derive(Clone)] pub struct State { - pub available_servers: ServerClients, pub redis_conn: RedisClient, pub algorithm: Algorithm, } @@ -35,17 +34,16 @@ impl State { pub async fn new(config: &SystemConfig) -> Result> { let servers = config.available_servers.split(',').collect::>(); - let available_servers: Vec = servers + let available_servers: Vec = servers .clone() .into_iter() - .map(Server::new) - .collect::, _>>()?; + .map(StaticServerData::new) + .collect::, _>>()?; let redis_conn = db::RedisClient::init_redis(&config.redis_url, available_servers.clone()).await?; Ok(State { - available_servers: ServerClients::new(available_servers), redis_conn, algorithm: config.algorithm.clone().into(), }) diff --git a/src/db/redis.rs b/src/db/redis.rs index 1b2ac22..52d0b0d 100644 --- a/src/db/redis.rs +++ b/src/db/redis.rs @@ -1,6 +1,11 @@ +use std::collections::HashMap; + use redis::{AsyncTypedCommands as _, cluster::ClusterClient, cluster_async::ClusterConnection}; -use crate::{error::Error, middleware::Server}; +use crate::{ + error::Error, + middleware::{ServerClient, StaticServerData}, +}; #[derive(Clone)] pub struct RedisClient(ClusterConnection); @@ -8,31 +13,163 @@ pub struct RedisClient(ClusterConnection); impl RedisClient { pub async fn init_redis( redis_url: &str, - available_servers: Vec, + available_servers: Vec, ) -> Result> { let nodes = vec![redis_url]; let client = ClusterClient::new(nodes)?; - let mut connection = client.get_async_connection().await?; + let connection = client.get_async_connection().await?; + + let mut client = Self(connection); + + // Preload server data into respective Redis keys + for server in available_servers.clone().into_iter() { + client.update_server_url(server.url.as_str()).await?; - // Preload server urls into Redis - for (server_index, server) in available_servers.clone().into_iter().enumerate() { - connection - .set(format!("server_{}", server_index), server.url.as_str()) - .await?; + client + .update_server_weight(server.url.as_str(), server.weight) + .await? } - Ok(Self(connection)) + Ok(client) } + // Basic Commands + + /// Set a key-value pair in Redis. pub async fn set(&mut self, key: &str, value: &str) -> Result<(), Error> { Ok(self.0.set(key, value).await?) } + /// Get the value associated with a key from Redis. pub async fn get(&mut self, key: &str) -> Result, Error> { Ok(self.0.get(key).await?) } + /// Delete a key from Redis. pub async fn delete(&mut self, key: &str) -> Result { Ok(self.0.del(key).await?) } + + // Server Data Commands + + /// Update the data of a server in Redis. + pub async fn update_server_url(&mut self, value: &str) -> Result<(), Error> { + Ok(self.0.rpush("server_url", value).await.map(|_| ())?) + } + + /// Get all server data from Redis. + pub async fn get_all_server_url(&mut self) -> Result, Error> { + self.0 + .lrange("server_url", 0, -1) + .await? + .into_iter() + .map(|v| Ok(StaticServerData::from_json(v)?.into())) + .collect::, _>>() + } + + // Server Load Commands + + /// Update the load of a server in Redis. + pub async fn update_server_load(&mut self, key: &str, value: u32) -> Result<(), Error> { + Ok(self.0.hset("server_load", key, value).await.map(|_| ())?) + } + + /// Get the load of a server from Redis. + pub async fn get_server_load(&mut self, key: &str) -> Result, Error> { + self.0 + .hget("server_load", key) + .await + .map_err(Error::RedisError)? + .map(|d| d.parse::().map_err(Error::ParseIntError)) + .transpose() + } + + /// Get all server load data from Redis. + pub async fn get_all_server_load(&mut self) -> Result, Error> { + self.0 + .hgetall("server_load") + .await? + .into_iter() + .map(|(k, v)| Ok((k, v.parse::().map_err(Error::ParseIntError)?))) + .collect::, _>>() + } + + // Latency Commands + + /// Update the mean latency record of a server in Redis. + pub async fn update_server_latency_record( + &mut self, + key: &str, + value: u128, + ) -> Result<(), Error> { + Ok(self.0.rpush(key, value).await.map(|_| ())?) + } + + pub async fn get_server_latency_record(&mut self, key: &str) -> Result, Error> { + self.0 + .lrange(key, 0, -1) + .await? + .into_iter() + .map(|v| Ok(v.parse()?)) + .collect::, _>>() + } + + /// Get the latency record of all servers in Redis. + pub async fn get_servers_latency_record( + &mut self, + ) -> Result>, Error> { + let server_client = self.get_all_server_url().await?; + + let mut res: HashMap> = HashMap::new(); + + for server in server_client { + let latencies = self.get_server_latency_record(server.url.as_str()).await?; + res.insert(server.url.to_string(), latencies); + } + Ok(res) + } + + /// Update the mean latency of a server in Redis. + pub async fn update_server_mean_latency( + &mut self, + key: &str, + value: u128, + ) -> Result<(), Error> { + Ok(self + .0 + .hset("server_latency", key, value) + .await + .map(|_| ())?) + } + + /// Get the mean latency of all servers in Redis. + pub async fn get_all_server_mean_latency(&mut self) -> Result, Error> { + self.0 + .hgetall("server_latency") + .await? + .into_iter() + .map(|(k, v)| Ok((k, v.parse::().map_err(Error::ParseIntError)?))) + .collect::, _>>() + } + + // Weights Commands + + /// Update the weight of a server in Redis. + pub async fn update_server_weight(&mut self, key: &str, value: u32) -> Result<(), Error> { + Ok(self + .0 + .hset("server_weights", key, value) + .await + .map(|_| ())?) + } + + /// Get the weights of all servers in Redis. + pub async fn get_all_server_weights(&mut self) -> Result, Error> { + self.0 + .hgetall("server_weights") + .await? + .into_iter() + .map(|(k, v)| Ok((k, v.parse::().map_err(Error::ParseIntError)?))) + .collect::, _>>() + } } diff --git a/src/error.rs b/src/error.rs index 4b53696..2d46bf5 100644 --- a/src/error.rs +++ b/src/error.rs @@ -1,3 +1,5 @@ +use std::string::ParseError; + use axum::response::{IntoResponse, Response}; use reqwest::StatusCode; @@ -22,6 +24,12 @@ pub enum Error { InvalidResponse, #[error("No Server Available")] NoServerAvailable, + #[error("Parse Error")] + ParseIntError(#[from] std::num::ParseIntError), + #[error("Parse Error")] + ParseError(#[from] ParseError), + #[error("Serialization Error")] + SerializationError(#[from] serde_json::Error), } impl IntoResponse for Error { @@ -32,7 +40,11 @@ impl IntoResponse for Error { | Error::Other(_) | Error::InvalidResponse | Error::RedisError(_) - | Error::NoServerAvailable => { + | Error::NoServerAvailable + | Error::ParseIntError(_) + | Error::ParseError(_) + | Error::SerializationError(_) => { + // TODO: Log error (StatusCode::INTERNAL_SERVER_ERROR, "Internal Server Error").into_response() } Error::Unauthorized => (StatusCode::UNAUTHORIZED, self).into_response(), diff --git a/src/main.rs b/src/main.rs index 1f1e5f6..956ba2f 100644 --- a/src/main.rs +++ b/src/main.rs @@ -2,27 +2,20 @@ use std::{net::SocketAddr, str::FromStr as _}; -use axum::{Router, routing::get}; -use tokio::task::JoinHandle; -use tower_http::{ - cors::{AllowHeaders, AllowMethods, AllowOrigin, CorsLayer}, - trace::TraceLayer, -}; use tracing::Level; use crate::{ + app::App, config::{State, SystemConfig}, - middleware::request_route, - servers::health::status, - services::server_status_worker, }; pub mod algorithms; +mod app; pub mod config; pub mod db; pub mod error; mod middleware; -mod servers; +mod route; mod services; #[tokio::main] @@ -36,78 +29,11 @@ async fn main() -> Result<(), Box> { let state = State::new(&config).await?; - let server = Router::new() - .route("/status", get(status)) - .layer( - CorsLayer::new() - .allow_headers(AllowHeaders::any()) - .allow_origin(AllowOrigin::any()) - .allow_methods(AllowMethods::any()), - ) - .layer(TraceLayer::new_for_http()) - // add rate limitter middleware mechanism - .layer(axum::middleware::from_fn_with_state( - state.clone(), - request_route, - )) - .with_state(state.clone()); - let addr = SocketAddr::from(([0, 0, 0, 0], config.port)); - - tracing::info!("Listening on {}", addr); - + tracing::info!("Listening on: {}", addr); let listener = tokio::net::TcpListener::bind(addr).await?; - let main = tokio::spawn(async move { axum::serve(listener, server).await }); - - let server_status_background_worker = tokio::spawn(async move { - let _: () = server_status_worker(state.clone().available_servers.available_servers).await; - Ok(()) - }); - - // let latency_tracker_background_worker = tokio::spawn(async move { - // let _: () = latency_tracker_worker(state.clone().available_servers.available_servers).await; - // Ok(()) - // }); - - let app = App::new( - main, - server_status_background_worker, - // latency_tracker_background_worker, - ); + let app = App::setup(state, listener).await?; app.start().await } - -type JoinHandleWrapper = JoinHandle>; - -/// Application struct to hold main and background worker tasks -struct App { - main: JoinHandleWrapper, - background_worker: JoinHandleWrapper, - // latency_tracker_background_worker: JoinHandleWrapper, -} - -impl App { - fn new( - main: JoinHandleWrapper, - background_worker: JoinHandleWrapper, - // latency_tracker_background_worker: JoinHandleWrapper, - ) -> Self { - Self { - main, - background_worker, - // latency_tracker_background_worker, - } - } - - async fn start(self) -> Result<(), Box> { - match tokio::try_join!(self.main, self.background_worker) { - Ok(_) => Ok(()), - Err(err) => Err(err)?, - } - } -} - -// background worker checking servers health status -// Load balancing algo? ref: https://www.cloudflare.com/learning/performance/types-of-load-balancing-algorithms/ diff --git a/src/middleware/mod.rs b/src/middleware/mod.rs index 3e389af..10fbf13 100644 --- a/src/middleware/mod.rs +++ b/src/middleware/mod.rs @@ -12,11 +12,11 @@ use crate::{config::State as AppState, error::Error}; mod server; -pub use server::{Server, ServerClients}; +pub use server::{ServerClient, StaticServerData}; /// Middleware function to route requests to appropriate servers pub async fn request_route( - State(state): State, + State(mut state): State, req: Request, next: Next, ) -> Result { @@ -35,20 +35,29 @@ pub async fn request_route( let route = parts.uri.to_string(); - let start_time = std::time::Instant::now(); - - let mut server: Server = state - .available_servers - .selected_server(state.algorithm) + let server_client = state + .algorithm + .select_server(state.redis_conn.clone()) .await?; - let response = server - .handle_request(parts.method, route.trim_start_matches('/'), json_body) + let start_time = std::time::Instant::now(); + + let response = server_client + .handle_request( + parts.method, + route.trim_start_matches('/'), + json_body, + state.redis_conn.clone(), + ) .await?; let latency = start_time.elapsed().as_millis(); - server.update_latencies(latency); + // TODO: move to background + state + .redis_conn + .update_server_latency_record(server_client.url.as_str(), latency) + .await?; Ok(response.into_response()) } diff --git a/src/middleware/server.rs b/src/middleware/server.rs index 8ec3137..2f165fb 100644 --- a/src/middleware/server.rs +++ b/src/middleware/server.rs @@ -1,122 +1,36 @@ -use std::{ - str::FromStr, - sync::{ - Arc, - atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering}, - }, -}; +use std::str::FromStr as _; use axum::response::{IntoResponse, Response}; use reqwest::{Method, Response as ReqwestResponse, StatusCode, Url}; +use serde::{Deserialize, Serialize}; -use crate::{algorithms::Algorithm, error::Error}; +use crate::{db::RedisClient, error::Error}; -/// Represents a collection of server clients for load balancing #[derive(Clone)] -pub struct ServerClients { - pub available_servers: Vec, -} - -impl ServerClients { - pub fn new(available_servers: Vec) -> Self { - Self { available_servers } - } - - /// Selects a server based on a load balancing algorithm - pub async fn selected_server(&self, algorithm: Algorithm) -> Result { - algorithm - .select_server(&self.available_servers) - .await - .inspect(|s| { - s.load.fetch_add(1, Ordering::Acquire); - }) - } -} - -#[derive(Clone)] -pub struct Server { +pub struct ServerClient { pub url: Url, pub client: reqwest::Client, - load: Arc, - weight: u32, - pub mean_latency: Arc, - pub latencies: Vec, - latencies_updated: Arc, } -impl Server { - pub fn new(url_and_weight: &str) -> anyhow::Result { - let (url, weight) = url_and_weight - .split_once('$') - .ok_or_else(|| anyhow::anyhow!("Invalid server format, expected 'url$weight'"))?; - - let weight = weight - .parse::() - .map_err(|_| anyhow::anyhow!("Invalid weight, expected a positive integer"))?; - - let url = Url::from_str(url)?; - - Ok(Self { - url, - client: Default::default(), - load: Arc::new(AtomicU32::new(0)), - weight, - mean_latency: Arc::new(AtomicU64::new(0)), - latencies: Vec::new(), - latencies_updated: Arc::new(AtomicBool::new(false)), - }) - } - - /// Returns the current load of the server - pub fn load(&self) -> u32 { - self.load.load(Ordering::Relaxed) - } - - /// Returns the weight of the server - pub fn weight(&self) -> u32 { - self.weight - } - - pub fn update_latencies(&mut self, latency: u128) { - if self.latencies.len() >= 20 { - // TODO: make it customisable - self.latencies.remove(0); - } - self.latencies.push(latency); - } - - pub fn latency_updated(&self) -> bool { - self.latencies_updated.load(Ordering::Relaxed) - } - - pub fn latency_update_status(&self, b: bool) { - self.latencies_updated.store(b, Ordering::Relaxed) - } - - pub fn mean_latency(&self) -> u64 { - self.mean_latency.load(Ordering::Relaxed) - } - +impl ServerClient { /// Handles incoming requests and forwards them to the server pub async fn handle_request( &self, method: Method, route: &str, body: Option, + mut redis_conn: RedisClient, ) -> Result { - // TODO: What if the request fails is the load count reduced? - match method { - Method::GET => self.get_request(route, body).await.inspect(|_| { - self.load.fetch_sub(1, Ordering::Release); - }), - Method::POST => self.post_request(route, body).await.inspect(|_| { - self.load.fetch_sub(1, Ordering::Release); - }), - _ => { - self.load.fetch_sub(1, Ordering::Release); - Err(Error::MethodNotAllowed) - } - } + let result = match method { + Method::GET => self.get_request(route, body).await, + Method::POST => self.post_request(route, body).await, + _ => return Err(Error::MethodNotAllowed), + }; + + // Update load once, regardless of success or failure + redis_conn.update_server_load(self.url.as_str(), 1).await?; + + result } /// Sends a POST request to the server @@ -176,6 +90,45 @@ impl Server { } } +impl From for ServerClient { + fn from(value: StaticServerData) -> Self { + Self { + url: value.url, + client: reqwest::Client::new(), + } + } +} + +#[derive(Serialize, Deserialize, Clone)] +pub struct StaticServerData { + pub url: Url, + pub weight: u32, +} + +impl StaticServerData { + pub fn from_json(data: String) -> Result { + serde_json::from_str(&data).map_err(Error::SerializationError) + } + + pub fn new(url_and_weight: &str) -> anyhow::Result { + let (url, weight) = url_and_weight + .split_once('$') + .ok_or_else(|| anyhow::anyhow!("Invalid server format, expected 'url$weight'"))?; + + let weight = weight + .parse::() + .map_err(|_| anyhow::anyhow!("Invalid weight, expected a positive integer"))?; + + let url = Url::from_str(url)?; + + Ok(Self { url, weight }) + } + + pub fn static_data(self) -> anyhow::Result { + Ok(serde_json::to_string(&self)?) + } +} + pub struct ApiResponse { status: StatusCode, message: String, diff --git a/src/servers/health.rs b/src/route/health.rs similarity index 100% rename from src/servers/health.rs rename to src/route/health.rs diff --git a/src/servers/mod.rs b/src/route/mod.rs similarity index 100% rename from src/servers/mod.rs rename to src/route/mod.rs diff --git a/src/services/latency_tracker_worker.rs b/src/services/latency_tracker_worker.rs index 538e1fe..7af4092 100644 --- a/src/services/latency_tracker_worker.rs +++ b/src/services/latency_tracker_worker.rs @@ -1,27 +1,24 @@ -use std::sync::atomic::Ordering; +use crate::db::RedisClient; -use crate::middleware::Server; - -pub async fn _latency_tracker_worker(available_servers: Vec) { +pub async fn latency_tracker_worker(redis_conn: RedisClient) { loop { - _check(&available_servers); + check(redis_conn.clone()).await; tokio::time::sleep(std::time::Duration::from_millis(100)).await; } } -fn _check(servers: &[Server]) { - for server in servers { - if server.latency_updated() { - server.latency_update_status(false); - let mean_latency = _mean_latency(&server.latencies); - server - .mean_latency - .store(mean_latency as u64, Ordering::Relaxed); +async fn check(mut redis_conn: RedisClient) { + if let Ok(data) = redis_conn.get_servers_latency_record().await { + for (url, latencies) in data { + let mean_latency = mean_latency(latencies); + _ = redis_conn + .update_server_mean_latency(&url, mean_latency) + .await; } } } -fn _mean_latency(latencies: &[u128]) -> u128 { +fn mean_latency(latencies: Vec) -> u128 { if latencies.is_empty() { 0 } else { diff --git a/src/services/mod.rs b/src/services/mod.rs index 533afac..9f4dd73 100644 --- a/src/services/mod.rs +++ b/src/services/mod.rs @@ -1,5 +1,5 @@ mod latency_tracker_worker; mod server_status_worker; -// pub use latency_tracker_worker::latency_tracker_worker; +pub use latency_tracker_worker::latency_tracker_worker; pub use server_status_worker::server_status_worker; diff --git a/src/services/server_status_worker.rs b/src/services/server_status_worker.rs index d264199..f2f6fc5 100644 --- a/src/services/server_status_worker.rs +++ b/src/services/server_status_worker.rs @@ -1,9 +1,9 @@ -use crate::middleware::Server; +use crate::db::RedisClient; /// Background worker that periodically checks the status of available servers -pub async fn server_status_worker(available_servers: Vec) { +pub async fn server_status_worker(redis_conn: RedisClient) { loop { - if let Err(failing_servers) = server_status(available_servers.clone()).await { + if let Err(failing_servers) = server_status(redis_conn.clone()).await { // TODO: remove them from the list of available servers tracing::warn!("Failing servers: {:#?}", failing_servers); } @@ -11,13 +11,17 @@ pub async fn server_status_worker(available_servers: Vec) { } } -async fn server_status(available_servers: Vec) -> Result<(), Vec> { +async fn server_status(mut redis_conn: RedisClient) -> Result<(), Vec> { let mut failing_servers = Vec::new(); - for server in available_servers { - if !server.is_available().await { - failing_servers.push(server.url.to_string()) + + if let Ok(data) = redis_conn.get_all_server_url().await { + for server in data { + if !server.is_available().await { + failing_servers.push(server.url.to_string()); + } } } + if failing_servers.is_empty() { Ok(()) } else {