llmproxy 0.2.2

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

    // Spawn background task to log weights periodically
    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] // Ignored: requires mock HTTP server for /v1/models endpoint
    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] // Ignored: requires mock HTTP server for /v1/models endpoint
    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(),
        };

        // First registration
        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();

        // Second registration
        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");
    }
}