fix: more graceful handling of missing usage in provider response (#907)

This commit is contained in:
Alice Hau
2025-01-29 20:01:05 -05:00
committed by GitHub
parent c96819bee0
commit 6051021a8d
9 changed files with 61 additions and 12 deletions
+9 -2
View File
@@ -5,7 +5,7 @@ use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::time::Duration;
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage};
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
use super::errors::ProviderError;
use super::formats::openai::{create_request, get_usage, response_to_message};
use super::oauth;
@@ -243,7 +243,14 @@ impl Provider for DatabricksProvider {
// Parse response
let message = response_to_message(response.clone())?;
let usage = get_usage(&response)?;
let usage = match get_usage(&response) {
Ok(usage) => usage,
Err(ProviderError::UsageError(e)) => {
tracing::warn!("Failed to get usage data: {}", e);
Usage::default()
}
Err(e) => return Err(e),
};
let model = get_model(&response);
super::utils::emit_debug_trace(self, &payload, &response, &usage);
+3
View File
@@ -19,6 +19,9 @@ pub enum ProviderError {
#[error("Execution error: {0}")]
ExecutionError(String),
#[error("Usage data error: {0}")]
UsageError(String),
}
impl From<anyhow::Error> for ProviderError {
@@ -1,6 +1,7 @@
use crate::message::{Message, MessageContent};
use crate::model::ModelConfig;
use crate::providers::base::Usage;
use crate::providers::errors::ProviderError;
use anyhow::{anyhow, Result};
use mcp_core::content::Content;
use mcp_core::role::Role;
@@ -201,6 +202,10 @@ pub fn get_usage(data: &Value) -> Result<Usage> {
Ok(Usage::new(input_tokens, output_tokens, total_tokens))
} else {
tracing::warn!(
"Failed to get usage data: {}",
ProviderError::UsageError("No usage data found in response".to_string())
);
// If no usage data, return None for all values
Ok(Usage::new(None, None, None))
}
@@ -1,6 +1,7 @@
use crate::message::{Message, MessageContent};
use crate::model::ModelConfig;
use crate::providers::base::Usage;
use crate::providers::errors::ProviderError;
use crate::providers::utils::{is_valid_function_name, sanitize_function_name};
use anyhow::Result;
use mcp_core::content::Content;
@@ -254,6 +255,10 @@ pub fn get_usage(data: &Value) -> Result<Usage> {
.map(|v| v as i32);
Ok(Usage::new(input_tokens, output_tokens, total_tokens))
} else {
tracing::warn!(
"Failed to get usage data: {}",
ProviderError::UsageError("No usage data found in response".to_string())
);
// If no usage data, return None for all values
Ok(Usage::new(None, None, None))
}
+3 -2
View File
@@ -1,6 +1,7 @@
use crate::message::{Message, MessageContent};
use crate::model::ModelConfig;
use crate::providers::base::Usage;
use crate::providers::errors::ProviderError;
use crate::providers::utils::{
convert_image, is_valid_function_name, sanitize_function_name, ImageFormat,
};
@@ -221,10 +222,10 @@ pub fn response_to_message(response: Value) -> anyhow::Result<Message> {
})
}
pub fn get_usage(data: &Value) -> anyhow::Result<Usage> {
pub fn get_usage(data: &Value) -> Result<Usage, ProviderError> {
let usage = data
.get("usage")
.ok_or_else(|| anyhow!("No usage data in response"))?;
.ok_or_else(|| ProviderError::UsageError("No usage data in response".to_string()))?;
let input_tokens = usage
.get("prompt_tokens")
+9 -2
View File
@@ -1,7 +1,7 @@
use super::errors::ProviderError;
use crate::message::Message;
use crate::model::ModelConfig;
use crate::providers::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage};
use crate::providers::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
use crate::providers::formats::openai::{create_request, get_usage, response_to_message};
use crate::providers::utils::get_model;
use anyhow::Result;
@@ -137,7 +137,14 @@ impl Provider for GroqProvider {
let response = self.post(payload.clone()).await?;
let message = response_to_message(response.clone())?;
let usage = get_usage(&response)?;
let usage = match get_usage(&response) {
Ok(usage) => usage,
Err(ProviderError::UsageError(e)) => {
tracing::warn!("Failed to get usage data: {}", e);
Usage::default()
}
Err(e) => return Err(e),
};
let model = get_model(&response);
super::utils::emit_debug_trace(self, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage)))
+9 -2
View File
@@ -1,4 +1,4 @@
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage};
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
use super::errors::ProviderError;
use super::utils::{get_model, handle_response_openai_compat};
use crate::message::Message;
@@ -104,7 +104,14 @@ impl Provider for OllamaProvider {
// Parse response
let message = response_to_message(response.clone())?;
let usage = get_usage(&response)?;
let usage = match get_usage(&response) {
Ok(usage) => usage,
Err(ProviderError::UsageError(e)) => {
tracing::warn!("Failed to get usage data: {}", e);
Usage::default()
}
Err(e) => return Err(e),
};
let model = get_model(&response);
super::utils::emit_debug_trace(self, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage)))
+9 -2
View File
@@ -4,7 +4,7 @@ use reqwest::Client;
use serde_json::Value;
use std::time::Duration;
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage};
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
use super::errors::ProviderError;
use super::formats::openai::{create_request, get_usage, response_to_message};
use super::utils::{emit_debug_trace, get_model, handle_response_openai_compat, ImageFormat};
@@ -115,7 +115,14 @@ impl Provider for OpenAiProvider {
// Parse response
let message = response_to_message(response.clone())?;
let usage = get_usage(&response)?;
let usage = match get_usage(&response) {
Ok(usage) => usage,
Err(ProviderError::UsageError(e)) => {
tracing::warn!("Failed to get usage data: {}", e);
Usage::default()
}
Err(e) => return Err(e),
};
let model = get_model(&response);
emit_debug_trace(self, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage)))
+9 -2
View File
@@ -4,7 +4,7 @@ use reqwest::Client;
use serde_json::{json, Value};
use std::time::Duration;
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage};
use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
use super::errors::ProviderError;
use super::utils::{emit_debug_trace, get_model, handle_response_openai_compat};
use crate::message::Message;
@@ -203,7 +203,14 @@ impl Provider for OpenRouterProvider {
// Parse response
let message = response_to_message(response.clone())?;
let usage = get_usage(&response)?;
let usage = match get_usage(&response) {
Ok(usage) => usage,
Err(ProviderError::UsageError(e)) => {
tracing::warn!("Failed to get usage data: {}", e);
Usage::default()
}
Err(e) => return Err(e),
};
let model = get_model(&response);
emit_debug_trace(self, &payload, &response, &usage);
Ok((message, ProviderUsage::new(model, usage)))