llmproxy 0.2.2

A simple HTTP proxy server for llm api requests
Documentation
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");

    // Check if any servers are registered
    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);

    // Extract request parts and model name
    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,
    };

    // Find candidate servers for the model
    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();
    }

    // Select best server using adaptive load balancing
    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;

    // Find the server index for metrics update
    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
    );

    // Build proxy request
    let new_req = match build_proxy_request(parts, body_bytes, &target_addr) {
        Ok(req) => req,
        Err(error_response) => return *error_response,
    };

    // Send request and track timing
    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;

    // Update server metrics
    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);

    // Return response
    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()
        }
    }
}