use anyhow::{anyhow, Context, Result};
use clap::Args;
use futures::StreamExt;
use indicatif::{ProgressBar, ProgressStyle};
use llm_test_bench_core::config::ConfigLoader;
use llm_test_bench_core::providers::{
AnthropicProvider, CompletionRequest, OpenAIProvider, Provider, ProviderError,
};
use std::path::PathBuf;
use std::time::Instant;
use crate::output::{display_error, display_response, OutputFormat, StreamingOutput};
#[derive(Args, Debug)]
pub struct TestArgs {
pub provider: String,
#[arg(short, long)]
pub prompt: String,
#[arg(short, long)]
pub model: Option<String>,
#[arg(short, long)]
pub temperature: Option<f32>,
#[arg(long)]
pub max_tokens: Option<u32>,
#[arg(long)]
pub top_p: Option<f32>,
#[arg(long)]
pub stop: Option<Vec<String>>,
#[arg(short, long)]
pub stream: bool,
#[arg(short, long, value_enum, default_value = "pretty")]
pub output_format: OutputFormat,
#[arg(short, long)]
pub config: Option<PathBuf>,
}
fn create_provider(
provider_name: &str,
config_path: &Option<PathBuf>,
) -> Result<Box<dyn Provider>> {
let mut loader = ConfigLoader::new();
if let Some(path) = config_path {
loader = loader.with_file(path);
}
let config = loader.load().context("Failed to load configuration")?;
let provider_config = config
.providers
.get(provider_name)
.ok_or_else(|| anyhow!("Provider '{}' not found in configuration", provider_name))?;
let api_key = std::env::var(&provider_config.api_key_env)
.with_context(|| format!("Environment variable '{}' not set", provider_config.api_key_env))?;
match provider_name {
"openai" => {
let provider = OpenAIProvider::with_base_url(api_key, provider_config.base_url.clone())
.context("Failed to create OpenAI provider")?;
Ok(Box::new(provider))
}
"anthropic" => {
let provider = AnthropicProvider::with_base_url(api_key, provider_config.base_url.clone());
Ok(Box::new(provider))
}
_ => Err(anyhow!("Unknown provider: {}", provider_name)),
}
}
fn build_completion_request(
args: &TestArgs,
provider: &Box<dyn Provider>,
) -> Result<CompletionRequest> {
let mut loader = ConfigLoader::new();
if let Some(ref path) = args.config {
loader = loader.with_file(path);
}
let config = loader.load()?;
let provider_config = config.providers.get(&args.provider).unwrap();
let model = args.model.clone().unwrap_or_else(|| provider_config.default_model.clone());
let supported_models = provider.supported_models();
if !supported_models.iter().any(|m| m.id == model) {
return Err(anyhow!(
"Model '{}' is not supported by provider '{}'\nSupported models: {}",
model,
args.provider,
supported_models
.iter()
.map(|m| m.id.as_str())
.collect::<Vec<_>>()
.join(", ")
));
}
Ok(CompletionRequest {
prompt: args.prompt.clone(),
model,
temperature: args.temperature,
max_tokens: args.max_tokens.map(|t| t as usize),
top_p: args.top_p,
stop: args.stop.clone(),
stream: args.stream,
})
}
async fn execute_non_streaming(
provider: &Box<dyn Provider>,
request: CompletionRequest,
output_format: OutputFormat,
) -> Result<()> {
let spinner = if output_format == OutputFormat::Pretty {
let pb = ProgressBar::new_spinner();
pb.set_style(
ProgressStyle::default_spinner()
.template("{spinner:.cyan} {msg}")
.unwrap()
.tick_strings(&["⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"]),
);
pb.set_message(format!("Requesting completion from {}...", provider.name()));
pb.enable_steady_tick(std::time::Duration::from_millis(100));
Some(pb)
} else {
None
};
let start = Instant::now();
let response = provider.complete(request).await.map_err(map_provider_error)?;
let elapsed = start.elapsed();
if let Some(pb) = spinner {
pb.finish_and_clear();
}
display_response(&response, output_format)?;
if output_format == OutputFormat::Pretty {
println!("⏱️ Response time: {:.2}s", elapsed.as_secs_f64());
}
Ok(())
}
async fn execute_streaming(
provider: &Box<dyn Provider>,
request: CompletionRequest,
output_format: OutputFormat,
) -> Result<()> {
let model = request.model.clone();
let mut stream = provider.stream(request).await.map_err(map_provider_error)?;
let mut output = StreamingOutput::new(output_format);
output.display_header(&model);
let mut is_first = true;
while let Some(chunk_result) = stream.next().await {
let chunk = chunk_result.map_err(map_provider_error)?;
if !chunk.is_empty() {
output.display_chunk(&chunk, is_first)?;
is_first = false;
}
}
output.display_footer(None);
Ok(())
}
fn map_provider_error(err: ProviderError) -> anyhow::Error {
match err {
ProviderError::InvalidApiKey => {
anyhow::anyhow!("Invalid API key\n\nSet the appropriate environment variable (OPENAI_API_KEY or ANTHROPIC_API_KEY)")
}
ProviderError::RateLimitExceeded { retry_after } => {
if let Some(duration) = retry_after {
anyhow::anyhow!("Rate limit exceeded\n\nRetry after {} seconds", duration.as_secs())
} else {
anyhow::anyhow!("Rate limit exceeded\n\nWait a moment and try again")
}
}
ProviderError::ContextLengthExceeded { tokens, max } => {
anyhow::anyhow!(
"Prompt too long: {} tokens (max: {})\n\nReduce prompt length or use a model with larger context",
tokens, max
)
}
ProviderError::NetworkError(e) => {
anyhow::anyhow!("Network error: {}\n\nCheck your internet connection", e)
}
ProviderError::ModelNotFound { model } => {
anyhow::anyhow!("Model not found: {}\n\nUse --help to see supported models", model)
}
ProviderError::Timeout(duration) => {
anyhow::anyhow!("Request timeout after {:?}\n\nTry again or use a shorter prompt", duration)
}
other => anyhow::anyhow!("{}", other),
}
}
pub async fn execute(args: TestArgs, verbose: bool) -> Result<()> {
if verbose {
tracing::info!("Test command starting with args: {:?}", args);
}
if let Some(temp) = args.temperature {
if !(0.0..=2.0).contains(&temp) {
return Err(anyhow!("Temperature must be between 0.0 and 2.0"));
}
}
if let Some(top_p) = args.top_p {
if !(0.0..=1.0).contains(&top_p) {
return Err(anyhow!("Top-p must be between 0.0 and 1.0"));
}
}
let provider = create_provider(&args.provider, &args.config)
.with_context(|| format!("Failed to initialize {} provider", args.provider))?;
let request = build_completion_request(&args, &provider)?;
if verbose {
tracing::info!("Request: model={}, stream={}", request.model, args.stream);
}
if args.stream {
execute_streaming(&provider, request, args.output_format).await?;
} else {
execute_non_streaming(&provider, request, args.output_format).await?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_temperature_validation() {
assert!((0.0..=2.0).contains(&0.7));
assert!((0.0..=2.0).contains(&0.0));
assert!((0.0..=2.0).contains(&2.0));
assert!(!(0.0..=2.0).contains(&2.1));
assert!(!(0.0..=2.0).contains(&-0.1));
}
#[test]
fn test_top_p_validation() {
assert!((0.0..=1.0).contains(&0.9));
assert!((0.0..=1.0).contains(&0.0));
assert!((0.0..=1.0).contains(&1.0));
assert!(!(0.0..=1.0).contains(&1.1));
assert!(!(0.0..=1.0).contains(&-0.1));
}
}