llmproxy 0.2.2

A simple HTTP proxy server for llm api requests
Documentation
use crate::state::{AppState, ProxyServer};
use hyper::Response;
use std::time::Duration;
use tracing;

/// Calculate UCB (Upper Confidence Bound) score for a server
/// Used for multi-armed bandit load balancing
///
/// UCB formula: avg_reward + c * sqrt(2 * ln(total) / count)
///   - avg_reward: average reward from recent requests (higher is better)
///   - exploration bonus: encourages trying less-explored servers
///   - c: exploration parameter (typically 1.0-2.0)
pub fn calculate_ucb_score(server: &ProxyServer, total_requests: u64, c: f64) -> f64 {
    // Never-tried servers get infinite score (try them first)
    if server.metrics.total_requests == 0 {
        return f64::INFINITY;
    }

    // Calculate average reward from sliding window
    let avg_reward = if server.metrics.recent_rewards.is_empty() {
        // No data yet, use a neutral reward based on default response time
        1000.0 / server.metrics.avg_response_time_ms
    } else {
        let sum: f64 = server.metrics.recent_rewards.iter().sum();
        sum / server.metrics.recent_rewards.len() as f64
    };

    // Apply failure penalty (reduce score by 50% per recent failure)
    let failure_penalty = 0.5_f64.powi(server.metrics.recent_failures as i32);
    let adjusted_reward = avg_reward * failure_penalty;

    // Calculate exploration bonus
    // More requests means we're more confident, so exploration bonus decreases
    let exploration_bonus = if total_requests > 0 {
        c * (2.0 * (total_requests as f64).ln() / server.metrics.total_requests as f64).sqrt()
    } else {
        0.0
    };

    adjusted_reward + exploration_bonus
}

/// Calculate weight for a server based on performance metrics
pub fn calculate_server_weight(server: &ProxyServer) -> f64 {
    let base_weight = 1000.0 / server.metrics.avg_response_time_ms;

    // Penalize servers with recent failures (reduce weight by 50% per failure)
    let failure_penalty = 0.5_f64.powi(server.metrics.recent_failures as i32);

    base_weight * failure_penalty
}

/// Update server metrics based on request timing and result
pub fn update_metrics(
    server: &mut ProxyServer,
    elapsed_ms: f64,
    response_result: &Result<Response<hyper::body::Incoming>, hyper_util::client::legacy::Error>,
) {
    let metrics = &mut server.metrics;

    // Update exponential moving average (alpha = 0.2 for smoothing)
    let alpha = 0.2;
    metrics.avg_response_time_ms =
        alpha * elapsed_ms + (1.0 - alpha) * metrics.avg_response_time_ms;

    metrics.total_requests += 1;
    metrics.last_request_time = Some(std::time::SystemTime::now());

    // --- Update UCB sliding window ---
    // Calculate reward: 1000.0 / response_time (higher is better)
    // Scale to ~1000 to make rewards easier to interpret
    let reward = 1000.0 / elapsed_ms.max(1.0); // Avoid division by zero

    // Add reward to sliding window
    metrics.recent_rewards.push_back(reward);

    // Remove old rewards if window is full
    if metrics.recent_rewards.len() > metrics.window_size {
        metrics.recent_rewards.pop_front();
    }

    match response_result {
        Ok(response) => {
            if response.status().is_success() {
                // Success: decay recent failures
                metrics.recent_failures = metrics.recent_failures.saturating_sub(1);
                tracing::debug!(
                    "Request succeeded in {:.2}ms, updated avg to {:.2}ms, reward: {:.2}",
                    elapsed_ms,
                    metrics.avg_response_time_ms,
                    reward
                );
            } else {
                // HTTP error: increment failures
                metrics.recent_failures = (metrics.recent_failures + 1).min(10);
                tracing::warn!(
                    "Request returned HTTP {} in {:.2}ms",
                    response.status(),
                    elapsed_ms
                );
            }
        }
        Err(_) => {
            // Network/connection error: increment failures
            metrics.recent_failures = (metrics.recent_failures + 1).min(10);
            tracing::error!("Request failed in {:.2}ms", elapsed_ms);
        }
    }
}

/// Background task that periodically logs load balancing weights
pub async fn log_weights_periodically(state: AppState) {
    let mut interval = tokio::time::interval(Duration::from_secs(60));
    interval.tick().await; // Skip the first immediate tick

    // Track previous request counts for QPS calculation
    // Key: server address, Value: (requests, timestamp)
    let mut previous_stats: std::collections::HashMap<String, (u64, std::time::Instant)> =
        std::collections::HashMap::new();

    loop {
        interval.tick().await;
        let now = std::time::Instant::now();

        let servers = state.servers.lock().await;
        if servers.is_empty() {
            continue;
        }

        // Group servers by model
        let mut models: std::collections::HashMap<String, Vec<&ProxyServer>> =
            std::collections::HashMap::new();
        for server in servers.iter() {
            models
                .entry(server.model_name.clone())
                .or_default()
                .push(server);
        }

        // Log weights and UCB scores for each model
        for (model_name, model_servers) in models {
            if model_servers.len() == 1 {
                // Skip logging for single-server models
                continue;
            }

            // Calculate total requests for UCB formula
            let total_requests: u64 = model_servers.iter().map(|s| s.metrics.total_requests).sum();

            // Calculate weights (for backward compatibility)
            let weights: Vec<f64> = model_servers
                .iter()
                .map(|s| calculate_server_weight(s))
                .collect();

            let total_weight: f64 = weights.iter().sum();
            let normalized_weights: Vec<f64> = if total_weight > 0.0 {
                weights.iter().map(|w| w / total_weight * 100.0).collect()
            } else {
                let uniform = 100.0 / model_servers.len() as f64;
                vec![uniform; model_servers.len()]
            };

            // Calculate UCB scores
            const UCB_C: f64 = 1.5; // Match the value in load_balancer.rs
            let ucb_scores: Vec<f64> = model_servers
                .iter()
                .map(|s| calculate_ucb_score(s, total_requests, UCB_C))
                .collect();

            // Calculate QPS for each server and total
            let mut total_qps = 0.0;
            let mut server_qps_values = Vec::new();

            for server in model_servers.iter() {
                let current_requests = server.metrics.total_requests;
                let qps = if let Some((prev_requests, prev_time)) = previous_stats.get(&server.addr)
                {
                    let elapsed_secs = now.duration_since(*prev_time).as_secs_f64();
                    if elapsed_secs > 0.0 {
                        (current_requests - prev_requests) as f64 / elapsed_secs
                    } else {
                        0.0
                    }
                } else {
                    0.0 // First time seeing this server
                };

                server_qps_values.push(qps);
                total_qps += qps;

                // Update previous stats
                previous_stats.insert(server.addr.clone(), (current_requests, now));
            }

            // Build log message with UCB scores and QPS
            let mut weight_info = Vec::new();
            for (i, server) in model_servers.iter().enumerate() {
                let avg_reward = if server.metrics.recent_rewards.is_empty() {
                    0.0
                } else {
                    let sum: f64 = server.metrics.recent_rewards.iter().sum();
                    sum / server.metrics.recent_rewards.len() as f64
                };

                weight_info.push(format!(
                    "  {} -> {:.1}% (avg: {:.0}ms, reqs: {}, qps: {:.1}, failures: {}, ucb: {:.2}, reward: {:.2})",
                    server.addr,
                    normalized_weights[i],
                    server.metrics.avg_response_time_ms,
                    server.metrics.total_requests,
                    server_qps_values[i],
                    server.metrics.recent_failures,
                    ucb_scores[i],
                    avg_reward
                ));
            }

            tracing::info!(
                "Load balancing metrics for model '{}' (UCB algorithm) [Total QPS: {:.1}]:\n{}",
                model_name,
                total_qps,
                weight_info.join("\n")
            );
        }
    }
}