use crate::load_balancer::{find_servers_for_model, select_server_adaptive};
use crate::models::{ResponseStatus, ServerResponse};
use crate::proxy::model_extractor::extract_model_name;
use crate::proxy::request_builder::build_proxy_request;
use crate::state::AppState;
use axum::{
extract::{Request, State},
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
use tracing;
pub async fn proxy_request_handler(
State(state): State<AppState>,
original_req: Request,
) -> Response {
tracing::trace!(?original_req, "Received proxy request");
let servers_guard = state.servers.lock().await;
if servers_guard.is_empty() {
tracing::warn!("No vLLM servers registered.");
return (
StatusCode::SERVICE_UNAVAILABLE,
Json(ServerResponse {
status: ResponseStatus::Error,
message: "No vLLM servers registered".to_string(),
}),
)
.into_response();
}
drop(servers_guard);
let (parts, body) = original_req.into_parts();
let (model_name, body_bytes) = match extract_model_name(body).await {
Ok(result) => result,
Err(error_response) => return *error_response,
};
let servers_guard = state.servers.lock().await;
let candidate_servers = find_servers_for_model(&servers_guard, &model_name);
if candidate_servers.is_empty() {
tracing::warn!("No server registered for model: {model_name}");
return (
StatusCode::BAD_REQUEST,
Json(ServerResponse {
status: ResponseStatus::Error,
message: format!("No server registered for model: {model_name}"),
}),
)
.into_response();
}
let selected_server = select_server_adaptive(&candidate_servers);
let target_addr = selected_server.addr.clone();
let selected_model = selected_server.model_name.clone();
let avg_response_time = selected_server.metrics.avg_response_time_ms;
let server_index = servers_guard
.iter()
.position(|s| s.model_name == selected_model && s.addr == target_addr)
.expect("Selected server must exist in list");
drop(servers_guard);
tracing::debug!(
"Selected server: {} for model {} (avg response time: {:.2}ms)",
target_addr,
model_name,
avg_response_time
);
let new_req = match build_proxy_request(parts, body_bytes, &target_addr) {
Ok(req) => req,
Err(error_response) => return *error_response,
};
let start_time = std::time::Instant::now();
let response_result = state.http_client.request(new_req).await;
let elapsed_ms = start_time.elapsed().as_millis() as f64;
let mut servers_guard = state.servers.lock().await;
if let Some(server) = servers_guard.get_mut(server_index) {
crate::metrics::update_metrics(server, elapsed_ms, &response_result);
}
drop(servers_guard);
match response_result {
Ok(response) => {
tracing::debug!(status = ?response.status(), "Received response from target");
response.into_response()
}
Err(err) => {
tracing::error!("Error forwarding request to {}: {}", target_addr, err);
(
StatusCode::BAD_GATEWAY,
Json(ServerResponse {
status: ResponseStatus::Error,
message: format!("Error forwarding request: {}", err),
}),
)
.into_response()
}
}
}