llmproxy 0.2.2

A simple HTTP proxy server for llm api requests
Documentation
use crate::models::{
    ModelInfo, ModelsListResponse, RegisterRequest, ResponseStatus, ServerResponse,
};
use crate::state::{AppState, EndpointModelsCache, ProxyServer, ServerMetrics};
use axum::{extract::State, http::StatusCode, response::IntoResponse, Json};
use hyper_util::client::legacy::Client;
use tracing;

pub async fn register_server(
    State(state): State<AppState>,
    Json(payload): Json<RegisterRequest>,
) -> impl IntoResponse {
    // Validate address format
    if payload.addr.trim().is_empty() || !payload.addr.contains(':') {
        tracing::warn!(
            "Invalid address provided for registration: {}",
            payload.addr
        );
        return (
            StatusCode::BAD_REQUEST,
            Json(ServerResponse {
                status: ResponseStatus::Error,
                message: "Invalid address format. Expected host:port".to_string(),
            }),
        );
    }

    let server_addr = payload.addr.trim().to_string();

    // Discover models from endpoint
    tracing::info!("Discovering models from endpoint: {}", server_addr);
    let discovered_models = match discover_models(&state.http_client, &server_addr).await {
        Ok(models) => models,
        Err(e) => {
            tracing::error!("Discovery failed for {}: {}", server_addr, e);
            return (
                StatusCode::BAD_GATEWAY,
                Json(ServerResponse {
                    status: ResponseStatus::Error,
                    message: format!("Failed to discover models: {}", e),
                }),
            );
        }
    };

    // Register all discovered models
    let mut servers = state.servers.lock().await;
    let mut models_cache = state.models_cache.lock().await;
    let mut newly_registered = Vec::new();
    let mut already_registered = Vec::new();

    for model in &discovered_models {
        let model_name = model.id.clone();

        // Check for duplicate
        if servers
            .iter()
            .any(|s| s.model_name == model_name && s.addr == server_addr)
        {
            already_registered.push(model_name);
            continue;
        }

        // Register new model
        servers.push(ProxyServer {
            model_name: model_name.clone(),
            addr: server_addr.clone(),
            metrics: ServerMetrics::default(),
        });
        newly_registered.push(model_name.clone());
        tracing::info!("Registered model '{}' at {}", model_name, server_addr);
    }

    // Update cache
    if let Some(cache_entry) = models_cache.iter_mut().find(|c| c.endpoint == server_addr) {
        cache_entry.models = discovered_models.clone();
        cache_entry.discovered_at = std::time::SystemTime::now();
    } else {
        models_cache.push(EndpointModelsCache {
            endpoint: server_addr.clone(),
            models: discovered_models.clone(),
            discovered_at: std::time::SystemTime::now(),
        });
    }

    // Drop locks
    drop(servers);
    drop(models_cache);

    // Build response message
    let message = if !newly_registered.is_empty() {
        if already_registered.is_empty() {
            format!(
                "Registered {} model(s): {}",
                newly_registered.len(),
                newly_registered.join(", ")
            )
        } else {
            format!(
                "Registered {} new model(s): {} (Already registered: {})",
                newly_registered.len(),
                newly_registered.join(", "),
                already_registered.join(", ")
            )
        }
    } else {
        format!(
            "All {} model(s) already registered: {}",
            already_registered.len(),
            already_registered.join(", ")
        )
    };

    let (status_code, response_status) = if !newly_registered.is_empty() {
        (StatusCode::CREATED, ResponseStatus::Success)
    } else {
        (StatusCode::OK, ResponseStatus::Warning)
    };

    (
        status_code,
        Json(ServerResponse {
            status: response_status,
            message,
        }),
    )
}

async fn discover_models(
    _http_client: &Client<hyper_util::client::legacy::connect::HttpConnector, axum::body::Body>,
    addr: &str,
) -> Result<Vec<ModelInfo>, String> {
    // Use reqwest for simpler HTTP GET with JSON parsing
    let url = format!("http://{}/v1/models", addr);

    let reqwest_client = reqwest::Client::new();
    let response = reqwest_client
        .get(&url)
        .send()
        .await
        .map_err(|e| format!("Failed to connect: {}", e))?;

    if !response.status().is_success() {
        return Err(format!("Endpoint returned HTTP {}", response.status()));
    }

    let models_response: ModelsListResponse = response
        .json()
        .await
        .map_err(|e| format!("Invalid JSON response: {}", e))?;

    if models_response.data.is_empty() {
        return Err("No models available at this endpoint".to_string());
    }

    Ok(models_response.data)
}