diff --git a/lib/crates/fabro-agent/src/cli.rs b/lib/crates/fabro-agent/src/cli.rs index fb7a4d96e..0a43e95fb 100644 --- a/lib/crates/fabro-agent/src/cli.rs +++ b/lib/crates/fabro-agent/src/cli.rs @@ -484,7 +484,7 @@ pub async fn run_with_args_and_client( model } else { Catalog::builtin() - .default_for_provider(&provider.to_string()) + .default_for_provider(<&str>::from(provider)) .map(|model| model.id.clone()) .ok_or_else(|| { anyhow::anyhow!( diff --git a/lib/crates/fabro-agent/tests/it/guardrails.rs b/lib/crates/fabro-agent/tests/it/guardrails.rs index 663f2b620..95b94b38a 100644 --- a/lib/crates/fabro-agent/tests/it/guardrails.rs +++ b/lib/crates/fabro-agent/tests/it/guardrails.rs @@ -4,9 +4,8 @@ use fabro_model::{Catalog, Provider}; #[test] fn profile_context_window_matches_catalog_for_default_models() { for &provider in Provider::ALL { - let provider_str: &str = provider.into(); let catalog_info = Catalog::builtin() - .default_for_provider(provider_str) + .default_for_provider(<&str>::from(provider)) .cloned() .unwrap_or_else(|| panic!("no default model for {provider:?} in catalog")); let model = &catalog_info.id; diff --git a/lib/crates/fabro-model/Cargo.toml b/lib/crates/fabro-model/Cargo.toml index 2b304b421..155be3a3f 100644 --- a/lib/crates/fabro-model/Cargo.toml +++ b/lib/crates/fabro-model/Cargo.toml @@ -20,4 +20,4 @@ serde_json.workspace = true strum.workspace = true [dev-dependencies] -insta.workspace = true \ No newline at end of file +insta.workspace = true diff --git a/lib/crates/fabro-model/src/catalog.rs b/lib/crates/fabro-model/src/catalog.rs index 9fe9e1964..cdada4a06 100644 --- a/lib/crates/fabro-model/src/catalog.rs +++ b/lib/crates/fabro-model/src/catalog.rs @@ -80,7 +80,7 @@ impl Catalog { #[must_use] pub fn default_from_env(&self) -> &Model { let provider = Provider::default_from_env(); - self.default_for_provider(provider.to_string().as_str()) + self.default_for_provider(<&str>::from(provider)) .unwrap_or_else(|| self.default_model()) } @@ -89,7 +89,7 @@ impl Catalog { #[must_use] pub fn default_for_configured(&self, configured: &[Provider]) -> &Model { let provider = Provider::default_for_configured(configured); - self.default_for_provider(provider.to_string().as_str()) + self.default_for_provider(<&str>::from(provider)) .unwrap_or_else(|| self.default_model()) } diff --git a/lib/crates/fabro-server/src/server.rs b/lib/crates/fabro-server/src/server.rs index 0fe52a578..942ea8106 100644 --- a/lib/crates/fabro-server/src/server.rs +++ b/lib/crates/fabro-server/src/server.rs @@ -54,7 +54,7 @@ use fabro_llm::types::{ ContentPart, FinishReason, Message as LlmMessage, Request as LlmRequest, Role, ToolChoice, ToolDefinition, }; -use fabro_model::{BilledModelUsage, BilledTokenCounts, Catalog, ModelTestMode, Provider}; +use fabro_model::{BilledModelUsage, BilledTokenCounts, Catalog, ModelTestMode}; use fabro_redact::redact_jsonl_line; use fabro_sandbox::daytona::{self, DaytonaSandbox}; use fabro_sandbox::reconnect::reconnect; diff --git a/lib/crates/fabro-server/src/server/handler/models.rs b/lib/crates/fabro-server/src/server/handler/models.rs index a4c190385..e3374fa6a 100644 --- a/lib/crates/fabro-server/src/server/handler/models.rs +++ b/lib/crates/fabro-server/src/server/handler/models.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use super::super::{ ApiError, AppState, FromStr, HashSet, IntoResponse, Json, MAX_PAGE_OFFSET, ModelTestMode, Path, - Provider, Query, RequiredUser, Response, Router, State, StatusCode, auth_issue_message, + Query, RequiredUser, Response, Router, State, StatusCode, auth_issue_message, default_page_limit, error, get, post, run_model_test, }; @@ -116,14 +116,12 @@ async fn test_model( .into_response(); } }; - if let Some((_, issue)) = llm_result + if let Some((provider_enum, issue)) = llm_result .auth_issues .iter() - .find(|(provider, _)| provider.to_string() == info.provider.as_str()) + .find(|(provider, _)| info.provider == <&str>::from(*provider)) { - let provider_enum = - Provider::from_str(info.provider.as_str()).unwrap_or(Provider::Anthropic); - return ApiError::bad_request(auth_issue_message(provider_enum, issue)).into_response(); + return ApiError::bad_request(auth_issue_message(*provider_enum, issue)).into_response(); } let provider_name = info.provider.as_str(); if !llm_result.client.provider_names().contains(&provider_name) { diff --git a/lib/crates/fabro-workflow/src/operations/start.rs b/lib/crates/fabro-workflow/src/operations/start.rs index f1f5b60b0..0e05b49bd 100644 --- a/lib/crates/fabro-workflow/src/operations/start.rs +++ b/lib/crates/fabro-workflow/src/operations/start.rs @@ -534,7 +534,7 @@ fn resolve_fallback_chain( .or_default() .push(model_ref.to_string()); } - Catalog::builtin().build_fallback_chain(&provider.to_string(), model, &by_provider) + Catalog::builtin().build_fallback_chain(<&str>::from(provider), model, &by_provider) } fn runtime_mcp_server(settings: &ResolvedMcpServerSettings) -> McpServerSettings { diff --git a/lib/crates/fabro-workflow/src/outcome.rs b/lib/crates/fabro-workflow/src/outcome.rs index 78885c9d6..3e18a52c9 100644 --- a/lib/crates/fabro-workflow/src/outcome.rs +++ b/lib/crates/fabro-workflow/src/outcome.rs @@ -27,8 +27,7 @@ pub fn billed_model_usage_from_llm( speed, }; let tokens = token_counts_from_llm_usage(usage); - let provider_str = provider.to_string(); - let facts = billing_facts_for_stage_usage(&provider_str, &tokens); + let facts = billing_facts_for_stage_usage(model.provider.as_str(), &tokens); let input = ModelBillingInput { usage: ModelUsage { model: model.clone(), @@ -39,7 +38,7 @@ pub fn billed_model_usage_from_llm( let total_usd_micros = Catalog::builtin() .get(model_id) - .filter(|candidate| candidate.provider == provider_str.as_str()) + .filter(|candidate| candidate.provider == model.provider) .and_then(|candidate| candidate.pricing_for(speed)) .and_then(|pricing| pricing.bill(&input)) .map(|amount| amount.0); diff --git a/lib/crates/fabro-workflow/src/run_materialization.rs b/lib/crates/fabro-workflow/src/run_materialization.rs index d911808d1..920152a31 100644 --- a/lib/crates/fabro-workflow/src/run_materialization.rs +++ b/lib/crates/fabro-workflow/src/run_materialization.rs @@ -37,8 +37,11 @@ pub fn materialize_run( let model = configured_model.or(graph_model).unwrap_or_else(|| { provider .as_deref() - .and_then(|value| value.parse::().ok()) - .and_then(|provider| catalog.default_for_provider(&provider.to_string())) + .and_then(|value| { + // Validate the provider string is known before looking up its default. + value.parse::().ok()?; + catalog.default_for_provider(value) + }) .unwrap_or_else(|| catalog.default_for_configured(configured_providers)) .id .clone()