Introduce Provider enum and ModelId type for compile-time provider safety

Replace scattered provider string literals ("anthropic", "openai", "gemini")
with a Provider enum and ModelId struct in the llm crate. This prevents
bugs like routing an OpenAI model to Anthropic's API (the bug fixed in
3263d0c) by making the provider identity a compile-time checked value.

Key changes:
- Add Provider enum (Anthropic, OpenAi, Gemini) with as_str/Display/FromStr
- Add ModelId struct bundling Provider + model name
- Replace WebFetchSummarizer's separate model+provider fields with ModelId
- Replace BaseProfile.id: &'static str with BaseProfile.provider: Provider
- Replace ProviderProfile::id() -> &str with provider() -> Provider
- Parse --provider CLI strings to Provider early via FromStr
- Update AgentBackend and CliBackend to use Provider instead of String/Option

Serialization boundaries (Request.provider, Response.provider, Client HashMap
keys) remain as strings, converted via provider.as_str().

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-02-28 03:29:30 -05:00
parent 3263d0c21e
commit 950a2d06a7
20 changed files with 349 additions and 263 deletions

View file

@ -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<Client>) -> Option<crate::tools::WebFetchSummarizer> {
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<Client>) -> Box<dyn ProviderProfile> {
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<Client>) -> Option<crate::tools::WebFetchSummarizer> {
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<Client>) -> Box<dyn ProviderProfile> {
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<dyn ProviderProfile> = 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<dyn ProviderProfile> = 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");

View file

@ -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,

View file

@ -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");
}

View file

@ -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");
}

View file

@ -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,
}

View file

@ -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");
}

View file

@ -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<String> {
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");

View file

@ -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");
}

View file

@ -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)

View file

@ -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<dyn ProviderAdapter>) -> 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![])
}

View file

@ -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<String>,
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<WebFetchSummarizer>) -> 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<WebFetchSummarizer>) -> 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<dyn ProviderAdapter> = Arc::new(
let default_provider: Arc<dyn llm::provider::ProviderAdapter> = 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<dyn ProviderAdapter> = Arc::new(
// "anthropic" provider has the model we actually want.
let target_provider: Arc<dyn llm::provider::ProviderAdapter> = 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));

View file

@ -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<dyn ProviderProfile> {
fn build_profile(provider: Provider, model: &str, client: &Client) -> Box<dyn ProviderProfile> {
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<dyn ProviderProfile> = {
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<dyn ProviderProfile> = {
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<dyn ProviderProfile> = 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 [<anthropic_ $scenario>]() {
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;
[<scenario_ $scenario>](&mut session, tmp.path()).await;
}
@ -103,7 +99,7 @@ macro_rules! provider_tests {
#[ignore = "requires LLM API keys"]
async fn [<openai_ $scenario>]() {
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;
[<scenario_ $scenario>](&mut session, tmp.path()).await;
}
@ -112,7 +108,7 @@ macro_rules! provider_tests {
#[ignore = "requires LLM API keys"]
async fn [<gemini_ $scenario>]() {
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;
[<scenario_ $scenario>](&mut session, tmp.path()).await;
}
@ -149,7 +145,7 @@ macro_rules! anthropic_gemini_tests {
#[ignore = "requires LLM API keys"]
async fn [<anthropic_ $scenario>]() {
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;
[<scenario_ $scenario>](&mut session, tmp.path()).await;
}
@ -158,7 +154,7 @@ macro_rules! anthropic_gemini_tests {
#[ignore = "requires LLM API keys"]
async fn [<gemini_ $scenario>]() {
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;
[<scenario_ $scenario>](&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

View file

@ -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<String>,
provider: Provider,
verbose: u8,
styles: &'static Styles,
sessions: Mutex<HashMap<String, Session>>,
@ -33,7 +34,7 @@ impl AgentBackend {
#[must_use]
pub fn new(
model: String,
provider: Option<String>,
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<dyn ProviderProfile>,
"gemini" => Arc::new(GeminiProfile::new(&factory_model)) as Arc<dyn ProviderProfile>,
_ => Arc::new(AnthropicProfile::new(&factory_model)) as Arc<dyn ProviderProfile>,
}
let child_profile: Arc<dyn ProviderProfile> = 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<dyn ProviderProfile> {
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,
);

View file

@ -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<CliResponse> {
}
/// Parse CLI output, choosing the right parser based on provider.
pub fn parse_cli_response(provider: &str, output: &str) -> Option<CliResponse> {
pub fn parse_cli_response(provider: Provider, output: &str) -> Option<CliResponse> {
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::<Provider>().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,

View file

@ -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::<Provider>())
.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)))
}

View file

@ -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::<Provider>())
.transpose()
.map_err(|e| anyhow::anyhow!("{e}"))?
.unwrap_or(Provider::Anthropic);
// Build registry factory
let factory = move |interviewer: Arc<dyn Interviewer>| {
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,
)))

View file

@ -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",
)

View file

@ -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<dyn attractor::handler::codergen::CodergenBackend>)
@ -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<dyn agent::ExecutionEnvironment> = 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<dyn agent::ExecutionEnvironment> = 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<dyn agent::ExecutionEnvironment> = 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<dyn agent::ExecutionEnvironment> = 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<dyn agent::ExecutionEnvironment> = 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<dyn agent::ExecutionEnvironment> = 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<dyn agent::ExecutionEnvironment> = 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<dyn agent::ExecutionEnvironment> = 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<dyn agent::ExecutionEnvironment> = 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);

View file

@ -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};

View file

@ -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<Self, Self::Err> {
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<String>) -> 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 =