use std::time::Instant;
use futures::future::join_all;
use crate::{
LLMProvider,
chat::{ChatMessage, Tool},
completion::CompletionRequest,
error::LLMError,
};
use super::ScoringFn;
#[derive(Debug)]
pub struct ParallelEvalResult {
pub text: String,
pub score: f32,
pub time_ms: u128,
pub provider_id: String,
}
pub struct ParallelEvaluator {
providers: Vec<(String, Box<dyn LLMProvider>)>,
scoring_fns: Vec<Box<ScoringFn>>,
include_timing: bool,
}
impl ParallelEvaluator {
pub fn new(providers: Vec<(String, Box<dyn LLMProvider>)>) -> Self {
Self {
providers,
scoring_fns: Vec::new(),
include_timing: true,
}
}
pub fn scoring<F>(mut self, f: F) -> Self
where
F: Fn(&str) -> f32 + Send + Sync + 'static,
{
self.scoring_fns.push(Box::new(f));
self
}
pub fn include_timing(mut self, include: bool) -> Self {
self.include_timing = include;
self
}
pub async fn evaluate_chat_parallel(
&self,
messages: &[ChatMessage],
) -> Result<Vec<ParallelEvalResult>, LLMError> {
let futures = self
.providers
.iter()
.map(|(id, provider)| {
let id = id.clone();
let messages = messages.to_vec();
async move {
let start = Instant::now();
let result = provider.chat(&messages, None).await;
let elapsed = start.elapsed().as_millis();
(id, result, elapsed)
}
})
.collect::<Vec<_>>();
let results = join_all(futures).await;
let mut eval_results = Vec::new();
for (id, result, elapsed) in results {
match result {
Ok(response) => {
let text = response.text().unwrap_or_default();
let score = self.compute_score(&text);
eval_results.push(ParallelEvalResult {
text,
score,
time_ms: elapsed,
provider_id: id,
});
}
Err(e) => {
eprintln!("Error from provider {id}: {e}");
}
}
}
Ok(eval_results)
}
pub async fn evaluate_chat_with_tools_parallel(
&self,
messages: &[ChatMessage],
tools: Option<&[Tool]>,
) -> Result<Vec<ParallelEvalResult>, LLMError> {
let futures = self
.providers
.iter()
.map(|(id, provider)| {
let id = id.clone();
let messages = messages.to_vec();
async move {
let start = Instant::now();
let result = provider.chat_with_tools(&messages, tools, None).await;
let elapsed = start.elapsed().as_millis();
(id, result, elapsed)
}
})
.collect::<Vec<_>>();
let results = join_all(futures).await;
let mut eval_results = Vec::new();
for (id, result, elapsed) in results {
match result {
Ok(response) => {
let text = response.text().unwrap_or_default();
let score = self.compute_score(&text);
eval_results.push(ParallelEvalResult {
text,
score,
time_ms: elapsed,
provider_id: id,
});
}
Err(e) => {
eprintln!("Error from provider {id}: {e}");
}
}
}
Ok(eval_results)
}
pub async fn evaluate_completion_parallel(
&self,
request: &CompletionRequest,
) -> Result<Vec<ParallelEvalResult>, LLMError> {
let futures = self
.providers
.iter()
.map(|(id, provider)| {
let id = id.clone();
let request = request.clone();
async move {
let start = Instant::now();
let result = provider.complete(&request, None).await;
let elapsed = start.elapsed().as_millis();
(id, result, elapsed)
}
})
.collect::<Vec<_>>();
let results = join_all(futures).await;
let mut eval_results = Vec::new();
for (id, result, elapsed) in results {
match result {
Ok(response) => {
let score = self.compute_score(&response.text);
eval_results.push(ParallelEvalResult {
text: response.text,
score,
time_ms: elapsed,
provider_id: id,
});
}
Err(e) => {
eprintln!("Error from provider {id}: {e}");
}
}
}
Ok(eval_results)
}
pub fn best_response<'a>(
&self,
results: &'a [ParallelEvalResult],
) -> Option<&'a ParallelEvalResult> {
if results.is_empty() {
return None;
}
results.iter().max_by(|a, b| {
a.score
.partial_cmp(&b.score)
.unwrap_or(std::cmp::Ordering::Equal)
})
}
fn compute_score(&self, response: &str) -> f32 {
let mut total = 0.0;
for sc in &self.scoring_fns {
total += sc(response);
}
total
}
}