use crate::state::{AppState, ProxyServer};
use hyper::Response;
use std::time::Duration;
use tracing;
pub fn calculate_ucb_score(server: &ProxyServer, total_requests: u64, c: f64) -> f64 {
if server.metrics.total_requests == 0 {
return f64::INFINITY;
}
let avg_reward = if server.metrics.recent_rewards.is_empty() {
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
};
let failure_penalty = 0.5_f64.powi(server.metrics.recent_failures as i32);
let adjusted_reward = avg_reward * failure_penalty;
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
}
pub fn calculate_server_weight(server: &ProxyServer) -> f64 {
let base_weight = 1000.0 / server.metrics.avg_response_time_ms;
let failure_penalty = 0.5_f64.powi(server.metrics.recent_failures as i32);
base_weight * failure_penalty
}
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;
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());
let reward = 1000.0 / elapsed_ms.max(1.0);
metrics.recent_rewards.push_back(reward);
if metrics.recent_rewards.len() > metrics.window_size {
metrics.recent_rewards.pop_front();
}
match response_result {
Ok(response) => {
if response.status().is_success() {
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 {
metrics.recent_failures = (metrics.recent_failures + 1).min(10);
tracing::warn!(
"Request returned HTTP {} in {:.2}ms",
response.status(),
elapsed_ms
);
}
}
Err(_) => {
metrics.recent_failures = (metrics.recent_failures + 1).min(10);
tracing::error!("Request failed in {:.2}ms", elapsed_ms);
}
}
}
pub async fn log_weights_periodically(state: AppState) {
let mut interval = tokio::time::interval(Duration::from_secs(60));
interval.tick().await;
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;
}
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);
}
for (model_name, model_servers) in models {
if model_servers.len() == 1 {
continue;
}
let total_requests: u64 = model_servers.iter().map(|s| s.metrics.total_requests).sum();
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()]
};
const UCB_C: f64 = 1.5; let ucb_scores: Vec<f64> = model_servers
.iter()
.map(|s| calculate_ucb_score(s, total_requests, UCB_C))
.collect();
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 };
server_qps_values.push(qps);
total_qps += qps;
previous_stats.insert(server.addr.clone(), (current_requests, now));
}
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")
);
}
}
}