diff --git a/Readme.md b/Readme.md index f69118a..54222c1 100644 --- a/Readme.md +++ b/Readme.md @@ -12,7 +12,7 @@ Prerequisites: Create a `.env` file at the project root (example): ```bash -AVAILABLE_SERVERS=http://localhost:3001,http://localhost:3002 +AVAILABLE_SERVERS=http://localhost:3001$4,http://localhost:3002$8 PORT=8080 ``` @@ -34,7 +34,7 @@ The server binds to the configured PORT. This project reads configuration from environment variables. The important variables are: -`AVAILABLE_SERVERS` — comma-separated list of backend base URLs (e.g. http://host:port). +`AVAILABLE_SERVERS` — comma-separated list of backend base URLs and their weights (e.g. http://host:port$weight). `PORT` — port to bind the load balancer to. diff --git a/src/algorithms/least_connection.rs b/src/algorithms/least_connection.rs new file mode 100644 index 0000000..5b9e53b --- /dev/null +++ b/src/algorithms/least_connection.rs @@ -0,0 +1,10 @@ +use crate::{error::Error, middleware::Server}; + +pub async fn least_connection(available_servers: &[Server]) -> Result { + let server = available_servers + .iter() + .min_by_key(|server| server.load()) + .ok_or_else(|| Error::NoServerAvailable)?; + + Ok(server.clone()) +} diff --git a/src/algorithms/mod.rs b/src/algorithms/mod.rs new file mode 100644 index 0000000..e0d7ae0 --- /dev/null +++ b/src/algorithms/mod.rs @@ -0,0 +1,44 @@ +use crate::{error::Error, middleware::Server}; + +mod least_connection; +mod resource_based; +mod weighted_least_connection; +mod weighted_response_time; + +#[derive(Clone, Default)] +pub enum Algorithm { + #[default] + LeastConnection, + ResourceBased, + WeightedLeastConnection, + WeightedResponseTime, +} + +impl From for Algorithm { + fn from(algorithm: String) -> Self { + match algorithm.as_str() { + "least_connection" => Algorithm::LeastConnection, + "resource_based" => Algorithm::ResourceBased, + "weighted_least_connection" => Algorithm::WeightedLeastConnection, + "weighted_response_time" => Algorithm::WeightedResponseTime, + _ => Algorithm::default(), + } + } +} + +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, + Algorithm::WeightedLeastConnection => { + weighted_least_connection::weighted_least_connection(available_servers).await + } + Algorithm::WeightedResponseTime => { + weighted_response_time::weighted_response_time(available_servers).await + } + } + } +} diff --git a/src/algorithms/resource_based.rs b/src/algorithms/resource_based.rs new file mode 100644 index 0000000..64e47ac --- /dev/null +++ b/src/algorithms/resource_based.rs @@ -0,0 +1,6 @@ +use crate::{error::Error, middleware::Server}; + +// TODO: Implement resource-based load balancing algorithm +pub async fn resource_based(available_servers: &[Server]) -> Result { + Ok(available_servers[0].clone()) +} diff --git a/src/algorithms/weighted_least_connection.rs b/src/algorithms/weighted_least_connection.rs new file mode 100644 index 0000000..84f2f87 --- /dev/null +++ b/src/algorithms/weighted_least_connection.rs @@ -0,0 +1,10 @@ +use crate::{error::Error, middleware::Server}; + +pub async fn weighted_least_connection(available_servers: &[Server]) -> Result { + let server = available_servers + .iter() + .min_by_key(|server| server.load() / server.weight()) + .ok_or_else(|| Error::NoServerAvailable)?; + + Ok(server.clone()) +} diff --git a/src/algorithms/weighted_response_time.rs b/src/algorithms/weighted_response_time.rs new file mode 100644 index 0000000..0d289b2 --- /dev/null +++ b/src/algorithms/weighted_response_time.rs @@ -0,0 +1,10 @@ +use crate::{error::Error, middleware::Server}; + +pub async fn weighted_response_time(available_servers: &[Server]) -> Result { + let server = available_servers + .iter() + .min_by_key(|server| server.mean_latency()) + .ok_or_else(|| Error::NoServerAvailable)?; + + Ok(server.clone()) +} diff --git a/src/config.rs b/src/config.rs index 977c0b7..8c7dbcf 100644 --- a/src/config.rs +++ b/src/config.rs @@ -1,15 +1,18 @@ use serde::Deserialize; use crate::{ + algorithms::Algorithm, db::{self, RedisClient}, middleware::{Server, ServerClients}, }; -#[derive(Deserialize, Debug)] +#[derive(Deserialize)] pub struct SystemConfig { pub available_servers: String, // TODO: This should be hosted in redis pub port: u16, pub redis_url: String, + pub algorithm: String, + pub trace_level: String, } impl SystemConfig { @@ -25,6 +28,7 @@ impl SystemConfig { pub struct State { pub available_servers: ServerClients, pub redis_conn: RedisClient, + pub algorithm: Algorithm, } impl State { @@ -43,6 +47,7 @@ impl State { Ok(State { available_servers: ServerClients::new(available_servers), redis_conn, + algorithm: config.algorithm.clone().into(), }) } } diff --git a/src/error.rs b/src/error.rs index dc84168..4b53696 100644 --- a/src/error.rs +++ b/src/error.rs @@ -20,6 +20,8 @@ pub enum Error { InvalidUrl, #[error("Invalid Response")] InvalidResponse, + #[error("No Server Available")] + NoServerAvailable, } impl IntoResponse for Error { @@ -29,7 +31,8 @@ impl IntoResponse for Error { Error::InternalServerError | Error::Other(_) | Error::InvalidResponse - | Error::RedisError(_) => { + | Error::RedisError(_) + | Error::NoServerAvailable => { (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 a1400f4..1f1e5f6 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,6 +1,6 @@ #![deny(clippy::disallowed_methods)] -use std::net::SocketAddr; +use std::{net::SocketAddr, str::FromStr as _}; use axum::{Router, routing::get}; use tokio::task::JoinHandle; @@ -14,9 +14,10 @@ use crate::{ config::{State, SystemConfig}, middleware::request_route, servers::health::status, - services::server_worker, + services::server_status_worker, }; +pub mod algorithms; pub mod config; pub mod db; pub mod error; @@ -26,13 +27,13 @@ mod services; #[tokio::main] async fn main() -> Result<(), Box> { + let config = SystemConfig::from_env()?; + tracing_subscriber::fmt() - .with_max_level(Level::INFO) + .with_max_level(Level::from_str(&config.trace_level)?) .pretty() .init(); - let config = SystemConfig::from_env()?; - let state = State::new(&config).await?; let server = Router::new() @@ -59,12 +60,21 @@ async fn main() -> Result<(), Box> { let main = tokio::spawn(async move { axum::serve(listener, server).await }); - let background_worker = tokio::spawn(async move { - let _: () = server_worker(state.clone().available_servers.available_servers).await; + let server_status_background_worker = tokio::spawn(async move { + let _: () = server_status_worker(state.clone().available_servers.available_servers).await; Ok(()) }); - let app = App::new(main, background_worker); + // 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, + ); app.start().await } @@ -75,13 +85,19 @@ type JoinHandleWrapper = JoinHandle>; struct App { main: JoinHandleWrapper, background_worker: JoinHandleWrapper, + // latency_tracker_background_worker: JoinHandleWrapper, } impl App { - fn new(main: JoinHandleWrapper, background_worker: JoinHandleWrapper) -> Self { + fn new( + main: JoinHandleWrapper, + background_worker: JoinHandleWrapper, + // latency_tracker_background_worker: JoinHandleWrapper, + ) -> Self { Self { main, background_worker, + // latency_tracker_background_worker, } } diff --git a/src/middleware/mod.rs b/src/middleware/mod.rs index b09a76a..3e389af 100644 --- a/src/middleware/mod.rs +++ b/src/middleware/mod.rs @@ -3,11 +3,12 @@ use axum::{ extract::State, http::Request, middleware::Next, + response::IntoResponse, }; use futures_util::stream::StreamExt; use serde_json::Value; -use crate::{config::State as AppState, error::Error, middleware::server::ApiResponse}; +use crate::{config::State as AppState, error::Error}; mod server; @@ -17,8 +18,12 @@ pub use server::{Server, ServerClients}; pub async fn request_route( State(state): State, req: Request, - _next: Next, -) -> Result { + next: Next, +) -> Result { + if req.uri().path().starts_with("/status") { + return Ok(next.run(req).await); + } + let (parts, body) = req.into_parts(); tracing::info!("New Request Received"); @@ -30,11 +35,22 @@ pub async fn request_route( let route = parts.uri.to_string(); - state + let start_time = std::time::Instant::now(); + + let mut server: Server = state .available_servers - .choiced_server() + .selected_server(state.algorithm) + .await?; + + let response = server .handle_request(parts.method, route.trim_start_matches('/'), json_body) - .await + .await?; + + let latency = start_time.elapsed().as_millis(); + + server.update_latencies(latency); + + Ok(response.into_response()) } struct BodyBytes(Bytes); diff --git a/src/middleware/server.rs b/src/middleware/server.rs index d6352bd..8ec3137 100644 --- a/src/middleware/server.rs +++ b/src/middleware/server.rs @@ -1,12 +1,18 @@ -use std::str::FromStr; +use std::{ + str::FromStr, + sync::{ + Arc, + atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering}, + }, +}; use axum::response::{IntoResponse, Response}; use reqwest::{Method, Response as ReqwestResponse, StatusCode, Url}; -use crate::error::Error; +use crate::{algorithms::Algorithm, error::Error}; /// Represents a collection of server clients for load balancing -#[derive(Clone, Debug)] +#[derive(Clone)] pub struct ServerClients { pub available_servers: Vec, } @@ -17,26 +23,80 @@ impl ServerClients { } /// Selects a server based on a load balancing algorithm - pub fn choiced_server(&self) -> Server { - // implement algorithm to select server here! - self.available_servers[0].clone() // placeholder for now + 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, Debug)] +#[derive(Clone)] pub struct Server { 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: &str) -> anyhow::Result { + 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: Url::from_str(url)?, + 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) + } + /// Handles incoming requests and forwards them to the server pub async fn handle_request( &self, @@ -44,10 +104,18 @@ impl Server { route: &str, body: Option, ) -> Result { + // TODO: What if the request fails is the load count reduced? match method { - Method::GET => self.get_request(route, body).await, - Method::POST => self.post_request(route, body).await, - _ => Err(Error::MethodNotAllowed), + 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) + } } } diff --git a/src/services/latency_tracker_worker.rs b/src/services/latency_tracker_worker.rs new file mode 100644 index 0000000..538e1fe --- /dev/null +++ b/src/services/latency_tracker_worker.rs @@ -0,0 +1,30 @@ +use std::sync::atomic::Ordering; + +use crate::middleware::Server; + +pub async fn _latency_tracker_worker(available_servers: Vec) { + loop { + _check(&available_servers); + 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); + } + } +} + +fn _mean_latency(latencies: &[u128]) -> u128 { + if latencies.is_empty() { + 0 + } else { + latencies.iter().sum::() / latencies.len() as u128 + } +} diff --git a/src/services/mod.rs b/src/services/mod.rs index 7a2628e..533afac 100644 --- a/src/services/mod.rs +++ b/src/services/mod.rs @@ -1,3 +1,5 @@ -mod background_worker; +mod latency_tracker_worker; +mod server_status_worker; -pub use background_worker::server_worker; +// pub use latency_tracker_worker::latency_tracker_worker; +pub use server_status_worker::server_status_worker; diff --git a/src/services/background_worker.rs b/src/services/server_status_worker.rs similarity index 92% rename from src/services/background_worker.rs rename to src/services/server_status_worker.rs index 6a7156c..d264199 100644 --- a/src/services/background_worker.rs +++ b/src/services/server_status_worker.rs @@ -1,7 +1,7 @@ use crate::middleware::Server; /// Background worker that periodically checks the status of available servers -pub async fn server_worker(available_servers: Vec) { +pub async fn server_status_worker(available_servers: Vec) { loop { if let Err(failing_servers) = server_status(available_servers.clone()).await { // TODO: remove them from the list of available servers