Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions Readme.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
```

Expand All @@ -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.

Expand Down
10 changes: 10 additions & 0 deletions src/algorithms/least_connection.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
use crate::{error::Error, middleware::Server};

pub async fn least_connection(available_servers: &[Server]) -> Result<Server, Error> {
let server = available_servers
.iter()
.min_by_key(|server| server.load())
.ok_or_else(|| Error::NoServerAvailable)?;

Ok(server.clone())
}
44 changes: 44 additions & 0 deletions src/algorithms/mod.rs
Original file line number Diff line number Diff line change
@@ -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<String> 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<Server, Error> {
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
}
}
}
}
6 changes: 6 additions & 0 deletions src/algorithms/resource_based.rs
Original file line number Diff line number Diff line change
@@ -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<Server, Error> {
Ok(available_servers[0].clone())
}
10 changes: 10 additions & 0 deletions src/algorithms/weighted_least_connection.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
use crate::{error::Error, middleware::Server};

pub async fn weighted_least_connection(available_servers: &[Server]) -> Result<Server, Error> {
let server = available_servers
.iter()
.min_by_key(|server| server.load() / server.weight())
.ok_or_else(|| Error::NoServerAvailable)?;

Ok(server.clone())
}
10 changes: 10 additions & 0 deletions src/algorithms/weighted_response_time.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
use crate::{error::Error, middleware::Server};

pub async fn weighted_response_time(available_servers: &[Server]) -> Result<Server, Error> {
let server = available_servers
.iter()
.min_by_key(|server| server.mean_latency())
.ok_or_else(|| Error::NoServerAvailable)?;

Ok(server.clone())
}
7 changes: 6 additions & 1 deletion src/config.rs
Original file line number Diff line number Diff line change
@@ -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 {
Expand All @@ -25,6 +28,7 @@ impl SystemConfig {
pub struct State {
pub available_servers: ServerClients,
pub redis_conn: RedisClient,
pub algorithm: Algorithm,
}

impl State {
Expand All @@ -43,6 +47,7 @@ impl State {
Ok(State {
available_servers: ServerClients::new(available_servers),
redis_conn,
algorithm: config.algorithm.clone().into(),
})
}
}
5 changes: 4 additions & 1 deletion src/error.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@ pub enum Error {
InvalidUrl,
#[error("Invalid Response")]
InvalidResponse,
#[error("No Server Available")]
NoServerAvailable,
}

impl IntoResponse for Error {
Expand All @@ -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(),
Expand Down
34 changes: 25 additions & 9 deletions src/main.rs
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -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;
Expand All @@ -26,13 +27,13 @@ mod services;

#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
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()
Expand All @@ -59,12 +60,21 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {

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
}
Expand All @@ -75,13 +85,19 @@ type JoinHandleWrapper = JoinHandle<Result<(), std::io::Error>>;
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,
}
}

Expand Down
28 changes: 22 additions & 6 deletions src/middleware/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand All @@ -17,8 +18,12 @@ pub use server::{Server, ServerClients};
pub async fn request_route(
State(state): State<AppState>,
req: Request<axum::body::Body>,
_next: Next,
) -> Result<ApiResponse, Error> {
next: Next,
) -> Result<impl IntoResponse, Error> {
if req.uri().path().starts_with("/status") {
return Ok(next.run(req).await);
}

let (parts, body) = req.into_parts();

tracing::info!("New Request Received");
Expand All @@ -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);
Expand Down
Loading