Feat: Let providers configure a fast model for summarization (#4228)

This commit is contained in:
David Katz
2025-08-21 17:41:33 -04:00
committed by GitHub
parent a121acd5e7
commit 72f4ebc640
33 changed files with 345 additions and 166 deletions
+2 -1
View File
@@ -537,8 +537,9 @@ mod tests {
goose::providers::base::ProviderMetadata::empty() goose::providers::base::ProviderMetadata::empty()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[rmcp::model::Tool], _tools: &[rmcp::model::Tool],
@@ -221,8 +221,9 @@ mod tests {
self.model_config.clone() self.model_config.clone()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[Tool], _tools: &[Tool],
+3 -2
View File
@@ -44,7 +44,7 @@ pub async fn summarize_messages(
// Send the request to the provider and fetch the response // Send the request to the provider and fetch the response
let (mut response, mut provider_usage) = provider let (mut response, mut provider_usage) = provider
.complete(&system_prompt, &summarization_request, &[]) .complete_fast(&system_prompt, &summarization_request, &[])
.await?; .await?;
// Set role to user as it will be used in following conversation as user content // Set role to user as it will be used in following conversation as user content
@@ -87,8 +87,9 @@ mod tests {
self.model_config.clone() self.model_config.clone()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[Tool], _tools: &[Tool],
+56 -3
View File
@@ -28,6 +28,7 @@ static MODEL_SPECIFIC_LIMITS: Lazy<Vec<(&'static str, usize)>> = Lazy::new(|| {
// anthropic - all 200k // anthropic - all 200k
("claude", 200_000), ("claude", 200_000),
// google // google
("gemini-1.5-flash", 1_000_000),
("gemini-1", 128_000), ("gemini-1", 128_000),
("gemini-2", 1_000_000), ("gemini-2", 1_000_000),
("gemma-3-27b", 128_000), ("gemma-3-27b", 128_000),
@@ -72,6 +73,7 @@ pub struct ModelConfig {
pub max_tokens: Option<i32>, pub max_tokens: Option<i32>,
pub toolshim: bool, pub toolshim: bool,
pub toolshim_model: Option<String>, pub toolshim_model: Option<String>,
pub fast_model: Option<String>,
} }
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
@@ -89,7 +91,7 @@ impl ModelConfig {
model_name: String, model_name: String,
context_env_var: Option<&str>, context_env_var: Option<&str>,
) -> Result<Self, ConfigError> { ) -> Result<Self, ConfigError> {
let context_limit = Self::parse_context_limit(&model_name, context_env_var)?; let context_limit = Self::parse_context_limit(&model_name, None, context_env_var)?;
let temperature = Self::parse_temperature()?; let temperature = Self::parse_temperature()?;
let toolshim = Self::parse_toolshim()?; let toolshim = Self::parse_toolshim()?;
let toolshim_model = Self::parse_toolshim_model()?; let toolshim_model = Self::parse_toolshim_model()?;
@@ -101,13 +103,16 @@ impl ModelConfig {
max_tokens: None, max_tokens: None,
toolshim, toolshim,
toolshim_model, toolshim_model,
fast_model: None,
}) })
} }
fn parse_context_limit( fn parse_context_limit(
model_name: &str, model_name: &str,
fast_model: Option<&str>,
custom_env_var: Option<&str>, custom_env_var: Option<&str>,
) -> Result<Option<usize>, ConfigError> { ) -> Result<Option<usize>, ConfigError> {
// First check if there's an explicit environment variable override
if let Some(env_var) = custom_env_var { if let Some(env_var) = custom_env_var {
if let Ok(val) = std::env::var(env_var) { if let Ok(val) = std::env::var(env_var) {
return Self::validate_context_limit(&val, env_var).map(Some); return Self::validate_context_limit(&val, env_var).map(Some);
@@ -116,7 +121,24 @@ impl ModelConfig {
if let Ok(val) = std::env::var("GOOSE_CONTEXT_LIMIT") { if let Ok(val) = std::env::var("GOOSE_CONTEXT_LIMIT") {
return Self::validate_context_limit(&val, "GOOSE_CONTEXT_LIMIT").map(Some); return Self::validate_context_limit(&val, "GOOSE_CONTEXT_LIMIT").map(Some);
} }
Ok(Self::get_model_specific_limit(model_name))
// Get the model's limit
let model_limit = Self::get_model_specific_limit(model_name);
// If there's a fast_model, get its limit and use the minimum
if let Some(fast_model_name) = fast_model {
let fast_model_limit = Self::get_model_specific_limit(fast_model_name);
// Return the minimum of both limits (if both exist)
match (model_limit, fast_model_limit) {
(Some(m), Some(f)) => Ok(Some(m.min(f))),
(Some(m), None) => Ok(Some(m)),
(None, Some(f)) => Ok(Some(f)),
(None, None) => Ok(None),
}
} else {
Ok(model_limit)
}
} }
fn validate_context_limit(val: &str, env_var: &str) -> Result<usize, ConfigError> { fn validate_context_limit(val: &str, env_var: &str) -> Result<usize, ConfigError> {
@@ -231,8 +253,39 @@ impl ModelConfig {
self self
} }
pub fn with_fast(mut self, fast_model: String) -> Self {
self.fast_model = Some(fast_model);
self
}
pub fn use_fast_model(&self) -> Self {
if let Some(fast_model) = &self.fast_model {
let mut config = self.clone();
config.model_name = fast_model.clone();
config
} else {
self.clone()
}
}
pub fn context_limit(&self) -> usize { pub fn context_limit(&self) -> usize {
self.context_limit.unwrap_or(DEFAULT_CONTEXT_LIMIT) // If we have an explicit context limit set, use it
if let Some(limit) = self.context_limit {
return limit;
}
// Otherwise, get the model's default limit
let main_limit =
Self::get_model_specific_limit(&self.model_name).unwrap_or(DEFAULT_CONTEXT_LIMIT);
// If we have a fast_model, also check its limit and use the minimum
if let Some(fast_model) = &self.fast_model {
let fast_limit =
Self::get_model_specific_limit(fast_model).unwrap_or(DEFAULT_CONTEXT_LIMIT);
main_limit.min(fast_limit)
} else {
main_limit
}
} }
pub fn new_or_fail(model_name: &str) -> ModelConfig { pub fn new_or_fail(model_name: &str) -> ModelConfig {
@@ -292,8 +292,9 @@ mod tests {
self.model_config.clone() self.model_config.clone()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[Tool], _tools: &[Tool],
+11 -7
View File
@@ -23,6 +23,7 @@ use crate::providers::retry::ProviderRetry;
use rmcp::model::Tool; use rmcp::model::Tool;
const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-4-0"; const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-4-0";
const ANTHROPIC_DEFAULT_FAST_MODEL: &str = "claude-3-7-sonnet-latest";
const ANTHROPIC_KNOWN_MODELS: &[&str] = &[ const ANTHROPIC_KNOWN_MODELS: &[&str] = &[
"claude-sonnet-4-0", "claude-sonnet-4-0",
"claude-sonnet-4-20250514", "claude-sonnet-4-20250514",
@@ -50,6 +51,8 @@ impl_provider_default!(AnthropicProvider);
impl AnthropicProvider { impl AnthropicProvider {
pub fn from_env(model: ModelConfig) -> Result<Self> { pub fn from_env(model: ModelConfig) -> Result<Self> {
let model = model.with_fast(ANTHROPIC_DEFAULT_FAST_MODEL.to_string());
let config = crate::config::Config::global(); let config = crate::config::Config::global();
let api_key: String = config.get_secret("ANTHROPIC_API_KEY")?; let api_key: String = config.get_secret("ANTHROPIC_API_KEY")?;
let host: String = config let host: String = config
@@ -179,16 +182,17 @@ impl Provider for AnthropicProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request(&self.model, system, messages, tools)?; let payload = create_request(model_config, system, messages, tools)?;
let response = self let response = self
.with_retry(|| async { self.post(&payload).await }) .with_retry(|| async { self.post(&payload).await })
@@ -201,9 +205,9 @@ impl Provider for AnthropicProvider {
tracing::debug!("🔍 Anthropic non-streaming parsed usage: input_tokens={:?}, output_tokens={:?}, total_tokens={:?}", tracing::debug!("🔍 Anthropic non-streaming parsed usage: input_tokens={:?}, output_tokens={:?}, total_tokens={:?}",
usage.input_tokens, usage.output_tokens, usage.total_tokens); usage.input_tokens, usage.output_tokens, usage.total_tokens);
let model = get_model(&json_response); let response_model = get_model(&json_response);
emit_debug_trace(&self.model, &payload, &json_response, &usage); emit_debug_trace(&self.model, &payload, &json_response, &usage);
let provider_usage = ProviderUsage::new(model, usage); let provider_usage = ProviderUsage::new(response_model, usage);
tracing::debug!( tracing::debug!(
"🔍 Anthropic non-streaming returning ProviderUsage: {:?}", "🔍 Anthropic non-streaming returning ProviderUsage: {:?}",
provider_usage provider_usage
@@ -271,7 +275,7 @@ impl Provider for AnthropicProvider {
let stream = response.bytes_stream().map_err(io::Error::other); let stream = response.bytes_stream().map_err(io::Error::other);
let model_config = self.model.clone(); let model = self.model.clone();
Ok(Box::pin(try_stream! { Ok(Box::pin(try_stream! {
let stream_reader = StreamReader::new(stream); let stream_reader = StreamReader::new(stream);
let framed = tokio_util::codec::FramedRead::new(stream_reader, tokio_util::codec::LinesCodec::new()).map_err(anyhow::Error::from); let framed = tokio_util::codec::FramedRead::new(stream_reader, tokio_util::codec::LinesCodec::new()).map_err(anyhow::Error::from);
@@ -280,7 +284,7 @@ impl Provider for AnthropicProvider {
pin!(message_stream); pin!(message_stream);
while let Some(message) = futures::StreamExt::next(&mut message_stream).await { while let Some(message) = futures::StreamExt::next(&mut message_stream).await {
let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?; let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?;
emit_debug_trace(&model_config, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default()); emit_debug_trace(&model, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default());
yield (message, usage); yield (message, usage);
} }
})) }))
+7 -6
View File
@@ -135,16 +135,17 @@ impl Provider for AzureProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request(&self.model, system, messages, tools, &ImageFormat::OpenAi)?; let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?;
let response = self let response = self
.with_retry(|| async { .with_retry(|| async {
let payload_clone = payload.clone(); let payload_clone = payload.clone();
@@ -157,8 +158,8 @@ impl Provider for AzureProvider {
tracing::debug!("Failed to get usage data"); tracing::debug!("Failed to get usage data");
Usage::default() Usage::default()
}); });
let model = get_model(&response); let response_model = get_model(&response);
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
} }
+30 -15
View File
@@ -317,25 +317,40 @@ pub trait Provider: Send + Sync {
where where
Self: Sized; Self: Sized;
/// Generate the next message using the configured model and other parameters // Internal implementation of complete, used by complete_fast and complete
/// // Providers should override this to implement their actual completion logic
/// # Arguments async fn complete_with_model(
/// * `system` - The system prompt that guides the model's behavior &self,
/// * `messages` - The conversation history as a sequence of messages model_config: &ModelConfig,
/// * `tools` - Optional list of tools the model can use system: &str,
/// messages: &[Message],
/// # Returns tools: &[Tool],
/// A tuple containing the model's response message and provider usage statistics ) -> Result<(Message, ProviderUsage), ProviderError>;
///
/// # Errors // Default implementation: use the provider's configured model
/// ProviderError
/// - It's important to raise ContextLengthExceeded correctly since agent handles it
async fn complete( async fn complete(
&self, &self,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError>; ) -> Result<(Message, ProviderUsage), ProviderError> {
let model_config = self.get_model_config();
self.complete_with_model(&model_config, system, messages, tools)
.await
}
// Check if a fast model is configured, otherwise fall back to regular model
async fn complete_fast(
&self,
system: &str,
messages: &[Message],
tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
let model_config = self.get_model_config();
let fast_config = model_config.use_fast_model();
self.complete_with_model(&fast_config, system, messages, tools)
.await
}
/// Get the model config from the provider /// Get the model config from the provider
fn get_model_config(&self) -> ModelConfig; fn get_model_config(&self) -> ModelConfig;
@@ -418,7 +433,7 @@ pub trait Provider: Send + Sync {
let prompt = self.create_session_name_prompt(&context); let prompt = self.create_session_name_prompt(&context);
let message = Message::user().with_text(&prompt); let message = Message::user().with_text(&prompt);
let result = self let result = self
.complete( .complete_fast(
"Reply with only a description in four words or less", "Reply with only a description in four words or less",
&[message], &[message],
&[], &[],
+4 -3
View File
@@ -152,16 +152,17 @@ impl Provider for BedrockProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let model_name = &self.model.model_name; let model_name = model_config.model_name.clone();
let (bedrock_message, bedrock_usage) = self let (bedrock_message, bedrock_usage) = self
.with_retry(|| self.converse(system, messages, tools)) .with_retry(|| self.converse(system, messages, tools))
+6 -5
View File
@@ -474,11 +474,12 @@ impl Provider for ClaudeCodeProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
@@ -495,7 +496,7 @@ impl Provider for ClaudeCodeProvider {
// Create a dummy payload for debug tracing // Create a dummy payload for debug tracing
let payload = json!({ let payload = json!({
"command": self.command, "command": self.command,
"model": self.model.model_name, "model": model_config.model_name,
"system": system, "system": system,
"messages": messages.len() "messages": messages.len()
}); });
@@ -505,11 +506,11 @@ impl Provider for ClaudeCodeProvider {
"usage": usage "usage": usage
}); });
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
Ok(( Ok((
message, message,
ProviderUsage::new(self.model.model_name.clone(), usage), ProviderUsage::new(model_config.model_name.clone(), usage),
)) ))
} }
} }
+6 -5
View File
@@ -407,11 +407,12 @@ impl Provider for CursorAgentProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
@@ -428,7 +429,7 @@ impl Provider for CursorAgentProvider {
// Create a dummy payload for debug tracing // Create a dummy payload for debug tracing
let payload = json!({ let payload = json!({
"command": self.command, "command": self.command,
"model": self.model.model_name, "model": model_config.model_name,
"system": system, "system": system,
"messages": messages.len() "messages": messages.len()
}); });
@@ -438,11 +439,11 @@ impl Provider for CursorAgentProvider {
"usage": usage "usage": usage
}); });
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
Ok(( Ok((
message, message,
ProviderUsage::new(self.model.model_name.clone(), usage), ProviderUsage::new(model_config.model_name.clone(), usage),
)) ))
} }
} }
+55 -18
View File
@@ -37,6 +37,7 @@ const DEFAULT_SCOPES: &[&str] = &["all-apis", "offline_access"];
const DEFAULT_TIMEOUT_SECS: u64 = 600; const DEFAULT_TIMEOUT_SECS: u64 = 600;
pub const DATABRICKS_DEFAULT_MODEL: &str = "databricks-claude-3-7-sonnet"; pub const DATABRICKS_DEFAULT_MODEL: &str = "databricks-claude-3-7-sonnet";
const DATABRICKS_DEFAULT_FAST_MODEL: &str = "gemini-1-5-flash";
pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[ pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[
"databricks-meta-llama-3-3-70b-instruct", "databricks-meta-llama-3-3-70b-instruct",
"databricks-meta-llama-3-1-405b-instruct", "databricks-meta-llama-3-1-405b-instruct",
@@ -137,13 +138,41 @@ impl DatabricksProvider {
let api_client = let api_client =
ApiClient::with_timeout(host, auth_method, Duration::from_secs(DEFAULT_TIMEOUT_SECS))?; ApiClient::with_timeout(host, auth_method, Duration::from_secs(DEFAULT_TIMEOUT_SECS))?;
Ok(Self { // Create the provider without the fast model first
let mut provider = Self {
api_client, api_client,
auth, auth,
model, model: model.clone(),
image_format: ImageFormat::OpenAi, image_format: ImageFormat::OpenAi,
retry_config, retry_config,
}) };
// Check if the default fast model exists in the workspace
let model_with_fast = tokio::task::block_in_place(|| {
tokio::runtime::Handle::current().block_on(async {
if let Ok(Some(models)) = provider.fetch_supported_models().await {
if models.contains(&DATABRICKS_DEFAULT_FAST_MODEL.to_string()) {
tracing::debug!(
"Found {} in Databricks workspace, setting as fast model",
DATABRICKS_DEFAULT_FAST_MODEL
);
model.with_fast(DATABRICKS_DEFAULT_FAST_MODEL.to_string())
} else {
tracing::debug!(
"{} not found in Databricks workspace, not setting fast model",
DATABRICKS_DEFAULT_FAST_MODEL
);
model
}
} else {
tracing::debug!("Could not fetch Databricks models, not setting fast model");
model
}
})
});
provider.model = model_with_fast;
Ok(provider)
} }
fn load_retry_config(config: &crate::config::Config) -> RetryConfig { fn load_retry_config(config: &crate::config::Config) -> RetryConfig {
@@ -195,17 +224,18 @@ impl DatabricksProvider {
}) })
} }
fn get_endpoint_path(&self, is_embedding: bool) -> String { fn get_endpoint_path(&self, model_name: &str, is_embedding: bool) -> String {
if is_embedding { if is_embedding {
"serving-endpoints/text-embedding-3-small/invocations".to_string() "serving-endpoints/text-embedding-3-small/invocations".to_string()
} else { } else {
format!("serving-endpoints/{}/invocations", self.model.model_name) format!("serving-endpoints/{}/invocations", model_name)
} }
} }
async fn post(&self, payload: Value) -> Result<Value, ProviderError> { async fn post(&self, payload: Value, model_name: Option<&str>) -> Result<Value, ProviderError> {
let is_embedding = payload.get("input").is_some() && payload.get("messages").is_none(); let is_embedding = payload.get("input").is_some() && payload.get("messages").is_none();
let path = self.get_endpoint_path(is_embedding); let model_to_use = model_name.unwrap_or(&self.model.model_name);
let path = self.get_endpoint_path(model_to_use, is_embedding);
let response = self.api_client.response_post(&path, &payload).await?; let response = self.api_client.response_post(&path, &payload).await?;
handle_response_openai_compat(response).await handle_response_openai_compat(response).await
@@ -238,32 +268,36 @@ impl Provider for DatabricksProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let mut payload = create_request(&self.model, system, messages, tools, &self.image_format)?; let mut payload =
create_request(model_config, system, messages, tools, &self.image_format)?;
payload payload
.as_object_mut() .as_object_mut()
.expect("payload should have model key") .expect("payload should have model key")
.remove("model"); .remove("model");
let response = self.with_retry(|| self.post(payload.clone())).await?; let response = self
.with_retry(|| self.post(payload.clone(), Some(&model_config.model_name)))
.await?;
let message = response_to_message(&response)?; let message = response_to_message(&response)?;
let usage = response.get("usage").map(get_usage).unwrap_or_else(|| { let usage = response.get("usage").map(get_usage).unwrap_or_else(|| {
tracing::debug!("Failed to get usage data"); tracing::debug!("Failed to get usage data");
Usage::default() Usage::default()
}); });
let model = get_model(&response); let response_model = get_model(&response);
super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); super::utils::emit_debug_trace(&self.model, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
async fn stream( async fn stream(
@@ -272,7 +306,10 @@ impl Provider for DatabricksProvider {
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<MessageStream, ProviderError> { ) -> Result<MessageStream, ProviderError> {
let mut payload = create_request(&self.model, system, messages, tools, &self.image_format)?; let model_config = self.model.clone();
let mut payload =
create_request(&model_config, system, messages, tools, &self.image_format)?;
payload payload
.as_object_mut() .as_object_mut()
.expect("payload should have model key") .expect("payload should have model key")
@@ -283,7 +320,7 @@ impl Provider for DatabricksProvider {
.unwrap() .unwrap()
.insert("stream".to_string(), Value::Bool(true)); .insert("stream".to_string(), Value::Bool(true));
let path = self.get_endpoint_path(false); let path = self.get_endpoint_path(&model_config.model_name, false);
let response = self let response = self
.with_retry(|| async { .with_retry(|| async {
let resp = self.api_client.response_post(&path, &payload).await?; let resp = self.api_client.response_post(&path, &payload).await?;
@@ -299,8 +336,8 @@ impl Provider for DatabricksProvider {
.await?; .await?;
let stream = response.bytes_stream().map_err(io::Error::other); let stream = response.bytes_stream().map_err(io::Error::other);
let model_config = self.model.clone();
let model = self.model.clone();
Ok(Box::pin(try_stream! { Ok(Box::pin(try_stream! {
let stream_reader = StreamReader::new(stream); let stream_reader = StreamReader::new(stream);
let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from); let framed = FramedRead::new(stream_reader, LinesCodec::new()).map_err(anyhow::Error::from);
@@ -309,7 +346,7 @@ impl Provider for DatabricksProvider {
pin!(message_stream); pin!(message_stream);
while let Some(message) = message_stream.next().await { while let Some(message) = message_stream.next().await {
let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?; let (message, usage) = message.map_err(|e| ProviderError::RequestFailed(format!("Stream decode error: {}", e)))?;
super::utils::emit_debug_trace(&model_config, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default()); super::utils::emit_debug_trace(&model, &payload, &message, &usage.as_ref().map(|f| f.usage).unwrap_or_default());
yield (message, usage); yield (message, usage);
} }
})) }))
@@ -408,7 +445,7 @@ impl EmbeddingCapable for DatabricksProvider {
"input": texts, "input": texts,
}); });
let response = self.with_retry(|| self.post(request.clone())).await?; let response = self.with_retry(|| self.post(request.clone(), None)).await?;
let embeddings = response["data"] let embeddings = response["data"]
.as_array() .as_array()
+2 -1
View File
@@ -202,8 +202,9 @@ mod tests {
self.model_config.clone() self.model_config.clone()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[Tool], _tools: &[Tool],
@@ -1045,6 +1045,7 @@ mod tests {
max_tokens: Some(1024), max_tokens: Some(1024),
toolshim: false, toolshim: false,
toolshim_model: None, toolshim_model: None,
fast_model: None,
}; };
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
let obj = request.as_object().unwrap(); let obj = request.as_object().unwrap();
@@ -1076,6 +1077,7 @@ mod tests {
max_tokens: Some(1024), max_tokens: Some(1024),
toolshim: false, toolshim: false,
toolshim_model: None, toolshim_model: None,
fast_model: None,
}; };
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
let obj = request.as_object().unwrap(); let obj = request.as_object().unwrap();
@@ -1108,6 +1110,7 @@ mod tests {
max_tokens: Some(1024), max_tokens: Some(1024),
toolshim: false, toolshim: false,
toolshim_model: None, toolshim_model: None,
fast_model: None,
}; };
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
let obj = request.as_object().unwrap(); let obj = request.as_object().unwrap();
@@ -1077,6 +1077,7 @@ mod tests {
max_tokens: Some(1024), max_tokens: Some(1024),
toolshim: false, toolshim: false,
toolshim_model: None, toolshim_model: None,
fast_model: None,
}; };
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
let obj = request.as_object().unwrap(); let obj = request.as_object().unwrap();
@@ -1108,6 +1109,7 @@ mod tests {
max_tokens: Some(1024), max_tokens: Some(1024),
toolshim: false, toolshim: false,
toolshim_model: None, toolshim_model: None,
fast_model: None,
}; };
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
let obj = request.as_object().unwrap(); let obj = request.as_object().unwrap();
@@ -1140,6 +1142,7 @@ mod tests {
max_tokens: Some(1024), max_tokens: Some(1024),
toolshim: false, toolshim: false,
toolshim_model: None, toolshim_model: None,
fast_model: None,
}; };
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?; let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
let obj = request.as_object().unwrap(); let obj = request.as_object().unwrap();
+5 -4
View File
@@ -512,23 +512,24 @@ impl Provider for GcpVertexAIProvider {
/// * `messages` - Array of previous messages in the conversation /// * `messages` - Array of previous messages in the conversation
/// * `tools` - Array of available tools for the model /// * `tools` - Array of available tools for the model
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
// Create request and context // Create request and context
let (request, context) = create_request(&self.model, system, messages, tools)?; let (request, context) = create_request(model_config, system, messages, tools)?;
// Send request and process response // Send request and process response
let response = self.post(&request, &context).await?; let response = self.post(&request, &context).await?;
let usage = get_usage(&response, &context)?; let usage = get_usage(&response, &context)?;
emit_debug_trace(&self.model, &request, &response, &usage); emit_debug_trace(model_config, &request, &response, &usage);
// Convert response to message // Convert response to message
let message = response_to_message(response, context)?; let message = response_to_message(response, context)?;
+4 -3
View File
@@ -319,11 +319,12 @@ impl Provider for GeminiCliProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
@@ -350,7 +351,7 @@ impl Provider for GeminiCliProvider {
"usage": usage "usage": usage
}); });
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
Ok(( Ok((
message, message,
+7 -6
View File
@@ -401,16 +401,17 @@ impl Provider for GithubCopilotProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request(&self.model, system, messages, tools, &ImageFormat::OpenAi)?; let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?;
// Make request with retry // Make request with retry
let response = self let response = self
@@ -426,9 +427,9 @@ impl Provider for GithubCopilotProvider {
tracing::debug!("Failed to get usage data"); tracing::debug!("Failed to get usage data");
Usage::default() Usage::default()
}); });
let model = get_model(&response); let response_model = get_model(&response);
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
/// Fetch supported models from GitHub Copliot; returns Err on failure, Ok(None) if not present /// Fetch supported models from GitHub Copliot; returns Err on failure, Ok(None) if not present
+14 -10
View File
@@ -14,6 +14,7 @@ use serde_json::Value;
pub const GOOGLE_API_HOST: &str = "https://generativelanguage.googleapis.com"; pub const GOOGLE_API_HOST: &str = "https://generativelanguage.googleapis.com";
pub const GOOGLE_DEFAULT_MODEL: &str = "gemini-2.5-flash"; pub const GOOGLE_DEFAULT_MODEL: &str = "gemini-2.5-flash";
pub const GOOGLE_DEFAULT_FAST_MODEL: &str = "gemini-1.5-flash";
pub const GOOGLE_KNOWN_MODELS: &[&str] = &[ pub const GOOGLE_KNOWN_MODELS: &[&str] = &[
// Gemini 2.5 models (latest generation) // Gemini 2.5 models (latest generation)
"gemini-2.5-pro", "gemini-2.5-pro",
@@ -55,6 +56,8 @@ impl_provider_default!(GoogleProvider);
impl GoogleProvider { impl GoogleProvider {
pub fn from_env(model: ModelConfig) -> Result<Self> { pub fn from_env(model: ModelConfig) -> Result<Self> {
let model = model.with_fast(GOOGLE_DEFAULT_FAST_MODEL.to_string());
let config = crate::config::Config::global(); let config = crate::config::Config::global();
let api_key: String = config.get_secret("GOOGLE_API_KEY")?; let api_key: String = config.get_secret("GOOGLE_API_KEY")?;
let host: String = config let host: String = config
@@ -72,8 +75,8 @@ impl GoogleProvider {
Ok(Self { api_client, model }) Ok(Self { api_client, model })
} }
async fn post(&self, payload: &Value) -> Result<Value, ProviderError> { async fn post(&self, model_name: &str, payload: &Value) -> Result<Value, ProviderError> {
let path = format!("v1beta/models/{}:generateContent", self.model.model_name); let path = format!("v1beta/models/{}:generateContent", model_name);
let response = self.api_client.response_post(&path, payload).await?; let response = self.api_client.response_post(&path, payload).await?;
handle_response_google_compat(response).await handle_response_google_compat(response).await
} }
@@ -101,34 +104,35 @@ impl Provider for GoogleProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request(&self.model, system, messages, tools)?; let payload = create_request(model_config, system, messages, tools)?;
// Make request // Make request
let response = self let response = self
.with_retry(|| async { .with_retry(|| async {
let payload_clone = payload.clone(); let payload_clone = payload.clone();
self.post(&payload_clone).await self.post(&model_config.model_name, &payload_clone).await
}) })
.await?; .await?;
// Parse response // Parse response
let message = response_to_message(unescape_json_values(&response))?; let message = response_to_message(unescape_json_values(&response))?;
let usage = get_usage(&response)?; let usage = get_usage(&response)?;
let model = match response.get("modelVersion") { let response_model = match response.get("modelVersion") {
Some(model_version) => model_version.as_str().unwrap_or_default().to_string(), Some(model_version) => model_version.as_str().unwrap_or_default().to_string(),
None => self.model.model_name.clone(), None => model_config.model_name.clone(),
}; };
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
let provider_usage = ProviderUsage::new(model, usage); let provider_usage = ProviderUsage::new(response_model, usage);
Ok((message, provider_usage)) Ok((message, provider_usage))
} }
+7 -6
View File
@@ -77,17 +77,18 @@ impl Provider for GroqProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request( let payload = create_request(
&self.model, model_config,
system, system,
messages, messages,
tools, tools,
@@ -101,9 +102,9 @@ impl Provider for GroqProvider {
tracing::debug!("Failed to get usage data"); tracing::debug!("Failed to get usage data");
Usage::default() Usage::default()
}); });
let model = get_model(&response); let response_model = get_model(&response);
super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); super::utils::emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
/// Fetch supported models from Groq; returns Err on failure, Ok(None) if no models found /// Fetch supported models from Groq; returns Err on failure, Ok(None) if no models found
+6 -3
View File
@@ -326,8 +326,9 @@ impl Provider for LeadWorkerProvider {
self.lead_provider.get_model_config() self.lead_provider.get_model_config()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
@@ -475,8 +476,9 @@ mod tests {
self.model_config.clone() self.model_config.clone()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[Tool], _tools: &[Tool],
@@ -635,8 +637,9 @@ mod tests {
self.model_config.clone() self.model_config.clone()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[Tool], _tools: &[Tool],
+6 -5
View File
@@ -161,14 +161,15 @@ impl Provider for LiteLLMProvider {
} }
#[tracing::instrument(skip_all, name = "provider_complete")] #[tracing::instrument(skip_all, name = "provider_complete")]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let mut payload = super::formats::openai::create_request( let mut payload = super::formats::openai::create_request(
&self.model, model_config,
system, system,
messages, messages,
tools, tools,
@@ -188,9 +189,9 @@ impl Provider for LiteLLMProvider {
let message = super::formats::openai::response_to_message(&response)?; let message = super::formats::openai::response_to_message(&response)?;
let usage = super::formats::openai::get_usage(&response); let usage = super::formats::openai::get_usage(&response);
let model = get_model(&response); let response_model = get_model(&response);
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
fn supports_embeddings(&self) -> bool { fn supports_embeddings(&self) -> bool {
+6 -5
View File
@@ -165,11 +165,12 @@ impl Provider for OllamaProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
@@ -197,9 +198,9 @@ impl Provider for OllamaProvider {
tracing::debug!("Failed to get usage data"); tracing::debug!("Failed to get usage data");
Usage::default() Usage::default()
}); });
let model = get_model(&response); let response_model = get_model(&response);
super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); super::utils::emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
/// Generate a session name based on the conversation history /// Generate a session name based on the conversation history
+7 -3
View File
@@ -29,6 +29,7 @@ use crate::providers::formats::openai::response_to_streaming_message;
use rmcp::model::Tool; use rmcp::model::Tool;
pub const OPEN_AI_DEFAULT_MODEL: &str = "gpt-4o"; pub const OPEN_AI_DEFAULT_MODEL: &str = "gpt-4o";
pub const OPEN_AI_DEFAULT_FAST_MODEL: &str = "gpt-4o-mini";
pub const OPEN_AI_KNOWN_MODELS: &[(&str, usize)] = &[ pub const OPEN_AI_KNOWN_MODELS: &[(&str, usize)] = &[
("gpt-4o", 128_000), ("gpt-4o", 128_000),
("gpt-4o-mini", 128_000), ("gpt-4o-mini", 128_000),
@@ -59,6 +60,8 @@ impl_provider_default!(OpenAiProvider);
impl OpenAiProvider { impl OpenAiProvider {
pub fn from_env(model: ModelConfig) -> Result<Self> { pub fn from_env(model: ModelConfig) -> Result<Self> {
let model = model.with_fast(OPEN_AI_DEFAULT_FAST_MODEL.to_string());
let config = crate::config::Config::global(); let config = crate::config::Config::global();
let api_key: String = config.get_secret("OPENAI_API_KEY")?; let api_key: String = config.get_secret("OPENAI_API_KEY")?;
let host: String = config let host: String = config
@@ -193,16 +196,17 @@ impl Provider for OpenAiProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request(&self.model, system, messages, tools, &ImageFormat::OpenAi)?; let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?;
let json_response = self.post(&payload).await?; let json_response = self.post(&payload).await?;
+6 -5
View File
@@ -238,11 +238,12 @@ impl Provider for OpenRouterProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
@@ -264,9 +265,9 @@ impl Provider for OpenRouterProvider {
tracing::debug!("Failed to get usage data"); tracing::debug!("Failed to get usage data");
Usage::default() Usage::default()
}); });
let model = get_model(&response); let response_model = get_model(&response);
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
/// Fetch supported models from OpenRouter API (only models with tool support) /// Fetch supported models from OpenRouter API (only models with tool support)
+4 -3
View File
@@ -280,16 +280,17 @@ impl Provider for SageMakerTgiProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let model_name = &self.model.model_name; let model_name = &model_config.model_name;
let request_payload = self.create_tgi_request(system, messages).map_err(|e| { let request_payload = self.create_tgi_request(system, messages).map_err(|e| {
ProviderError::RequestFailed(format!("Failed to create request: {}", e)) ProviderError::RequestFailed(format!("Failed to create request: {}", e))
+7 -6
View File
@@ -299,16 +299,17 @@ impl Provider for SnowflakeProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request(&self.model, system, messages, tools)?; let payload = create_request(model_config, system, messages, tools)?;
let response = self let response = self
.with_retry(|| async { .with_retry(|| async {
@@ -320,9 +321,9 @@ impl Provider for SnowflakeProvider {
// Parse response // Parse response
let message = response_to_message(&response)?; let message = response_to_message(&response)?;
let usage = get_usage(&response)?; let usage = get_usage(&response)?;
let model = get_model(&response); let response_model = get_model(&response);
super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); super::utils::emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
} }
+4 -2
View File
@@ -112,8 +112,9 @@ impl Provider for TestProvider {
) )
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
@@ -188,8 +189,9 @@ mod tests {
) )
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[Tool], _tools: &[Tool],
+13 -23
View File
@@ -1,4 +1,4 @@
use anyhow::{Error, Result}; use anyhow::Result;
use async_trait::async_trait; use async_trait::async_trait;
use serde_json::Value; use serde_json::Value;
@@ -113,23 +113,6 @@ impl TetrateProvider {
} }
} }
fn create_request_based_on_model(
provider: &TetrateProvider,
system: &str,
messages: &[Message],
tools: &[Tool],
) -> anyhow::Result<Value, Error> {
let payload = create_request(
&provider.model,
system,
messages,
tools,
&super::utils::ImageFormat::OpenAi,
)?;
Ok(payload)
}
#[async_trait] #[async_trait]
impl Provider for TetrateProvider { impl Provider for TetrateProvider {
fn metadata() -> ProviderMetadata { fn metadata() -> ProviderMetadata {
@@ -157,17 +140,24 @@ impl Provider for TetrateProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
// Create the base payload // Create the base payload using the provided model_config
let payload = create_request_based_on_model(self, system, messages, tools)?; let payload = create_request(
model_config,
system,
messages,
tools,
&super::utils::ImageFormat::OpenAi,
)?;
// Make request // Make request
let response = self let response = self
@@ -184,7 +174,7 @@ impl Provider for TetrateProvider {
Usage::default() Usage::default()
}); });
let model = get_model(&response); let model = get_model(&response);
emit_debug_trace(&self.model, &payload, &response, &usage); emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(model, usage)))
} }
+8 -7
View File
@@ -246,12 +246,13 @@ impl Provider for VeniceProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(_system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
_system: &str, model_config: &ModelConfig,
system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
@@ -259,10 +260,10 @@ impl Provider for VeniceProvider {
let mut formatted_messages = Vec::new(); let mut formatted_messages = Vec::new();
// Add the system message if present // Add the system message if present
if !_system.is_empty() { if !system.is_empty() {
formatted_messages.push(json!({ formatted_messages.push(json!({
"role": "system", "role": "system",
"content": _system "content": system
})); }));
} }
@@ -391,7 +392,7 @@ impl Provider for VeniceProvider {
// Build Venice-specific payload // Build Venice-specific payload
let mut payload = json!({ let mut payload = json!({
"model": strip_flags(&self.model.model_name), "model": strip_flags(&model_config.model_name),
"messages": formatted_messages, "messages": formatted_messages,
"stream": false, "stream": false,
"temperature": 0.7, "temperature": 0.7,
@@ -470,7 +471,7 @@ impl Provider for VeniceProvider {
return Ok(( return Ok((
message, message,
ProviderUsage::new( ProviderUsage::new(
strip_flags(&self.model.model_name).to_string(), strip_flags(&model_config.model_name).to_string(),
Usage::default(), Usage::default(),
), ),
)); ));
+7 -6
View File
@@ -93,17 +93,18 @@ impl Provider for XaiProvider {
} }
#[tracing::instrument( #[tracing::instrument(
skip(self, system, messages, tools), skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens) fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)] )]
async fn complete( async fn complete_with_model(
&self, &self,
model_config: &ModelConfig,
system: &str, system: &str,
messages: &[Message], messages: &[Message],
tools: &[Tool], tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request( let payload = create_request(
&self.model, model_config,
system, system,
messages, messages,
tools, tools,
@@ -117,8 +118,8 @@ impl Provider for XaiProvider {
tracing::debug!("Failed to get usage data"); tracing::debug!("Failed to get usage data");
Usage::default() Usage::default()
}); });
let model = get_model(&response); let response_model = get_model(&response);
super::utils::emit_debug_trace(&self.model, &payload, &response, &usage); super::utils::emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage))) Ok((message, ProviderUsage::new(response_model, usage)))
} }
} }
+2 -1
View File
@@ -1390,8 +1390,9 @@ mod tests {
self.model_config.clone() self.model_config.clone()
} }
async fn complete( async fn complete_with_model(
&self, &self,
_model_config: &ModelConfig,
_system: &str, _system: &str,
_messages: &[Message], _messages: &[Message],
_tools: &[Tool], _tools: &[Tool],
+40
View File
@@ -592,6 +592,16 @@ mod final_output_tool_tests {
ProviderUsage::new("mock".to_string(), Usage::default()), ProviderUsage::new("mock".to_string(), Usage::default()),
)) ))
} }
async fn complete_with_model(
&self,
_model_config: &ModelConfig,
system: &str,
messages: &[Message],
tools: &[Tool],
) -> anyhow::Result<(Message, ProviderUsage), ProviderError> {
self.complete(system, messages, tools).await
}
} }
let agent = Agent::new(); let agent = Agent::new();
@@ -713,6 +723,16 @@ mod final_output_tool_tests {
) -> Result<(Message, ProviderUsage), ProviderError> { ) -> Result<(Message, ProviderUsage), ProviderError> {
Err(ProviderError::NotImplemented("Not implemented".to_string())) Err(ProviderError::NotImplemented("Not implemented".to_string()))
} }
async fn complete_with_model(
&self,
_model_config: &ModelConfig,
system: &str,
messages: &[Message],
tools: &[Tool],
) -> anyhow::Result<(Message, ProviderUsage), ProviderError> {
self.complete(system, messages, tools).await
}
} }
let agent = Agent::new(); let agent = Agent::new();
@@ -829,6 +849,16 @@ mod retry_tests {
)) ))
} }
} }
async fn complete_with_model(
&self,
_model_config: &ModelConfig,
system: &str,
messages: &[Message],
tools: &[Tool],
) -> anyhow::Result<(Message, ProviderUsage), ProviderError> {
self.complete(system, messages, tools).await
}
} }
#[tokio::test] #[tokio::test]
@@ -1002,6 +1032,16 @@ mod max_turns_tests {
Ok((message, usage)) Ok((message, usage))
} }
async fn complete_with_model(
&self,
_model_config: &ModelConfig,
system_prompt: &str,
messages: &[Message],
tools: &[Tool],
) -> anyhow::Result<(Message, ProviderUsage), ProviderError> {
self.complete(system_prompt, messages, tools).await
}
fn get_model_config(&self) -> ModelConfig { fn get_model_config(&self) -> ModelConfig {
ModelConfig::new("mock-model").unwrap() ModelConfig::new("mock-model").unwrap()
} }