diff --git a/Cargo.lock b/Cargo.lock index ff698ec11..ff407fa63 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2015,6 +2015,7 @@ dependencies = [ name = "fabro-model" version = "0.221.0-nightly.1" dependencies = [ + "chrono", "fabro-static", "insta", "serde", diff --git a/lib/crates/fabro-agent/src/cli.rs b/lib/crates/fabro-agent/src/cli.rs index 1276bba26..fb7a4d96e 100644 --- a/lib/crates/fabro-agent/src/cli.rs +++ b/lib/crates/fabro-agent/src/cli.rs @@ -205,18 +205,18 @@ fn build_tool_approval( } fn summarizer_model_id(provider: Provider) -> ModelHandle { + let model = match provider { + Provider::OpenAi | Provider::OpenAiCompatible => "gpt-4o-mini", + Provider::Gemini => "gemini-2.0-flash", + Provider::Anthropic => "claude-haiku-4-5", + Provider::Kimi => "kimi-k2.5", + Provider::Zai => "glm-4.7", + Provider::Minimax => "minimax-m2.5", + Provider::Inception => "mercury", + }; ModelHandle::ByName { - provider, - model: match provider { - Provider::OpenAi | Provider::OpenAiCompatible => "gpt-4o-mini", - Provider::Gemini => "gemini-2.0-flash", - Provider::Anthropic => "claude-haiku-4-5", - Provider::Kimi => "kimi-k2.5", - Provider::Zai => "glm-4.7", - Provider::Minimax => "minimax-m2.5", - Provider::Inception => "mercury", - } - .to_string(), + provider: provider.into(), + model: model.to_string(), } } @@ -484,7 +484,7 @@ pub async fn run_with_args_and_client( model } else { Catalog::builtin() - .default_for_provider(provider) + .default_for_provider(&provider.to_string()) .map(|model| model.id.clone()) .ok_or_else(|| { anyhow::anyhow!( diff --git a/lib/crates/fabro-agent/src/tools.rs b/lib/crates/fabro-agent/src/tools.rs index 95eb55e34..08f73e5d2 100644 --- a/lib/crates/fabro-agent/src/tools.rs +++ b/lib/crates/fabro-agent/src/tools.rs @@ -1292,7 +1292,7 @@ mod tests { let summarizer = WebFetchSummarizer { client, model_id: ModelHandle::ByName { - provider: fabro_model::Provider::Anthropic, + provider: fabro_model::ProviderId::from("anthropic"), model: "mock-model".to_string(), }, }; @@ -1393,7 +1393,7 @@ mod tests { let summarizer = WebFetchSummarizer { client, model_id: ModelHandle::ByName { - provider: fabro_model::Provider::Anthropic, + provider: fabro_model::ProviderId::from("anthropic"), model: "target-model".to_string(), }, }; diff --git a/lib/crates/fabro-agent/tests/it/guardrails.rs b/lib/crates/fabro-agent/tests/it/guardrails.rs index 5c2c89723..663f2b620 100644 --- a/lib/crates/fabro-agent/tests/it/guardrails.rs +++ b/lib/crates/fabro-agent/tests/it/guardrails.rs @@ -4,8 +4,9 @@ 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) + .default_for_provider(provider_str) .cloned() .unwrap_or_else(|| panic!("no default model for {provider:?} in catalog")); let model = &catalog_info.id; diff --git a/lib/crates/fabro-agent/tests/it/parity_matrix.rs b/lib/crates/fabro-agent/tests/it/parity_matrix.rs index 8b8630e4d..971947969 100644 --- a/lib/crates/fabro-agent/tests/it/parity_matrix.rs +++ b/lib/crates/fabro-agent/tests/it/parity_matrix.rs @@ -17,7 +17,7 @@ use fabro_auth::EnvCredentialSource; use fabro_llm::client::Client; use fabro_llm::provider::{Provider, ProviderAdapter}; use fabro_llm::providers::OpenAiAdapter; -use fabro_model::ModelHandle; +use fabro_model::{ModelHandle, ProviderId}; use fabro_test::{TwinScenario, TwinScenarios, TwinToolCall, twin_openai}; use tokio::sync::Mutex as AsyncMutex; @@ -35,15 +35,15 @@ fn summarizer_model_id(provider: Provider) -> ModelHandle { | Provider::Minimax | Provider::Inception | Provider::OpenAiCompatible => ModelHandle::ByName { - provider: Provider::OpenAi, + provider: ProviderId::from("openai"), model: "gpt-5.4-mini".to_string(), }, Provider::Gemini => ModelHandle::ByName { - provider: Provider::Gemini, + provider: ProviderId::from("gemini"), model: "gemini-3-flash-preview".to_string(), }, Provider::Anthropic => ModelHandle::ByName { - provider: Provider::Anthropic, + provider: ProviderId::from("anthropic"), model: "claude-haiku-4-5".to_string(), }, } diff --git a/lib/crates/fabro-api/tests/model_round_trip.rs b/lib/crates/fabro-api/tests/model_round_trip.rs index 314f35103..836fafb30 100644 --- a/lib/crates/fabro-api/tests/model_round_trip.rs +++ b/lib/crates/fabro-api/tests/model_round_trip.rs @@ -1,7 +1,7 @@ use std::any::{TypeId, type_name}; use fabro_api::types::Model as ApiModel; -use fabro_model::{Model, ModelCosts, ModelFeatures, ModelLimits, Provider}; +use fabro_model::{Model, ModelCosts, ModelFeatures, ModelLimits, ProviderId}; #[test] fn model_reuses_canonical_type() { @@ -12,7 +12,7 @@ fn model_reuses_canonical_type() { fn model_json_matches_openapi_shape() { let model = Model { id: "claude-opus-4-7".to_string(), - provider: Provider::Anthropic, + provider: ProviderId::from("anthropic"), family: "claude-4".to_string(), display_name: "Claude Opus 4.7".to_string(), limits: ModelLimits { diff --git a/lib/crates/fabro-cli/src/commands/model.rs b/lib/crates/fabro-cli/src/commands/model.rs index e66fbdbfd..3095c9778 100644 --- a/lib/crates/fabro-cli/src/commands/model.rs +++ b/lib/crates/fabro-cli/src/commands/model.rs @@ -2,7 +2,7 @@ use anyhow::{Context, Result, bail}; use cli_table::format::{Border, Justify, Separator}; use cli_table::{Cell, CellStruct, Color, Style, Table}; use fabro_api::types as api_types; -use fabro_model::{Catalog, Model, ModelTestMode, Provider}; +use fabro_model::{Catalog, Model, ModelTestMode}; use fabro_util::terminal::Styles; use serde::Serialize; @@ -21,7 +21,7 @@ enum ModelTestResultKind { #[derive(Serialize)] struct ModelTestRow { model: String, - provider: Provider, + provider: String, result: ModelTestResultKind, #[serde(skip_serializing_if = "Option::is_none")] detail: Option, @@ -104,6 +104,7 @@ fn model_row(model: &Model, use_color: bool) -> Vec { model.id.clone().cell().bold(use_color), model .provider + .to_string() .cell() .foreground_color(color_if(use_color, Color::Ansi256(8))), aliases @@ -160,21 +161,21 @@ fn model_test_row_from_status(model: &Model, status: &str, result_color: Color) match result_color { Color::Green => ModelTestRow { model: model.id.clone(), - provider: model.provider, + provider: model.provider.to_string(), result: ModelTestResultKind::Pass, detail: None, error: None, }, Color::Yellow => ModelTestRow { model: model.id.clone(), - provider: model.provider, + provider: model.provider.to_string(), result: ModelTestResultKind::Skip, detail: Some(trimmed.to_string()), error: None, }, _ => ModelTestRow { model: model.id.clone(), - provider: model.provider, + provider: model.provider.to_string(), result: ModelTestResultKind::Fail, detail: None, error: Some( @@ -276,7 +277,7 @@ async fn test_models_via_server( for info in &unconfigured { skipped += 1; - let provider_name = info.provider.display_name().to_string(); + let provider_name = info.provider.to_string(); if !skipped_providers.contains(&provider_name) { skipped_providers.push(provider_name); } @@ -437,33 +438,33 @@ mod tests { server_client::Client::new_no_proxy(api_url).unwrap() } - fn test_model_json(id: &str, provider: Provider) -> serde_json::Value { + fn test_model_json(id: &str, provider: &str) -> serde_json::Value { serde_json::to_value(Model { - id: id.to_string(), - provider, - family: "test".to_string(), - display_name: format!("{id} display"), - limits: ModelLimits { + id: id.to_string(), + provider: fabro_model::ProviderId::from(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 { + training: None, + knowledge_cutoff: None, + features: ModelFeatures { tools: true, vision: false, reasoning: false, effort: false, }, - costs: ModelCosts { + 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, - configured: false, + aliases: vec!["tm".to_string()], + default: false, + configured: false, }) .unwrap() } @@ -640,7 +641,7 @@ mod tests { .header("Content-Type", "application/json") .body( serde_json::json!({ - "data": [test_model_json("test-model", Provider::Anthropic)], + "data": [test_model_json("test-model", "anthropic")], "meta": { "has_more": false } }) .to_string(), @@ -654,7 +655,7 @@ mod tests { mock.assert_async().await; assert_eq!(models.len(), 1); assert_eq!(models[0].id, "test-model"); - assert_eq!(models[0].provider, Provider::Anthropic); + assert_eq!(models[0].provider, "anthropic"); } #[tokio::test] @@ -671,7 +672,7 @@ mod tests { .header("Content-Type", "application/json") .body( serde_json::json!({ - "data": [test_model_json("model-a", Provider::Anthropic)], + "data": [test_model_json("model-a", "anthropic")], "meta": { "has_more": false } }) .to_string(), @@ -700,7 +701,7 @@ mod tests { .header("Content-Type", "application/json") .body( serde_json::json!({ - "data": [test_model_json("claude-sonnet-4-5", Provider::Anthropic)], + "data": [test_model_json("claude-sonnet-4-5", "anthropic")], "meta": { "has_more": false } }) .to_string(), @@ -729,7 +730,7 @@ mod tests { .header("Content-Type", "application/json") .body( serde_json::json!({ - "data": [test_model_json("model-a", Provider::Anthropic)], + "data": [test_model_json("model-a", "anthropic")], "meta": { "has_more": true } }) .to_string(), @@ -746,7 +747,7 @@ mod tests { .header("Content-Type", "application/json") .body( serde_json::json!({ - "data": [test_model_json("model-b", Provider::OpenAi)], + "data": [test_model_json("model-b", "openai")], "meta": { "has_more": false } }) .to_string(), diff --git a/lib/crates/fabro-cli/src/commands/run/overrides.rs b/lib/crates/fabro-cli/src/commands/run/overrides.rs index c1b73177b..7746d486c 100644 --- a/lib/crates/fabro-cli/src/commands/run/overrides.rs +++ b/lib/crates/fabro-cli/src/commands/run/overrides.rs @@ -39,6 +39,7 @@ fn model_from_args(model: Option<&str>, provider: Option<&str>) -> Option Resul .await .context("failed to create LLM client")?; + let provider_str = provider.to_string(); let probe_model = Catalog::builtin() - .probe_for_provider(provider) + .probe_for_provider(&provider_str) .map_or_else(|| format!("unknown-{provider}"), |model| model.id.clone()); let params = GenerateParams::new(probe_model, Arc::new(client)) - .provider(<&'static str>::from(provider)) + .provider(&provider_str) .prompt("Say OK") .max_tokens(16); diff --git a/lib/crates/fabro-cli/tests/it/cmd/model_test.rs b/lib/crates/fabro-cli/tests/it/cmd/model_test.rs index 5707c21f1..f2211ec21 100644 --- a/lib/crates/fabro-cli/tests/it/cmd/model_test.rs +++ b/lib/crates/fabro-cli/tests/it/cmd/model_test.rs @@ -250,7 +250,7 @@ fn model_test_skipped_footer_sources_from_listing() { String::from_utf8_lossy(&output.stderr) ); let stderr = String::from_utf8_lossy(&output.stderr); - assert!(stderr.contains("Skipped 1 model(s) (no credentials: OpenAI)")); + assert!(stderr.contains("Skipped 1 model(s) (no credentials: openai)")); } #[test] diff --git a/lib/crates/fabro-config/src/builders.rs b/lib/crates/fabro-config/src/builders.rs index 4b09c407c..38d2f2cd7 100644 --- a/lib/crates/fabro-config/src/builders.rs +++ b/lib/crates/fabro-config/src/builders.rs @@ -497,6 +497,7 @@ command = ["demo-mcp"] provider: Some(InterpString::parse("openai")), name: Some(InterpString::parse("gpt-5")), fallbacks: Vec::new(), + controls: None, }), execution: Some(RunExecutionLayer { mode: Some(RunMode::DryRun), diff --git a/lib/crates/fabro-config/src/layers/combine.rs b/lib/crates/fabro-config/src/layers/combine.rs index 58a0f2a03..3b0d24a24 100644 --- a/lib/crates/fabro-config/src/layers/combine.rs +++ b/lib/crates/fabro-config/src/layers/combine.rs @@ -13,10 +13,14 @@ use fabro_types::settings::{Duration, InterpString, Size}; use super::LogFilter; use super::cli::{CliAuthLayer, CliLoggingLayer, CliTargetLayer}; use super::features::FeaturesLayer; +use super::llm::{ + CredentialRef, ModelControlsLayer, ModelCostTableLayer, ModelFeaturesLayer, ModelLimitsLayer, +}; use super::run::{ DaytonaSnapshotLayer, HookAgentMarker, HookEntry, HookTlsMode, InterviewProviderLayer, LocalSandboxLayer, ModelRefOrSplice, NotificationProviderLayer, RunArtifactsLayer, - RunCheckpointLayer, RunGoalLayer, RunPrepareLayer, ScmGitHubLayer, StringOrSplice, + RunCheckpointLayer, RunGoalLayer, RunModelControlsLayer, RunPrepareLayer, ScmGitHubLayer, + StringOrSplice, }; use super::server::{ ObjectStoreLocalLayer, ObjectStoreS3Layer, ServerApiLayer, ServerAuthGithubLayer, @@ -57,6 +61,7 @@ macro_rules! impl_combine_or_option { impl_combine_or_option!( String, bool, + f64, u16, u32, u64, @@ -84,9 +89,15 @@ impl_combine_or_option!( LogFilter, ); -impl Combine for Option> { - fn combine(self, other: Self) -> Self { - self.or(other) +impl Combine for Vec { + fn combine(self, _other: Self) -> Self { + self + } +} + +impl Combine for Vec { + fn combine(self, _other: Self) -> Self { + self } } @@ -123,9 +134,14 @@ impl_combine_self!( DaytonaSnapshotLayer, InterviewProviderLayer, LocalSandboxLayer, + ModelControlsLayer, + ModelCostTableLayer, + ModelFeaturesLayer, + ModelLimitsLayer, NotificationProviderLayer, RunArtifactsLayer, RunGoalLayer, + RunModelControlsLayer, RunPrepareLayer, ScmGitHubLayer, ObjectStoreLocalLayer, diff --git a/lib/crates/fabro-config/src/layers/llm.rs b/lib/crates/fabro-config/src/layers/llm.rs new file mode 100644 index 000000000..c6c9d5893 --- /dev/null +++ b/lib/crates/fabro-config/src/layers/llm.rs @@ -0,0 +1,232 @@ +//! Sparse `[llm]` settings layer: provider and model catalog data. + +use std::collections::BTreeMap; + +use serde::de::Error as _; +use serde::{Deserialize, Serialize}; + +use super::maps::MergeMap; + +/// Deserialize `knowledge_cutoff` from either a TOML date or a string. +/// +/// When TOML source contains an unquoted `2025-01-01`, the `toml` crate +/// intermediate `Value` representation stores it as a `Datetime`. +/// When it's quoted `"2025-01-01"`, it's a string. We accept both. +fn deserialize_knowledge_cutoff<'de, D>(deserializer: D) -> Result, D::Error> +where + D: serde::Deserializer<'de>, +{ + // Deserialize as a generic TOML value first, then coerce to string. + let opt: Option = Option::deserialize(deserializer)?; + match opt { + None => Ok(None), + Some(toml::Value::String(s)) => Ok(Some(s)), + Some(toml::Value::Datetime(dt)) => Ok(Some(dt.to_string())), + Some(other) => Err(D::Error::custom(format!( + "expected a date string or TOML date for knowledge_cutoff, got {other}" + ))), + } +} + +/// Top-level `[llm]` settings layer. +/// +/// This only contains `providers` and `models` subtrees. +/// Legacy keys like `provider` or `model` at `[llm]` level should be caught +/// by the parse-time migration hint, not parsed here. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] +#[serde(deny_unknown_fields)] +pub struct LlmLayer { + /// `[llm.providers.]` — merge-by-key across layers. + #[serde(default, skip_serializing_if = "MergeMap::is_empty")] + pub providers: MergeMap, + /// `[llm.models.]` — merge-by-key across layers. + #[serde(default, skip_serializing_if = "MergeMap::is_empty")] + pub models: MergeMap, +} + +/// `[llm.providers.]` — a single provider's settings. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] +#[serde(deny_unknown_fields)] +pub struct ProviderSettingsLayer { + /// Human-readable display name. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub display_name: Option, + /// Adapter key (e.g. "anthropic", "openai", "openai_compatible"). + /// Validated against the adapter registry at catalog build time. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub adapter: Option, + /// Base URL for API requests. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub base_url: Option, + /// Ordered credential references. Replaces as whole array across layers. + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub credentials: Vec, + /// Priority for default provider selection. Higher wins. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub priority: Option, + /// Whether this provider is available for runtime selection. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub enabled: Option, + /// Alternative names for this provider. + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub aliases: Vec, +} + +/// `[llm.models.]` — a single model's settings. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] +#[serde(deny_unknown_fields)] +pub struct ModelSettingsLayer { + /// Provider ID this model belongs to. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub provider: Option, + /// The model identifier sent to the provider API. + /// When omitted, defaults to the catalog model ID. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub api_id: Option, + /// Human-readable display name. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub display_name: Option, + /// Model family (e.g. "claude-4", "gpt-5"). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub family: Option, + /// Knowledge cutoff date (YYYY-MM-DD string). + #[serde( + default, + skip_serializing_if = "Option::is_none", + deserialize_with = "deserialize_knowledge_cutoff" + )] + pub knowledge_cutoff: Option, + /// Whether this is the default model for its provider. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub default: Option, + /// Whether this model is available for selection. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub enabled: Option, + /// Alternative names for this model. + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub aliases: Vec, + /// Estimated output tokens per second. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub estimated_output_tps: Option, + /// Model limits. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub limits: Option, + /// Model feature flags. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub features: Option, + /// Base cost rates. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub costs: Option, + /// Supported control values. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub controls: Option, +} + +/// Model context window and output limits. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ModelLimitsLayer { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub context_window: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_output: Option, +} + +/// Model feature flags. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ModelFeaturesLayer { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tools: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub vision: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reasoning: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub effort: Option, +} + +/// Cost rates in USD per million tokens. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct CostRatesLayer { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub input_cost_per_mtok: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub output_cost_per_mtok: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cache_input_cost_per_mtok: Option, +} + +/// Model cost table: base rates plus optional per-speed overrides. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ModelCostTableLayer { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub input_cost_per_mtok: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub output_cost_per_mtok: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cache_input_cost_per_mtok: Option, + /// Per-speed cost overrides. Keys are speed names (e.g. "fast"). + #[serde(default, skip_serializing_if = "BTreeMap::is_empty")] + pub speed: BTreeMap, +} + +/// Model control allow-lists. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct ModelControlsLayer { + /// Allowed reasoning effort values (e.g. `["low", "medium", "high"]`). + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub reasoning_effort: Vec, + /// Additional speed values beyond standard (e.g. `["fast"]`). + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub speed: Vec, +} + +/// A typed credential reference. Only `credential:` and `env:` +/// are valid. Literal secrets fail deserialization. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum CredentialRef { + /// `credential:` — read from fabro-vault. + Credential(String), + /// `env:` — read from process environment, then vault fallback. + Env(String), +} + +impl std::fmt::Display for CredentialRef { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Credential(id) => write!(f, "credential:{id}"), + Self::Env(name) => write!(f, "env:{name}"), + } + } +} + +impl Serialize for CredentialRef { + fn serialize(&self, serializer: S) -> Result { + serializer.serialize_str(&self.to_string()) + } +} + +impl<'de> Deserialize<'de> for CredentialRef { + fn deserialize>(deserializer: D) -> Result { + let raw = String::deserialize(deserializer)?; + if let Some(id) = raw.strip_prefix("credential:") { + if id.is_empty() { + return Err(D::Error::custom("credential: ref must have a non-empty ID")); + } + return Ok(Self::Credential(id.to_string())); + } + if let Some(name) = raw.strip_prefix("env:") { + if name.is_empty() { + return Err(D::Error::custom("env: ref must have a non-empty name")); + } + return Ok(Self::Env(name.to_string())); + } + Err(D::Error::custom(format!( + "invalid credential reference '{raw}': must start with 'credential:' or 'env:'" + ))) + } +} diff --git a/lib/crates/fabro-config/src/layers/mod.rs b/lib/crates/fabro-config/src/layers/mod.rs index 60af5e79e..2ab1dc222 100644 --- a/lib/crates/fabro-config/src/layers/mod.rs +++ b/lib/crates/fabro-config/src/layers/mod.rs @@ -1,6 +1,7 @@ mod cli; mod combine; mod features; +mod llm; mod log_filter; mod maps; mod project; @@ -16,6 +17,10 @@ pub use cli::{ }; pub(crate) use combine::Combine; pub use features::FeaturesLayer; +pub use llm::{ + CostRatesLayer, CredentialRef, LlmLayer, ModelControlsLayer, ModelCostTableLayer, + ModelFeaturesLayer, ModelLimitsLayer, ModelSettingsLayer, ProviderSettingsLayer, +}; pub use log_filter::LogFilter; pub use maps::{MergeMap, ReplaceMap, StickyMap}; pub use project::ProjectLayer; @@ -24,8 +29,9 @@ pub use run::{ GitAuthorLayer, HookAgentMarker, HookEntry, HookTlsMode, InterviewProviderLayer, InterviewsLayer, LocalSandboxLayer, McpEntryLayer, ModelRefOrSplice, NotificationProviderLayer, NotificationRouteLayer, PrepareStep, RunAgentLayer, RunArtifactsLayer, RunCheckpointLayer, - RunExecutionLayer, RunGitLayer, RunGoalLayer, RunLayer, RunModelLayer, RunPrepareLayer, - RunPullRequestLayer, RunSandboxLayer, RunScmLayer, ScmGitHubLayer, StringOrSplice, + RunExecutionLayer, RunGitLayer, RunGoalLayer, RunLayer, RunModelControlsLayer, RunModelLayer, + RunPrepareLayer, RunPullRequestLayer, RunSandboxLayer, RunScmLayer, ScmGitHubLayer, + StringOrSplice, }; pub use server::{ DiscordIntegrationLayer, GithubIntegrationLayer, IntegrationWebhooksLayer, diff --git a/lib/crates/fabro-config/src/layers/run.rs b/lib/crates/fabro-config/src/layers/run.rs index 8786315aa..7f02ea7a8 100644 --- a/lib/crates/fabro-config/src/layers/run.rs +++ b/lib/crates/fabro-config/src/layers/run.rs @@ -107,6 +107,21 @@ pub struct RunModelLayer { #[serde(default, skip_serializing_if = "Vec::is_empty")] #[option(default = "[]", value_type = "array")] pub fallbacks: Vec, + /// Default model controls for runs. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub controls: Option, +} + +/// `[run.model.controls]` — run-level default model controls. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +pub struct RunModelControlsLayer { + /// Default reasoning effort for runs. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reasoning_effort: Option, + /// Default speed for runs. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub speed: Option, } /// A single `fallbacks` entry: either a parsed `ModelRef` or the splice marker. diff --git a/lib/crates/fabro-config/src/layers/settings.rs b/lib/crates/fabro-config/src/layers/settings.rs index 4d3401c79..50e81e03c 100644 --- a/lib/crates/fabro-config/src/layers/settings.rs +++ b/lib/crates/fabro-config/src/layers/settings.rs @@ -11,6 +11,7 @@ use serde::{Deserialize, Serialize}; use super::cli::CliLayer; use super::features::FeaturesLayer; +use super::llm::LlmLayer; use super::project::ProjectLayer; use super::run::RunLayer; use super::server::ServerLayer; @@ -29,6 +30,8 @@ pub(crate) struct SettingsLayer { #[serde(default, skip_serializing_if = "Option::is_none")] pub run: Option, #[serde(default, skip_serializing_if = "Option::is_none")] + pub llm: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] pub cli: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub server: Option, diff --git a/lib/crates/fabro-config/src/lib.rs b/lib/crates/fabro-config/src/lib.rs index 1394f08b9..114a03f93 100644 --- a/lib/crates/fabro-config/src/lib.rs +++ b/lib/crates/fabro-config/src/lib.rs @@ -37,19 +37,22 @@ pub use fabro_util::path::expand_tilde; pub use home::Home; pub use layers::{ CliAuthLayer, CliExecAgentLayer, CliExecLayer, CliExecModelLayer, CliLayer, CliLoggingLayer, - CliOutputLayer, CliTargetLayer, CliUpdatesLayer, DaytonaDockerfileLayer, DaytonaSandboxLayer, - DaytonaSnapshotLayer, DiscordIntegrationLayer, DockerSandboxLayer, FeaturesLayer, - GitAuthorLayer, GithubIntegrationLayer, HookAgentMarker, HookEntry, HookTlsMode, - IntegrationWebhooksLayer, InterviewProviderLayer, InterviewsLayer, LocalSandboxLayer, - LogFilter, McpEntryLayer, MergeMap, ModelRefOrSplice, NotificationProviderLayer, - NotificationRouteLayer, ObjectStoreLocalLayer, ObjectStoreS3Layer, PrepareStep, ProjectLayer, - ReplaceMap, RunAgentLayer, RunArtifactsLayer, RunCheckpointLayer, RunExecutionLayer, - RunGitLayer, RunGoalLayer, RunLayer, RunModelLayer, RunPrepareLayer, RunPullRequestLayer, - RunSandboxLayer, RunScmLayer, ScmGitHubLayer, ServerApiLayer, ServerArtifactsLayer, - ServerAuthGithubLayer, ServerAuthLayer, ServerIntegrationsLayer, ServerIpAllowlistLayer, - ServerIpAllowlistOverrideLayer, ServerLayer, ServerListenLayer, ServerLoggingLayer, - ServerSchedulerLayer, ServerSlateDbLayer, ServerStorageLayer, ServerWebLayer, - SlackIntegrationLayer, StickyMap, StringOrSplice, TeamsIntegrationLayer, WorkflowLayer, + CliOutputLayer, CliTargetLayer, CliUpdatesLayer, CostRatesLayer, CredentialRef, + DaytonaDockerfileLayer, DaytonaSandboxLayer, DaytonaSnapshotLayer, DiscordIntegrationLayer, + DockerSandboxLayer, FeaturesLayer, GitAuthorLayer, GithubIntegrationLayer, HookAgentMarker, + HookEntry, HookTlsMode, IntegrationWebhooksLayer, InterviewProviderLayer, InterviewsLayer, + LlmLayer, LocalSandboxLayer, LogFilter, McpEntryLayer, MergeMap, ModelControlsLayer, + ModelCostTableLayer, ModelFeaturesLayer, ModelLimitsLayer, ModelRefOrSplice, + ModelSettingsLayer, NotificationProviderLayer, NotificationRouteLayer, ObjectStoreLocalLayer, + ObjectStoreS3Layer, PrepareStep, ProjectLayer, ProviderSettingsLayer, ReplaceMap, + RunAgentLayer, RunArtifactsLayer, RunCheckpointLayer, RunExecutionLayer, RunGitLayer, + RunGoalLayer, RunLayer, RunModelControlsLayer, RunModelLayer, RunPrepareLayer, + RunPullRequestLayer, RunSandboxLayer, RunScmLayer, ScmGitHubLayer, ServerApiLayer, + ServerArtifactsLayer, ServerAuthGithubLayer, ServerAuthLayer, ServerIntegrationsLayer, + ServerIpAllowlistLayer, ServerIpAllowlistOverrideLayer, ServerLayer, ServerListenLayer, + ServerLoggingLayer, ServerSchedulerLayer, ServerSlateDbLayer, ServerStorageLayer, + ServerWebLayer, SlackIntegrationLayer, StickyMap, StringOrSplice, TeamsIntegrationLayer, + WorkflowLayer, }; pub(crate) use layers::{Combine, SettingsLayer}; pub use logging::{resolve_log_destination, resolve_log_destination_with_env}; diff --git a/lib/crates/fabro-config/src/parse.rs b/lib/crates/fabro-config/src/parse.rs index 50d842b31..1ba258afc 100644 --- a/lib/crates/fabro-config/src/parse.rs +++ b/lib/crates/fabro-config/src/parse.rs @@ -5,7 +5,7 @@ use crate::SettingsLayer; const CURRENT_VERSION: u32 = 1; const ALLOWED_TOP_LEVEL_KEYS: &[&str] = &[ - "_version", "project", "workflow", "run", "cli", "server", "features", + "_version", "project", "workflow", "run", "llm", "cli", "server", "features", ]; #[derive(Debug, Clone, PartialEq, Eq)] @@ -26,7 +26,7 @@ impl fmt::Display for ParseError { } else { write!( f, - "unknown top-level settings key `{key}`: expected one of `_version`, `project`, `workflow`, `run`, `cli`, `server`, `features`" + "unknown top-level settings key `{key}`: expected one of `_version`, `project`, `workflow`, `run`, `llm`, `cli`, `server`, `features`" ) } } @@ -98,7 +98,7 @@ fn rename_hint(key: &str) -> Option { "goal" | "goal_file" | "work_dir" | "directory" => "move to `[run]`", "graph" => "move to `[workflow]`", "labels" => "move to `[run.metadata]`", - "llm" => "rename to `[run.model]`", + // "llm" is now a valid top-level key for provider/model catalog settings. "vars" => "rename to `[run.inputs]`", "setup" => "rename to `[run.prepare]`", "sandbox" => "move under `[run.sandbox]`", diff --git a/lib/crates/fabro-config/src/tests/llm_settings.rs b/lib/crates/fabro-config/src/tests/llm_settings.rs new file mode 100644 index 000000000..73d3e7de0 --- /dev/null +++ b/lib/crates/fabro-config/src/tests/llm_settings.rs @@ -0,0 +1,254 @@ +use crate::{CredentialRef, SettingsLayer}; + +#[test] +fn parses_llm_provider_settings() { + let input = r#" +_version = 1 + +[llm.providers.kimi] +display_name = "Kimi" +adapter = "openai_compatible" +base_url = "https://api.moonshot.ai/v1" +credentials = ["credential:kimi", "env:KIMI_API_KEY"] +priority = 60 +enabled = true +aliases = ["moonshot"] +"#; + + let layer: SettingsLayer = input.parse().unwrap(); + let llm = layer.llm.unwrap(); + let kimi = llm.providers.get("kimi").unwrap(); + + assert_eq!(kimi.display_name.as_deref(), Some("Kimi")); + assert_eq!(kimi.adapter.as_deref(), Some("openai_compatible")); + assert_eq!(kimi.base_url.as_deref(), Some("https://api.moonshot.ai/v1")); + assert_eq!(kimi.credentials.len(), 2); + assert_eq!( + kimi.credentials[0], + CredentialRef::Credential("kimi".to_string()) + ); + assert_eq!( + kimi.credentials[1], + CredentialRef::Env("KIMI_API_KEY".to_string()) + ); + assert_eq!(kimi.priority, Some(60)); + assert_eq!(kimi.enabled, Some(true)); + assert_eq!(kimi.aliases, vec!["moonshot"]); +} + +#[test] +fn parses_llm_model_settings() { + let input = r#" +_version = 1 + +[llm.models."kimi-k2.5"] +provider = "kimi" +api_id = "kimi-k2.5" +display_name = "Kimi K2.5" +family = "kimi" +knowledge_cutoff = 2025-01-01 +default = true +enabled = true +aliases = ["kimi"] +estimated_output_tps = 50.0 + +[llm.models."kimi-k2.5".limits] +context_window = 262144 +max_output = 32768 + +[llm.models."kimi-k2.5".features] +tools = true +vision = false +reasoning = true +effort = false + +[llm.models."kimi-k2.5".costs] +input_cost_per_mtok = 0.60 +output_cost_per_mtok = 2.50 +cache_input_cost_per_mtok = 0.15 + +[llm.models."kimi-k2.5".controls] +reasoning_effort = ["low", "medium", "high"] +"#; + + let layer: SettingsLayer = input.parse().unwrap(); + let llm = layer.llm.unwrap(); + let model = llm.models.get("kimi-k2.5").unwrap(); + + assert_eq!(model.provider.as_deref(), Some("kimi")); + assert_eq!(model.api_id.as_deref(), Some("kimi-k2.5")); + assert_eq!(model.display_name.as_deref(), Some("Kimi K2.5")); + assert_eq!(model.family.as_deref(), Some("kimi")); + assert_eq!(model.default, Some(true)); + assert_eq!(model.enabled, Some(true)); + assert_eq!(model.aliases, vec!["kimi"]); + assert_eq!(model.estimated_output_tps, Some(50.0)); + + let limits = model.limits.as_ref().unwrap(); + assert_eq!(limits.context_window, Some(262_144)); + assert_eq!(limits.max_output, Some(32_768)); + + let features = model.features.as_ref().unwrap(); + assert_eq!(features.tools, Some(true)); + assert_eq!(features.vision, Some(false)); + assert_eq!(features.reasoning, Some(true)); + assert_eq!(features.effort, Some(false)); + + let costs = model.costs.as_ref().unwrap(); + assert_eq!(costs.input_cost_per_mtok, Some(0.60)); + assert_eq!(costs.output_cost_per_mtok, Some(2.50)); + assert_eq!(costs.cache_input_cost_per_mtok, Some(0.15)); + + let controls = model.controls.as_ref().unwrap(); + assert_eq!(controls.reasoning_effort, vec!["low", "medium", "high"]); +} + +#[test] +fn parses_model_speed_costs() { + let input = r#" +_version = 1 + +[llm.models."claude-opus-4-6".costs] +input_cost_per_mtok = 5.0 +output_cost_per_mtok = 25.0 +cache_input_cost_per_mtok = 0.5 + +[llm.models."claude-opus-4-6".costs.speed.fast] +input_cost_per_mtok = 30.0 +output_cost_per_mtok = 150.0 +cache_input_cost_per_mtok = 3.0 + +[llm.models."claude-opus-4-6".controls] +reasoning_effort = ["low", "medium", "high"] +speed = ["fast"] +"#; + + let layer: SettingsLayer = input.parse().unwrap(); + let model = layer + .llm + .unwrap() + .models + .into_inner() + .remove("claude-opus-4-6") + .unwrap(); + + let costs = model.costs.unwrap(); + assert_eq!(costs.input_cost_per_mtok, Some(5.0)); + let fast_costs = &costs.speed["fast"]; + assert_eq!(fast_costs.input_cost_per_mtok, Some(30.0)); + assert_eq!(fast_costs.output_cost_per_mtok, Some(150.0)); + + let controls = model.controls.unwrap(); + assert_eq!(controls.speed, vec!["fast"]); +} + +#[test] +fn rejects_literal_credential_secret() { + let input = r#" +_version = 1 + +[llm.providers.custom] +adapter = "openai_compatible" +credentials = ["sk-secret-key-literal"] +"#; + + let err = input.parse::().unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("credential:"), + "error should mention credential: prefix, got: {msg}" + ); +} + +#[test] +fn rejects_empty_credential_id() { + let input = r#" +_version = 1 + +[llm.providers.custom] +credentials = ["credential:"] +"#; + + let err = input.parse::().unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("non-empty"), + "error should mention non-empty, got: {msg}" + ); +} + +#[test] +fn rejects_empty_env_name() { + let input = r#" +_version = 1 + +[llm.providers.custom] +credentials = ["env:"] +"#; + + let err = input.parse::().unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("non-empty"), + "error should mention non-empty, got: {msg}" + ); +} + +#[test] +fn llm_top_level_key_accepted() { + let input = r#" +_version = 1 + +[llm.providers.test] +adapter = "openai_compatible" +"#; + + let layer: SettingsLayer = input.parse().unwrap(); + assert!(layer.llm.is_some()); +} + +#[test] +fn rejects_unknown_field_under_llm_providers() { + let input = r#" +_version = 1 + +[llm.providers.test] +adapter = "openai_compatible" +unknown_field = "value" +"#; + + let err = input.parse::().unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("unknown"), + "error should mention unknown field, got: {msg}" + ); +} + +#[test] +fn credential_ref_display() { + assert_eq!( + CredentialRef::Credential("kimi".to_string()).to_string(), + "credential:kimi" + ); + assert_eq!( + CredentialRef::Env("KIMI_API_KEY".to_string()).to_string(), + "env:KIMI_API_KEY" + ); +} + +#[test] +fn run_model_controls_parse() { + let input = r#" +_version = 1 + +[run.model.controls] +reasoning_effort = "high" +speed = "fast" +"#; + + let layer: SettingsLayer = input.parse().unwrap(); + let controls = layer.run.unwrap().model.unwrap().controls.unwrap(); + assert_eq!(controls.reasoning_effort.as_deref(), Some("high")); + assert_eq!(controls.speed.as_deref(), Some("fast")); +} diff --git a/lib/crates/fabro-config/src/tests/mod.rs b/lib/crates/fabro-config/src/tests/mod.rs index 593af6efc..94e279ef8 100644 --- a/lib/crates/fabro-config/src/tests/mod.rs +++ b/lib/crates/fabro-config/src/tests/mod.rs @@ -1,5 +1,6 @@ mod combine; mod defaults; +mod llm_settings; mod log_filter; mod resolve_cli; mod resolve_features; diff --git a/lib/crates/fabro-llm/src/model_test.rs b/lib/crates/fabro-llm/src/model_test.rs index 3d98f5016..4d3882d76 100644 --- a/lib/crates/fabro-llm/src/model_test.rs +++ b/lib/crates/fabro-llm/src/model_test.rs @@ -55,7 +55,7 @@ pub async fn run_model_test( async fn run_basic_test(info: &Model, client: Arc) -> ModelTestOutcome { let params = GenerateParams::new(&info.id, client) - .provider(<&'static str>::from(info.provider)) + .provider(info.provider.as_str()) .prompt("Say OK") .max_tokens(16); @@ -123,7 +123,7 @@ fn build_deep_test_params(info: &Model, client: Arc) -> Option::from(info.provider)) + .provider(info.provider.as_str()) .prompt( "Use the add tool twice: first add 15 and 27, then add that result to 42. \ Finally, tell me whether the grand total is even or odd and why.", @@ -159,7 +159,7 @@ fn validate_deep_result(result: &GenerateResult) -> Result<(), String> { mod tests { use std::collections::HashMap; - use fabro_model::{ModelCosts, ModelFeatures, ModelLimits, Provider}; + use fabro_model::{ModelCosts, ModelFeatures, ModelLimits}; use super::*; use crate::types::{FinishReason, Message, Response, StepResult, TokenCounts, ToolResult}; @@ -167,7 +167,7 @@ mod tests { fn test_model_with(features: ModelFeatures) -> Model { Model { id: "test-model".to_string(), - provider: Provider::Anthropic, + provider: fabro_model::ProviderId::from("anthropic"), family: "test".to_string(), display_name: "Test Model".to_string(), limits: ModelLimits { diff --git a/lib/crates/fabro-llm/src/types.rs b/lib/crates/fabro-llm/src/types.rs index 58e785207..b0160dc91 100644 --- a/lib/crates/fabro-llm/src/types.rs +++ b/lib/crates/fabro-llm/src/types.rs @@ -411,28 +411,7 @@ pub struct RateLimitInfo { // --- 3.8 ReasoningEffort --- -#[derive( - Debug, - Clone, - Copy, - PartialEq, - Eq, - Hash, - Serialize, - Deserialize, - strum::Display, - strum::EnumString, - strum::IntoStaticStr, -)] -#[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] -pub enum ReasoningEffort { - Low, - Medium, - High, - XHigh, - Max, -} +pub use fabro_model::ReasoningEffort; // --- 3.6 Request --- diff --git a/lib/crates/fabro-model/Cargo.toml b/lib/crates/fabro-model/Cargo.toml index 70fc3a44c..2b304b421 100644 --- a/lib/crates/fabro-model/Cargo.toml +++ b/lib/crates/fabro-model/Cargo.toml @@ -13,10 +13,11 @@ doctest = false workspace = true [dependencies] +chrono.workspace = true fabro-static.workspace = true serde.workspace = true serde_json.workspace = true strum.workspace = true [dev-dependencies] -insta.workspace = true +insta.workspace = true \ No newline at end of file diff --git a/lib/crates/fabro-model/src/adapter.rs b/lib/crates/fabro-model/src/adapter.rs new file mode 100644 index 000000000..bc5fe446b --- /dev/null +++ b/lib/crates/fabro-model/src/adapter.rs @@ -0,0 +1,133 @@ +use crate::billing::Speed; +use crate::reasoning_effort::ReasoningEffort; + +/// Identifies the kind of agent profile an adapter's models use. +/// +/// This is an internal dispatch key, not a settings field. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum AgentProfileKind { + Anthropic, + OpenAi, + Gemini, +} + +/// How an API key is sent with requests. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ApiKeyHeaderPolicy { + Bearer, + Custom { name: &'static str }, +} + +/// Control capabilities declared by an adapter. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct AdapterControlCapabilities { + pub native_reasoning_effort: &'static [ReasoningEffort], + pub additional_speeds: &'static [Speed], +} + +/// Static metadata for a provider adapter. +/// +/// This is Rust-owned code, not settings data. It describes behavioral +/// contracts that adapters implement. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct AdapterMetadata { + pub key: &'static str, + pub default_profile: AgentProfileKind, + pub api_key_header: ApiKeyHeaderPolicy, + pub controls: AdapterControlCapabilities, +} + +/// Built-in adapter metadata registry. +/// +/// New adapters require Rust code; new providers using existing adapters +/// only require settings data. +pub fn builtin_adapter_metadata() -> &'static [AdapterMetadata] { + static METADATA: &[AdapterMetadata] = &[ + AdapterMetadata { + key: "anthropic", + default_profile: AgentProfileKind::Anthropic, + api_key_header: ApiKeyHeaderPolicy::Custom { name: "x-api-key" }, + controls: AdapterControlCapabilities { + native_reasoning_effort: &[ + ReasoningEffort::Low, + ReasoningEffort::Medium, + ReasoningEffort::High, + ], + additional_speeds: &[Speed::Fast], + }, + }, + AdapterMetadata { + key: "openai", + default_profile: AgentProfileKind::OpenAi, + api_key_header: ApiKeyHeaderPolicy::Bearer, + controls: AdapterControlCapabilities { + native_reasoning_effort: &[ + ReasoningEffort::Low, + ReasoningEffort::Medium, + ReasoningEffort::High, + ], + additional_speeds: &[], + }, + }, + AdapterMetadata { + key: "gemini", + default_profile: AgentProfileKind::Gemini, + api_key_header: ApiKeyHeaderPolicy::Bearer, + controls: AdapterControlCapabilities { + native_reasoning_effort: &[ + ReasoningEffort::Low, + ReasoningEffort::Medium, + ReasoningEffort::High, + ], + additional_speeds: &[], + }, + }, + AdapterMetadata { + key: "openai_compatible", + default_profile: AgentProfileKind::OpenAi, + api_key_header: ApiKeyHeaderPolicy::Bearer, + controls: AdapterControlCapabilities { + native_reasoning_effort: &[], + additional_speeds: &[], + }, + }, + ]; + METADATA +} + +/// Look up adapter metadata by key. +#[must_use] +pub fn adapter_metadata(key: &str) -> Option<&'static AdapterMetadata> { + builtin_adapter_metadata().iter().find(|m| m.key == key) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn builtin_metadata_has_four_adapters() { + assert_eq!(builtin_adapter_metadata().len(), 4); + } + + #[test] + fn lookup_by_key() { + let anthropic = adapter_metadata("anthropic").unwrap(); + assert_eq!(anthropic.default_profile, AgentProfileKind::Anthropic); + assert_eq!(anthropic.api_key_header, ApiKeyHeaderPolicy::Custom { + name: "x-api-key", + }); + } + + #[test] + fn lookup_unknown_key() { + assert!(adapter_metadata("unknown").is_none()); + } + + #[test] + fn openai_compatible_has_empty_controls() { + let compat = adapter_metadata("openai_compatible").unwrap(); + assert!(compat.controls.native_reasoning_effort.is_empty()); + assert!(compat.controls.additional_speeds.is_empty()); + } +} diff --git a/lib/crates/fabro-model/src/billing.rs b/lib/crates/fabro-model/src/billing.rs index f793473de..b8b933fbc 100644 --- a/lib/crates/fabro-model/src/billing.rs +++ b/lib/crates/fabro-model/src/billing.rs @@ -1,7 +1,7 @@ use serde::{Deserialize, Serialize}; use strum::{Display, EnumString, IntoStaticStr}; -use crate::{Model, Provider}; +use crate::{Model, ProviderId}; const TOKENS_PER_MTOK: i128 = 1_000_000; const ANTHROPIC_FAST_MODE_MULTIPLIER_NUMERATOR: i64 = 6; @@ -117,7 +117,7 @@ pub enum Speed { #[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] pub struct ModelRef { - pub provider: Provider, + pub provider: ProviderId, pub model_id: String, #[serde(default, skip_serializing_if = "Option::is_none")] pub speed: Option, @@ -271,16 +271,17 @@ pub enum ModelBillingFacts { impl ModelBillingFacts { #[must_use] - pub fn for_provider(provider: Provider) -> Self { + pub fn for_provider(provider: &str) -> Self { match provider { - Provider::OpenAi => Self::OpenAi(OpenAiBillingFacts::default()), - Provider::OpenAiCompatible => Self::OpenAiCompatible(OpenAiBillingFacts::default()), - Provider::Anthropic => Self::Anthropic(AnthropicBillingFacts::default()), - Provider::Gemini => Self::Gemini(GeminiBillingFacts::default()), - Provider::Kimi => Self::Kimi(OpenAiBillingFacts::default()), - Provider::Zai => Self::Zai(OpenAiBillingFacts::default()), - Provider::Minimax => Self::Minimax(OpenAiBillingFacts::default()), - Provider::Inception => Self::Inception(OpenAiBillingFacts::default()), + "openai" => Self::OpenAi(OpenAiBillingFacts::default()), + "anthropic" => Self::Anthropic(AnthropicBillingFacts::default()), + "gemini" => Self::Gemini(GeminiBillingFacts::default()), + "kimi" => Self::Kimi(OpenAiBillingFacts::default()), + "zai" => Self::Zai(OpenAiBillingFacts::default()), + "minimax" => Self::Minimax(OpenAiBillingFacts::default()), + "inception" => Self::Inception(OpenAiBillingFacts::default()), + // Unknown providers and openai_compatible get OpenAI-compatible billing facts + _ => Self::OpenAiCompatible(OpenAiBillingFacts::default()), } } } @@ -361,7 +362,7 @@ impl Model { #[must_use] pub fn billing_model_ref(&self, speed: Option) -> ModelRef { ModelRef { - provider: self.provider, + provider: self.provider.clone(), model_id: self.id.clone(), speed, } @@ -379,8 +380,10 @@ impl Model { .cache_input_cost_per_mtok .map(PricePerMTok::from_usd); - let (input, output, cached_input) = match (self.provider, speed) { - (Provider::Anthropic, Some(Speed::Fast)) + let provider_str = self.provider.as_str(); + + let (input, output, cached_input) = match (provider_str, speed) { + ("anthropic", Some(Speed::Fast)) if self.id == "claude-opus-4-7" || self.id == "claude-opus-4-6" => { ( @@ -404,20 +407,13 @@ impl Model { _ => return None, }; - let policy = match self.provider { - Provider::OpenAi => ModelPricingPolicy::OpenAi(OpenAiModelPricing { + let policy = match provider_str { + "openai" => ModelPricingPolicy::OpenAi(OpenAiModelPricing { input, cached_input, output, }), - Provider::OpenAiCompatible => { - ModelPricingPolicy::OpenAiCompatible(OpenAiModelPricing { - input, - cached_input, - output, - }) - } - Provider::Anthropic => ModelPricingPolicy::Anthropic(AnthropicModelPricing { + "anthropic" => ModelPricingPolicy::Anthropic(AnthropicModelPricing { input, cache_read: cached_input, cache_write_5m: Some(input.multiply_ratio( @@ -430,28 +426,34 @@ impl Model { )), output, }), - Provider::Gemini => ModelPricingPolicy::Gemini(GeminiModelPricing { + "gemini" => ModelPricingPolicy::Gemini(GeminiModelPricing { input, output, cached_input, storage: None, }), - Provider::Kimi => ModelPricingPolicy::Kimi(OpenAiModelPricing { + "kimi" => ModelPricingPolicy::Kimi(OpenAiModelPricing { input, cached_input, output, }), - Provider::Zai => ModelPricingPolicy::Zai(OpenAiModelPricing { + "zai" => ModelPricingPolicy::Zai(OpenAiModelPricing { input, cached_input, output, }), - Provider::Minimax => ModelPricingPolicy::Minimax(OpenAiModelPricing { + "minimax" => ModelPricingPolicy::Minimax(OpenAiModelPricing { input, cached_input, output, }), - Provider::Inception => ModelPricingPolicy::Inception(OpenAiModelPricing { + "inception" => ModelPricingPolicy::Inception(OpenAiModelPricing { + input, + cached_input, + output, + }), + // Unknown providers get OpenAI-compatible pricing + _ => ModelPricingPolicy::OpenAiCompatible(OpenAiModelPricing { input, cached_input, output, @@ -573,13 +575,13 @@ fn bill_gemini( #[cfg(test)] mod tests { use super::*; - use crate::Catalog; + use crate::{Catalog, ProviderId}; #[test] fn openai_pricing_bills_cached_input_and_reasoning_output() { let pricing = ModelPricing { model: ModelRef { - provider: Provider::OpenAi, + provider: ProviderId::from("openai"), model_id: "gpt-5.4".to_string(), speed: None, }, @@ -632,7 +634,7 @@ mod tests { fn anthropic_billing_supports_distinct_cache_write_buckets() { let pricing = ModelPricing { model: ModelRef { - provider: Provider::Anthropic, + provider: ProviderId::from("anthropic"), model_id: "claude-opus-4-6".to_string(), speed: Some(Speed::Fast), }, @@ -678,7 +680,7 @@ mod tests { fn gemini_billing_requires_storage_pricing_when_storage_facts_exist() { let pricing = ModelPricing { model: ModelRef { - provider: Provider::Gemini, + provider: ProviderId::from("gemini"), model_id: "gemini-3.1-pro-preview".to_string(), speed: None, }, diff --git a/lib/crates/fabro-model/src/catalog.rs b/lib/crates/fabro-model/src/catalog.rs index 5cd208250..9fe9e1964 100644 --- a/lib/crates/fabro-model/src/catalog.rs +++ b/lib/crates/fabro-model/src/catalog.rs @@ -50,7 +50,7 @@ impl Catalog { /// List all models, optionally filtered by provider. #[must_use] - pub fn list(&self, provider: Option) -> Vec<&Model> { + pub fn list(&self, provider: Option<&str>) -> Vec<&Model> { match provider { None => self.models.iter().collect(), Some(p) => self.models.iter().filter(|m| m.provider == p).collect(), @@ -71,7 +71,7 @@ impl Catalog { /// The default model for a specific provider. #[must_use] - pub fn default_for_provider(&self, p: Provider) -> Option<&Model> { + pub fn default_for_provider(&self, p: &str) -> Option<&Model> { self.models.iter().find(|m| m.provider == p && m.default) } @@ -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) + self.default_for_provider(provider.to_string().as_str()) .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) + self.default_for_provider(provider.to_string().as_str()) .unwrap_or_else(|| self.default_model()) } @@ -97,9 +97,9 @@ impl Catalog { /// connectivity checks. Falls back to the provider's default when no /// explicit override is configured. #[must_use] - pub fn probe_for_provider(&self, p: Provider) -> Option<&Model> { + pub fn probe_for_provider(&self, p: &str) -> Option<&Model> { let override_id: Option<&str> = match p { - Provider::OpenAi => Some("gpt-5.4-mini"), + "openai" => Some("gpt-5.4-mini"), _ => None, }; if let Some(id) = override_id { @@ -117,7 +117,7 @@ impl Catalog { /// `features.reasoning`. Among matches, picks the closest by /// `costs.input_cost_per_mtok` (absolute diff). #[must_use] - pub fn closest(&self, target: Provider, reference: &Model) -> Option<&Model> { + pub fn closest(&self, target: &str, reference: &Model) -> Option<&Model> { self.models .iter() .filter(|m| { @@ -144,7 +144,7 @@ impl Catalog { #[must_use] pub fn build_fallback_chain( &self, - primary: Provider, + primary: &str, model: &str, fallbacks: &HashMap>, ) -> Vec { @@ -152,18 +152,18 @@ impl Catalog { return Vec::new(); }; - let Some(fallback_providers) = fallbacks.get(<&'static str>::from(primary)) else { + let Some(fallback_providers) = fallbacks.get(primary) else { return Vec::new(); }; fallback_providers .iter() .filter_map(|provider_str| { - let provider = provider_str.parse::().ok()?; - self.closest(provider, reference).map(|m| FallbackTarget { - provider: provider_str.clone(), - model: m.id.clone(), - }) + self.closest(provider_str, reference) + .map(|m| FallbackTarget { + provider: provider_str.clone(), + model: m.id.clone(), + }) }) .collect() } @@ -175,6 +175,7 @@ mod tests { use super::*; use crate::provider::Provider; + use crate::provider_id::ProviderId; // ---- Catalog struct tests ---- @@ -203,15 +204,15 @@ mod tests { #[test] fn builtin_list_by_provider() { - let anthropic = Catalog::builtin().list(Some(Provider::Anthropic)); + let anthropic = Catalog::builtin().list(Some("anthropic")); assert!(!anthropic.is_empty()); - assert!(anthropic.iter().all(|m| m.provider == Provider::Anthropic)); + assert!(anthropic.iter().all(|m| m.provider == "anthropic")); } #[test] fn builtin_list_unknown_provider_empty() { // OpenAiCompatible has no catalog models - let models = Catalog::builtin().list(Some(Provider::OpenAiCompatible)); + let models = Catalog::builtin().list(Some("openai_compatible")); assert!(models.is_empty()); } @@ -224,61 +225,47 @@ mod tests { #[test] fn builtin_default_for_provider() { let m = Catalog::builtin() - .default_for_provider(Provider::Anthropic) + .default_for_provider("anthropic") .unwrap(); assert_eq!(m.id, "claude-sonnet-4-6"); assert!(m.default); - let m = Catalog::builtin() - .default_for_provider(Provider::OpenAi) - .unwrap(); + let m = Catalog::builtin().default_for_provider("openai").unwrap(); assert_eq!(m.id, "gpt-5.4"); - let m = Catalog::builtin() - .default_for_provider(Provider::Gemini) - .unwrap(); + let m = Catalog::builtin().default_for_provider("gemini").unwrap(); assert_eq!(m.id, "gemini-3.1-pro-preview"); } #[test] fn builtin_probe_openai_returns_override() { - let m = Catalog::builtin() - .probe_for_provider(Provider::OpenAi) - .unwrap(); + let m = Catalog::builtin().probe_for_provider("openai").unwrap(); assert_eq!(m.id, "gpt-5.4-mini"); } #[test] fn builtin_probe_anthropic_returns_default() { - let m = Catalog::builtin() - .probe_for_provider(Provider::Anthropic) - .unwrap(); + let m = Catalog::builtin().probe_for_provider("anthropic").unwrap(); assert_eq!(m.id, "claude-sonnet-4-6"); } #[test] fn builtin_probe_gemini_returns_default() { - let m = Catalog::builtin() - .probe_for_provider(Provider::Gemini) - .unwrap(); + let m = Catalog::builtin().probe_for_provider("gemini").unwrap(); assert_eq!(m.id, "gemini-3.1-pro-preview"); } #[test] fn builtin_closest_opus_to_gemini() { let opus = Catalog::builtin().get("claude-opus-4-6").unwrap(); - let result = Catalog::builtin().closest(Provider::Gemini, opus).unwrap(); + let result = Catalog::builtin().closest("gemini", opus).unwrap(); assert_eq!(result.id, "gemini-3.1-pro-preview"); } #[test] fn builtin_closest_no_match() { let haiku = Catalog::builtin().get("claude-haiku-4-5").unwrap(); - assert!( - Catalog::builtin() - .closest(Provider::OpenAi, haiku) - .is_none() - ); + assert!(Catalog::builtin().closest("openai", haiku).is_none()); } #[test] @@ -287,11 +274,8 @@ mod tests { "gemini".to_string(), "openai".to_string(), ])]); - let chain = Catalog::builtin().build_fallback_chain( - Provider::Anthropic, - "claude-opus-4-6", - &fallbacks, - ); + let chain = + Catalog::builtin().build_fallback_chain("anthropic", "claude-opus-4-6", &fallbacks); assert_eq!(chain.len(), 2); assert_eq!(chain[0].provider, "gemini"); assert_eq!(chain[0].model, "gemini-3.1-pro-preview"); @@ -302,19 +286,15 @@ mod tests { #[test] fn builtin_build_fallback_chain_unknown_model() { let fallbacks = HashMap::from([("anthropic".to_string(), vec!["gemini".to_string()])]); - let chain = - Catalog::builtin().build_fallback_chain(Provider::Anthropic, "unknown-xyz", &fallbacks); + let chain = Catalog::builtin().build_fallback_chain("anthropic", "unknown-xyz", &fallbacks); assert!(chain.is_empty()); } #[test] fn builtin_build_fallback_chain_provider_not_in_map() { let fallbacks = HashMap::from([("openai".to_string(), vec!["anthropic".to_string()])]); - let chain = Catalog::builtin().build_fallback_chain( - Provider::Anthropic, - "claude-opus-4-6", - &fallbacks, - ); + let chain = + Catalog::builtin().build_fallback_chain("anthropic", "claude-opus-4-6", &fallbacks); assert!(chain.is_empty()); } @@ -324,11 +304,8 @@ mod tests { "openai".to_string(), "kimi".to_string(), ])]); - let chain = Catalog::builtin().build_fallback_chain( - Provider::Anthropic, - "claude-haiku-4-5", - &fallbacks, - ); + let chain = + Catalog::builtin().build_fallback_chain("anthropic", "claude-haiku-4-5", &fallbacks); assert_eq!(chain.len(), 1); assert_eq!(chain[0].provider, "kimi"); assert_eq!(chain[0].model, "kimi-k2.5"); @@ -337,11 +314,8 @@ mod tests { #[test] fn builtin_build_fallback_chain_empty_map() { let fallbacks = HashMap::new(); - let chain = Catalog::builtin().build_fallback_chain( - Provider::Anthropic, - "claude-opus-4-6", - &fallbacks, - ); + let chain = + Catalog::builtin().build_fallback_chain("anthropic", "claude-opus-4-6", &fallbacks); assert!(chain.is_empty()); } @@ -351,7 +325,7 @@ mod tests { let models = vec![Model { id: "test-model".to_string(), - provider: Provider::Anthropic, + provider: ProviderId::from("anthropic"), family: "test".to_string(), display_name: "Test Model".to_string(), limits: ModelLimits { @@ -390,7 +364,8 @@ mod tests { #[test] fn every_provider_has_catalog_models() { for &provider in Provider::ALL { - let models = Catalog::builtin().list(Some(provider)); + let provider_str = provider.to_string(); + let models = Catalog::builtin().list(Some(&provider_str)); assert!( !models.is_empty(), "Provider {provider:?} has no models in catalog" @@ -401,8 +376,9 @@ mod tests { #[test] fn every_provider_has_exactly_one_default_model() { for &provider in Provider::ALL { + let provider_str = provider.to_string(); let defaults: Vec<_> = Catalog::builtin() - .list(Some(provider)) + .list(Some(&provider_str)) .into_iter() .filter(|m| m.default) .collect(); @@ -418,13 +394,12 @@ mod tests { } #[test] - fn catalog_providers_roundtrip_through_static_str() { + fn catalog_providers_roundtrip_through_provider_enum() { for model in Catalog::builtin().list(None) { - let roundtripped = Provider::from_str(<&'static str>::from(model.provider)); - assert_eq!( - roundtripped, - Ok(model.provider), - "catalog model '{}' provider {:?} does not roundtrip through IntoStaticStr", + let roundtripped = Provider::from_str(model.provider.as_str()); + assert!( + roundtripped.is_ok(), + "catalog model '{}' provider {:?} does not parse as a Provider enum", model.id, model.provider ); @@ -451,7 +426,9 @@ mod tests { insta::assert_debug_snapshot!(info, @r#" Model { id: "claude-opus-4-6", - provider: Anthropic, + provider: ProviderId( + "anthropic", + ), family: "claude-4", display_name: "Claude Opus 4.6", limits: ModelLimits { @@ -519,7 +496,9 @@ mod tests { insta::assert_debug_snapshot!(m, @r#" Model { id: "gemini-3.1-flash-lite-preview", - provider: Gemini, + provider: ProviderId( + "gemini", + ), family: "gemini-3", display_name: "Gemini 3.1 Flash Lite (Preview)", limits: ModelLimits { @@ -577,7 +556,9 @@ mod tests { insta::assert_debug_snapshot!(m, @r#" Model { id: "kimi-k2.5", - provider: Kimi, + provider: ProviderId( + "kimi", + ), family: "kimi-k2", display_name: "Kimi K2.5", limits: ModelLimits { @@ -627,13 +608,13 @@ mod tests { #[test] fn glm_4_7_in_catalog() { let m = Catalog::builtin().get("glm-4.7").unwrap(); - assert_eq!(m.provider, Provider::Zai); + assert_eq!(m.provider, "zai"); } #[test] fn minimax_m2_5_in_catalog() { let m = Catalog::builtin().get("minimax-m2.5").unwrap(); - assert_eq!(m.provider, Provider::Minimax); + assert_eq!(m.provider, "minimax"); } #[test] @@ -642,7 +623,9 @@ mod tests { insta::assert_debug_snapshot!(m, @r#" Model { id: "mercury-2", - provider: Inception, + provider: ProviderId( + "inception", + ), family: "mercury", display_name: "Mercury 2", limits: ModelLimits { @@ -691,7 +674,9 @@ mod tests { insta::assert_debug_snapshot!(m, @r#" Model { id: "gpt-5.4", - provider: OpenAi, + provider: ProviderId( + "openai", + ), family: "gpt-5", display_name: "GPT-5.4", limits: ModelLimits { @@ -742,7 +727,9 @@ mod tests { insta::assert_debug_snapshot!(m, @r#" Model { id: "gpt-5.4-pro", - provider: OpenAi, + provider: ProviderId( + "openai", + ), family: "gpt-5", display_name: "GPT-5.4 Pro", limits: ModelLimits { @@ -819,7 +806,9 @@ mod tests { insta::assert_debug_snapshot!(m, @r#" Model { id: "gpt-5.3-codex-spark", - provider: OpenAi, + provider: ProviderId( + "openai", + ), family: "gpt-5", display_name: "GPT-5.3 Codex Spark", limits: ModelLimits { @@ -870,23 +859,21 @@ mod tests { #[test] fn closest_model_sonnet_to_gemini() { let sonnet = Catalog::builtin().get("claude-sonnet-4-5").unwrap(); - let result = Catalog::builtin() - .closest(Provider::Gemini, sonnet) - .unwrap(); + let result = Catalog::builtin().closest("gemini", sonnet).unwrap(); assert_eq!(result.id, "gemini-3.1-pro-preview"); } #[test] fn closest_model_haiku_to_kimi() { let haiku = Catalog::builtin().get("claude-haiku-4-5").unwrap(); - let result = Catalog::builtin().closest(Provider::Kimi, haiku).unwrap(); + let result = Catalog::builtin().closest("kimi", haiku).unwrap(); assert_eq!(result.id, "kimi-k2.5"); } #[test] fn closest_model_no_capability_match() { let glm = Catalog::builtin().get("glm-4.7").unwrap(); - assert!(Catalog::builtin().closest(Provider::Gemini, glm).is_none()); + assert!(Catalog::builtin().closest("gemini", glm).is_none()); } // ---- Cost tests ---- diff --git a/lib/crates/fabro-model/src/lib.rs b/lib/crates/fabro-model/src/lib.rs index 9a95e5815..f333b287e 100644 --- a/lib/crates/fabro-model/src/lib.rs +++ b/lib/crates/fabro-model/src/lib.rs @@ -1,10 +1,17 @@ +pub mod adapter; pub mod billing; pub mod catalog; pub mod model_ref; pub mod model_test; pub mod provider; +pub mod provider_id; +pub mod reasoning_effort; pub mod types; +pub use adapter::{ + AdapterControlCapabilities, AdapterMetadata, AgentProfileKind, ApiKeyHeaderPolicy, + adapter_metadata, builtin_adapter_metadata, +}; pub use billing::{ AnthropicBillingFacts, AnthropicModelPricing, BilledModelUsage, BilledTokenCounts, GeminiBillingFacts, GeminiModelPricing, GeminiStoragePricing, GeminiStorageSegment, @@ -15,4 +22,6 @@ pub use catalog::{Catalog, FallbackTarget}; pub use model_ref::ModelHandle; pub use model_test::ModelTestMode; pub use provider::Provider; +pub use provider_id::{ModelId, ProviderId}; +pub use reasoning_effort::ReasoningEffort; pub use types::{Model, ModelCosts, ModelFeatures, ModelLimits}; diff --git a/lib/crates/fabro-model/src/model_ref.rs b/lib/crates/fabro-model/src/model_ref.rs index 96cbced9b..3b9aac245 100644 --- a/lib/crates/fabro-model/src/model_ref.rs +++ b/lib/crates/fabro-model/src/model_ref.rs @@ -1,7 +1,7 @@ use std::fmt; use std::sync::Arc; -use crate::provider::Provider; +use crate::provider_id::ProviderId; use crate::types::Model; /// A reference to a model — either a fully resolved `Model` or a @@ -12,7 +12,7 @@ pub enum ModelHandle { Resolved(Arc), /// An unresolved provider:model pair (e.g. from CLI input or config). ByName { - provider: Provider, + provider: ProviderId, model: String, }, } @@ -29,10 +29,10 @@ impl ModelHandle { /// The provider for this model. #[must_use] - pub fn provider(&self) -> Provider { + pub fn provider(&self) -> &ProviderId { match self { - Self::Resolved(m) => m.provider, - Self::ByName { provider, .. } => *provider, + Self::Resolved(m) => &m.provider, + Self::ByName { provider, .. } => provider, } } } @@ -64,7 +64,7 @@ mod tests { #[test] fn by_name_display() { let r = ModelHandle::ByName { - provider: Provider::Anthropic, + provider: ProviderId::from("anthropic"), model: "claude-opus-4-6".to_string(), }; assert_eq!(r.to_string(), "anthropic:claude-opus-4-6"); @@ -73,11 +73,11 @@ mod tests { #[test] fn by_name_accessors() { let r = ModelHandle::ByName { - provider: Provider::OpenAi, + provider: ProviderId::from("openai"), model: "gpt-5.4".to_string(), }; assert_eq!(r.model_id(), "gpt-5.4"); - assert_eq!(r.provider(), Provider::OpenAi); + assert_eq!(r.provider(), "openai"); } #[test] @@ -92,17 +92,17 @@ mod tests { let info = Catalog::builtin().get("gpt-5.4").unwrap().clone(); let r = ModelHandle::Resolved(Arc::new(info)); assert_eq!(r.model_id(), "gpt-5.4"); - assert_eq!(r.provider(), Provider::OpenAi); + assert_eq!(r.provider(), "openai"); } #[test] fn debug_format() { let r = ModelHandle::ByName { - provider: Provider::Gemini, + provider: ProviderId::from("gemini"), model: "gemini-3.1-pro-preview".to_string(), }; let debug = format!("{r:?}"); assert!(debug.contains("ByName")); - assert!(debug.contains("Gemini")); + assert!(debug.contains("gemini")); } } diff --git a/lib/crates/fabro-model/src/provider_id.rs b/lib/crates/fabro-model/src/provider_id.rs new file mode 100644 index 000000000..f35fcd824 --- /dev/null +++ b/lib/crates/fabro-model/src/provider_id.rs @@ -0,0 +1,200 @@ +use std::fmt; +use std::str::FromStr; + +use serde::{Deserialize, Serialize}; + +/// A string-backed provider identifier. +/// +/// Unlike the closed `Provider` enum, `ProviderId` can represent any +/// provider — built-in or user-defined through settings. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(transparent)] +pub struct ProviderId(String); + +impl ProviderId { + /// Create a new provider ID from a string. + #[must_use] + pub fn new(id: impl Into) -> Self { + Self(id.into()) + } + + /// The string value of this provider ID. + #[must_use] + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl fmt::Display for ProviderId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.0) + } +} + +impl FromStr for ProviderId { + type Err = std::convert::Infallible; + + fn from_str(s: &str) -> Result { + Ok(Self(s.to_string())) + } +} + +impl From<&str> for ProviderId { + fn from(s: &str) -> Self { + Self(s.to_string()) + } +} + +impl From for ProviderId { + fn from(s: String) -> Self { + Self(s) + } +} + +impl AsRef for ProviderId { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl PartialEq for ProviderId { + fn eq(&self, other: &str) -> bool { + self.0 == other + } +} + +impl PartialEq<&str> for ProviderId { + fn eq(&self, other: &&str) -> bool { + self.0 == *other + } +} + +/// Convert from the legacy `Provider` enum for migration compatibility. +impl From for ProviderId { + fn from(p: crate::Provider) -> Self { + Self(p.to_string()) + } +} + +/// A string-backed model identifier. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(transparent)] +pub struct ModelId(String); + +impl ModelId { + /// Create a new model ID from a string. + #[must_use] + pub fn new(id: impl Into) -> Self { + Self(id.into()) + } + + /// The string value of this model ID. + #[must_use] + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl fmt::Display for ModelId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.0) + } +} + +impl FromStr for ModelId { + type Err = std::convert::Infallible; + + fn from_str(s: &str) -> Result { + Ok(Self(s.to_string())) + } +} + +impl From<&str> for ModelId { + fn from(s: &str) -> Self { + Self(s.to_string()) + } +} + +impl From for ModelId { + fn from(s: String) -> Self { + Self(s) + } +} + +impl AsRef for ModelId { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl PartialEq for ModelId { + fn eq(&self, other: &str) -> bool { + self.0 == other + } +} + +impl PartialEq<&str> for ModelId { + fn eq(&self, other: &&str) -> bool { + self.0 == *other + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn provider_id_from_str() { + let id: ProviderId = "anthropic".parse().unwrap(); + assert_eq!(id.as_str(), "anthropic"); + } + + #[test] + fn provider_id_display() { + let id = ProviderId::new("openai"); + assert_eq!(id.to_string(), "openai"); + } + + #[test] + fn provider_id_serde_roundtrip() { + let id = ProviderId::new("kimi"); + let json = serde_json::to_string(&id).unwrap(); + assert_eq!(json, "\"kimi\""); + let parsed: ProviderId = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, id); + } + + #[test] + fn provider_id_from_legacy_provider() { + let id = ProviderId::from(crate::Provider::Anthropic); + assert_eq!(id.as_str(), "anthropic"); + } + + #[test] + fn provider_id_eq_str() { + let id = ProviderId::new("anthropic"); + assert_eq!(id, "anthropic"); + assert_eq!(id, *"anthropic"); + } + + #[test] + fn model_id_from_str() { + let id: ModelId = "claude-opus-4-6".parse().unwrap(); + assert_eq!(id.as_str(), "claude-opus-4-6"); + } + + #[test] + fn model_id_display() { + let id = ModelId::new("gpt-5.4"); + assert_eq!(id.to_string(), "gpt-5.4"); + } + + #[test] + fn model_id_serde_roundtrip() { + let id = ModelId::new("gemini-3.1-pro-preview"); + let json = serde_json::to_string(&id).unwrap(); + assert_eq!(json, "\"gemini-3.1-pro-preview\""); + let parsed: ModelId = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, id); + } +} diff --git a/lib/crates/fabro-model/src/reasoning_effort.rs b/lib/crates/fabro-model/src/reasoning_effort.rs new file mode 100644 index 000000000..798382f16 --- /dev/null +++ b/lib/crates/fabro-model/src/reasoning_effort.rs @@ -0,0 +1,78 @@ +use serde::{Deserialize, Serialize}; +use strum::{Display, EnumString, IntoStaticStr}; + +/// Reasoning effort level for models that support native effort control. +/// +/// Values are code-owned; adding a new level is a Rust change. +#[derive( + Debug, + Clone, + Copy, + PartialEq, + Eq, + Hash, + PartialOrd, + Ord, + Serialize, + Deserialize, + Display, + EnumString, + IntoStaticStr, +)] +#[serde(rename_all = "lowercase")] +#[strum(serialize_all = "lowercase")] +pub enum ReasoningEffort { + Low, + Medium, + High, + XHigh, + Max, +} + +#[cfg(test)] +mod tests { + use std::str::FromStr; + + use super::*; + + #[test] + fn from_str_round_trip() { + assert_eq!(ReasoningEffort::from_str("low"), Ok(ReasoningEffort::Low)); + assert_eq!( + ReasoningEffort::from_str("medium"), + Ok(ReasoningEffort::Medium) + ); + assert_eq!(ReasoningEffort::from_str("high"), Ok(ReasoningEffort::High)); + assert_eq!( + ReasoningEffort::from_str("xhigh"), + Ok(ReasoningEffort::XHigh) + ); + assert_eq!(ReasoningEffort::from_str("max"), Ok(ReasoningEffort::Max)); + assert_eq!(ReasoningEffort::XHigh.to_string(), "xhigh"); + assert_eq!(<&'static str>::from(ReasoningEffort::XHigh), "xhigh"); + assert_eq!(ReasoningEffort::Max.to_string(), "max"); + assert_eq!(<&'static str>::from(ReasoningEffort::Max), "max"); + } + + #[test] + fn from_str_rejects_unknown() { + assert!(ReasoningEffort::from_str("bogus").is_err()); + } + + #[test] + fn serde_roundtrip() { + let effort = ReasoningEffort::High; + let json = serde_json::to_string(&effort).unwrap(); + assert_eq!(json, "\"high\""); + let parsed: ReasoningEffort = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, effort); + } + + #[test] + fn ord_ordering() { + assert!(ReasoningEffort::Low < ReasoningEffort::Medium); + assert!(ReasoningEffort::Medium < ReasoningEffort::High); + assert!(ReasoningEffort::High < ReasoningEffort::XHigh); + assert!(ReasoningEffort::XHigh < ReasoningEffort::Max); + } +} diff --git a/lib/crates/fabro-model/src/types.rs b/lib/crates/fabro-model/src/types.rs index cd96efca8..dbd9c8f6e 100644 --- a/lib/crates/fabro-model/src/types.rs +++ b/lib/crates/fabro-model/src/types.rs @@ -1,6 +1,6 @@ use serde::{Deserialize, Serialize}; -use crate::provider::Provider; +use crate::provider_id::ProviderId; // --- 2.9 Model --- @@ -34,7 +34,7 @@ pub struct ModelCosts { #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct Model { pub id: String, - pub provider: Provider, + pub provider: ProviderId, pub family: String, pub display_name: String, pub limits: ModelLimits, @@ -58,8 +58,8 @@ impl Model { &self.id } - pub fn provider(&self) -> Provider { - self.provider + pub fn provider(&self) -> &ProviderId { + &self.provider } pub fn family(&self) -> &str { @@ -130,13 +130,12 @@ impl Model { #[cfg(test)] mod tests { use crate::catalog::Catalog; - use crate::provider::Provider; #[test] fn inherent_methods_return_correct_values() { let info = Catalog::builtin().get("claude-opus-4-7").unwrap(); assert_eq!(info.id(), "claude-opus-4-7"); - assert_eq!(info.provider(), Provider::Anthropic); + assert_eq!(info.provider(), "anthropic"); assert_eq!(info.family(), "claude-4"); assert_eq!(info.display_name(), "Claude Opus 4.7"); assert_eq!(info.context_window(), 1_000_000); diff --git a/lib/crates/fabro-server/src/diagnostics.rs b/lib/crates/fabro-server/src/diagnostics.rs index aa3cddc08..29a23ca85 100644 --- a/lib/crates/fabro-server/src/diagnostics.rs +++ b/lib/crates/fabro-server/src/diagnostics.rs @@ -148,7 +148,7 @@ async fn check_llm_providers(state: &AppState) -> CheckResult { fn probe_model(provider: Provider) -> String { Catalog::builtin() - .probe_for_provider(provider) + .probe_for_provider(&provider.to_string()) .map_or_else(|| format!("unknown-{provider}"), |m| m.id.clone()) } diff --git a/lib/crates/fabro-server/src/run_manifest.rs b/lib/crates/fabro-server/src/run_manifest.rs index 5aeea806b..c8c7b9fa2 100644 --- a/lib/crates/fabro-server/src/run_manifest.rs +++ b/lib/crates/fabro-server/src/run_manifest.rs @@ -313,6 +313,7 @@ fn manifest_args_overrides(args: Option<&types::ManifestArgs>) -> ManifestSettin provider: args.provider.as_deref().map(InterpString::parse), name: args.model.as_deref().map(InterpString::parse), fallbacks: Vec::new(), + controls: None, }); let local_worktree = args .worktree_mode diff --git a/lib/crates/fabro-server/src/server/handler/models.rs b/lib/crates/fabro-server/src/server/handler/models.rs index 00ad1449e..a4c190385 100644 --- a/lib/crates/fabro-server/src/server/handler/models.rs +++ b/lib/crates/fabro-server/src/server/handler/models.rs @@ -35,32 +35,21 @@ async fn list_models( State(state): State>, Query(params): Query, ) -> Response { - let provider = match params.provider.as_deref() { - Some(value) => match Provider::from_str(value) { - Ok(provider) => Some(provider), - Err(_) => { - return ApiError::new( - StatusCode::BAD_REQUEST, - format!("unknown provider: {value}"), - ) - .into_response(); - } - }, - None => None, - }; + let provider_filter = params.provider.as_deref(); let query = params.query.as_ref().map(|value| value.to_lowercase()); let limit = params.limit.clamp(1, 100) as usize; let offset = params.offset.min(MAX_PAGE_OFFSET) as usize; - let configured: HashSet = state + let configured: HashSet = state .llm_source .configured_providers() .await .into_iter() + .map(|p| p.to_string()) .collect(); let mut models = fabro_model::Catalog::builtin() - .list(provider) + .list(provider_filter) .into_iter() .filter(|model| match &query { Some(query) => { @@ -75,7 +64,7 @@ async fn list_models( }) .cloned() .map(|mut model| { - model.configured = configured.contains(&model.provider); + model.configured = configured.contains(model.provider.as_str()); model }) .collect::>(); @@ -130,11 +119,13 @@ async fn test_model( if let Some((_, issue)) = llm_result .auth_issues .iter() - .find(|(provider, _)| *provider == info.provider) + .find(|(provider, _)| provider.to_string() == info.provider.as_str()) { - return ApiError::bad_request(auth_issue_message(info.provider, issue)).into_response(); + 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(); } - let provider_name = <&'static str>::from(info.provider); + let provider_name = info.provider.as_str(); if !llm_result.client.provider_names().contains(&provider_name) { return Json(serde_json::json!({ "model_id": info.id, diff --git a/lib/crates/fabro-server/src/server/tests.rs b/lib/crates/fabro-server/src/server/tests.rs index 6078c362f..446ccbc11 100644 --- a/lib/crates/fabro-server/src/server/tests.rs +++ b/lib/crates/fabro-server/src/server/tests.rs @@ -2693,7 +2693,7 @@ async fn list_models_marks_configured_false_when_no_credential_material() { } #[tokio::test] -async fn list_models_invalid_provider_returns_400() { +async fn list_models_unknown_provider_returns_empty_list() { let app = test_app_with(); let req = Request::builder() @@ -2703,7 +2703,13 @@ async fn list_models_invalid_provider_returns_400() { .unwrap(); let response = app.oneshot(req).await.unwrap(); - assert_status!(response, StatusCode::BAD_REQUEST).await; + let body = checked_response!(response, StatusCode::OK).await; + let bytes = axum::body::to_bytes(body.into_body(), usize::MAX) + .await + .unwrap(); + let json: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + let data = json["data"].as_array().expect("data should be an array"); + assert!(data.is_empty(), "unknown provider should return empty list"); } #[tokio::test] diff --git a/lib/crates/fabro-server/tests/it/scenario/usage.rs b/lib/crates/fabro-server/tests/it/scenario/usage.rs index 976e0bf69..8e91f8c59 100644 --- a/lib/crates/fabro-server/tests/it/scenario/usage.rs +++ b/lib/crates/fabro-server/tests/it/scenario/usage.rs @@ -150,5 +150,5 @@ fn assert_non_llm_billing(billing: &serde_json::Value, expected_stage_ids: &[&st let total_runtime_secs = billing["totals"]["runtime_secs"] .as_f64() .expect("totals should include runtime_secs"); - assert_eq!(total_runtime_secs, runtime_secs); + assert!((total_runtime_secs - runtime_secs).abs() < f64::EPSILON); } diff --git a/lib/crates/fabro-workflow/src/operations/start.rs b/lib/crates/fabro-workflow/src/operations/start.rs index 4e623bcec..f1f5b60b0 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, model, &by_provider) + Catalog::builtin().build_fallback_chain(&provider.to_string(), 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 6fad164aa..78885c9d6 100644 --- a/lib/crates/fabro-workflow/src/outcome.rs +++ b/lib/crates/fabro-workflow/src/outcome.rs @@ -20,13 +20,15 @@ pub fn billed_model_usage_from_llm( usage: &LlmTokenCounts, ) -> BilledModelUsage { let speed = parse_speed(requested_speed); + let provider_id = fabro_model::ProviderId::from(provider); let model = ModelRef { - provider, + provider: provider_id, model_id: model_id.to_string(), speed, }; let tokens = token_counts_from_llm_usage(usage); - let facts = billing_facts_for_stage_usage(provider, &tokens); + let provider_str = provider.to_string(); + let facts = billing_facts_for_stage_usage(&provider_str, &tokens); let input = ModelBillingInput { usage: ModelUsage { model: model.clone(), @@ -37,7 +39,7 @@ pub fn billed_model_usage_from_llm( let total_usd_micros = Catalog::builtin() .get(model_id) - .filter(|candidate| candidate.provider == provider) + .filter(|candidate| candidate.provider == provider_str.as_str()) .and_then(|candidate| candidate.pricing_for(speed)) .and_then(|pricing| pricing.bill(&input)) .map(|amount| amount.0); @@ -144,9 +146,9 @@ fn token_counts_from_llm_usage(usage: &LlmTokenCounts) -> TokenCounts { usage.clone() } -fn billing_facts_for_stage_usage(provider: Provider, tokens: &TokenCounts) -> ModelBillingFacts { +fn billing_facts_for_stage_usage(provider: &str, tokens: &TokenCounts) -> ModelBillingFacts { match provider { - Provider::Anthropic => ModelBillingFacts::Anthropic(AnthropicBillingFacts { + "anthropic" => ModelBillingFacts::Anthropic(AnthropicBillingFacts { cache_write_5m_tokens: tokens.cache_write_tokens, cache_write_1h_tokens: 0, }), diff --git a/lib/crates/fabro-workflow/src/run_materialization.rs b/lib/crates/fabro-workflow/src/run_materialization.rs index 62806c3e6..d911808d1 100644 --- a/lib/crates/fabro-workflow/src/run_materialization.rs +++ b/lib/crates/fabro-workflow/src/run_materialization.rs @@ -38,7 +38,7 @@ pub fn materialize_run( provider .as_deref() .and_then(|value| value.parse::().ok()) - .and_then(|provider| catalog.default_for_provider(provider)) + .and_then(|provider| catalog.default_for_provider(&provider.to_string())) .unwrap_or_else(|| catalog.default_for_configured(configured_providers)) .id .clone()