Integrate pricing with canonical model (#6130)
This commit is contained in:
@@ -6,8 +6,6 @@ use goose_server::auth::check_token;
|
||||
use tower_http::cors::{Any, CorsLayer};
|
||||
use tracing::info;
|
||||
|
||||
use goose::providers::pricing::initialize_pricing_cache;
|
||||
|
||||
// Graceful shutdown signal
|
||||
#[cfg(unix)]
|
||||
async fn shutdown_signal() {
|
||||
@@ -32,13 +30,6 @@ pub async fn run() -> Result<()> {
|
||||
|
||||
let settings = configuration::Settings::new()?;
|
||||
|
||||
if let Err(e) = initialize_pricing_cache().await {
|
||||
tracing::warn!(
|
||||
"Failed to initialize pricing cache: {}. Pricing data may not be available.",
|
||||
e
|
||||
);
|
||||
}
|
||||
|
||||
let secret_key =
|
||||
std::env::var("GOOSE_SERVER__SECRET_KEY").unwrap_or_else(|_| "test".to_string());
|
||||
|
||||
|
||||
@@ -351,6 +351,7 @@ derive_utoipa!(Icon as IconSchema);
|
||||
super::routes::config_management::remove_custom_provider,
|
||||
super::routes::config_management::check_provider,
|
||||
super::routes::config_management::set_config_provider,
|
||||
super::routes::config_management::get_pricing,
|
||||
super::routes::agent::start_agent,
|
||||
super::routes::agent::resume_agent,
|
||||
super::routes::agent::get_tools,
|
||||
@@ -417,6 +418,9 @@ derive_utoipa!(Icon as IconSchema);
|
||||
super::routes::config_management::UpdateCustomProviderRequest,
|
||||
super::routes::config_management::CheckProviderRequest,
|
||||
super::routes::config_management::SetProviderRequest,
|
||||
super::routes::config_management::PricingQuery,
|
||||
super::routes::config_management::PricingResponse,
|
||||
super::routes::config_management::PricingData,
|
||||
super::routes::action_required::ConfirmToolActionRequest,
|
||||
super::routes::reply::ChatRequest,
|
||||
super::routes::session::ImportSessionRequest,
|
||||
|
||||
@@ -13,10 +13,8 @@ use goose::config::{Config, ConfigError};
|
||||
use goose::model::ModelConfig;
|
||||
use goose::providers::auto_detect::detect_provider_from_api_key;
|
||||
use goose::providers::base::{ProviderMetadata, ProviderType};
|
||||
use goose::providers::canonical::maybe_get_canonical_model;
|
||||
use goose::providers::create_with_default_model;
|
||||
use goose::providers::pricing::{
|
||||
get_all_pricing, get_model_pricing, parse_model_id, refresh_pricing,
|
||||
};
|
||||
use goose::providers::providers as get_providers;
|
||||
use goose::{
|
||||
agents::execute_commands, agents::ExtensionConfig, config::permission::PermissionLevel,
|
||||
@@ -470,7 +468,8 @@ pub struct PricingResponse {
|
||||
|
||||
#[derive(Deserialize, ToSchema)]
|
||||
pub struct PricingQuery {
|
||||
pub configured_only: bool,
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
@@ -484,84 +483,28 @@ pub struct PricingQuery {
|
||||
pub async fn get_pricing(
|
||||
Json(query): Json<PricingQuery>,
|
||||
) -> Result<Json<PricingResponse>, StatusCode> {
|
||||
let configured_only = query.configured_only;
|
||||
|
||||
// If refresh requested (configured_only = false), refresh the cache
|
||||
if !configured_only {
|
||||
if let Err(e) = refresh_pricing().await {
|
||||
tracing::error!("Failed to refresh pricing data: {}", e);
|
||||
}
|
||||
}
|
||||
let canonical_model =
|
||||
maybe_get_canonical_model(&query.provider, &query.model).ok_or(StatusCode::NOT_FOUND)?;
|
||||
|
||||
let mut pricing_data = Vec::new();
|
||||
|
||||
if !configured_only {
|
||||
// Get ALL pricing data from the cache
|
||||
let all_pricing = get_all_pricing().await;
|
||||
|
||||
for (provider, models) in all_pricing {
|
||||
for (model, pricing) in models {
|
||||
pricing_data.push(PricingData {
|
||||
provider: provider.clone(),
|
||||
model: model.clone(),
|
||||
input_token_cost: pricing.input_cost,
|
||||
output_token_cost: pricing.output_cost,
|
||||
currency: "$".to_string(),
|
||||
context_length: pricing.context_length,
|
||||
});
|
||||
}
|
||||
}
|
||||
} else {
|
||||
for (metadata, provider_type) in get_providers().await {
|
||||
// Skip unconfigured providers if filtering
|
||||
if !check_provider_configured(&metadata, provider_type) {
|
||||
continue;
|
||||
}
|
||||
|
||||
for model_info in &metadata.known_models {
|
||||
// Handle OpenRouter models specially - they store full provider/model names
|
||||
let (lookup_provider, lookup_model) = if metadata.name == "openrouter" {
|
||||
// For OpenRouter, parse the model name to extract real provider/model
|
||||
if let Some((provider, model)) = parse_model_id(&model_info.name) {
|
||||
(provider, model)
|
||||
} else {
|
||||
// Fallback if parsing fails
|
||||
(metadata.name.clone(), model_info.name.clone())
|
||||
}
|
||||
} else {
|
||||
// For other providers, use names as-is
|
||||
(metadata.name.clone(), model_info.name.clone())
|
||||
};
|
||||
|
||||
// Only get pricing from OpenRouter cache
|
||||
if let Some(pricing) = get_model_pricing(&lookup_provider, &lookup_model).await {
|
||||
pricing_data.push(PricingData {
|
||||
provider: metadata.name.clone(),
|
||||
model: model_info.name.clone(),
|
||||
input_token_cost: pricing.input_cost,
|
||||
output_token_cost: pricing.output_cost,
|
||||
currency: "$".to_string(),
|
||||
context_length: pricing.context_length,
|
||||
});
|
||||
}
|
||||
// No fallback to hardcoded prices
|
||||
}
|
||||
}
|
||||
if let (Some(input_cost), Some(output_cost)) = (
|
||||
canonical_model.pricing.prompt,
|
||||
canonical_model.pricing.completion,
|
||||
) {
|
||||
pricing_data.push(PricingData {
|
||||
provider: query.provider.clone(),
|
||||
model: query.model.clone(),
|
||||
input_token_cost: input_cost,
|
||||
output_token_cost: output_cost,
|
||||
currency: "$".to_string(),
|
||||
context_length: Some(canonical_model.context_length as u32),
|
||||
});
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
"Returning pricing for {} models{}",
|
||||
pricing_data.len(),
|
||||
if configured_only {
|
||||
" (configured providers only)"
|
||||
} else {
|
||||
" (all cached models)"
|
||||
}
|
||||
);
|
||||
|
||||
Ok(Json(PricingResponse {
|
||||
pricing: pricing_data,
|
||||
source: "openrouter".to_string(),
|
||||
source: "canonical".to_string(),
|
||||
}))
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user