diff --git a/crates/agent/src/cli.rs b/crates/agent/src/cli.rs index 2acab319a..b50082521 100644 --- a/crates/agent/src/cli.rs +++ b/crates/agent/src/cli.rs @@ -5,6 +5,7 @@ use crate::{ }; use clap::{Parser, ValueEnum}; use llm::client::Client; +use llm::provider::{ModelId, Provider}; use std::io::{IsTerminal, Write}; use std::path::PathBuf; use std::sync::{Arc, Mutex}; @@ -63,12 +64,11 @@ enum PermissionLevel { Full, } -fn default_model(provider: &str) -> &'static str { +fn default_model(provider: Provider) -> &'static str { match provider { - "openai" => "gpt-5.2-codex", - "gemini" => "gemini-3.1-pro-preview", - // anthropic and unknown providers - _ => "claude-opus-4-6", + Provider::OpenAi => "gpt-5.2-codex", + Provider::Gemini => "gemini-3.1-pro-preview", + Provider::Anthropic => "claude-opus-4-6", } } @@ -142,44 +142,38 @@ fn build_tool_approval( }) } -fn build_summarizer(provider: &str, llm_client: Option) -> Option { - let client = llm_client?; - let model = match provider { - "openai" => "gpt-4o-mini", - "gemini" => "gemini-2.0-flash", - // anthropic and unknown providers - _ => "claude-haiku-4-5-20251001", - }; - let provider_name = match provider { - "openai" => Some("openai".to_string()), - "gemini" => Some("gemini".to_string()), - _ => None, - }; - Some(crate::tools::WebFetchSummarizer { - client, - model: model.into(), - provider: provider_name, - }) -} - -fn build_profile(provider: &str, model: &str, llm_client: Option) -> Box { - let summarizer = build_summarizer(provider, llm_client); +fn summarizer_model_id(provider: Provider) -> ModelId { match provider { - "openai" => Box::new(OpenAiProfile::with_summarizer(model, summarizer)), - "gemini" => Box::new(GeminiProfile::with_summarizer(model, summarizer)), - // anthropic and unknown providers - _ => Box::new(AnthropicProfile::with_summarizer(model, summarizer)), + Provider::OpenAi => ModelId::new(Provider::OpenAi, "gpt-4o-mini"), + Provider::Gemini => ModelId::new(Provider::Gemini, "gemini-2.0-flash"), + Provider::Anthropic => ModelId::new(Provider::Anthropic, "claude-haiku-4-5-20251001"), } } -fn validate_api_key(provider: &str) -> bool { +fn build_summarizer(provider: Provider, llm_client: Option) -> Option { + let client = llm_client?; + Some(crate::tools::WebFetchSummarizer { + client, + model_id: summarizer_model_id(provider), + }) +} + +fn build_profile(provider: Provider, model: &str, llm_client: Option) -> Box { + let summarizer = build_summarizer(provider, llm_client); match provider { - "anthropic" => std::env::var("ANTHROPIC_API_KEY").is_ok(), - "openai" => std::env::var("OPENAI_API_KEY").is_ok(), - "gemini" => { + Provider::OpenAi => Box::new(OpenAiProfile::with_summarizer(model, summarizer)), + Provider::Gemini => Box::new(GeminiProfile::with_summarizer(model, summarizer)), + Provider::Anthropic => Box::new(AnthropicProfile::with_summarizer(model, summarizer)), + } +} + +fn validate_api_key(provider: Provider) -> bool { + match provider { + Provider::Anthropic => std::env::var("ANTHROPIC_API_KEY").is_ok(), + Provider::OpenAi => std::env::var("OPENAI_API_KEY").is_ok(), + Provider::Gemini => { std::env::var("GEMINI_API_KEY").is_ok() || std::env::var("GOOGLE_API_KEY").is_ok() } - _ => false, } } @@ -331,9 +325,12 @@ pub async fn run() -> anyhow::Result<()> { // Resolve color support once, leak to get 'static lifetime for use across threads let styles: &'static Styles = Box::leak(Box::new(Styles::detect_stderr())); + // Parse provider string to enum early for compile-time safety + let provider: Provider = cli.provider.parse().map_err(|e: String| anyhow::anyhow!("{e}"))?; + // Validate provider API key - if !validate_api_key(&cli.provider) { - anyhow::bail!("API key not set for provider '{}'", cli.provider); + if !validate_api_key(provider) { + anyhow::bail!("API key not set for provider '{provider}'"); } // Build LLM client @@ -351,12 +348,12 @@ pub async fn run() -> anyhow::Result<()> { let model = cli .model .as_deref() - .unwrap_or_else(|| default_model(&cli.provider)); + .unwrap_or_else(|| default_model(provider)); eprintln!( "{}Using model: {model}{}", styles.dim, styles.reset, ); - let mut profile = build_profile(&cli.provider, model, Some(client.clone())); + let mut profile = build_profile(provider, model, Some(client.clone())); // Build execution environment let cwd = std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")); @@ -379,16 +376,15 @@ pub async fn run() -> anyhow::Result<()> { )); let manager_for_callback = manager.clone(); let factory_client = client.clone(); - let factory_provider = cli.provider.clone(); let factory_model = model.to_string(); let factory_env = Arc::clone(&env); let factory_approval = config.tool_approval.clone(); let factory: SessionFactory = Arc::new(move || { - let child_summarizer = build_summarizer(&factory_provider, Some(factory_client.clone())); - let child_profile: Arc = match factory_provider.as_str() { - "openai" => Arc::new(OpenAiProfile::with_summarizer(&factory_model, child_summarizer)), - "gemini" => Arc::new(GeminiProfile::with_summarizer(&factory_model, child_summarizer)), - _ => Arc::new(AnthropicProfile::with_summarizer(&factory_model, child_summarizer)), + let child_summarizer = build_summarizer(provider, Some(factory_client.clone())); + let child_profile: Arc = match provider { + Provider::OpenAi => Arc::new(OpenAiProfile::with_summarizer(&factory_model, child_summarizer)), + Provider::Gemini => Arc::new(GeminiProfile::with_summarizer(&factory_model, child_summarizer)), + Provider::Anthropic => Arc::new(AnthropicProfile::with_summarizer(&factory_model, child_summarizer)), }; Session::new( factory_client.clone(), @@ -526,6 +522,7 @@ pub async fn run() -> anyhow::Result<()> { #[cfg(test)] mod tests { use super::*; + use llm::provider::Provider; use serde_json::json; static NO_COLOR: Styles = Styles::new(false); @@ -596,24 +593,17 @@ mod tests { #[test] fn default_model_anthropic() { - assert_eq!(default_model("anthropic"), "claude-opus-4-6"); + assert_eq!(default_model(Provider::Anthropic), "claude-opus-4-6"); } #[test] fn default_model_openai() { - assert_eq!(default_model("openai"), "gpt-5.2-codex"); + assert_eq!(default_model(Provider::OpenAi), "gpt-5.2-codex"); } #[test] fn default_model_gemini() { - assert_eq!(default_model("gemini"), "gemini-3.1-pro-preview"); - } - - // validate_api_key tests - - #[test] - fn validate_api_key_unknown_provider() { - assert!(!validate_api_key("unknown")); + assert_eq!(default_model(Provider::Gemini), "gemini-3.1-pro-preview"); } // build_tool_approval non-interactive tests @@ -650,27 +640,27 @@ mod tests { #[test] fn build_profile_anthropic() { - let profile = build_profile("anthropic", "model", None); - assert_eq!(profile.id(), "anthropic"); + let profile = build_profile(Provider::Anthropic, "model", None); + assert_eq!(profile.provider(), Provider::Anthropic); } #[test] fn build_profile_openai() { - let profile = build_profile("openai", "model", None); - assert_eq!(profile.id(), "openai"); + let profile = build_profile(Provider::OpenAi, "model", None); + assert_eq!(profile.provider(), Provider::OpenAi); } #[test] fn build_profile_gemini() { - let profile = build_profile("gemini", "model", None); - assert_eq!(profile.id(), "gemini"); + let profile = build_profile(Provider::Gemini, "model", None); + assert_eq!(profile.provider(), Provider::Gemini); } // subagent tool registration tests #[test] fn build_profile_can_register_subagent_tools() { - let mut profile = build_profile("anthropic", "model", None); + let mut profile = build_profile(Provider::Anthropic, "model", None); let manager = Arc::new(tokio::sync::Mutex::new(SubAgentManager::new(1))); let factory: SessionFactory = Arc::new(|| { panic!("factory should not be called in this test"); diff --git a/crates/agent/src/compaction.rs b/crates/agent/src/compaction.rs index 890aedf12..3662f207e 100644 --- a/crates/agent/src/compaction.rs +++ b/crates/agent/src/compaction.rs @@ -100,7 +100,7 @@ and conversational filler.{file_ops_section}" "Here is the conversation to summarize:\n\n{rendered}" )), ], - provider: Some(provider_profile.id().to_string()), + provider: Some(provider_profile.provider().as_str().to_string()), tools: None, tool_choice: None, response_format: None, diff --git a/crates/agent/src/profiles/anthropic.rs b/crates/agent/src/profiles/anthropic.rs index 77e6d03c1..ec8336b2d 100644 --- a/crates/agent/src/profiles/anthropic.rs +++ b/crates/agent/src/profiles/anthropic.rs @@ -6,6 +6,7 @@ use crate::provider_profile::{ProfileCapabilities, ProviderProfile}; use crate::skills::Skill; use crate::tool_registry::ToolRegistry; use crate::tools::{make_edit_file_tool, register_core_tools, WebFetchSummarizer}; +use llm::provider::Provider; use super::EnvContext; @@ -32,7 +33,7 @@ impl AnthropicProfile { Self { base: BaseProfile { - id: "anthropic", + provider: Provider::Anthropic, model: model.into(), registry, }, @@ -41,8 +42,8 @@ impl AnthropicProfile { } impl ProviderProfile for AnthropicProfile { - fn id(&self) -> &str { - self.base.id + fn provider(&self) -> Provider { + self.base.provider } fn model(&self) -> &str { @@ -191,7 +192,7 @@ mod tests { #[test] fn anthropic_profile_identity() { let profile = AnthropicProfile::new("claude-sonnet-4-20250514"); - assert_eq!(profile.id(), "anthropic"); + assert_eq!(profile.provider(), Provider::Anthropic); assert_eq!(profile.model(), "claude-sonnet-4-20250514"); } diff --git a/crates/agent/src/profiles/gemini.rs b/crates/agent/src/profiles/gemini.rs index 171aa06dd..dcae4f3f9 100644 --- a/crates/agent/src/profiles/gemini.rs +++ b/crates/agent/src/profiles/gemini.rs @@ -9,6 +9,7 @@ use crate::tools::{ make_edit_file_tool, make_list_dir_tool, make_read_many_files_tool, register_core_tools, WebFetchSummarizer, }; +use llm::provider::Provider; use super::EnvContext; @@ -34,7 +35,7 @@ impl GeminiProfile { Self { base: BaseProfile { - id: "gemini", + provider: Provider::Gemini, model: model.into(), registry, }, @@ -43,8 +44,8 @@ impl GeminiProfile { } impl ProviderProfile for GeminiProfile { - fn id(&self) -> &str { - self.base.id + fn provider(&self) -> Provider { + self.base.provider } fn model(&self) -> &str { @@ -221,7 +222,7 @@ mod tests { #[test] fn gemini_profile_identity() { let profile = GeminiProfile::new("gemini-2.0-flash"); - assert_eq!(profile.id(), "gemini"); + assert_eq!(profile.provider(), Provider::Gemini); assert_eq!(profile.model(), "gemini-2.0-flash"); } diff --git a/crates/agent/src/profiles/mod.rs b/crates/agent/src/profiles/mod.rs index d16765951..fd8263b46 100644 --- a/crates/agent/src/profiles/mod.rs +++ b/crates/agent/src/profiles/mod.rs @@ -9,13 +9,14 @@ pub use openai::OpenAiProfile; use crate::execution_env::ExecutionEnvironment; use crate::skills::{format_skills_prompt_section, Skill}; use crate::tool_registry::ToolRegistry; +use llm::provider::Provider; /// Common fields shared by all provider profiles. /// -/// Each concrete profile embeds this struct and delegates `id()`, `model()`, +/// Each concrete profile embeds this struct and delegates `provider()`, `model()`, /// `tool_registry()`, and `tool_registry_mut()` to it. pub struct BaseProfile { - pub id: &'static str, + pub provider: Provider, pub model: String, pub registry: ToolRegistry, } diff --git a/crates/agent/src/profiles/openai.rs b/crates/agent/src/profiles/openai.rs index 5d80ef852..836f9ab0c 100644 --- a/crates/agent/src/profiles/openai.rs +++ b/crates/agent/src/profiles/openai.rs @@ -7,6 +7,7 @@ use crate::skills::Skill; use crate::tool_registry::ToolRegistry; use crate::tools::{register_core_tools, WebFetchSummarizer}; use crate::v4a_patch::make_apply_patch_tool; +use llm::provider::Provider; use super::EnvContext; @@ -31,7 +32,7 @@ impl OpenAiProfile { Self { base: BaseProfile { - id: "openai", + provider: Provider::OpenAi, model: model.into(), registry, }, @@ -45,8 +46,8 @@ impl OpenAiProfile { } impl ProviderProfile for OpenAiProfile { - fn id(&self) -> &str { - self.base.id + fn provider(&self) -> Provider { + self.base.provider } fn model(&self) -> &str { @@ -194,7 +195,7 @@ mod tests { #[test] fn openai_profile_identity() { let profile = OpenAiProfile::new("o3-mini"); - assert_eq!(profile.id(), "openai"); + assert_eq!(profile.provider(), Provider::OpenAi); assert_eq!(profile.model(), "o3-mini"); } diff --git a/crates/agent/src/project_docs.rs b/crates/agent/src/project_docs.rs index 610db61ea..6ae4d3c04 100644 --- a/crates/agent/src/project_docs.rs +++ b/crates/agent/src/project_docs.rs @@ -1,4 +1,5 @@ use crate::execution_env::ExecutionEnvironment; +use llm::provider::Provider; const BUDGET_BYTES: usize = 32768; @@ -6,15 +7,14 @@ pub async fn discover_project_docs( env: &dyn ExecutionEnvironment, git_root: &str, working_dir: &str, - provider_id: &str, + provider: Provider, ) -> Vec { let directories = build_directory_walk(git_root, working_dir); - let candidate_filenames: Vec<&str> = match provider_id { - "anthropic" => vec!["AGENTS.md", "CLAUDE.md"], - "openai" => vec!["AGENTS.md", ".codex/instructions.md"], - "gemini" => vec!["AGENTS.md", "GEMINI.md"], - _ => vec!["AGENTS.md"], + let candidate_filenames: Vec<&str> = match provider { + Provider::Anthropic => vec!["AGENTS.md", "CLAUDE.md"], + Provider::OpenAi => vec!["AGENTS.md", ".codex/instructions.md"], + Provider::Gemini => vec!["AGENTS.md", "GEMINI.md"], }; let mut results = Vec::new(); @@ -99,7 +99,7 @@ mod tests { files, ..Default::default() }); - let docs = discover_project_docs(env.as_ref(), "/repo", "/repo", "anthropic").await; + let docs = discover_project_docs(env.as_ref(), "/repo", "/repo", Provider::Anthropic).await; assert_eq!(docs.len(), 1); assert_eq!(docs[0], "Agent instructions"); } @@ -120,7 +120,7 @@ mod tests { ..Default::default() }); let anthropic_docs = - discover_project_docs(env.as_ref(), "/repo", "/repo", "anthropic").await; + discover_project_docs(env.as_ref(), "/repo", "/repo", Provider::Anthropic).await; assert_eq!(anthropic_docs.len(), 2); assert_eq!(anthropic_docs[0], "agents"); assert_eq!(anthropic_docs[1], "claude"); @@ -129,7 +129,7 @@ mod tests { files: files.clone(), ..Default::default() }); - let openai_docs = discover_project_docs(env.as_ref(), "/repo", "/repo", "openai").await; + let openai_docs = discover_project_docs(env.as_ref(), "/repo", "/repo", Provider::OpenAi).await; assert_eq!(openai_docs.len(), 2); assert_eq!(openai_docs[0], "agents"); assert_eq!(openai_docs[1], "copilot"); @@ -138,7 +138,7 @@ mod tests { files, ..Default::default() }); - let gemini_docs = discover_project_docs(env.as_ref(), "/repo", "/repo", "gemini").await; + let gemini_docs = discover_project_docs(env.as_ref(), "/repo", "/repo", Provider::Gemini).await; assert_eq!(gemini_docs.len(), 2); assert_eq!(gemini_docs[0], "agents"); assert_eq!(gemini_docs[1], "gemini"); @@ -157,7 +157,7 @@ mod tests { files, ..Default::default() }); - let docs = discover_project_docs(env.as_ref(), "/repo", "/repo", "anthropic").await; + let docs = discover_project_docs(env.as_ref(), "/repo", "/repo", Provider::Anthropic).await; assert_eq!(docs.len(), 2); assert_eq!(docs[0], large_content); // Second doc should be truncated to fit remaining budget @@ -177,7 +177,7 @@ mod tests { ..Default::default() }); let docs = - discover_project_docs(env.as_ref(), "/repo", "/repo/src/app", "anthropic").await; + discover_project_docs(env.as_ref(), "/repo", "/repo/src/app", Provider::Anthropic).await; assert_eq!(docs.len(), 3); assert_eq!(docs[0], "root agents"); assert_eq!(docs[1], "src agents"); diff --git a/crates/agent/src/provider_profile.rs b/crates/agent/src/provider_profile.rs index 8cb3a1388..46a24cf58 100644 --- a/crates/agent/src/provider_profile.rs +++ b/crates/agent/src/provider_profile.rs @@ -7,6 +7,7 @@ use crate::subagent::{ }; use crate::tool_registry::ToolRegistry; use std::sync::Arc; +use llm::provider::Provider; use llm::types::ToolDefinition; /// Static capabilities of a provider profile. @@ -18,7 +19,7 @@ pub struct ProfileCapabilities { } pub trait ProviderProfile: Send + Sync { - fn id(&self) -> &str; + fn provider(&self) -> Provider; fn model(&self) -> &str; fn tool_registry(&self) -> &ToolRegistry; fn tool_registry_mut(&mut self) -> &mut ToolRegistry; @@ -81,11 +82,12 @@ pub trait ProviderProfile: Send + Sync { mod tests { use super::*; use crate::test_support::{MockExecutionEnvironment, TestProfile}; + use llm::provider::Provider; #[test] - fn profile_id_and_model() { + fn profile_provider_and_model() { let profile = TestProfile::new(); - assert_eq!(profile.id(), "mock"); + assert_eq!(profile.provider(), Provider::Anthropic); assert_eq!(profile.model(), "mock-model"); } diff --git a/crates/agent/src/session.rs b/crates/agent/src/session.rs index fb9b37d9f..95f4354a5 100644 --- a/crates/agent/src/session.rs +++ b/crates/agent/src/session.rs @@ -82,7 +82,7 @@ impl Session { self.execution_env.as_ref(), &doc_root, self.execution_env.working_directory(), - self.provider_profile.id(), + self.provider_profile.provider(), ) .await; @@ -371,7 +371,7 @@ impl Session { // Call LLM (streaming) with retry for transient errors let retry_emitter = self.event_emitter.clone(); let retry_session_id = self.id.clone(); - let retry_provider = self.provider_profile.id().to_string(); + let retry_provider = self.provider_profile.provider().as_str().to_string(); let retry_model = self.provider_profile.model().to_string(); let retry_policy = llm::types::RetryPolicy { max_retries: 3, @@ -612,7 +612,7 @@ impl Session { Request { model: self.provider_profile.model().to_string(), messages, - provider: Some(self.provider_profile.id().to_string()), + provider: Some(self.provider_profile.provider().as_str().to_string()), tools: if has_tools { Some(tools) } else { None }, tool_choice: if has_tools { Some(ToolChoice::Auto) diff --git a/crates/agent/src/test_support.rs b/crates/agent/src/test_support.rs index db735160a..f89f3af73 100644 --- a/crates/agent/src/test_support.rs +++ b/crates/agent/src/test_support.rs @@ -11,7 +11,7 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; use llm::client::Client; use llm::error::SdkError; -use llm::provider::{ProviderAdapter, StreamEventStream}; +use llm::provider::{Provider, ProviderAdapter, StreamEventStream}; use llm::types::{ContentPart, FinishReason, Message, Request, Response, StreamEvent, Usage}; use tokio_util::sync::CancellationToken; @@ -344,8 +344,8 @@ impl TestProfile { } impl ProviderProfile for TestProfile { - fn id(&self) -> &'static str { - "mock" + fn provider(&self) -> Provider { + Provider::Anthropic } fn model(&self) -> &'static str { @@ -488,7 +488,9 @@ pub fn text_response(text: &str) -> Response { pub async fn make_client(provider: Arc) -> Client { let mut providers = HashMap::new(); - providers.insert(provider.name().to_string(), provider); + providers.insert(provider.name().to_string(), provider.clone()); + // Also register under "anthropic" so TestProfile (Provider::Anthropic) routes correctly + providers.insert("anthropic".to_string(), provider); Client::new(providers, Some("mock".into()), vec![]) } diff --git a/crates/agent/src/tools.rs b/crates/agent/src/tools.rs index 2886ea9ea..6609af1a3 100644 --- a/crates/agent/src/tools.rs +++ b/crates/agent/src/tools.rs @@ -2,6 +2,7 @@ use crate::config::SessionConfig; use crate::execution_env::GrepOptions; use crate::tool_registry::RegisteredTool; use llm::client::Client; +use llm::provider::ModelId; use llm::types::{Message, Request, ToolDefinition}; use std::borrow::Cow; use std::fmt::Write; @@ -13,8 +14,7 @@ const MAX_WEB_FETCH_BYTES: usize = 100 * 1024; #[derive(Clone)] pub struct WebFetchSummarizer { pub client: Client, - pub model: String, - pub provider: Option, + pub model_id: ModelId, } /// Returns true if the input looks like it contains HTML markup. @@ -534,9 +534,9 @@ pub(crate) fn make_web_fetch_tool(summarizer: Option) -> Reg "Content from {url}:\n---\n{content}\n---\n\n{user_prompt}\n\nRespond concisely based only on the content above." ); let request = Request { - model: s.model.clone(), + model: s.model_id.model.clone(), messages: vec![Message::user(summarization_prompt)], - provider: s.provider.clone(), + provider: Some(s.model_id.provider.as_str().to_string()), tools: None, tool_choice: None, response_format: None, @@ -549,7 +549,7 @@ pub(crate) fn make_web_fetch_tool(summarizer: Option) -> Reg provider_options: None, }; let response = s.client.complete(&request).await.map_err(|e| { - format!("web_fetch summarization (model={}) failed: {e}", s.model) + format!("web_fetch summarization (model={}) failed: {e}", s.model_id.model) })?; Ok(response.text()) } @@ -984,8 +984,7 @@ mod tests { let client = make_client(provider).await; let summarizer = WebFetchSummarizer { client, - model: "mock-model".into(), - provider: None, + model_id: ModelId::new(llm::provider::Provider::Anthropic, "mock-model"), }; let tool = make_web_fetch_tool(Some(summarizer)); @@ -1035,10 +1034,9 @@ mod tests { async fn web_fetch_summarizer_routes_to_specified_provider() { use crate::test_support::{MockErrorProvider, MockLlmProvider, text_response}; use llm::error::{ProviderErrorDetail, ProviderErrorKind, SdkError}; - use llm::provider::ProviderAdapter; // "other_provider" is the default — it rejects all requests. - let default_provider: Arc = Arc::new( + let default_provider: Arc = Arc::new( MockErrorProvider { error: SdkError::Provider { kind: ProviderErrorKind::NotFound, @@ -1048,20 +1046,20 @@ mod tests { }, }, ); - // "target_provider" has the model we actually want. - let target_provider: Arc = Arc::new( + // "anthropic" provider has the model we actually want. + let target_provider: Arc = Arc::new( MockLlmProvider::new(vec![text_response("summarized content")]), ); let mut providers = HashMap::new(); providers.insert("other_provider".to_string(), default_provider); - providers.insert(target_provider.name().to_string(), target_provider); + // Register under "anthropic" so ModelId { provider: Anthropic, .. } routes here + providers.insert("anthropic".to_string(), target_provider); let client = Client::new(providers, Some("other_provider".into()), vec![]); let summarizer = WebFetchSummarizer { client, - model: "target-model".into(), - provider: Some("mock".into()), + model_id: ModelId::new(llm::provider::Provider::Anthropic, "target-model"), }; let tool = make_web_fetch_tool(Some(summarizer)); diff --git a/crates/agent/tests/parity_matrix.rs b/crates/agent/tests/parity_matrix.rs index b0dbf640f..8ed6d76cf 100644 --- a/crates/agent/tests/parity_matrix.rs +++ b/crates/agent/tests/parity_matrix.rs @@ -6,31 +6,33 @@ use agent::{ Session, SessionConfig, SubAgentManager, WebFetchSummarizer, }; use llm::client::Client; +use llm::provider::{ModelId, Provider}; -fn build_summarizer(provider: &str, client: &Client) -> WebFetchSummarizer { - let (summarizer_model, provider_name) = match provider { - "openai" => ("gpt-4o-mini", Some("openai".to_string())), - "gemini" => ("gemini-2.0-flash", Some("gemini".to_string())), - _ => ("claude-haiku-4-5-20251001", None), - }; +fn summarizer_model_id(provider: Provider) -> ModelId { + match provider { + Provider::OpenAi => ModelId::new(Provider::OpenAi, "gpt-4o-mini"), + Provider::Gemini => ModelId::new(Provider::Gemini, "gemini-2.0-flash"), + Provider::Anthropic => ModelId::new(Provider::Anthropic, "claude-haiku-4-5-20251001"), + } +} + +fn build_summarizer(provider: Provider, client: &Client) -> WebFetchSummarizer { WebFetchSummarizer { client: client.clone(), - model: summarizer_model.into(), - provider: provider_name, + model_id: summarizer_model_id(provider), } } -fn build_profile(provider: &str, model: &str, client: &Client) -> Box { +fn build_profile(provider: Provider, model: &str, client: &Client) -> Box { let summarizer = Some(build_summarizer(provider, client)); match provider { - "anthropic" => Box::new(AnthropicProfile::with_summarizer(model, summarizer)), - "openai" => Box::new(OpenAiProfile::with_summarizer(model, summarizer)), - "gemini" => Box::new(GeminiProfile::with_summarizer(model, summarizer)), - _ => panic!("unknown provider: {provider}"), + Provider::Anthropic => Box::new(AnthropicProfile::with_summarizer(model, summarizer)), + Provider::OpenAi => Box::new(OpenAiProfile::with_summarizer(model, summarizer)), + Provider::Gemini => Box::new(GeminiProfile::with_summarizer(model, summarizer)), } } -async fn make_session(provider: &str, model: &str, cwd: &Path) -> Session { +async fn make_session(provider: Provider, model: &str, cwd: &Path) -> Session { dotenvy::dotenv().ok(); let client = Client::from_env().await.expect("Client::from_env failed"); let mut profile = build_profile(provider, model, &client); @@ -39,31 +41,25 @@ async fn make_session(provider: &str, model: &str, cwd: &Path) -> Session { // Register subagent tools so spawn_agent / wait / send_input / close_agent are available let manager = Arc::new(tokio::sync::Mutex::new(SubAgentManager::new(3))); let factory_client = client.clone(); - let factory_provider: &str = provider; let factory_model: String = model.to_string(); let factory_cwd = cwd.to_path_buf(); - let factory: agent::subagent::SessionFactory = { - let provider = factory_provider.to_string(); - let model = factory_model; - Arc::new(move || { - let sub_profile: Arc = { - let summarizer = Some(build_summarizer(&provider, &factory_client)); - match provider.as_str() { - "anthropic" => Arc::new(AnthropicProfile::with_summarizer(&model, summarizer)), - "openai" => Arc::new(OpenAiProfile::with_summarizer(&model, summarizer)), - "gemini" => Arc::new(GeminiProfile::with_summarizer(&model, summarizer)), - _ => panic!("unknown provider: {provider}"), - } - }; - let sub_env = Arc::new(LocalExecutionEnvironment::new(factory_cwd.clone())); - Session::new( - factory_client.clone(), - sub_profile, - sub_env, - SessionConfig::default(), - ) - }) - }; + let factory: agent::subagent::SessionFactory = Arc::new(move || { + let sub_profile: Arc = { + let summarizer = Some(build_summarizer(provider, &factory_client)); + match provider { + Provider::Anthropic => Arc::new(AnthropicProfile::with_summarizer(&factory_model, summarizer)), + Provider::OpenAi => Arc::new(OpenAiProfile::with_summarizer(&factory_model, summarizer)), + Provider::Gemini => Arc::new(GeminiProfile::with_summarizer(&factory_model, summarizer)), + } + }; + let sub_env = Arc::new(LocalExecutionEnvironment::new(factory_cwd.clone())); + Session::new( + factory_client.clone(), + sub_profile, + sub_env, + SessionConfig::default(), + ) + }); profile.register_subagent_tools(manager, factory, 0); let profile: Arc = Arc::from(profile); @@ -75,7 +71,7 @@ async fn make_session(provider: &str, model: &str, cwd: &Path) -> Session { } async fn make_session_with_config( - provider: &str, + provider: Provider, model: &str, cwd: &Path, config: SessionConfig, @@ -94,7 +90,7 @@ macro_rules! provider_tests { #[ignore = "requires LLM API keys"] async fn []() { let tmp = tempfile::tempdir().expect("failed to create tempdir"); - let mut session = make_session("anthropic", "claude-haiku-4-5-20251001", tmp.path()).await; + let mut session = make_session(Provider::Anthropic, "claude-haiku-4-5-20251001", tmp.path()).await; session.initialize().await; [](&mut session, tmp.path()).await; } @@ -103,7 +99,7 @@ macro_rules! provider_tests { #[ignore = "requires LLM API keys"] async fn []() { let tmp = tempfile::tempdir().expect("failed to create tempdir"); - let mut session = make_session("openai", "gpt-4o-mini", tmp.path()).await; + let mut session = make_session(Provider::OpenAi, "gpt-4o-mini", tmp.path()).await; session.initialize().await; [](&mut session, tmp.path()).await; } @@ -112,7 +108,7 @@ macro_rules! provider_tests { #[ignore = "requires LLM API keys"] async fn []() { let tmp = tempfile::tempdir().expect("failed to create tempdir"); - let mut session = make_session("gemini", "gemini-2.5-flash", tmp.path()).await; + let mut session = make_session(Provider::Gemini, "gemini-2.5-flash", tmp.path()).await; session.initialize().await; [](&mut session, tmp.path()).await; } @@ -149,7 +145,7 @@ macro_rules! anthropic_gemini_tests { #[ignore = "requires LLM API keys"] async fn []() { let tmp = tempfile::tempdir().expect("failed to create tempdir"); - let mut session = make_session("anthropic", "claude-haiku-4-5-20251001", tmp.path()).await; + let mut session = make_session(Provider::Anthropic, "claude-haiku-4-5-20251001", tmp.path()).await; session.initialize().await; [](&mut session, tmp.path()).await; } @@ -158,7 +154,7 @@ macro_rules! anthropic_gemini_tests { #[ignore = "requires LLM API keys"] async fn []() { let tmp = tempfile::tempdir().expect("failed to create tempdir"); - let mut session = make_session("gemini", "gemini-2.5-flash", tmp.path()).await; + let mut session = make_session(Provider::Gemini, "gemini-2.5-flash", tmp.path()).await; session.initialize().await; [](&mut session, tmp.path()).await; } @@ -340,9 +336,9 @@ macro_rules! reasoning_effort_tests { }; } -reasoning_effort_tests!("anthropic", "claude-haiku-4-5-20251001", anthropic_reasoning_effort); +reasoning_effort_tests!(Provider::Anthropic, "claude-haiku-4-5-20251001", anthropic_reasoning_effort); // gpt-4o-mini does not support the reasoning.effort parameter, so no OpenAI test. -reasoning_effort_tests!("gemini", "gemini-2.5-flash", gemini_reasoning_effort); +reasoning_effort_tests!(Provider::Gemini, "gemini-2.5-flash", gemini_reasoning_effort); // --------------------------------------------------------------------------- // Scenario 12: subagent_spawn @@ -383,9 +379,9 @@ macro_rules! loop_detection_tests { }; } -loop_detection_tests!("anthropic", "claude-haiku-4-5-20251001", anthropic_loop_detection); -loop_detection_tests!("openai", "gpt-4o-mini", openai_loop_detection); -loop_detection_tests!("gemini", "gemini-2.5-flash", gemini_loop_detection); +loop_detection_tests!(Provider::Anthropic, "claude-haiku-4-5-20251001", anthropic_loop_detection); +loop_detection_tests!(Provider::OpenAi, "gpt-4o-mini", openai_loop_detection); +loop_detection_tests!(Provider::Gemini, "gemini-2.5-flash", gemini_loop_detection); // --------------------------------------------------------------------------- // Scenario 14: error_recovery diff --git a/crates/attractor/src/cli/backend.rs b/crates/attractor/src/cli/backend.rs index 7a1ab485e..aafc651ea 100644 --- a/crates/attractor/src/cli/backend.rs +++ b/crates/attractor/src/cli/backend.rs @@ -9,6 +9,7 @@ use agent::{ subagent::{SessionFactory, SubAgentManager}, }; use llm::client::Client; +use llm::provider::Provider; use util::terminal::Styles; use crate::context::Context; @@ -23,7 +24,7 @@ use crate::outcome::StageUsage; /// and reused so the LLM sees the full conversation history. pub struct AgentBackend { model: String, - provider: Option, + provider: Provider, verbose: u8, styles: &'static Styles, sessions: Mutex>, @@ -33,7 +34,7 @@ impl AgentBackend { #[must_use] pub fn new( model: String, - provider: Option, + provider: Provider, verbose: u8, styles: &'static Styles, ) -> Self { @@ -69,17 +70,14 @@ impl AgentBackend { // Build factory that creates child sessions WITHOUT subagent tools let factory_client = client.clone(); - let factory_provider = self.provider.clone(); + let factory_provider = self.provider; let factory_model = self.model.clone(); let factory_env = Arc::clone(execution_env); let factory: SessionFactory = Arc::new(move || { - let child_profile = { - let provider = factory_provider.as_deref().unwrap_or("anthropic"); - match provider { - "openai" => Arc::new(OpenAiProfile::new(&factory_model)) as Arc, - "gemini" => Arc::new(GeminiProfile::new(&factory_model)) as Arc, - _ => Arc::new(AnthropicProfile::new(&factory_model)) as Arc, - } + let child_profile: Arc = match factory_provider { + Provider::OpenAi => Arc::new(OpenAiProfile::new(&factory_model)), + Provider::Gemini => Arc::new(GeminiProfile::new(&factory_model)), + Provider::Anthropic => Arc::new(AnthropicProfile::new(&factory_model)), }; Session::new( factory_client.clone(), @@ -101,11 +99,10 @@ impl AgentBackend { } fn build_profile(&self) -> Box { - let provider = self.provider.as_deref().unwrap_or("anthropic"); - match provider { - "openai" => Box::new(OpenAiProfile::new(&self.model)), - "gemini" => Box::new(GeminiProfile::new(&self.model)), - _ => Box::new(AnthropicProfile::new(&self.model)), + match self.provider { + Provider::OpenAi => Box::new(OpenAiProfile::new(&self.model)), + Provider::Gemini => Box::new(GeminiProfile::new(&self.model)), + Provider::Anthropic => Box::new(AnthropicProfile::new(&self.model)), } } } @@ -125,8 +122,8 @@ impl CodergenBackend for AgentBackend { let model = node.llm_model().unwrap_or(&self.model); let provider = node .llm_provider() - .or(self.provider.as_deref()) - .map(String::from); + .map(String::from) + .or_else(|| Some(self.provider.as_str().to_string())); let request = llm::types::Request { model: model.to_string(), @@ -413,7 +410,7 @@ impl CodergenBackend for AgentBackend { let provider_used = serde_json::json!({ "mode": "agent_loop", - "provider": self.provider.as_deref().unwrap_or("anthropic"), + "provider": self.provider.as_str(), "model": &self.model, }); if let Ok(json) = serde_json::to_string_pretty(&provider_used) { @@ -459,12 +456,12 @@ mod tests { let styles = Box::leak(Box::new(Styles::new(false))); let backend = AgentBackend::new( "claude-opus-4-6".to_string(), - Some("openai".to_string()), + Provider::OpenAi, 2, styles, ); assert_eq!(backend.model, "claude-opus-4-6"); - assert_eq!(backend.provider.as_deref(), Some("openai")); + assert_eq!(backend.provider, Provider::OpenAi); assert_eq!(backend.verbose, 2); } @@ -473,7 +470,7 @@ mod tests { let styles = Box::leak(Box::new(Styles::new(false))); let backend = AgentBackend::new( "claude-opus-4-6".to_string(), - None, + Provider::Anthropic, 0, styles, ); @@ -485,7 +482,7 @@ mod tests { let styles = Box::leak(Box::new(Styles::new(false))); let backend = AgentBackend::new( "claude-opus-4-6".to_string(), - None, + Provider::Anthropic, 0, styles, ); diff --git a/crates/attractor/src/cli/cli_backend.rs b/crates/attractor/src/cli/cli_backend.rs index 8450df7bc..3f0ac9366 100644 --- a/crates/attractor/src/cli/cli_backend.rs +++ b/crates/attractor/src/cli/cli_backend.rs @@ -3,6 +3,7 @@ use std::sync::Arc; use agent::ExecutionEnvironment; use async_trait::async_trait; +use llm::provider::Provider; use crate::context::Context; use crate::error::AttractorError; @@ -25,24 +26,23 @@ pub fn is_cli_only_model(model: &str) -> bool { /// The `prompt_file` is the path to a file containing the prompt text, which /// will be shell-redirected into the command's stdin. #[must_use] -pub fn cli_command_for_provider(provider: &str, model: &str, prompt_file: &str) -> String { +pub fn cli_command_for_provider(provider: Provider, model: &str, prompt_file: &str) -> String { let model_flag = if model.is_empty() { String::new() } else { match provider { - "openai" => format!(" -m {model}"), - "gemini" => format!(" -m {model}"), - _ => format!(" --model {model}"), + Provider::OpenAi | Provider::Gemini => format!(" -m {model}"), + Provider::Anthropic => format!(" --model {model}"), } }; match provider { // --full-auto: sandboxed auto-execution, escalates on request - "openai" => format!("codex exec --json --full-auto{model_flag} < {prompt_file}"), + Provider::OpenAi => format!("codex exec --json --full-auto{model_flag} < {prompt_file}"), // --yolo: auto-approve all tool calls - "gemini" => format!("gemini -o json --yolo{model_flag} < {prompt_file}"), + Provider::Gemini => format!("gemini -o json --yolo{model_flag} < {prompt_file}"), // --dangerously-skip-permissions: bypass all permission checks (required for non-interactive use). // CLAUDECODE= unset to allow running inside a Claude Code session. - _ => format!("CLAUDECODE= claude -p --output-format stream-json --dangerously-skip-permissions{model_flag} < {prompt_file}"), + Provider::Anthropic => format!("CLAUDECODE= claude -p --output-format stream-json --dangerously-skip-permissions{model_flag} < {prompt_file}"), } } @@ -170,23 +170,23 @@ fn parse_gemini_json(output: &str) -> Option { } /// Parse CLI output, choosing the right parser based on provider. -pub fn parse_cli_response(provider: &str, output: &str) -> Option { +pub fn parse_cli_response(provider: Provider, output: &str) -> Option { match provider { - "openai" => parse_codex_ndjson(output), - "gemini" => parse_gemini_json(output), - _ => parse_claude_ndjson(output), + Provider::OpenAi => parse_codex_ndjson(output), + Provider::Gemini => parse_gemini_json(output), + Provider::Anthropic => parse_claude_ndjson(output), } } /// CLI backend that invokes external CLI tools (claude, codex, gemini) via `exec_command()`. pub struct CliBackend { model: String, - provider: String, + provider: Provider, } impl CliBackend { #[must_use] - pub fn new(model: String, provider: String) -> Self { + pub fn new(model: String, provider: Provider) -> Self { Self { model, provider } } @@ -267,13 +267,16 @@ impl CodergenBackend for CliBackend { // 3. Build and execute CLI command let model = node.llm_model().unwrap_or(&self.model); - let provider = node.llm_provider().unwrap_or(&self.provider); + let provider = node + .llm_provider() + .and_then(|s| s.parse::().ok()) + .unwrap_or(self.provider); let command = cli_command_for_provider(provider, model, prompt_path); let _ = tokio::fs::create_dir_all(stage_dir).await; let provider_used = serde_json::json!({ "mode": "cli", - "provider": provider, + "provider": provider.as_str(), "model": model, "command": &command, }); @@ -410,7 +413,7 @@ mod tests { #[test] fn cli_command_for_codex() { - let cmd = cli_command_for_provider("openai", "gpt-5.3-codex", "/tmp/prompt.txt"); + let cmd = cli_command_for_provider(Provider::OpenAi, "gpt-5.3-codex", "/tmp/prompt.txt"); assert!(cmd.starts_with("codex exec --json --full-auto")); assert!(cmd.contains("-m gpt-5.3-codex")); assert!(cmd.ends_with("< /tmp/prompt.txt")); @@ -418,7 +421,7 @@ mod tests { #[test] fn cli_command_for_claude() { - let cmd = cli_command_for_provider("anthropic", "claude-opus-4-6", "/tmp/prompt.txt"); + let cmd = cli_command_for_provider(Provider::Anthropic, "claude-opus-4-6", "/tmp/prompt.txt"); assert!(cmd.contains("claude -p")); assert!(cmd.contains("--dangerously-skip-permissions")); assert!(cmd.contains("--output-format stream-json")); @@ -427,27 +430,20 @@ mod tests { #[test] fn cli_command_for_gemini() { - let cmd = cli_command_for_provider("gemini", "gemini-3.1-pro", "/tmp/prompt.txt"); + let cmd = cli_command_for_provider(Provider::Gemini, "gemini-3.1-pro", "/tmp/prompt.txt"); assert!(cmd.starts_with("gemini -o json --yolo")); assert!(cmd.contains("-m gemini-3.1-pro")); } - #[test] - fn cli_command_defaults_to_claude() { - let cmd = cli_command_for_provider("unknown_provider", "some-model", "/tmp/prompt.txt"); - assert!(cmd.contains("claude ")); - assert!(cmd.contains("--dangerously-skip-permissions")); - } - #[test] fn cli_command_omits_model_when_empty() { - let cmd = cli_command_for_provider("openai", "", "/tmp/prompt.txt"); + let cmd = cli_command_for_provider(Provider::OpenAi, "", "/tmp/prompt.txt"); assert!(cmd.starts_with("codex exec --json --full-auto")); assert!(!cmd.contains("-m ")); - let cmd = cli_command_for_provider("anthropic", "", "/tmp/prompt.txt"); + let cmd = cli_command_for_provider(Provider::Anthropic, "", "/tmp/prompt.txt"); assert!(cmd.contains("--dangerously-skip-permissions")); assert!(!cmd.contains("--model ")); - let cmd = cli_command_for_provider("gemini", "", "/tmp/prompt.txt"); + let cmd = cli_command_for_provider(Provider::Gemini, "", "/tmp/prompt.txt"); assert!(cmd.contains("--yolo")); assert!(!cmd.contains("-m ")); } @@ -468,7 +464,7 @@ mod tests { let output = r#"{"type":"system","message":"Claude CLI v1.0"} {"type":"assistant","message":{"content":"thinking..."}} {"type":"result","result":"Here is the implementation.","usage":{"input_tokens":100,"output_tokens":50}}"#; - let response = parse_cli_response("anthropic", output).unwrap(); + let response = parse_cli_response(Provider::Anthropic, output).unwrap(); assert_eq!(response.text, "Here is the implementation."); assert_eq!(response.input_tokens, 100); assert_eq!(response.output_tokens, 50); @@ -478,7 +474,7 @@ mod tests { fn parse_claude_ndjson_uses_last_result() { let output = r#"{"type":"result","result":"first","usage":{"input_tokens":10,"output_tokens":5}} {"type":"result","result":"second","usage":{"input_tokens":20,"output_tokens":10}}"#; - let response = parse_cli_response("anthropic", output).unwrap(); + let response = parse_cli_response(Provider::Anthropic, output).unwrap(); assert_eq!(response.text, "second"); assert_eq!(response.input_tokens, 20); } @@ -487,13 +483,13 @@ mod tests { fn parse_claude_ndjson_returns_none_for_no_result() { let output = r#"{"type":"system","message":"hello"} {"type":"assistant","message":{"content":"no result line"}}"#; - assert!(parse_cli_response("anthropic", output).is_none()); + assert!(parse_cli_response(Provider::Anthropic, output).is_none()); } #[test] fn parse_gemini_json_extracts_text_and_usage() { let output = r#"{"session_id":"abc","response":"Gemini says hello","stats":{"models":{"gemini-2.5-flash":{"tokens":{"input":200,"candidates":80,"total":280}}}}}"#; - let response = parse_cli_response("gemini", output).unwrap(); + let response = parse_cli_response(Provider::Gemini, output).unwrap(); assert_eq!(response.text, "Gemini says hello"); assert_eq!(response.input_tokens, 200); assert_eq!(response.output_tokens, 80); @@ -502,7 +498,7 @@ mod tests { #[test] fn parse_gemini_json_handles_missing_stats() { let output = r#"{"response":"hello"}"#; - let response = parse_cli_response("gemini", output).unwrap(); + let response = parse_cli_response(Provider::Gemini, output).unwrap(); assert_eq!(response.text, "hello"); assert_eq!(response.input_tokens, 0); assert_eq!(response.output_tokens, 0); @@ -510,7 +506,7 @@ mod tests { #[test] fn parse_gemini_json_returns_none_for_invalid_json() { - assert!(parse_cli_response("gemini", "not json").is_none()); + assert!(parse_cli_response(Provider::Gemini, "not json").is_none()); } // -- Cycle 4: parse_cli_response — Codex NDJSON -- @@ -522,7 +518,7 @@ mod tests { {"type":"item.completed","item":{"id":"item_0","type":"reasoning","text":"thinking..."}} {"type":"item.completed","item":{"id":"item_1","type":"agent_message","text":"Fixed the bug."}} {"type":"turn.completed","usage":{"input_tokens":300,"output_tokens":150}}"#; - let response = parse_cli_response("openai", output).unwrap(); + let response = parse_cli_response(Provider::OpenAi, output).unwrap(); assert_eq!(response.text, "Fixed the bug."); assert_eq!(response.input_tokens, 300); assert_eq!(response.output_tokens, 150); @@ -531,14 +527,14 @@ mod tests { #[test] fn parse_codex_ndjson_handles_no_message() { let output = r#"{"type":"turn.completed","usage":{"input_tokens":10,"output_tokens":5}}"#; - let response = parse_cli_response("openai", output).unwrap(); + let response = parse_cli_response(Provider::OpenAi, output).unwrap(); assert_eq!(response.text, ""); assert_eq!(response.input_tokens, 10); } #[test] fn parse_codex_ndjson_returns_none_for_no_events() { - assert!(parse_cli_response("openai", "not json at all").is_none()); + assert!(parse_cli_response(Provider::OpenAi, "not json at all").is_none()); } // -- Cycle 5: Node::backend() accessor (tested here since the accessor is simple) -- @@ -567,7 +563,7 @@ mod tests { node.attrs .insert("backend".to_string(), AttrValue::String("cli".to_string())); - let cli_backend = CliBackend::new("model".into(), "anthropic".into()); + let cli_backend = CliBackend::new("model".into(), Provider::Anthropic); let router = BackendRouter::new( Box::new(StubBackend), cli_backend, @@ -579,7 +575,7 @@ mod tests { fn router_uses_api_by_default() { let node = Node::new("test"); - let cli_backend = CliBackend::new("model".into(), "anthropic".into()); + let cli_backend = CliBackend::new("model".into(), Provider::Anthropic); let router = BackendRouter::new( Box::new(StubBackend), cli_backend, @@ -595,7 +591,7 @@ mod tests { AttrValue::String("claude-opus-4-6".to_string()), ); - let cli_backend = CliBackend::new("model".into(), "anthropic".into()); + let cli_backend = CliBackend::new("model".into(), Provider::Anthropic); let router = BackendRouter::new( Box::new(StubBackend), cli_backend, diff --git a/crates/attractor/src/cli/run.rs b/crates/attractor/src/cli/run.rs index b7933f2aa..1368bf70e 100644 --- a/crates/attractor/src/cli/run.rs +++ b/crates/attractor/src/cli/run.rs @@ -18,6 +18,8 @@ use crate::outcome::StageStatus; use crate::pipeline::PipelineBuilder; use crate::validation::Severity; +use llm::provider::Provider; + use super::backend::AgentBackend; use super::cli_backend::{BackendRouter, CliBackend}; use super::task_config; @@ -399,6 +401,14 @@ pub async fn run_command(args: RunArgs, styles: &'static Styles) -> anyhow::Resu None => (model, provider), }; + // Parse provider string to enum (defaults to Anthropic) + let provider_enum: Provider = provider + .as_deref() + .map(|s| s.parse::()) + .transpose() + .map_err(|e| anyhow::anyhow!("{e}"))? + .unwrap_or(Provider::Anthropic); + // 7. Build engine let registry = default_registry(interviewer.clone(), || { if dry_run_mode { @@ -406,13 +416,13 @@ pub async fn run_command(args: RunArgs, styles: &'static Styles) -> anyhow::Resu } else { let api = AgentBackend::new( model.clone(), - provider.clone(), + provider_enum, args.verbose, styles, ); let cli = CliBackend::new( model.clone(), - provider.clone().unwrap_or_else(|| "anthropic".to_string()), + provider_enum, ); Some(Box::new(BackendRouter::new(Box::new(api), cli))) } diff --git a/crates/attractor/src/cli/serve.rs b/crates/attractor/src/cli/serve.rs index 930473ae8..fcd26147c 100644 --- a/crates/attractor/src/cli/serve.rs +++ b/crates/attractor/src/cli/serve.rs @@ -1,5 +1,6 @@ use std::sync::Arc; +use llm::provider::Provider; use util::terminal::Styles; use tokio::net::TcpListener; @@ -40,30 +41,37 @@ pub async fn serve_command(args: ServeArgs, styles: &'static Styles) -> anyhow:: }; // Resolve model/provider defaults - let provider = args.provider; - let model = args.model.unwrap_or_else(|| match provider.as_deref() { + let provider_str = args.provider; + let model = args.model.unwrap_or_else(|| match provider_str.as_deref() { Some("openai") => "gpt-5.2".to_string(), Some("gemini") => "gemini-3.1-pro-preview".to_string(), _ => "claude-opus-4-6".to_string(), }); // Resolve model alias through catalog - let (model, provider) = match llm::catalog::get_model_info(&model) { - Some(info) => (info.id, provider.or(Some(info.provider))), - None => (model, provider), + let (model, provider_str) = match llm::catalog::get_model_info(&model) { + Some(info) => (info.id, provider_str.or(Some(info.provider))), + None => (model, provider_str), }; + // Parse provider string to enum (defaults to Anthropic) + let provider_enum: Provider = provider_str + .as_deref() + .map(|s| s.parse::()) + .transpose() + .map_err(|e| anyhow::anyhow!("{e}"))? + .unwrap_or(Provider::Anthropic); + // Build registry factory let factory = move |interviewer: Arc| { let model = model.clone(); - let provider = provider.clone(); default_registry(interviewer, move || { if dry_run_mode { None } else { Some(Box::new(AgentBackend::new( model.clone(), - provider.clone(), + provider_enum, 0, styles, ))) diff --git a/crates/attractor/tests/daytona_integration.rs b/crates/attractor/tests/daytona_integration.rs index 024b9051f..cf3a3d3ae 100644 --- a/crates/attractor/tests/daytona_integration.rs +++ b/crates/attractor/tests/daytona_integration.rs @@ -20,6 +20,7 @@ use attractor::handler::exit::ExitHandler; use attractor::handler::start::StartHandler; use attractor::handler::{Handler, HandlerRegistry}; use attractor::outcome::{Outcome, StageStatus}; +use llm::provider::Provider; async fn create_env() -> DaytonaExecutionEnvironment { dotenvy::dotenv().ok(); @@ -319,7 +320,7 @@ use attractor::handler::codergen::{CodergenBackend, CodergenResult}; /// /// Installs the CLI tool in the sandbox, then runs the CliBackend against it. async fn run_daytona_cli_test( - provider: &str, + provider: Provider, model: &str, install_command: &str, ) { @@ -338,7 +339,7 @@ async fn run_daytona_cli_test( install_result.exit_code, install_result.stdout ); - let backend = CliBackend::new(model.to_string(), provider.to_string()); + let backend = CliBackend::new(model.to_string(), provider); let node = Node::new("daytona_cli_test"); let context = Context::new(); let emitter = Arc::new(EventEmitter::new()); @@ -392,7 +393,7 @@ async fn run_daytona_cli_test( #[ignore] // requires DAYTONA_API_KEY + Claude CLI auth async fn daytona_cli_claude() { run_daytona_cli_test( - "anthropic", + Provider::Anthropic, "haiku", "curl -fsSL https://claude.ai/install.sh | sh", ) @@ -403,7 +404,7 @@ async fn daytona_cli_claude() { #[ignore] // requires DAYTONA_API_KEY + OpenAI/Codex auth async fn daytona_cli_codex() { run_daytona_cli_test( - "openai", + Provider::OpenAi, "o4-mini", "npm install -g @openai/codex", ) @@ -414,7 +415,7 @@ async fn daytona_cli_codex() { #[ignore] // requires DAYTONA_API_KEY + Gemini auth async fn daytona_cli_gemini() { run_daytona_cli_test( - "gemini", + Provider::Gemini, "gemini-2.5-flash", "npm install -g @google/gemini-cli", ) diff --git a/crates/attractor/tests/integration.rs b/crates/attractor/tests/integration.rs index 9a766533d..e829f0547 100644 --- a/crates/attractor/tests/integration.rs +++ b/crates/attractor/tests/integration.rs @@ -26,6 +26,7 @@ use attractor::stylesheet::{apply_stylesheet, parse_stylesheet}; use attractor::transform::{StylesheetApplicationTransform, Transform, VariableExpansionTransform}; use attractor::cli::backend::AgentBackend; use attractor::handler::default_registry; +use llm::provider::Provider; use attractor::validation::{validate, validate_or_raise, Severity}; use util::terminal::Styles; @@ -6629,7 +6630,7 @@ async fn attractor_e2e_with_real_llm() { let registry = default_registry(interviewer, move || { Some(Box::new(AgentBackend::new( model.clone(), - None, + Provider::Anthropic, 0, &TEST_STYLES, )) as Box) @@ -7331,7 +7332,7 @@ async fn cli_backend_run_writes_prompt_and_calls_exec() { let claude_output = r#"{"type":"result","result":"I fixed the bug.","usage":{"input_tokens":500,"output_tokens":200}}"#; let test_env = Arc::new(CliTestEnv::new(claude_output)); let env: Arc = test_env.clone(); - let backend = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let backend = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let node = Node::new("fix_code"); let context = Context::new(); @@ -7376,7 +7377,7 @@ async fn cli_backend_run_detects_changed_files() { CliTestEnv::new(claude_output) .with_git_diff_after("src/main.rs\nsrc/lib.rs\n"), ); - let backend = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let backend = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let node = Node::new("implement"); let context = Context::new(); @@ -7401,7 +7402,7 @@ async fn cli_backend_run_with_codex_provider() { let codex_output = "{\"type\":\"item.completed\",\"item\":{\"id\":\"item_0\",\"type\":\"agent_message\",\"text\":\"Implemented the feature.\"}}\n{\"type\":\"turn.completed\",\"usage\":{\"input_tokens\":300,\"output_tokens\":150}}"; let test_env = Arc::new(CliTestEnv::new(codex_output)); let env: Arc = test_env.clone(); - let backend = CliBackend::new("gpt-5.3-codex".into(), "openai".into()); + let backend = CliBackend::new("gpt-5.3-codex".into(), Provider::OpenAi); let node = Node::new("implement"); let context = Context::new(); @@ -7459,7 +7460,7 @@ async fn cli_backend_run_fails_on_nonzero_exit() { } let failing_env: Arc = Arc::new(FailingCliEnv); - let backend = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let backend = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let node = Node::new("step"); let context = Context::new(); let emitter = Arc::new(EventEmitter::new()); @@ -7483,7 +7484,7 @@ async fn cli_backend_run_fails_on_nonzero_exit() { #[tokio::test] async fn cli_backend_run_fails_on_unparseable_output() { let env: Arc = Arc::new(CliTestEnv::new("this is not json at all")); - let backend = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let backend = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let node = Node::new("step"); let context = Context::new(); @@ -7507,7 +7508,7 @@ async fn cli_backend_run_uses_node_model_override() { let claude_output = r#"{"type":"result","result":"ok","usage":{"input_tokens":10,"output_tokens":5}}"#; let test_env = Arc::new(CliTestEnv::new(claude_output)); let env: Arc = test_env.clone(); - let backend = CliBackend::new("default-model".into(), "anthropic".into()); + let backend = CliBackend::new("default-model".into(), Provider::Anthropic); let mut node = Node::new("step"); node.attrs.insert("llm_model".to_string(), AttrValue::String("claude-sonnet-4-5".to_string())); @@ -7532,7 +7533,7 @@ async fn cli_backend_run_uses_node_provider_override() { let codex_output = "{\"type\":\"item.completed\",\"item\":{\"id\":\"item_0\",\"type\":\"agent_message\",\"text\":\"ok\"}}\n{\"type\":\"turn.completed\",\"usage\":{\"input_tokens\":10,\"output_tokens\":5}}"; let test_env = Arc::new(CliTestEnv::new(codex_output)); let env: Arc = test_env.clone(); - let backend = CliBackend::new("default-model".into(), "anthropic".into()); + let backend = CliBackend::new("default-model".into(), Provider::Anthropic); let mut node = Node::new("step"); node.attrs.insert("llm_provider".to_string(), AttrValue::String("openai".to_string())); @@ -7556,7 +7557,7 @@ async fn cli_backend_run_uses_node_provider_override() { async fn cli_backend_run_writes_provider_used_json() { let claude_output = r#"{"type":"result","result":"done","usage":{"input_tokens":10,"output_tokens":5}}"#; let env: Arc = Arc::new(CliTestEnv::new(claude_output)); - let backend = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let backend = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let node = Node::new("step"); let context = Context::new(); @@ -7587,7 +7588,7 @@ async fn backend_router_delegates_to_cli_for_cli_node() { let env: Arc = Arc::new(CliTestEnv::new(claude_output)); let api_backend = Box::new(MockCodergenBackend); // would return "Response for ..." - let cli = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let cli = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let router = BackendRouter::new(api_backend, cli); let mut node = Node::new("cli_step"); @@ -7616,7 +7617,7 @@ async fn backend_router_delegates_to_api_for_normal_node() { let env = local_env(); let api_backend = Box::new(MockCodergenBackend); - let cli = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let cli = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let router = BackendRouter::new(api_backend, cli); let mut node = Node::new("api_step"); @@ -7645,7 +7646,7 @@ async fn backend_router_delegates_to_cli_for_backend_attr() { let env: Arc = Arc::new(CliTestEnv::new(codex_output)); let api_backend = Box::new(MockCodergenBackend); - let cli = CliBackend::new("gpt-5.3-codex".into(), "openai".into()); + let cli = CliBackend::new("gpt-5.3-codex".into(), Provider::OpenAi); let router = BackendRouter::new(api_backend, cli); let mut node = Node::new("codex_step"); @@ -7705,7 +7706,7 @@ async fn full_pipeline_with_cli_backend_node() { // Build engine with BackendRouter let api = MockCodergenBackend; - let cli = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let cli = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let router = BackendRouter::new(Box::new(api), cli); let codergen_handler = CodergenHandler::new(Some(Box::new(router))); @@ -7715,7 +7716,7 @@ async fn full_pipeline_with_cli_backend_node() { registry.register("codergen", Box::new(CodergenHandler::new(Some(Box::new({ // Second BackendRouter for the "codergen" handler let api2 = MockCodergenBackend; - let cli2 = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let cli2 = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); BackendRouter::new(Box::new(api2), cli2) }))))); @@ -7792,14 +7793,14 @@ async fn stylesheet_backend_property_routes_to_cli() { // Run the pipeline let api = MockCodergenBackend; - let cli = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let cli = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let router = BackendRouter::new(Box::new(api), cli); let mut registry = HandlerRegistry::new(Box::new(CodergenHandler::new(Some(Box::new(router))))); registry.register("start", Box::new(StartHandler)); registry.register("exit", Box::new(ExitHandler)); let api2 = MockCodergenBackend; - let cli2 = CliBackend::new("claude-opus-4-6".into(), "anthropic".into()); + let cli2 = CliBackend::new("claude-opus-4-6".into(), Provider::Anthropic); let router2 = BackendRouter::new(Box::new(api2), cli2); registry.register("codergen", Box::new(CodergenHandler::new(Some(Box::new(router2))))); @@ -7827,9 +7828,9 @@ async fn stylesheet_backend_property_routes_to_cli() { use attractor::cli::cli_backend::parse_cli_response; /// Run a real CLI tool via LocalExecutionEnvironment and verify the full flow. -async fn run_real_cli_test(provider: &str, model: &str) { +async fn run_real_cli_test(provider: Provider, model: &str) { let env = local_env(); - let backend = CliBackend::new(model.to_string(), provider.to_string()); + let backend = CliBackend::new(model.to_string(), provider); let mut node = Node::new("real_cli_test"); node.attrs.insert("prompt".to_string(), AttrValue::String("What is 2+2? Reply with just the number.".to_string())); @@ -7862,25 +7863,25 @@ async fn run_real_cli_test(provider: &str, model: &str) { &std::fs::read_to_string(&provider_path).unwrap() ).unwrap(); assert_eq!(provider_json["mode"], "cli"); - assert_eq!(provider_json["provider"], provider); + assert_eq!(provider_json["provider"], provider.as_str()); } #[tokio::test] #[ignore] // requires `claude` CLI installed async fn real_cli_claude() { - run_real_cli_test("anthropic", "haiku").await; + run_real_cli_test(Provider::Anthropic, "haiku").await; } #[tokio::test] #[ignore] // requires `codex` CLI installed and OpenAI auth async fn real_cli_codex() { - run_real_cli_test("openai", "").await; + run_real_cli_test(Provider::OpenAi, "").await; } #[tokio::test] #[ignore] // requires `gemini` CLI installed and Google auth async fn real_cli_gemini() { - run_real_cli_test("gemini", "gemini-2.5-flash").await; + run_real_cli_test(Provider::Gemini, "gemini-2.5-flash").await; } /// Verify parse_cli_response works against real Claude CLI output captured from stream-json. @@ -7890,7 +7891,7 @@ fn parse_real_claude_stream_json() { let output = r#"{"type":"system","subtype":"init","cwd":"/tmp","session_id":"abc"} {"type":"assistant","message":{"content":[{"type":"text","text":"4"}]}} {"type":"result","subtype":"success","is_error":false,"duration_ms":2000,"num_turns":1,"result":"4","usage":{"input_tokens":9,"output_tokens":5}}"#; - let response = parse_cli_response("anthropic", output).unwrap(); + let response = parse_cli_response(Provider::Anthropic, output).unwrap(); assert_eq!(response.text, "4"); assert_eq!(response.input_tokens, 9); assert_eq!(response.output_tokens, 5); @@ -7905,7 +7906,7 @@ fn parse_real_codex_ndjson() { {"type":"item.completed","item":{"id":"item_0","type":"reasoning","text":"**Confirming simple numeric reply**"}} {"type":"item.completed","item":{"id":"item_1","type":"agent_message","text":"4"}} {"type":"turn.completed","usage":{"input_tokens":7999,"cached_input_tokens":7040,"output_tokens":33}}"#; - let response = parse_cli_response("openai", output).unwrap(); + let response = parse_cli_response(Provider::OpenAi, output).unwrap(); assert_eq!(response.text, "4"); assert_eq!(response.input_tokens, 7999); assert_eq!(response.output_tokens, 33); @@ -7916,7 +7917,7 @@ fn parse_real_codex_ndjson() { fn parse_real_gemini_json() { // Real output captured from: gemini "What is 2+2?" -m gemini-2.5-flash --sandbox -o json let output = r#"{"session_id":"abc","response":"4","stats":{"models":{"gemini-2.5-flash":{"api":{"totalRequests":1,"totalErrors":0,"totalLatencyMs":618},"tokens":{"input":123,"prompt":8911,"candidates":1,"total":8912,"cached":8788,"thoughts":0,"tool":0}}},"tools":{"totalCalls":0},"files":{"totalLinesAdded":0,"totalLinesRemoved":0}}}"#; - let response = parse_cli_response("gemini", output).unwrap(); + let response = parse_cli_response(Provider::Gemini, output).unwrap(); assert_eq!(response.text, "4"); assert_eq!(response.input_tokens, 123); assert_eq!(response.output_tokens, 1); diff --git a/crates/llm/src/lib.rs b/crates/llm/src/lib.rs index c7927b25b..747bb09fe 100644 --- a/crates/llm/src/lib.rs +++ b/crates/llm/src/lib.rs @@ -11,3 +11,4 @@ pub mod providers; // Re-export module-level default client helpers (Section 2.5). pub use generate::set_default_client; +pub use provider::{ModelId, Provider}; diff --git a/crates/llm/src/provider.rs b/crates/llm/src/provider.rs index 1825608ca..8eef74d1a 100644 --- a/crates/llm/src/provider.rs +++ b/crates/llm/src/provider.rs @@ -1,7 +1,87 @@ use crate::error::SdkError; use crate::types::{Request, Response, StreamEvent, ToolChoice}; use futures::Stream; +use serde::{Deserialize, Serialize}; +use std::fmt; use std::pin::Pin; +use std::str::FromStr; + +// --------------------------------------------------------------------------- +// Provider enum — compile-time safe provider identity +// --------------------------------------------------------------------------- + +/// Known LLM provider variants. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum Provider { + Anthropic, + OpenAi, + Gemini, +} + +impl Provider { + /// Stable lowercase string representation used in `Request.provider`, + /// adapter names, and other serialization boundaries. + #[must_use] + pub fn as_str(self) -> &'static str { + match self { + Self::Anthropic => "anthropic", + Self::OpenAi => "openai", + Self::Gemini => "gemini", + } + } +} + +impl fmt::Display for Provider { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} + +impl FromStr for Provider { + type Err = String; + + fn from_str(s: &str) -> Result { + match s { + "anthropic" => Ok(Self::Anthropic), + "openai" | "open_ai" => Ok(Self::OpenAi), + "gemini" => Ok(Self::Gemini), + other => Err(format!("unknown provider: {other}")), + } + } +} + +// --------------------------------------------------------------------------- +// ModelId — bundles a provider with a model name +// --------------------------------------------------------------------------- + +/// A model identifier that pairs a [`Provider`] with the provider-specific +/// model name (e.g. `"claude-opus-4-6"` or `"gpt-4o-mini"`). +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct ModelId { + pub provider: Provider, + pub model: String, +} + +impl ModelId { + #[must_use] + pub fn new(provider: Provider, model: impl Into) -> Self { + Self { + provider, + model: model.into(), + } + } +} + +impl fmt::Display for ModelId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}:{}", self.provider, self.model) + } +} + +// --------------------------------------------------------------------------- +// ProviderAdapter trait +// --------------------------------------------------------------------------- /// Async stream of `StreamEvents` returned by streaming providers. pub type StreamEventStream =