mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-06 02:48:25 +00:00
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:
parent
3263d0c21e
commit
950a2d06a7
20 changed files with 349 additions and 263 deletions
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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![])
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)))
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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 =
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue