use crate::handlers::{list_servers, register_server, test_server, unregister_server};
use crate::proxy::proxy_request_handler;
use crate::state::AppState;
use axum::{
routing::{get, post},
Router,
};
use hyper_util::{client::legacy::Client, rt::TokioExecutor};
use std::{net::SocketAddr, time::Duration};
pub async fn run(addr: SocketAddr) {
let http_client = Client::builder(TokioExecutor::new())
.pool_idle_timeout(Duration::from_secs(30))
.http2_only(false)
.build_http();
let state = AppState::new(http_client);
let state_for_logging = state.clone();
tokio::spawn(async move {
crate::metrics::log_weights_periodically(state_for_logging).await;
});
let app = app(state);
let listener = tokio::net::TcpListener::bind(addr).await.unwrap();
tracing::info!("Listening on {}", listener.local_addr().unwrap());
axum::serve(listener, app.into_make_service())
.await
.unwrap();
}
fn app(state: AppState) -> Router {
let api_routes = Router::new()
.route("/register", post(register_server))
.route("/unregister", post(unregister_server))
.route("/health", get(|| async { "OK" }))
.route("/list", get(list_servers))
.route("/test", post(test_server));
let proxy_router = Router::new().fallback(proxy_request_handler);
Router::new()
.merge(api_routes)
.merge(proxy_router)
.with_state(state)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::models::{RegisterRequest, ResponseStatus, ServerResponse};
use axum::{
body::Body,
http::{self, Request, StatusCode},
};
use tower::ServiceExt;
fn test_app_state() -> AppState {
let http_client = Client::builder(TokioExecutor::new()).build_http();
AppState::new(http_client)
}
#[tokio::test]
#[ignore] async fn test_register_server_ok() {
let state = test_app_state();
let app = app(state.clone());
let payload = RegisterRequest {
model_name: None,
addr: "localhost:8001".to_string(),
};
let response = app
.oneshot(
Request::builder()
.method(http::Method::POST)
.uri("/register")
.header(http::header::CONTENT_TYPE, mime::APPLICATION_JSON.as_ref())
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::CREATED);
let servers = state.servers.lock().await;
assert_eq!(servers.len(), 1);
assert_eq!(servers[0].model_name, "test_model");
assert_eq!(servers[0].addr, "localhost:8001");
}
#[tokio::test]
#[ignore] async fn test_register_server_already_exists() {
let state = test_app_state();
let app = app(state.clone());
let payload = RegisterRequest {
model_name: None,
addr: "localhost:8001".to_string(),
};
app.clone()
.oneshot(
Request::builder()
.method(http::Method::POST)
.uri("/register")
.header(http::header::CONTENT_TYPE, mime::APPLICATION_JSON.as_ref())
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap(),
)
.await
.unwrap();
let response = app
.oneshot(
Request::builder()
.method(http::Method::POST)
.uri("/register")
.header(http::header::CONTENT_TYPE, mime::APPLICATION_JSON.as_ref())
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let server_response: ServerResponse = serde_json::from_slice(&body).unwrap();
assert_eq!(server_response.status, ResponseStatus::Warning);
assert_eq!(server_response.message, "Server already registered");
}
}