diff --git a/Cargo.lock b/Cargo.lock index 608789855..4cc9286f5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1714,10 +1714,6 @@ dependencies = [ "async-trait", "base64", "bytes", - "clap", - "cli-table", - "dialoguer", - "fabro-api", "fabro-macros", "fabro-model", "fabro-test", @@ -1725,9 +1721,7 @@ dependencies = [ "futures", "http", "httpmock", - "indicatif", "insta", - "progenitor-client", "rand 0.8.5", "reqwest 0.13.2", "serde", diff --git a/lib/crates/fabro-cli/src/args.rs b/lib/crates/fabro-cli/src/args.rs index b455668c0..299a6c9e4 100644 --- a/lib/crates/fabro-cli/src/args.rs +++ b/lib/crates/fabro-cli/src/args.rs @@ -4,7 +4,6 @@ use std::path::PathBuf; use clap::{Args, Subcommand, ValueEnum}; use fabro_agent::cli::AgentArgs; use fabro_graphviz::render::GraphFormat; -use fabro_llm::cli::ModelsCommand; pub(crate) const LONG_VERSION: &str = concat!( env!("CARGO_PKG_VERSION"), @@ -729,6 +728,35 @@ impl RunsCommands { } } +#[derive(Subcommand)] +pub(crate) enum ModelsCommand { + /// List available models + List { + /// Filter by provider + #[arg(short, long)] + provider: Option, + + /// Search for models matching this string + #[arg(short, long)] + query: Option, + }, + + /// Test model availability by sending a simple prompt + Test { + /// Filter by provider + #[arg(short, long)] + provider: Option, + + /// Test a specific model + #[arg(short, long)] + model: Option, + + /// Run a multi-turn tool-use test (catches reasoning round-trip bugs) + #[arg(long)] + deep: bool, + }, +} + #[derive(Subcommand)] pub(crate) enum Commands { /// Run an agentic coding session diff --git a/lib/crates/fabro-cli/src/commands/model.rs b/lib/crates/fabro-cli/src/commands/model.rs index b46328d82..b12cba26c 100644 --- a/lib/crates/fabro-cli/src/commands/model.rs +++ b/lib/crates/fabro-cli/src/commands/model.rs @@ -1,10 +1,42 @@ -use anyhow::Result; -use fabro_llm::cli::{ModelsCommand, run_models}; +use anyhow::{Context, Result, anyhow, bail}; +use cli_table::format::{Border, Justify, Separator}; +use cli_table::{Cell, CellStruct, Color, Style, Table}; +use fabro_api::{self, types as api_types}; +use fabro_model::{Catalog, Model, Provider}; +use fabro_util::terminal::Styles; +use serde::Serialize; +use serde::de::DeserializeOwned; -use crate::args::GlobalArgs; +use crate::args::{GlobalArgs, ModelsCommand}; use crate::server_client; use crate::user_config; +#[derive(Serialize)] +#[serde(rename_all = "snake_case")] +enum ModelTestResultKind { + Pass, + Fail, + Skip, +} + +#[derive(Serialize)] +struct ModelTestRow { + model: String, + provider: Provider, + result: ModelTestResultKind, + #[serde(skip_serializing_if = "Option::is_none")] + detail: Option, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, +} + +#[derive(Serialize)] +struct ModelTestOutput { + results: Vec, + total: usize, + failures: u32, +} + pub(crate) async fn execute(command: Option, globals: &GlobalArgs) -> Result<()> { let cli_settings = user_config::load_user_settings_with_globals(globals)?; let client = match globals.server_url.as_deref() { @@ -20,3 +52,729 @@ pub(crate) async fn execute(command: Option, globals: &GlobalArgs run_models(command, client, globals.json).await } + +fn format_context_window(tokens: i64) -> String { + let rounded = ((tokens + 500) / 1_000) * 1_000; + if rounded >= 1_000_000 { + format!("{}m", rounded / 1_000_000) + } else if rounded >= 1_000 { + format!("{}k", rounded / 1_000) + } else { + tokens.to_string() + } +} + +fn format_cost(cost: Option) -> String { + match cost { + None => "-".to_string(), + Some(c) => format!("${c:.1}"), + } +} + +fn format_speed(tps: Option) -> String { + match tps { + None => "-".to_string(), + #[allow(clippy::cast_possible_truncation)] + Some(t) => format!("{} tok/s", t as i64), + } +} + +fn color_if(use_color: bool, color: Color) -> Option { + if use_color { Some(color) } else { None } +} + +fn color_choice(use_color: bool) -> cli_table::ColorChoice { + if use_color { + cli_table::ColorChoice::Auto + } else { + cli_table::ColorChoice::Never + } +} + +fn model_row(model: &Model, use_color: bool) -> Vec { + let aliases = model.aliases.join(", "); + let cost = format!( + "{} / {}", + format_cost(model.costs.input_cost_per_mtok), + format_cost(model.costs.output_cost_per_mtok), + ); + vec![ + model.id.clone().cell().bold(use_color), + model + .provider + .cell() + .foreground_color(color_if(use_color, Color::Ansi256(8))), + aliases + .cell() + .foreground_color(color_if(use_color, Color::Ansi256(8))), + format_context_window(model.limits.context_window) + .cell() + .justify(Justify::Right), + cost.cell().justify(Justify::Right), + format_speed(model.estimated_output_tps) + .cell() + .justify(Justify::Right) + .foreground_color(color_if(use_color, Color::Cyan)), + ] +} + +fn models_title(use_color: bool) -> Vec { + vec![ + "MODEL".cell().bold(use_color), + "PROVIDER".cell().bold(use_color), + "ALIASES".cell().bold(use_color), + "CONTEXT".cell().bold(use_color).justify(Justify::Right), + "COST".cell().bold(use_color).justify(Justify::Right), + "SPEED".cell().bold(use_color).justify(Justify::Right), + ] +} + +#[allow(clippy::print_stdout)] +fn print_models_table(models: &[Model], styles: &Styles) { + let use_color = styles.use_color; + let rows: Vec> = models + .iter() + .map(|model| model_row(model, use_color)) + .collect(); + let table = rows + .table() + .title(models_title(use_color)) + .color_choice(color_choice(use_color)) + .border(Border::builder().build()) + .separator(Separator::builder().build()); + println!("{}", table.display().unwrap()); +} + +fn model_test_row_from_status(model: &Model, status: &str, result_color: Color) -> ModelTestRow { + let trimmed = status.trim(); + match result_color { + Color::Green => ModelTestRow { + model: model.id.clone(), + provider: model.provider, + result: ModelTestResultKind::Pass, + detail: None, + error: None, + }, + Color::Yellow => ModelTestRow { + model: model.id.clone(), + provider: model.provider, + result: ModelTestResultKind::Skip, + detail: Some(trimmed.to_string()), + error: None, + }, + _ => ModelTestRow { + model: model.id.clone(), + provider: model.provider, + result: ModelTestResultKind::Fail, + detail: None, + error: Some( + trimmed + .strip_prefix("error: ") + .unwrap_or(trimmed) + .to_string(), + ), + }, + } +} + +fn map_api_error(err: progenitor_client::Error) -> anyhow::Error +where + E: serde::Serialize + std::fmt::Debug, +{ + match err { + progenitor_client::Error::ErrorResponse(response) => { + let status = response.status(); + if let Ok(value) = serde_json::to_value(response.into_inner()) { + if let Some(detail) = value + .get("errors") + .and_then(serde_json::Value::as_array) + .and_then(|errors| errors.first()) + .and_then(|entry| entry.get("detail")) + .and_then(serde_json::Value::as_str) + { + return anyhow!("{detail}"); + } + } + anyhow!("request failed with status {status}") + } + progenitor_client::Error::UnexpectedResponse(response) => { + anyhow!("request failed with status {}", response.status()) + } + other => anyhow!("{other}"), + } +} + +fn convert_type(value: TInput) -> Result +where + TInput: serde::Serialize, + TOutput: DeserializeOwned, +{ + serde_json::from_value(serde_json::to_value(value)?).map_err(Into::into) +} + +async fn fetch_models_from_server( + client: &fabro_api::Client, + provider: Option<&str>, + query: Option<&str>, +) -> Result> { + let mut offset = 0u64; + let mut models = Vec::new(); + + loop { + let mut request = client.list_models().page_limit(100u64).page_offset(offset); + if let Some(provider) = provider { + request = request.provider(provider.to_string()); + } + if let Some(query) = query { + request = request.query(query.to_string()); + } + + let response = request.send().await.map_err(map_api_error)?; + let parsed = response.into_inner(); + let count = parsed.data.len() as u64; + models.extend(convert_type::<_, Vec>(parsed.data)?); + if !parsed.meta.has_more { + break; + } + offset += count; + } + + Ok(models) +} + +async fn test_model_via_server( + client: &fabro_api::Client, + model_id: &str, + mode: Option, +) -> Result { + let mut request = client.test_model().id(model_id.to_string()); + if let Some(mode) = mode { + request = request.mode(mode); + } + let response = request.send().await.map_err(map_api_error)?; + Ok(response.into_inner()) +} + +#[allow(clippy::print_stdout, clippy::print_stderr)] +async fn test_models_via_server( + client: &fabro_api::Client, + provider: Option<&str>, + model: Option<&str>, + deep: bool, + styles: &Styles, + json_output: bool, +) -> Result<()> { + let request_mode = deep.then_some(api_types::ModelTestMode::Deep); + + let use_color = styles.use_color; + let mut title = models_title(use_color); + title.push("RESULT".cell().bold(use_color)); + + let mut rows: Vec> = Vec::new(); + let mut json_rows = Vec::new(); + let mut failures = 0u32; + if let Some(model_id) = model { + if !json_output { + eprint!("Testing {model_id}..."); + } + let result = test_model_via_server(client, model_id, request_mode).await; + if !json_output { + eprintln!(" done"); + } + + let (info, result_color, status) = match result { + Ok(resp) => { + let info = Catalog::builtin() + .get(&resp.model_id) + .cloned() + .with_context(|| { + format!("Unknown model returned by server: {}", resp.model_id) + })?; + if resp.status == api_types::ModelTestResultStatus::Ok { + (info, Color::Green, "ok".to_string()) + } else { + failures += 1; + let message = resp + .error_message + .unwrap_or_else(|| "unknown error".to_string()); + (info, Color::Red, format!("error: {message}")) + } + } + Err(err) if err.to_string().contains("Model not found") => { + bail!("Unknown model: {model_id}"); + } + Err(err) => { + let info = Catalog::builtin() + .get(model_id) + .cloned() + .with_context(|| format!("Unknown model: {model_id}"))?; + failures += 1; + (info, Color::Red, format!("error: {err}")) + } + }; + + let mut row = model_row(&info, use_color); + row.push( + status + .clone() + .cell() + .foreground_color(color_if(use_color, result_color)), + ); + rows.push(row); + json_rows.push(model_test_row_from_status(&info, &status, result_color)); + } else { + let models_to_test = fetch_models_from_server(client, provider, None).await?; + if models_to_test.is_empty() { + bail!("No models found"); + } + + for info in &models_to_test { + if !json_output { + eprint!("Testing {}...", info.id); + } + let result = test_model_via_server(client, &info.id, request_mode).await; + if !json_output { + eprintln!(" done"); + } + + let (result_color, status) = match result { + Ok(resp) if resp.status == api_types::ModelTestResultStatus::Ok => { + (Color::Green, "ok".to_string()) + } + Ok(resp) => { + failures += 1; + let message = resp + .error_message + .unwrap_or_else(|| "unknown error".to_string()); + (Color::Red, format!("error: {message}")) + } + Err(err) => { + failures += 1; + (Color::Red, format!("error: {err}")) + } + }; + + let mut row = model_row(info, use_color); + row.push( + status + .clone() + .cell() + .foreground_color(color_if(use_color, result_color)), + ); + rows.push(row); + json_rows.push(model_test_row_from_status(info, &status, result_color)); + } + } + + if json_output { + println!( + "{}", + serde_json::to_string_pretty(&ModelTestOutput { + total: json_rows.len(), + failures, + results: json_rows, + })? + ); + if failures > 0 { + bail!("{failures} model(s) failed"); + } + return Ok(()); + } + + let table = rows + .table() + .title(title) + .color_choice(color_choice(use_color)) + .border(Border::builder().build()) + .separator(Separator::builder().build()); + println!("{}", table.display()?); + + if failures > 0 { + bail!("{failures} model(s) failed"); + } + + Ok(()) +} + +#[allow(clippy::print_stdout)] +async fn run_models( + command: Option, + client: fabro_api::Client, + json_output: bool, +) -> Result<()> { + let command = command.unwrap_or(ModelsCommand::List { + provider: None, + query: None, + }); + + let styles = Styles::detect_stdout(); + + match command { + ModelsCommand::List { provider, query } => { + let models = + fetch_models_from_server(&client, provider.as_deref(), query.as_deref()).await?; + + if json_output { + println!("{}", serde_json::to_string_pretty(&models)?); + } else { + print_models_table(&models, &styles); + } + } + ModelsCommand::Test { + provider, + model, + deep, + } => { + test_models_via_server( + &client, + provider.as_deref(), + model.as_deref(), + deep, + &styles, + json_output, + ) + .await?; + } + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use fabro_model::{ModelCosts, ModelFeatures, ModelLimits}; + + fn test_http_client() -> reqwest::Client { + reqwest::Client::builder().no_proxy().build().unwrap() + } + + fn test_api_client(base_url: &str) -> fabro_api::Client { + fabro_api::Client::new_with_client(base_url, test_http_client()) + } + + fn test_model_json(id: &str, provider: Provider) -> serde_json::Value { + serde_json::to_value(Model { + id: id.to_string(), + provider, + family: "test".to_string(), + display_name: format!("{id} display"), + limits: ModelLimits { + context_window: 128_000, + max_output: Some(4096), + }, + training: None, + knowledge_cutoff: None, + features: ModelFeatures { + tools: true, + vision: false, + reasoning: false, + effort: false, + }, + costs: ModelCosts { + input_cost_per_mtok: Some(1.0), + output_cost_per_mtok: Some(2.0), + cache_input_cost_per_mtok: None, + }, + estimated_output_tps: Some(100.0), + aliases: vec!["tm".to_string()], + default: false, + }) + .unwrap() + } + + #[test] + fn format_context_window_millions() { + assert_eq!(format_context_window(1_000_000), "1m"); + } + + #[test] + fn format_context_window_thousands() { + assert_eq!(format_context_window(128_000), "128k"); + } + + #[test] + fn format_context_window_small() { + assert_eq!(format_context_window(400), "400"); + } + + #[test] + fn format_context_window_rounds_up() { + assert_eq!(format_context_window(1500), "2k"); + } + + #[test] + fn format_context_window_rounds_down() { + assert_eq!(format_context_window(1499), "1k"); + } + + #[test] + fn format_context_window_zero() { + assert_eq!(format_context_window(0), "0"); + } + + #[test] + fn format_cost_none() { + assert_eq!(format_cost(None), "-"); + } + + #[test] + fn format_cost_some() { + assert_eq!(format_cost(Some(3.0)), "$3.0"); + } + + #[test] + fn format_cost_fractional() { + assert_eq!(format_cost(Some(15.75)), "$15.8"); + } + + #[test] + fn format_speed_none() { + assert_eq!(format_speed(None), "-"); + } + + #[test] + fn format_speed_some() { + assert_eq!(format_speed(Some(85.5)), "85 tok/s"); + } + + #[tokio::test] + async fn test_model_via_server_parses_ok() { + let server = httpmock::MockServer::start_async().await; + server + .mock_async(|when, then| { + when.method("POST").path("/api/v1/models/test-model/test"); + then.status(200) + .header("Content-Type", "application/json") + .body( + serde_json::json!({ + "model_id": "test-model", + "status": "ok" + }) + .to_string(), + ); + }) + .await; + + let client = test_api_client(&server.url("")); + let response = test_model_via_server(&client, "test-model", None) + .await + .unwrap(); + + assert_eq!(response.status, api_types::ModelTestResultStatus::Ok); + assert!(response.error_message.is_none()); + } + + #[tokio::test] + async fn test_model_via_server_passes_mode_and_parses_error() { + let server = httpmock::MockServer::start_async().await; + server + .mock_async(|when, then| { + when.method("POST") + .path("/api/v1/models/test-model/test") + .query_param("mode", "deep"); + then.status(200) + .header("Content-Type", "application/json") + .body( + serde_json::json!({ + "model_id": "test-model", + "status": "error", + "error_message": "timeout" + }) + .to_string(), + ); + }) + .await; + + let client = test_api_client(&server.url("")); + let response = + test_model_via_server(&client, "test-model", Some(api_types::ModelTestMode::Deep)) + .await + .unwrap(); + + assert_eq!(response.status, api_types::ModelTestResultStatus::Error); + assert_eq!(response.error_message.as_deref(), Some("timeout")); + } + + #[tokio::test] + async fn test_model_via_server_404() { + let server = httpmock::MockServer::start_async().await; + server + .mock_async(|when, then| { + when.method("POST").path("/api/v1/models/bad-model/test"); + then.status(404) + .header("Content-Type", "application/json") + .body( + serde_json::json!({ + "errors": [{"status": "404", "title": "Not Found", "detail": "Model not found"}] + }) + .to_string(), + ); + }) + .await; + + let client = test_api_client(&server.url("")); + let result = test_model_via_server(&client, "bad-model", None).await; + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("Model not found")); + } + + #[tokio::test] + async fn fetch_models_from_server_parses_response() { + let server = httpmock::MockServer::start_async().await; + let mock = server + .mock_async(|when, then| { + when.method("GET") + .path("/api/v1/models") + .query_param("page[limit]", "100") + .query_param("page[offset]", "0"); + then.status(200) + .header("Content-Type", "application/json") + .body( + serde_json::json!({ + "data": [test_model_json("test-model", Provider::Anthropic)], + "meta": { "has_more": false } + }) + .to_string(), + ); + }) + .await; + + let client = test_api_client(&server.url("")); + let models = fetch_models_from_server(&client, None, None).await.unwrap(); + + mock.assert_async().await; + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "test-model"); + assert_eq!(models[0].provider, Provider::Anthropic); + } + + #[tokio::test] + async fn fetch_models_from_server_filters_by_provider() { + let server = httpmock::MockServer::start_async().await; + server + .mock_async(|when, then| { + when.method("GET") + .path("/api/v1/models") + .query_param("page[limit]", "100") + .query_param("page[offset]", "0") + .query_param("provider", "anthropic"); + then.status(200) + .header("Content-Type", "application/json") + .body( + serde_json::json!({ + "data": [test_model_json("model-a", Provider::Anthropic)], + "meta": { "has_more": false } + }) + .to_string(), + ); + }) + .await; + + let client = test_api_client(&server.url("")); + let models = fetch_models_from_server(&client, Some("anthropic"), None) + .await + .unwrap(); + + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "model-a"); + } + + #[tokio::test] + async fn fetch_models_from_server_passes_query_param() { + let server = httpmock::MockServer::start_async().await; + let mock = server + .mock_async(|when, then| { + when.method("GET") + .path("/api/v1/models") + .query_param("page[limit]", "100") + .query_param("page[offset]", "0") + .query_param("query", "sonnet"); + then.status(200) + .header("Content-Type", "application/json") + .body( + serde_json::json!({ + "data": [test_model_json("claude-sonnet-4-5", Provider::Anthropic)], + "meta": { "has_more": false } + }) + .to_string(), + ); + }) + .await; + + let client = test_api_client(&server.url("")); + let models = fetch_models_from_server(&client, None, Some("sonnet")) + .await + .unwrap(); + + mock.assert_async().await; + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "claude-sonnet-4-5"); + } + + #[tokio::test] + async fn fetch_models_from_server_follows_pagination() { + let server = httpmock::MockServer::start_async().await; + let first_page = server + .mock_async(|when, then| { + when.method("GET") + .path("/api/v1/models") + .query_param("page[limit]", "100") + .query_param("page[offset]", "0"); + then.status(200) + .header("Content-Type", "application/json") + .body( + serde_json::json!({ + "data": [test_model_json("model-a", Provider::Anthropic)], + "meta": { "has_more": true } + }) + .to_string(), + ); + }) + .await; + let second_page = server + .mock_async(|when, then| { + when.method("GET") + .path("/api/v1/models") + .query_param("page[limit]", "100") + .query_param("page[offset]", "1"); + then.status(200) + .header("Content-Type", "application/json") + .body( + serde_json::json!({ + "data": [test_model_json("model-b", Provider::OpenAi)], + "meta": { "has_more": false } + }) + .to_string(), + ); + }) + .await; + + let client = test_api_client(&server.url("")); + let models = fetch_models_from_server(&client, None, None).await.unwrap(); + + first_page.assert_async().await; + second_page.assert_async().await; + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "model-a"); + assert_eq!(models[1].id, "model-b"); + } + + #[tokio::test] + async fn fetch_models_from_server_error_on_failure() { + let server = httpmock::MockServer::start_async().await; + server + .mock_async(|when, then| { + when.method("GET") + .path("/api/v1/models") + .query_param("page[limit]", "100") + .query_param("page[offset]", "0"); + then.status(500).body("internal error"); + }) + .await; + + let client = test_api_client(&server.url("")); + let result = fetch_models_from_server(&client, None, None).await; + assert!(result.is_err()); + } +} diff --git a/lib/crates/fabro-llm/Cargo.toml b/lib/crates/fabro-llm/Cargo.toml index aa19f37ac..7c62ca633 100644 --- a/lib/crates/fabro-llm/Cargo.toml +++ b/lib/crates/fabro-llm/Cargo.toml @@ -31,13 +31,7 @@ reqwest.workspace = true base64.workspace = true bytes.workspace = true tokio-util.workspace = true -cli-table.workspace = true -indicatif.workspace = true -clap.workspace = true -dialoguer.workspace = true tracing.workspace = true -progenitor-client = "0.13" -fabro-api = { path = "../fabro-api" } fabro-model = { path = "../fabro-model" } fabro-util = { path = "../fabro-util" } diff --git a/lib/crates/fabro-llm/src/cli.rs b/lib/crates/fabro-llm/src/cli.rs deleted file mode 100644 index 2465e91e4..000000000 --- a/lib/crates/fabro-llm/src/cli.rs +++ /dev/null @@ -1,1745 +0,0 @@ -use std::io::{self, IsTerminal, Read, Write}; - -use dialoguer::console::Term; -use dialoguer::theme::ColorfulTheme; - -use anyhow::{Context, Result, anyhow, bail}; -use clap::{Args, Subcommand}; -use cli_table::format::{Border, Justify, Separator}; -use cli_table::{Cell, CellStruct, Color, Style, Table}; -use futures::StreamExt; -use serde::Serialize; -use serde::de::DeserializeOwned; -use tokio::task; - -use fabro_api::{self, types as api_types}; -use fabro_util::terminal::Styles; - -use fabro_model::{Catalog, Model, Provider}; - -use crate::generate::{self, GenerateParams}; -use crate::types::{Message, StreamEvent, Usage}; - -pub struct ServerConnection { - pub client: reqwest::Client, - pub base_url: String, -} - -#[derive(Serialize)] -#[serde(rename_all = "snake_case")] -enum ModelTestResultKind { - Pass, - Fail, - Skip, -} - -#[derive(Serialize)] -struct ModelTestRow { - model: String, - provider: Provider, - result: ModelTestResultKind, - #[serde(skip_serializing_if = "Option::is_none")] - detail: Option, - #[serde(skip_serializing_if = "Option::is_none")] - error: Option, -} - -#[derive(Serialize)] -struct ModelTestOutput { - results: Vec, - total: usize, - failures: u32, -} - -#[derive(Args)] -pub struct PromptArgs { - /// The prompt text (also accepts stdin) - pub prompt: Option, - - /// Model to use - #[arg(short, long)] - pub model: Option, - - /// System prompt - #[arg(short, long)] - pub system: Option, - - /// Do not stream output - #[arg(long)] - pub no_stream: bool, - - /// Show token usage - #[arg(short, long)] - pub usage: bool, - - /// JSON schema for structured output (inline JSON string) - #[arg(short = 'S', long)] - pub schema: Option, - - /// key=value options (temperature, `max_tokens`, `top_p`) - #[arg(short, long, value_parser = parse_option)] - pub option: Vec<(String, String)>, -} - -#[derive(Subcommand)] -pub enum ModelsCommand { - /// List available models - List { - /// Filter by provider - #[arg(short, long)] - provider: Option, - - /// Search for models matching this string - #[arg(short, long)] - query: Option, - }, - - /// Test model availability by sending a simple prompt - Test { - /// Filter by provider - #[arg(short, long)] - provider: Option, - - /// Test a specific model - #[arg(short, long)] - model: Option, - - /// Run a multi-turn tool-use test (catches reasoning round-trip bugs) - #[arg(long)] - deep: bool, - }, -} - -fn parse_option(s: &str) -> Result<(String, String), String> { - let (key, value) = s - .split_once('=') - .ok_or_else(|| format!("expected key=value, got {s}"))?; - Ok((key.to_string(), value.to_string())) -} - -fn format_context_window(tokens: i64) -> String { - let rounded = ((tokens + 500) / 1_000) * 1_000; - if rounded >= 1_000_000 { - format!("{}m", rounded / 1_000_000) - } else if rounded >= 1_000 { - format!("{}k", rounded / 1_000) - } else { - tokens.to_string() - } -} - -fn format_cost(cost: Option) -> String { - match cost { - None => "-".to_string(), - Some(c) => format!("${c:.1}"), - } -} - -fn format_speed(tps: Option) -> String { - match tps { - None => "-".to_string(), - #[allow(clippy::cast_possible_truncation)] // f64-to-integer: fractional loss is fine - Some(t) => format!("{} tok/s", t as i64), - } -} - -fn color_if(use_color: bool, color: Color) -> Option { - if use_color { Some(color) } else { None } -} - -fn color_choice(use_color: bool) -> cli_table::ColorChoice { - if use_color { - cli_table::ColorChoice::Auto - } else { - cli_table::ColorChoice::Never - } -} - -fn model_row(model: &Model, use_color: bool) -> Vec { - let aliases = model.aliases.join(", "); - let cost = format!( - "{} / {}", - format_cost(model.costs.input_cost_per_mtok), - format_cost(model.costs.output_cost_per_mtok), - ); - vec![ - model.id.clone().cell().bold(use_color), - model - .provider - .cell() - .foreground_color(color_if(use_color, Color::Ansi256(8))), - aliases - .cell() - .foreground_color(color_if(use_color, Color::Ansi256(8))), - format_context_window(model.limits.context_window) - .cell() - .justify(Justify::Right), - cost.cell().justify(Justify::Right), - format_speed(model.estimated_output_tps) - .cell() - .justify(Justify::Right) - .foreground_color(color_if(use_color, Color::Cyan)), - ] -} - -fn models_title(use_color: bool) -> Vec { - vec![ - "MODEL".cell().bold(use_color), - "PROVIDER".cell().bold(use_color), - "ALIASES".cell().bold(use_color), - "CONTEXT".cell().bold(use_color).justify(Justify::Right), - "COST".cell().bold(use_color).justify(Justify::Right), - "SPEED".cell().bold(use_color).justify(Justify::Right), - ] -} - -#[allow(clippy::print_stdout)] -fn print_models_table(models: &[Model], s: &Styles) { - let use_color = s.use_color; - let rows: Vec> = models.iter().map(|m| model_row(m, use_color)).collect(); - let table = rows - .table() - .title(models_title(use_color)) - .color_choice(color_choice(use_color)) - .border(Border::builder().build()) - .separator(Separator::builder().build()); - println!("{}", table.display().unwrap()); -} - -fn read_stdin_prompt() -> Option { - let stdin = io::stdin(); - if stdin.is_terminal() { - return None; - } - let mut buf = String::new(); - stdin.lock().read_to_string(&mut buf).ok()?; - let trimmed = buf.trim(); - if trimmed.is_empty() { - None - } else { - Some(trimmed.to_string()) - } -} - -fn resolve_prompt(arg: Option, stdin: Option) -> Result { - match (stdin, arg) { - (Some(s), Some(a)) => Ok(format!("{s}\n{a}")), - (Some(s), None) => Ok(s), - (None, Some(a)) => Ok(a), - (None, None) => { - bail!("Error: no prompt provided. Pass a prompt as an argument or pipe text via stdin.") - } - } -} - -/// Returns (`model_id`, provider) from the catalog, falling back to the first catalog model. -fn resolve_model(model_arg: Option) -> (String, Option) { - let raw = model_arg.unwrap_or_else(|| { - Catalog::builtin() - .list(None) - .first() - .map_or_else(|| "claude-sonnet-4-5".to_string(), |m| m.id.clone()) - }); - match Catalog::builtin().get(&raw) { - Some(info) => (info.id.clone(), Some(info.provider.to_string())), - None => (raw, None), - } -} - -fn apply_options( - mut params: GenerateParams, - options: &[(String, String)], -) -> Result { - let mut provider_opts = serde_json::Map::new(); - - for (key, value) in options { - match key.as_str() { - "temperature" => { - let v: f64 = value - .parse() - .with_context(|| format!("invalid temperature value: {value}"))?; - params = params.temperature(v); - } - "max_tokens" => { - let v: i64 = value - .parse() - .with_context(|| format!("invalid max_tokens value: {value}"))?; - params = params.max_tokens(v); - } - "top_p" => { - let v: f64 = value - .parse() - .with_context(|| format!("invalid top_p value: {value}"))?; - params = params.top_p(v); - } - _ => { - provider_opts.insert(key.clone(), serde_json::Value::String(value.clone())); - } - } - } - - if !provider_opts.is_empty() { - params = params.provider_options(serde_json::Value::Object(provider_opts)); - } - - Ok(params) -} - -#[allow(clippy::print_stderr)] -fn print_usage(usage: &Usage) { - eprintln!( - "Tokens: {} input, {} output, {} total", - usage.input_tokens, usage.output_tokens, usage.total_tokens - ); -} - -#[derive(Args)] -pub struct ChatArgs { - /// Model to use - #[arg(short, long)] - pub model: Option, - - /// System prompt - #[arg(short, long)] - pub system: Option, -} - -#[allow(clippy::print_stdout, clippy::print_stderr)] -pub async fn run_chat(args: ChatArgs) -> Result<()> { - let (model_id, provider) = resolve_model(args.model); - eprintln!("Using model: {model_id}"); - - let mut messages: Vec = Vec::new(); - let is_tty = io::stdin().is_terminal(); - - loop { - let line = if is_tty { - let result = task::spawn_blocking(|| { - dialoguer::Input::::with_theme(&ColorfulTheme::default()) - .with_prompt(">") - .interact_on(&Term::stderr()) - }) - .await?; - match result { - Ok(line) => line, - Err(_) => break, - } - } else { - eprint!("> "); - io::stderr().flush()?; - let mut buf = String::new(); - if io::stdin().read_line(&mut buf)? == 0 { - break; - } - buf.trim_end().to_string() - }; - - let trimmed = line.trim(); - if trimmed.is_empty() { - continue; - } - - messages.push(Message::user(trimmed)); - - let mut params = GenerateParams::new(&model_id) - .messages(messages.clone()) - .max_tokens(4096); - if let Some(ref p) = provider { - params = params.provider(p); - } - if let Some(ref sys) = args.system { - params = params.system(sys); - } - - let mut stream_result = generate::stream(params).await?; - let mut full_text = String::new(); - while let Some(event) = stream_result.next().await { - if let StreamEvent::TextDelta { delta, .. } = event? { - print!("{delta}"); - full_text.push_str(&delta); - } - } - println!(); - - messages.push(Message::assistant(full_text)); - } - - Ok(()) -} - -#[allow(clippy::print_stdout, clippy::print_stderr)] -pub async fn run_prompt(args: PromptArgs, json_output: bool) -> Result<()> { - let stdin_prompt = read_stdin_prompt(); - let prompt_text = resolve_prompt(args.prompt, stdin_prompt)?; - let (model_id, provider) = resolve_model(args.model); - - if !json_output { - eprintln!("Using model: {model_id}"); - } - - let mut params = GenerateParams::new(&model_id).prompt(&prompt_text); - if let Some(p) = provider { - params = params.provider(&p); - } - if let Some(sys) = args.system { - params = params.system(&sys); - } - params = apply_options(params, &args.option)?; - - let schema: Option = match &args.schema { - Some(s) => Some(serde_json::from_str(s).context("--schema must be valid JSON")?), - None => None, - }; - - match (args.no_stream, schema) { - (true, Some(schema)) => { - let result = generate::generate_object(params, schema).await?; - let object = result.output.as_ref().unwrap_or(&serde_json::Value::Null); - println!("{}", serde_json::to_string_pretty(object)?); - if args.usage { - print_usage(&result.usage); - } - } - (true, None) => { - let result = generate::generate(params).await?; - if json_output { - println!( - "{}", - serde_json::to_string_pretty(&serde_json::json!({ - "response": result.text(), - "model": model_id, - "usage": result.usage, - }))? - ); - } else { - print!("{}", result.text()); - if args.usage { - print_usage(&result.usage); - } - } - } - (false, Some(schema)) => { - let mut stream_result = generate::stream_object(params, schema).await?; - while let Some(event) = stream_result.next().await { - event?; - } - if let Some(object) = stream_result.object() { - println!("{}", serde_json::to_string_pretty(object)?); - } - } - (false, None) => { - let mut stream_result = generate::stream(params).await?; - let mut full_text = String::new(); - while let Some(event) = stream_result.next().await { - if let StreamEvent::TextDelta { delta, .. } = event? { - if json_output { - full_text.push_str(&delta); - } else { - print!("{delta}"); - } - } - } - if json_output { - let usage = stream_result - .response() - .map(|response| response.usage.clone()); - let mut value = serde_json::Map::new(); - value.insert("response".to_string(), full_text.into()); - value.insert("model".to_string(), model_id.into()); - if let Some(usage) = usage { - value.insert("usage".to_string(), serde_json::to_value(usage)?); - } - println!("{}", serde_json::to_string_pretty(&value)?); - } else { - println!(); - if args.usage { - if let Some(response) = stream_result.response() { - print_usage(&response.usage); - } - } - } - } - } - - Ok(()) -} - -#[allow(clippy::print_stdout, clippy::print_stderr)] -pub async fn run_prompt_via_server( - args: PromptArgs, - server: &ServerConnection, - json_output: bool, -) -> Result<()> { - let stdin_prompt = read_stdin_prompt(); - let prompt_text = resolve_prompt(args.prompt, stdin_prompt)?; - - // Extract known options - let mut temperature: Option = None; - let mut max_tokens: Option = None; - let mut top_p: Option = None; - for (key, value) in &args.option { - match key.as_str() { - "temperature" => { - temperature = Some( - value - .parse() - .with_context(|| format!("invalid temperature value: {value}"))?, - ); - } - "max_tokens" => { - max_tokens = Some( - value - .parse() - .with_context(|| format!("invalid max_tokens value: {value}"))?, - ); - } - "top_p" => { - top_p = Some( - value - .parse() - .with_context(|| format!("invalid top_p value: {value}"))?, - ); - } - _ => {} - } - } - - let schema: Option = match &args.schema { - Some(s) => Some(serde_json::from_str(s).context("--schema must be valid JSON")?), - None => None, - }; - - // Force non-streaming for structured output - let use_stream = !args.no_stream && schema.is_none(); - - let mut body = serde_json::json!({ - "messages": [{"role": "user", "content": [{"kind": "text", "data": prompt_text}]}], - "stream": use_stream, - }); - if let Some(ref model) = args.model { - body["model"] = serde_json::Value::String(model.clone()); - } - if let Some(ref system) = args.system { - body["system"] = serde_json::Value::String(system.clone()); - } - if let Some(ref schema) = schema { - body["schema"] = schema.clone(); - } - if let Some(t) = temperature { - body["temperature"] = serde_json::json!(t); - } - if let Some(m) = max_tokens { - body["max_tokens"] = serde_json::json!(m); - } - if let Some(t) = top_p { - body["top_p"] = serde_json::json!(t); - } - - let url = format!("{}/completions", server.base_url); - - if use_stream { - let response = server - .client - .post(&url) - .json(&body) - .send() - .await - .with_context(|| format!("Failed to connect to server at {}", server.base_url))?; - - let status = response.status(); - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - bail!("Server returned {status}: {text}"); - } - - let show_usage = args.usage; - let mut output_usage: Option = None; - let mut output_model = args.model.clone(); - let mut full_text = String::new(); - - parse_sse_frames(response, |event_type, data| { - if event_type == "stream_event" { - if let Ok(event) = serde_json::from_str::(data) { - match event { - StreamEvent::TextDelta { delta, .. } => { - if json_output { - full_text.push_str(&delta); - } else { - print!("{delta}"); - let _ = io::stdout().flush(); - } - } - StreamEvent::Finish { - usage, response, .. - } => { - output_usage = Some(usage); - if output_model.is_none() { - output_model = Some(response.model.clone()); - } - } - StreamEvent::Error { error, .. } => { - bail!("Server error: {error}"); - } - _ => {} - } - } - } - Ok(true) - }) - .await?; - if json_output { - let mut value = serde_json::Map::new(); - value.insert("response".to_string(), full_text.into()); - if let Some(model) = output_model { - value.insert("model".to_string(), model.into()); - } - if let Some(usage) = output_usage { - value.insert("usage".to_string(), serde_json::to_value(usage)?); - } - println!("{}", serde_json::to_string_pretty(&value)?); - } else { - println!(); - if show_usage { - if let Some(usage) = output_usage { - print_usage(&usage); - } - } - } - } else { - // Non-streaming - let response = server - .client - .post(&url) - .json(&body) - .send() - .await - .with_context(|| format!("Failed to connect to server at {}", server.base_url))?; - - let status = response.status(); - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - bail!("Server returned {status}: {text}"); - } - - let result: serde_json::Value = response - .json() - .await - .context("Failed to parse completion response")?; - - if schema.is_some() { - if let Some(output) = result.get("output") { - println!("{}", serde_json::to_string_pretty(output)?); - } else { - // Extract text from message.content parts - print_message_text(&result["message"]); - } - } else if json_output { - let mut value = serde_json::Map::new(); - value.insert( - "response".to_string(), - extract_message_text(&result["message"]).into(), - ); - let model = result - .get("model") - .and_then(serde_json::Value::as_str) - .map(ToOwned::to_owned) - .or(args.model.clone()); - if let Some(model) = model { - value.insert("model".to_string(), model.into()); - } - if let Some(usage) = result.get("usage") { - value.insert("usage".to_string(), usage.clone()); - } - println!("{}", serde_json::to_string_pretty(&value)?); - } else { - print_message_text(&result["message"]); - } - - if args.usage && !json_output { - let input = result["usage"]["input_tokens"].as_i64().unwrap_or(0); - let output = result["usage"]["output_tokens"].as_i64().unwrap_or(0); - eprintln!( - "Tokens: {} input, {} output, {} total", - input, - output, - input + output - ); - } - } - - Ok(()) -} - -/// Extract and print text from a CompletionMessage JSON value. -#[allow(clippy::print_stdout)] -fn print_message_text(message: &serde_json::Value) { - print!("{}", extract_message_text(message)); -} - -fn extract_message_text(message: &serde_json::Value) -> String { - let mut text = String::new(); - if let Some(content) = message["content"].as_array() { - for part in content { - if part["kind"].as_str() == Some("text") { - if let Some(part_text) = part["data"].as_str() { - text.push_str(part_text); - } - } - } - } - text -} - -fn model_test_row_from_status(model: &Model, status: &str, result_color: Color) -> ModelTestRow { - let trimmed = status.trim(); - match result_color { - Color::Green => ModelTestRow { - model: model.id.clone(), - provider: model.provider, - result: ModelTestResultKind::Pass, - detail: None, - error: None, - }, - Color::Yellow => ModelTestRow { - model: model.id.clone(), - provider: model.provider, - result: ModelTestResultKind::Skip, - detail: Some( - trimmed - .strip_prefix("deep: skipped (") - .and_then(|rest| rest.strip_suffix(')')) - .or_else(|| { - trimmed - .strip_prefix("deep: ok (") - .and_then(|rest| rest.strip_suffix(')')) - }) - .unwrap_or(trimmed) - .to_string(), - ), - error: None, - }, - _ => ModelTestRow { - model: model.id.clone(), - provider: model.provider, - result: ModelTestResultKind::Fail, - detail: None, - error: Some( - trimmed - .strip_prefix("error: ") - .or_else(|| trimmed.strip_prefix("deep: error: ")) - .or_else(|| { - trimmed - .strip_prefix("deep: fail (") - .and_then(|rest| rest.strip_suffix(')')) - }) - .unwrap_or(trimmed) - .to_string(), - ), - }, - } -} - -/// Parse SSE frames from a server response, calling `on_frame` for each complete frame. -/// -/// Each frame provides `(event_type, data)`. -async fn parse_sse_frames( - response: reqwest::Response, - mut on_frame: impl FnMut(&str, &str) -> Result, -) -> Result<()> { - let mut stream = response.bytes_stream(); - let mut buffer = String::new(); - - while let Some(chunk) = stream.next().await { - let chunk = chunk.context("Error reading stream")?; - buffer.push_str(&String::from_utf8_lossy(&chunk)); - - while let Some(pos) = buffer.find("\n\n") { - let frame = buffer[..pos].to_string(); - buffer = buffer[pos + 2..].to_string(); - - let mut event_type = String::new(); - let mut data = String::new(); - for line in frame.lines() { - if let Some(val) = line.strip_prefix("event: ") { - event_type = val.to_string(); - } else if let Some(val) = line.strip_prefix("data: ") { - data = val.to_string(); - } - } - - if !on_frame(&event_type, &data)? { - return Ok(()); - } - } - } - Ok(()) -} - -/// Stream session SSE events, printing text deltas to stdout in real-time. -#[allow(clippy::print_stdout)] -async fn stream_session_text(response: reqwest::Response) -> Result<()> { - parse_sse_frames(response, |event_type, data| { - match event_type { - "content_delta" => { - if let Ok(parsed) = serde_json::from_str::(data) { - if let Some(delta) = parsed["delta"].as_str() { - print!("{delta}"); - let _ = io::stdout().flush(); - } - } - } - "done" => { - println!(); - return Ok(false); - } - "error" => { - if let Ok(parsed) = serde_json::from_str::(data) { - let msg = parsed["message"].as_str().unwrap_or("Unknown error"); - bail!("Server error: {msg}"); - } - } - _ => {} - } - Ok(true) - }) - .await -} - -#[allow(clippy::print_stderr)] -pub async fn run_chat_via_server(args: ChatArgs, server: &ServerConnection) -> Result<()> { - let is_tty = io::stdin().is_terminal(); - let mut session_id: Option = None; - - loop { - let line = if is_tty { - let result = task::spawn_blocking(|| { - dialoguer::Input::::with_theme(&ColorfulTheme::default()) - .with_prompt(">") - .interact_on(&Term::stderr()) - }) - .await?; - match result { - Ok(line) => line, - Err(_) => break, - } - } else { - eprint!("> "); - io::stderr().flush()?; - let mut buf = String::new(); - if io::stdin().read_line(&mut buf)? == 0 { - break; - } - buf.trim_end().to_string() - }; - - let trimmed = line.trim(); - if trimmed.is_empty() { - continue; - } - - if let Some(sid) = &session_id { - // Subsequent messages: send message - let body = serde_json::json!({ "content": trimmed }); - let url = format!("{}/sessions/{sid}/messages", server.base_url); - let response = server - .client - .post(&url) - .json(&body) - .send() - .await - .with_context(|| format!("Failed to connect to server at {}", server.base_url))?; - - let status = response.status(); - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - bail!("Server returned {status}: {text}"); - } - - // Stream events - let events_url = format!("{}/sessions/{sid}/events", server.base_url); - let events_response = server - .client - .get(&events_url) - .send() - .await - .context("Failed to connect to event stream")?; - - if !events_response.status().is_success() { - let text = events_response.text().await.unwrap_or_default(); - bail!("Event stream returned error: {text}"); - } - - stream_session_text(events_response).await?; - } else { - // First message: create session - let mut body = serde_json::json!({ "content": trimmed }); - if let Some(ref model) = args.model { - body["model"] = serde_json::Value::String(model.clone()); - } - if let Some(ref system) = args.system { - body["system"] = serde_json::Value::String(system.clone()); - } - - let url = format!("{}/sessions", server.base_url); - let response = server - .client - .post(&url) - .json(&body) - .send() - .await - .with_context(|| format!("Failed to connect to server at {}", server.base_url))?; - - let status = response.status(); - if !status.is_success() { - let text = response.text().await.unwrap_or_default(); - bail!("Server returned {status}: {text}"); - } - - let create_resp: serde_json::Value = response - .json() - .await - .context("Failed to parse session creation response")?; - - let sid = create_resp["id"] - .as_str() - .context("Missing session id in response")? - .to_string(); - let model_id = create_resp["model"]["id"].as_str().unwrap_or("unknown"); - eprintln!("Using model: {model_id}"); - - // Stream events - let events_url = format!("{}/sessions/{sid}/events", server.base_url); - let events_response = server - .client - .get(&events_url) - .send() - .await - .context("Failed to connect to event stream")?; - - if !events_response.status().is_success() { - let text = events_response.text().await.unwrap_or_default(); - bail!("Event stream returned error: {text}"); - } - - stream_session_text(events_response).await?; - session_id = Some(sid); - } - } - - Ok(()) -} - -fn map_api_error(err: progenitor_client::Error) -> anyhow::Error -where - E: serde::Serialize + std::fmt::Debug, -{ - match err { - progenitor_client::Error::ErrorResponse(response) => { - if let Ok(value) = serde_json::to_value(response.into_inner()) { - if let Some(detail) = value - .get("errors") - .and_then(serde_json::Value::as_array) - .and_then(|errors| errors.first()) - .and_then(|entry| entry.get("detail")) - .and_then(serde_json::Value::as_str) - { - return anyhow!("{detail}"); - } - } - anyhow!("request failed") - } - progenitor_client::Error::UnexpectedResponse(response) => { - anyhow!("request failed with status {}", response.status()) - } - other => anyhow!("{other}"), - } -} - -fn convert_type(value: TInput) -> Result -where - TInput: serde::Serialize, - TOutput: DeserializeOwned, -{ - serde_json::from_value(serde_json::to_value(value)?).map_err(Into::into) -} - -async fn fetch_models_from_server( - client: &fabro_api::Client, - provider: Option<&str>, - query: Option<&str>, -) -> Result> { - let mut offset = 0u64; - let mut models = Vec::new(); - - loop { - let mut request = client.list_models().page_limit(100u64).page_offset(offset); - if let Some(provider) = provider { - request = request.provider(provider.to_string()); - } - if let Some(query) = query { - request = request.query(query.to_string()); - } - - let response = request.send().await.map_err(map_api_error)?; - let parsed = response.into_inner(); - let count = parsed.data.len() as u64; - models.extend(convert_type::<_, Vec>(parsed.data)?); - if !parsed.meta.has_more { - break; - } - offset += count; - } - - Ok(models) -} - -async fn test_model_via_server( - client: &fabro_api::Client, - model_id: &str, - mode: Option, -) -> Result { - let mut request = client.test_model().id(model_id.to_string()); - if let Some(mode) = mode { - request = request.mode(mode); - } - let response = request.send().await.map_err(map_api_error)?; - Ok(response.into_inner()) -} - -#[allow(clippy::print_stdout, clippy::print_stderr)] -async fn test_models_via_server( - client: &fabro_api::Client, - provider: Option<&str>, - model: Option<&str>, - deep: bool, - s: &Styles, - json_output: bool, -) -> Result<()> { - let request_mode = deep.then_some(api_types::ModelTestMode::Deep); - - let use_color = s.use_color; - let mut title = models_title(use_color); - title.push("RESULT".cell().bold(use_color)); - - let mut rows: Vec> = Vec::new(); - let mut json_rows = Vec::new(); - let mut failures = 0u32; - if let Some(model_id) = model { - if !json_output { - eprint!("Testing {model_id}..."); - } - let result = test_model_via_server(client, model_id, request_mode).await; - if !json_output { - eprintln!(" done"); - } - - let (info, result_color, status) = match result { - Ok(resp) => { - let info = Catalog::builtin() - .get(&resp.model_id) - .cloned() - .with_context(|| { - format!("Unknown model returned by server: {}", resp.model_id) - })?; - if resp.status == api_types::ModelTestResultStatus::Ok { - (info, Color::Green, "ok".to_string()) - } else { - failures += 1; - let message = resp - .error_message - .unwrap_or_else(|| "unknown error".to_string()); - (info, Color::Red, format!("error: {message}")) - } - } - Err(err) if err.to_string().contains("Model not found") => { - bail!("Unknown model: {model_id}"); - } - Err(err) => { - let info = Catalog::builtin() - .get(model_id) - .cloned() - .with_context(|| format!("Unknown model: {model_id}"))?; - failures += 1; - (info, Color::Red, format!("error: {err}")) - } - }; - - let mut row = model_row(&info, use_color); - row.push( - status - .clone() - .cell() - .foreground_color(color_if(use_color, result_color)), - ); - rows.push(row); - json_rows.push(model_test_row_from_status(&info, &status, result_color)); - } else { - let models_to_test = fetch_models_from_server(client, provider, None).await?; - if models_to_test.is_empty() { - bail!("No models found"); - } - - for info in &models_to_test { - if !json_output { - eprint!("Testing {}...", info.id); - } - let result = test_model_via_server(client, &info.id, request_mode).await; - if !json_output { - eprintln!(" done"); - } - - let (result_color, status) = match result { - Ok(resp) if resp.status == api_types::ModelTestResultStatus::Ok => { - (Color::Green, "ok".to_string()) - } - Ok(resp) => { - failures += 1; - let message = resp - .error_message - .unwrap_or_else(|| "unknown error".to_string()); - (Color::Red, format!("error: {message}")) - } - Err(err) => { - failures += 1; - (Color::Red, format!("error: {err}")) - } - }; - - let mut row = model_row(info, use_color); - row.push( - status - .clone() - .cell() - .foreground_color(color_if(use_color, result_color)), - ); - rows.push(row); - json_rows.push(model_test_row_from_status(info, &status, result_color)); - } - } - - if json_output { - println!( - "{}", - serde_json::to_string_pretty(&ModelTestOutput { - total: json_rows.len(), - failures, - results: json_rows, - })? - ); - if failures > 0 { - bail!("{failures} model(s) failed"); - } - return Ok(()); - } - - let table = rows - .table() - .title(title) - .color_choice(color_choice(use_color)) - .border(Border::builder().build()) - .separator(Separator::builder().build()); - println!("{}", table.display()?); - - if failures > 0 { - bail!("{failures} model(s) failed"); - } - - Ok(()) -} - -#[allow(clippy::print_stdout)] -pub async fn run_models( - command: Option, - client: fabro_api::Client, - json_output: bool, -) -> Result<()> { - let command = command.unwrap_or(ModelsCommand::List { - provider: None, - query: None, - }); - - let styles = Styles::detect_stdout(); - - match command { - ModelsCommand::List { provider, query } => { - let models = - fetch_models_from_server(&client, provider.as_deref(), query.as_deref()).await?; - - if json_output { - println!("{}", serde_json::to_string_pretty(&models)?); - } else { - print_models_table(&models, &styles); - } - } - ModelsCommand::Test { - provider, - model, - deep, - } => { - test_models_via_server( - &client, - provider.as_deref(), - model.as_deref(), - deep, - &styles, - json_output, - ) - .await?; - } - } - - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - use fabro_model::{ModelCosts, ModelFeatures, ModelLimits, Provider}; - - fn test_http_client() -> reqwest::Client { - reqwest::Client::builder().no_proxy().build().unwrap() - } - - fn test_api_client(base_url: &str) -> fabro_api::Client { - fabro_api::Client::new_with_client(base_url, test_http_client()) - } - - fn test_model_json(id: &str, provider: Provider) -> serde_json::Value { - serde_json::to_value(Model { - id: id.to_string(), - provider, - family: "test".to_string(), - display_name: format!("{id} display"), - limits: ModelLimits { - context_window: 128_000, - max_output: Some(4096), - }, - training: None, - knowledge_cutoff: None, - features: ModelFeatures { - tools: true, - vision: false, - reasoning: false, - effort: false, - }, - costs: ModelCosts { - input_cost_per_mtok: Some(1.0), - output_cost_per_mtok: Some(2.0), - cache_input_cost_per_mtok: None, - }, - estimated_output_tps: Some(100.0), - aliases: vec!["tm".to_string()], - default: false, - }) - .unwrap() - } - - // --- parse_option --- - - #[test] - fn parse_option_valid() { - let (k, v) = parse_option("temperature=0.7").unwrap(); - assert_eq!(k, "temperature"); - assert_eq!(v, "0.7"); - } - - #[test] - fn parse_option_value_with_equals() { - let (k, v) = parse_option("key=a=b").unwrap(); - assert_eq!(k, "key"); - assert_eq!(v, "a=b"); - } - - #[test] - fn parse_option_no_equals() { - assert!(parse_option("nope").is_err()); - } - - // --- format_context_window --- - - #[test] - fn format_context_window_millions() { - assert_eq!(format_context_window(1_000_000), "1m"); - } - - #[test] - fn format_context_window_thousands() { - assert_eq!(format_context_window(128_000), "128k"); - } - - #[test] - fn format_context_window_small() { - assert_eq!(format_context_window(400), "400"); - } - - #[test] - fn format_context_window_rounds_up() { - // 1500 rounds to 2000 -> "2k" - assert_eq!(format_context_window(1500), "2k"); - } - - #[test] - fn format_context_window_rounds_down() { - // 1499 rounds to 1000 -> "1k" - assert_eq!(format_context_window(1499), "1k"); - } - - #[test] - fn format_context_window_zero() { - assert_eq!(format_context_window(0), "0"); - } - - // --- format_cost --- - - #[test] - fn format_cost_none() { - assert_eq!(format_cost(None), "-"); - } - - #[test] - fn format_cost_some() { - assert_eq!(format_cost(Some(3.0)), "$3.0"); - } - - #[test] - fn format_cost_fractional() { - assert_eq!(format_cost(Some(15.75)), "$15.8"); - } - - // --- format_speed --- - - #[test] - fn format_speed_none() { - assert_eq!(format_speed(None), "-"); - } - - #[test] - fn format_speed_some() { - assert_eq!(format_speed(Some(85.5)), "85 tok/s"); - } - - // --- resolve_prompt --- - - #[test] - fn resolve_prompt_arg_only() { - let result = resolve_prompt(Some("hello".into()), None).unwrap(); - assert_eq!(result, "hello"); - } - - #[test] - fn resolve_prompt_stdin_only() { - let result = resolve_prompt(None, Some("piped".into())).unwrap(); - assert_eq!(result, "piped"); - } - - #[test] - fn resolve_prompt_both_concatenates() { - let result = resolve_prompt(Some("arg".into()), Some("stdin".into())).unwrap(); - assert_eq!(result, "stdin\narg"); - } - - #[test] - fn resolve_prompt_neither_errors() { - assert!(resolve_prompt(None, None).is_err()); - } - - // --- resolve_model --- - - #[test] - fn resolve_model_explicit_known() { - let (model, provider) = resolve_model(Some("claude-sonnet-4-5".into())); - assert_eq!(model, "claude-sonnet-4-5"); - assert_eq!(provider, Some("anthropic".to_string())); - } - - #[test] - fn resolve_model_explicit_unknown() { - let (model, provider) = resolve_model(Some("nonexistent-model-xyz".into())); - assert_eq!(model, "nonexistent-model-xyz"); - assert_eq!(provider, None); - } - - #[test] - fn resolve_model_none_uses_default() { - let (model, provider) = resolve_model(None); - // Should return some valid model from catalog - assert!(!model.is_empty()); - assert!(provider.is_some()); - } - - // --- apply_options --- - - #[test] - fn apply_options_temperature() { - let params = GenerateParams::new("test-model"); - let result = apply_options(params, &[("temperature".into(), "0.7".into())]).unwrap(); - assert_eq!(result.temperature, Some(0.7)); - } - - #[test] - fn apply_options_max_tokens() { - let params = GenerateParams::new("test-model"); - let result = apply_options(params, &[("max_tokens".into(), "4096".into())]).unwrap(); - assert_eq!(result.max_tokens, Some(4096)); - } - - #[test] - fn apply_options_top_p() { - let params = GenerateParams::new("test-model"); - let result = apply_options(params, &[("top_p".into(), "0.9".into())]).unwrap(); - assert_eq!(result.top_p, Some(0.9)); - } - - #[test] - fn apply_options_unknown_key_goes_to_provider_opts() { - let params = GenerateParams::new("test-model"); - let result = apply_options(params, &[("custom_key".into(), "custom_val".into())]).unwrap(); - let opts = result.provider_options.unwrap(); - assert_eq!(opts["custom_key"], "custom_val"); - } - - #[test] - fn apply_options_invalid_temperature_errors() { - let params = GenerateParams::new("test-model"); - assert!(apply_options(params, &[("temperature".into(), "not_a_number".into())]).is_err()); - } - - #[test] - fn apply_options_invalid_max_tokens_errors() { - let params = GenerateParams::new("test-model"); - assert!(apply_options(params, &[("max_tokens".into(), "abc".into())]).is_err()); - } - - #[test] - fn apply_options_empty() { - let params = GenerateParams::new("test-model"); - let result = apply_options(params, &[]).unwrap(); - assert_eq!(result.temperature, None); - assert_eq!(result.max_tokens, None); - assert_eq!(result.provider_options, None); - } - - // --- test_model_via_server --- - - #[tokio::test] - async fn test_model_via_server_parses_ok() { - let server = httpmock::MockServer::start_async().await; - server - .mock_async(|when, then| { - when.method("POST").path("/api/v1/models/test-model/test"); - then.status(200) - .header("Content-Type", "application/json") - .body( - serde_json::json!({ - "model_id": "test-model", - "status": "ok" - }) - .to_string(), - ); - }) - .await; - - let client = test_api_client(&server.url("")); - let resp = test_model_via_server(&client, "test-model", None) - .await - .unwrap(); - - assert_eq!(resp.status, api_types::ModelTestResultStatus::Ok); - assert!(resp.error_message.is_none()); - } - - #[tokio::test] - async fn test_model_via_server_passes_mode_and_parses_error() { - let server = httpmock::MockServer::start_async().await; - server - .mock_async(|when, then| { - when.method("POST") - .path("/api/v1/models/test-model/test") - .query_param("mode", "deep"); - then.status(200) - .header("Content-Type", "application/json") - .body( - serde_json::json!({ - "model_id": "test-model", - "status": "error", - "error_message": "timeout" - }) - .to_string(), - ); - }) - .await; - - let client = test_api_client(&server.url("")); - let resp = - test_model_via_server(&client, "test-model", Some(api_types::ModelTestMode::Deep)) - .await - .unwrap(); - - assert_eq!(resp.status, api_types::ModelTestResultStatus::Error); - assert_eq!(resp.error_message.as_deref(), Some("timeout")); - } - - #[tokio::test] - async fn test_model_via_server_404() { - let server = httpmock::MockServer::start_async().await; - server - .mock_async(|when, then| { - when.method("POST").path("/api/v1/models/bad-model/test"); - then.status(404) - .header("Content-Type", "application/json") - .body(serde_json::json!({ - "errors": [{"status": "404", "title": "Not Found", "detail": "Model not found"}] - }).to_string()); - }) - .await; - - let client = test_api_client(&server.url("")); - let result = test_model_via_server(&client, "bad-model", None).await; - assert!(result.is_err()); - assert!(result.unwrap_err().to_string().contains("Model not found")); - } - - // --- fetch_models_from_server --- - - #[tokio::test] - async fn fetch_models_from_server_parses_response() { - let server = httpmock::MockServer::start_async().await; - let mock = server - .mock_async(|when, then| { - when.method("GET") - .path("/api/v1/models") - .query_param("page[limit]", "100") - .query_param("page[offset]", "0"); - then.status(200) - .header("Content-Type", "application/json") - .body( - serde_json::json!({ - "data": [test_model_json("test-model", Provider::Anthropic)], - "meta": { "has_more": false } - }) - .to_string(), - ); - }) - .await; - - let client = test_api_client(&server.url("")); - let models = fetch_models_from_server(&client, None, None).await.unwrap(); - - mock.assert_async().await; - assert_eq!(models.len(), 1); - assert_eq!(models[0].id, "test-model"); - assert_eq!(models[0].provider, fabro_model::Provider::Anthropic); - } - - #[tokio::test] - async fn fetch_models_from_server_filters_by_provider() { - let server = httpmock::MockServer::start_async().await; - server - .mock_async(|when, then| { - when.method("GET") - .path("/api/v1/models") - .query_param("page[limit]", "100") - .query_param("page[offset]", "0") - .query_param("provider", "anthropic"); - then.status(200) - .header("Content-Type", "application/json") - .body( - serde_json::json!({ - "data": [test_model_json("model-a", Provider::Anthropic)], - "meta": { "has_more": false } - }) - .to_string(), - ); - }) - .await; - - let client = test_api_client(&server.url("")); - let models = fetch_models_from_server(&client, Some("anthropic"), None) - .await - .unwrap(); - - assert_eq!(models.len(), 1); - assert_eq!(models[0].id, "model-a"); - } - - #[tokio::test] - async fn fetch_models_from_server_passes_query_param() { - let server = httpmock::MockServer::start_async().await; - let mock = server - .mock_async(|when, then| { - when.method("GET") - .path("/api/v1/models") - .query_param("page[limit]", "100") - .query_param("page[offset]", "0") - .query_param("query", "sonnet"); - then.status(200) - .header("Content-Type", "application/json") - .body( - serde_json::json!({ - "data": [test_model_json("claude-sonnet-4-5", Provider::Anthropic)], - "meta": { "has_more": false } - }) - .to_string(), - ); - }) - .await; - - let client = test_api_client(&server.url("")); - let models = fetch_models_from_server(&client, None, Some("sonnet")) - .await - .unwrap(); - - mock.assert_async().await; - assert_eq!(models.len(), 1); - assert_eq!(models[0].id, "claude-sonnet-4-5"); - } - - #[tokio::test] - async fn fetch_models_from_server_follows_pagination() { - let server = httpmock::MockServer::start_async().await; - let first_page = server - .mock_async(|when, then| { - when.method("GET") - .path("/api/v1/models") - .query_param("page[limit]", "100") - .query_param("page[offset]", "0"); - then.status(200) - .header("Content-Type", "application/json") - .body( - serde_json::json!({ - "data": [test_model_json("model-a", Provider::Anthropic)], - "meta": { "has_more": true } - }) - .to_string(), - ); - }) - .await; - let second_page = server - .mock_async(|when, then| { - when.method("GET") - .path("/api/v1/models") - .query_param("page[limit]", "100") - .query_param("page[offset]", "1"); - then.status(200) - .header("Content-Type", "application/json") - .body( - serde_json::json!({ - "data": [test_model_json("model-b", Provider::OpenAi)], - "meta": { "has_more": false } - }) - .to_string(), - ); - }) - .await; - - let client = test_api_client(&server.url("")); - let models = fetch_models_from_server(&client, None, None).await.unwrap(); - - first_page.assert_async().await; - second_page.assert_async().await; - assert_eq!(models.len(), 2); - assert_eq!(models[0].id, "model-a"); - assert_eq!(models[1].id, "model-b"); - } - - #[tokio::test] - async fn fetch_models_from_server_error_on_failure() { - let server = httpmock::MockServer::start_async().await; - server - .mock_async(|when, then| { - when.method("GET") - .path("/api/v1/models") - .query_param("page[limit]", "100") - .query_param("page[offset]", "0"); - then.status(500).body("internal error"); - }) - .await; - - let client = test_api_client(&server.url("")); - let result = fetch_models_from_server(&client, None, None).await; - assert!(result.is_err()); - } - - // --- run_prompt_via_server --- - - #[tokio::test] - async fn run_prompt_via_server_non_streaming() { - let mock_server = httpmock::MockServer::start_async().await; - let mock = mock_server - .mock_async(|when, then| { - when.method("POST").path("/completions"); - then.status(200) - .header("Content-Type", "application/json") - .body( - serde_json::json!({ - "id": "msg_123", - "model": "test-model", - "content": "Hello world", - "stop_reason": "end_turn", - "usage": {"input_tokens": 10, "output_tokens": 5} - }) - .to_string(), - ); - }) - .await; - - let server = ServerConnection { - client: test_http_client(), - base_url: mock_server.url(""), - }; - - let args = PromptArgs { - prompt: Some("Hello".into()), - model: Some("test-model".into()), - system: None, - no_stream: true, - usage: false, - schema: None, - option: vec![], - }; - - let result = run_prompt_via_server(args, &server, false).await; - assert!(result.is_ok()); - mock.assert_async().await; - } - - #[tokio::test] - async fn run_prompt_via_server_streaming() { - let mock_server = httpmock::MockServer::start_async().await; - let sse_body = "\ -event: stream_event\n\ -data: {\"type\":\"text_delta\",\"delta\":\"Hi\",\"text_id\":null}\n\ -\n\ -event: stream_event\n\ -data: {\"type\":\"finish\",\"finish_reason\":\"stop\",\"usage\":{\"input_tokens\":5,\"output_tokens\":2,\"total_tokens\":7},\"response\":{\"id\":\"r1\",\"model\":\"test\",\"provider\":\"test\",\"message\":{\"role\":\"assistant\",\"content\":[{\"kind\":\"text\",\"data\":\"Hi\"}],\"name\":null,\"tool_call_id\":null},\"finish_reason\":\"stop\",\"usage\":{\"input_tokens\":5,\"output_tokens\":2,\"total_tokens\":7},\"raw\":null,\"warnings\":[],\"rate_limit\":null}}\n\ -\n"; - - let mock = mock_server - .mock_async(|when, then| { - when.method("POST").path("/completions"); - then.status(200) - .header("Content-Type", "text/event-stream") - .body(sse_body); - }) - .await; - - let server = ServerConnection { - client: test_http_client(), - base_url: mock_server.url(""), - }; - - let args = PromptArgs { - prompt: Some("Hello".into()), - model: Some("test-model".into()), - system: None, - no_stream: false, - usage: false, - schema: None, - option: vec![], - }; - - let result = run_prompt_via_server(args, &server, false).await; - assert!(result.is_ok()); - mock.assert_async().await; - } -} diff --git a/lib/crates/fabro-llm/src/lib.rs b/lib/crates/fabro-llm/src/lib.rs index 91f10b6ae..c7b7a51bf 100644 --- a/lib/crates/fabro-llm/src/lib.rs +++ b/lib/crates/fabro-llm/src/lib.rs @@ -1,4 +1,3 @@ -pub mod cli; pub mod client; pub mod error; pub mod generate; diff --git a/lib/crates/fabro-llm/src/model_test.rs b/lib/crates/fabro-llm/src/model_test.rs index 390f54e91..1b02771d6 100644 --- a/lib/crates/fabro-llm/src/model_test.rs +++ b/lib/crates/fabro-llm/src/model_test.rs @@ -5,8 +5,7 @@ use tokio::time; use crate::generate::{self, GenerateParams}; use crate::tools::Tool; -use crate::types::GenerateResult; -use crate::types::ReasoningEffort; +use crate::types::{GenerateResult, ReasoningEffort}; use fabro_model::Model; #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] @@ -198,6 +197,7 @@ fn validate_deep_result(result: &GenerateResult) -> Result<(), String> { #[cfg(test)] mod tests { use super::*; + use crate::types::ToolResult; use crate::types::{FinishReason, Message, Response, StepResult, Usage}; use fabro_model::{ModelCosts, ModelFeatures, ModelLimits, Provider}; @@ -259,10 +259,7 @@ mod tests { #[test] fn validate_deep_result_does_not_fail_only_for_missing_reasoning() { - let tool_results = vec![crate::types::ToolResult::success( - "call_1", - serde_json::json!(42), - )]; + let tool_results = vec![ToolResult::success("call_1", serde_json::json!(42))]; let first_step = StepResult { response: response_with_text("tool step"), tool_results: tool_results.clone(), diff --git a/lib/crates/fabro-llm/src/providers/http_api.rs b/lib/crates/fabro-llm/src/providers/http_api.rs index d63bde0a5..30cfe0d81 100644 --- a/lib/crates/fabro-llm/src/providers/http_api.rs +++ b/lib/crates/fabro-llm/src/providers/http_api.rs @@ -21,11 +21,11 @@ impl HttpApi { fn build_client(timeout: AdapterTimeout) -> reqwest::Client { #[cfg(test)] { - return reqwest::Client::builder() + reqwest::Client::builder() .connect_timeout(Duration::from_secs_f64(timeout.connect)) .no_proxy() .build() - .unwrap_or_default(); + .unwrap_or_default() } #[cfg(not(test))] {