diff --git a/crates/goose-provider-types/src/conversation/token_usage.rs b/crates/goose-provider-types/src/conversation/token_usage.rs index d62a4d3701d5..56d1430d0e83 100644 --- a/crates/goose-provider-types/src/conversation/token_usage.rs +++ b/crates/goose-provider-types/src/conversation/token_usage.rs @@ -5,6 +5,8 @@ use utoipa::ToSchema; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ProviderUsage { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub provider: Option, pub model: String, pub usage: Usage, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -32,12 +34,23 @@ pub struct DraftStats { impl ProviderUsage { pub fn new(model: String, usage: Usage) -> Self { Self { + provider: None, model, usage, stats: None, } } + pub fn with_provider(mut self, provider: impl Into) -> Self { + self.provider = Some(provider.into()); + self + } + + pub fn with_model(mut self, model: impl Into) -> Self { + self.model = model.into(); + self + } + pub fn with_stats(mut self, stats: ProviderStats) -> Self { self.stats = Some(stats); self @@ -47,6 +60,7 @@ impl ProviderUsage { /// Uses the model from this ProviderUsage pub fn combine_with(&self, other: &ProviderUsage) -> ProviderUsage { ProviderUsage { + provider: self.provider.clone().or_else(|| other.provider.clone()), model: self.model.clone(), usage: self.usage + other.usage, stats: self.stats.clone().or_else(|| other.stats.clone()), diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 5365f1b9bde9..ca742e13b344 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -35,7 +35,8 @@ use crate::providers::inventory::{ }; use crate::scheduler_trait::SchedulerTrait; use crate::session::{ - EnabledExtensionsState, ExtensionData, ExtensionState, Session, SessionManager, SessionType, + EnabledExtensionsState, ExtensionData, ExtensionState, ProviderUsageSnapshotState, Session, + SessionManager, SessionType, }; use crate::source_roots::SourceRoot; use crate::utils::sanitize_unicode_tags; @@ -836,6 +837,13 @@ pub(super) struct UsageUpdates { pub(super) standard: UsageUpdate, } +fn provider_usage_meta(session: &Session) -> Option { + let provider_usage = ProviderUsageSnapshotState::from_extension_data(&session.extension_data)?; + serde_json::json!({ "goose": { "providerUsage": provider_usage } }) + .as_object() + .cloned() +} + pub(super) fn build_usage_updates(session: &Session) -> Option { let used = session.usage.total_tokens.unwrap_or(0).max(0) as u64; let ctx_limit = session.model_config.as_ref()?.context_limit() as u64; @@ -859,6 +867,9 @@ pub(super) fn build_usage_updates(session: &Session) -> Option { if let Some(amount) = session.accumulated_cost { standard = standard.cost(Cost::new(amount, "USD")); } + if let Some(meta) = provider_usage_meta(session) { + standard = standard.meta(meta); + } standard }, }) @@ -3888,6 +3899,37 @@ print(\"hello, world\") assert_eq!(updates.standard.size, 258_000); } + #[test] + fn test_build_usage_update_includes_provider_usage_snapshot_meta() { + let mut session = make_session_with_usage( + TokenUsage::new(Some(80), Some(40), Some(120)), + TokenUsage::default(), + ); + session.model_config = Some( + goose_providers::model::ModelConfig::new("test-model") + .with_context_limit(Some(258_000)), + ); + let mut provider_usage = ProviderUsageSnapshotState::default(); + provider_usage.add_usage( + "test-provider", + "test-model", + "2026-07-02T10:00:00Z".to_string(), + TokenUsage::new(Some(10), Some(5), Some(15)), + ); + provider_usage + .to_extension_data(&mut session.extension_data) + .unwrap(); + + let updates = build_usage_updates(&session).expect("usage updates should be present"); + let standard = serde_json::to_value(updates.standard).unwrap(); + let entry = &standard["_meta"]["goose"]["providerUsage"]["entries"][0]; + + assert_eq!(entry["providerId"], "test-provider"); + assert_eq!(entry["modelId"], "test-model"); + assert_eq!(entry["inputTokens"], 10); + assert_eq!(entry["outputTokens"], 5); + } + #[test] fn test_build_usage_update_requires_model_config() { let session = make_session_with_usage( diff --git a/crates/goose/src/acp/server/providers.rs b/crates/goose/src/acp/server/providers.rs index d1a496137f0c..d8e14300cff4 100644 --- a/crates/goose/src/acp/server/providers.rs +++ b/crates/goose/src/acp/server/providers.rs @@ -72,6 +72,47 @@ fn inventory_entry_to_dto(entry: ProviderInventoryEntry) -> ProviderInventoryEnt } } +fn provider_price_per_million(price: Option) -> Option { + price.map(|price| price * 1_000_000.0) +} + +struct ModelPricing { + input: Option, + output: Option, + cache_read: Option, + cache_write: Option, + currency: String, + using_canonical: bool, +} + +fn resolve_model_pricing( + provider: Option<&crate::providers::base::ModelInfo>, + canonical: Option<&crate::providers::canonical::CanonicalModel>, +) -> ModelPricing { + if let Some(provider) = provider + .filter(|model| model.input_token_cost.is_some() || model.output_token_cost.is_some()) + { + return ModelPricing { + input: provider_price_per_million(provider.input_token_cost), + output: provider_price_per_million(provider.output_token_cost), + cache_read: None, + cache_write: None, + currency: provider.currency.clone().unwrap_or_else(|| "$".to_string()), + using_canonical: false, + }; + } + + ModelPricing { + input: canonical.and_then(|model| model.cost.input), + output: canonical.and_then(|model| model.cost.output), + cache_read: canonical.and_then(|model| model.cost.cache_read), + cache_write: canonical.and_then(|model| model.cost.cache_write), + currency: "$".to_string(), + using_canonical: canonical + .is_some_and(|model| model.cost.input.is_some() || model.cost.output.is_some()), + } +} + fn provider_config_key_to_dto(key: crate::providers::base::ConfigKey) -> ProviderConfigKey { ProviderConfigKey { name: key.name, @@ -1022,24 +1063,128 @@ impl GooseAcpAgent { ) -> Result { use goose_providers::model::ModelConfig; - let model_info = - crate::providers::canonical::maybe_get_canonical_model(&req.provider, &req.model).map( - |canonical_model| CanonicalModelInfoDto { - provider: req.provider.clone(), - model: req.model.clone(), - context_limit: canonical_model.limit.context, - max_output_tokens: canonical_model.limit.output, - reasoning: canonical_model - .reasoning - .unwrap_or_else(|| ModelConfig::new(&req.model).is_reasoning_model()), - input_token_cost: canonical_model.cost.input, - output_token_cost: canonical_model.cost.output, - cache_read_token_cost: canonical_model.cost.cache_read, - cache_write_token_cost: canonical_model.cost.cache_write, - currency: "$".to_string(), - }, + // TODO: Rename this legacy method to provider model info. It now returns + // provider-declared metadata first and falls back to canonical metadata. + let provider_model = crate::providers::get_from_registry(&req.provider) + .await + .ok() + .and_then(|entry| entry.model_info(&req.model)); + + let canonical_model = + crate::providers::canonical::maybe_get_canonical_model(&req.provider, &req.model); + + if provider_model.is_none() && canonical_model.is_none() { + return Ok(CanonicalModelInfoResponse { model_info: None }); + } + + let provider = provider_model.as_ref(); + let canonical = canonical_model.as_ref(); + let pricing = resolve_model_pricing(provider, canonical); + if pricing.using_canonical { + warn!( + provider = %req.provider, + model = %req.model, + "Using canonical pricing because provider model pricing is not configured" ); + } + + Ok(CanonicalModelInfoResponse { + model_info: Some(CanonicalModelInfoDto { + provider: req.provider.clone(), + model: req.model.clone(), + context_limit: provider + .and_then(|model| (model.context_limit > 0).then_some(model.context_limit)) + .or_else(|| canonical.map(|model| model.limit.context)) + .unwrap_or_default(), + max_output_tokens: canonical.and_then(|model| model.limit.output), + reasoning: provider.map(|model| model.reasoning).unwrap_or(false) + || canonical + .and_then(|model| model.reasoning) + .unwrap_or_else(|| ModelConfig::new(&req.model).is_reasoning_model()), + input_token_cost: pricing.input, + output_token_cost: pricing.output, + cache_read_token_cost: pricing.cache_read, + cache_write_token_cost: pricing.cache_write, + currency: pricing.currency, + }), + }) + } +} + +#[cfg(test)] +mod tests { + use super::{provider_price_per_million, resolve_model_pricing}; + use crate::providers::base::ModelInfo; + use crate::providers::canonical::{CanonicalModel, Limit, Modalities, Pricing}; + + #[test] + fn provider_price_per_million_normalizes_per_token_price() { + assert_eq!(provider_price_per_million(Some(0.000003)), Some(3.0)); + assert_eq!(provider_price_per_million(None), None); + } + + fn provider_model(input: Option, output: Option) -> ModelInfo { + ModelInfo { + name: "provider-model".to_string(), + input_token_cost: input, + output_token_cost: output, + currency: Some("RUB".to_string()), + ..ModelInfo::new("provider-model", 128_000) + } + } + + fn canonical_model() -> CanonicalModel { + CanonicalModel { + id: "provider/model".to_string(), + name: "Canonical Model".to_string(), + family: None, + attachment: None, + reasoning: None, + thinking_mode: None, + tool_call: false, + temperature: None, + knowledge: None, + release_date: None, + last_updated: None, + modalities: Modalities::default(), + open_weights: None, + cost: Pricing { + input: Some(5.0), + output: Some(15.0), + cache_read: Some(1.0), + cache_write: Some(2.0), + }, + limit: Limit::default(), + } + } + + #[test] + fn provider_pricing_does_not_mix_canonical_fallback_fields() { + let provider = provider_model(Some(0.0005), None); + let canonical = canonical_model(); + + let pricing = resolve_model_pricing(Some(&provider), Some(&canonical)); + + assert_eq!(pricing.input, Some(500.0)); + assert_eq!(pricing.output, None); + assert_eq!(pricing.cache_read, None); + assert_eq!(pricing.cache_write, None); + assert_eq!(pricing.currency, "RUB"); + assert!(!pricing.using_canonical); + } + + #[test] + fn canonical_pricing_is_used_only_when_provider_pricing_is_absent() { + let provider = provider_model(None, None); + let canonical = canonical_model(); + + let pricing = resolve_model_pricing(Some(&provider), Some(&canonical)); - Ok(CanonicalModelInfoResponse { model_info }) + assert_eq!(pricing.input, Some(5.0)); + assert_eq!(pricing.output, Some(15.0)); + assert_eq!(pricing.cache_read, Some(1.0)); + assert_eq!(pricing.cache_write, Some(2.0)); + assert_eq!(pricing.currency, "$"); + assert!(pricing.using_canonical); } } diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index d4a8f592e906..6f586414ccca 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1756,8 +1756,10 @@ impl Agent { ); let compact_model_config = self.model_config_for_session(&session_config.id).await?; + let compact_provider = self.provider().await?; + let compact_provider_name = compact_provider.get_name().to_string(); match compact_messages( - self.provider().await?.as_ref(), + compact_provider.as_ref(), &compact_model_config, &session_config.id, &conversation_to_compact, @@ -1766,6 +1768,9 @@ impl Agent { .await { Ok((compacted_conversation, summarization_usage)) => { + let summarization_usage = summarization_usage + .with_provider(compact_provider_name) + .with_model(compact_model_config.model_name.clone()); session_manager.replace_conversation(&session_config.id, &compacted_conversation).await?; self.update_session_metrics(&session_config.id, session_config.schedule_id.clone(), &summarization_usage, true).await?; @@ -1834,7 +1839,7 @@ impl Agent { .ok() .and_then(|model_info| model_info.resolved_model) .map(|resolved_model| InferenceMetadata { - provider: provider_name, + provider: provider_name.clone(), requested_model, resolved_model: Some(resolved_model), }); @@ -2033,9 +2038,12 @@ impl Agent { Ok((response, usage)) => { compaction_attempts = 0; - if let Some(ref usage) = usage { - self.update_session_metrics(&session_config.id, session_config.schedule_id.clone(), usage, false).await?; - yield AgentEvent::Usage(usage.clone()); + if let Some(usage) = usage { + let usage = usage + .with_provider(provider_name.clone()) + .with_model(model_config.model_name.clone()); + self.update_session_metrics(&session_config.id, session_config.schedule_id.clone(), &usage, false).await?; + yield AgentEvent::Usage(usage); } if let Some(response) = response { @@ -2421,6 +2429,9 @@ impl Agent { .await { Ok((compacted_conversation, usage)) => { + let usage = usage + .with_provider(provider_name.clone()) + .with_model(model_config.model_name.clone()); session_manager.replace_conversation(&session_config.id, &compacted_conversation).await?; self.update_session_metrics(&session_config.id, session_config.schedule_id.clone(), &usage, true).await?; conversation = compacted_conversation; diff --git a/crates/goose/src/agents/execute_commands.rs b/crates/goose/src/agents/execute_commands.rs index 84e9152b0382..e071125ebeb3 100644 --- a/crates/goose/src/agents/execute_commands.rs +++ b/crates/goose/src/agents/execute_commands.rs @@ -157,8 +157,10 @@ impl Agent { .ok_or_else(|| anyhow!("Session has no conversation"))?; let model_config = self.model_config_for_session(session_id).await?; + let provider = self.provider().await?; + let provider_name = provider.get_name().to_string(); let (compacted_conversation, usage) = compact_messages( - self.provider().await?.as_ref(), + provider.as_ref(), &model_config, session_id, &conversation, @@ -170,6 +172,9 @@ impl Agent { .replace_conversation(session_id, &compacted_conversation) .await?; + let usage = usage + .with_provider(provider_name) + .with_model(model_config.model_name); self.update_session_metrics(session_id, session.schedule_id, &usage, true) .await?; @@ -510,6 +515,7 @@ fn user_only_assistant_text(text: impl Into) -> Message { mod tests { use super::*; use crate::conversation::message::MessageContent; + use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; #[test] fn parse_slash_command_splits_on_literal_space() { @@ -561,6 +567,21 @@ mod tests { )); } + #[test] + fn compact_usage_is_attributed_to_request_identity() { + let usage = ProviderUsage::new( + "resolved-model-from-provider".to_string(), + Usage::new(Some(10), Some(2), Some(12)), + ); + + let usage = usage + .with_provider("custom-provider") + .with_model("configured-model"); + + assert_eq!(usage.provider.as_deref(), Some("custom-provider")); + assert_eq!(usage.model, "configured-model"); + } + #[test] fn status_is_registered_as_a_builtin_command() { assert!(list_commands() diff --git a/crates/goose/src/agents/reply_parts.rs b/crates/goose/src/agents/reply_parts.rs index c1c366888c02..02a79cf6ba56 100644 --- a/crates/goose/src/agents/reply_parts.rs +++ b/crates/goose/src/agents/reply_parts.rs @@ -21,6 +21,7 @@ use crate::providers::toolshim::{ augment_message_with_selected_tool_interpreter, convert_tool_messages_to_text, modify_system_prompt_for_tool_json, sanitize_residual_markers, }; +use crate::session::ProviderUsageSnapshotState; use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; use goose_providers::model::ModelConfig; use rmcp::model::Tool; @@ -529,9 +530,10 @@ impl Agent { let accumulated_usage = session.accumulated_usage + usage.usage; - let accumulated_cost = session - .provider_name + let accumulated_cost = usage + .provider .as_deref() + .or(session.provider_name.as_deref()) .and_then(|pn| self.accumulate_cost(session.accumulated_cost, usage, pn)) .or(session.accumulated_cost); @@ -543,14 +545,28 @@ impl Agent { usage.usage }; - manager + let mut update = manager .update(session_id) .schedule_id(schedule_id) .usage(current_usage) .accumulated_usage(accumulated_usage) - .accumulated_cost(accumulated_cost) - .apply() - .await?; + .accumulated_cost(accumulated_cost); + + if let Some(provider_name) = usage + .provider + .as_deref() + .or(session.provider_name.as_deref()) + { + update = update.extension_data(ProviderUsageSnapshotState::extension_data_with_usage( + &session.extension_data, + provider_name, + &usage.model, + chrono::Utc::now().to_rfc3339(), + usage.usage, + )?); + } + + update.apply().await?; Ok(()) } @@ -619,6 +635,7 @@ mod tests { use crate::conversation::message::Message; use crate::providers::base::Provider; use crate::session::session_manager::SessionType; + use crate::session::ExtensionState; use async_trait::async_trait; use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; use goose_providers::model::ModelConfig; @@ -712,6 +729,57 @@ mod tests { Ok(()) } + #[tokio::test] + async fn update_session_metrics_prefers_usage_provider_for_snapshot() -> anyhow::Result<()> { + let agent = crate::agents::Agent::new(); + + let session = agent + .config + .session_manager + .create_session( + std::env::current_dir().unwrap(), + "test-provider-usage-attribution".to_string(), + SessionType::Hidden, + GooseMode::default(), + ) + .await?; + + agent + .config + .session_manager + .update(&session.id) + .provider_name("session-provider") + .apply() + .await?; + + let usage = ProviderUsage::new( + "resolved-model-from-provider".to_string(), + Usage::new(Some(10), Some(2), Some(12)), + ) + .with_provider("openai") + .with_model("gpt-4o"); + agent + .update_session_metrics(&session.id, None, &usage, false) + .await?; + + let session = agent + .config + .session_manager + .get_session(&session.id, false) + .await?; + let snapshot = crate::session::ProviderUsageSnapshotState::from_extension_data( + &session.extension_data, + ) + .expect("provider usage snapshot"); + + assert_eq!(snapshot.entries.len(), 1); + assert_eq!(snapshot.entries[0].provider_id, "openai"); + assert_eq!(snapshot.entries[0].model_id, "gpt-4o"); + assert!(session.accumulated_cost.is_some()); + + Ok(()) + } + #[tokio::test] async fn test_stream_error_propagation() { use futures::StreamExt; diff --git a/crates/goose/src/providers/provider_registry.rs b/crates/goose/src/providers/provider_registry.rs index 8447a758b390..93f07c02e6dd 100644 --- a/crates/goose/src/providers/provider_registry.rs +++ b/crates/goose/src/providers/provider_registry.rs @@ -46,6 +46,14 @@ impl ProviderEntry { self.supports_inventory_refresh } + pub fn model_info(&self, model_name: &str) -> Option { + self.metadata + .known_models + .iter() + .find(|model| model.name.eq_ignore_ascii_case(model_name)) + .cloned() + } + pub fn inventory_identity(&self) -> Result { (self.inventory_identity)() } @@ -414,4 +422,42 @@ mod tests { assert!(!entry.inventory_configured()); } + + #[test] + fn provider_entry_model_info_returns_declared_pricing() { + let mut config = test_config(); + config.models = vec![ModelInfo { + name: "gpt-5.5".to_string(), + input_token_cost: Some(500.0), + output_token_cost: Some(3000.0), + currency: Some("RUB".to_string()), + reasoning: true, + ..ModelInfo::new("gpt-5.5", 128_000) + }]; + + let mut registry = ProviderRegistry::new(None); + registry.register_with_name::( + &config, + ProviderType::Declarative, + false, + |_| unreachable!("constructor is not used by this test"), + || Ok(InventoryIdentityInput::new("custom_hf", "openai")), + ); + + let model_info = registry + .entries + .get("custom_hf") + .and_then(|entry| entry.model_info("GPT-5.5")) + .expect("declared model info should be present"); + + assert_eq!( + ( + model_info.input_token_cost, + model_info.output_token_cost, + model_info.currency.as_deref() + ), + (Some(500.0), Some(3000.0), Some("RUB")) + ); + assert!(model_info.reasoning); + } } diff --git a/crates/goose/src/session/extension_data.rs b/crates/goose/src/session/extension_data.rs index b286a3e220b7..5bbfff1eb0e5 100644 --- a/crates/goose/src/session/extension_data.rs +++ b/crates/goose/src/session/extension_data.rs @@ -6,6 +6,7 @@ use crate::config::extensions::is_extension_available; use crate::config::ExtensionConfig; use crate::session::SessionManager; use anyhow::Result; +use goose_providers::conversation::token_usage::Usage; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::HashMap; @@ -143,6 +144,81 @@ impl EnabledExtensionsState { } } +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq)] +#[serde(rename_all = "camelCase")] +pub struct ProviderUsageSnapshotState { + pub entries: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +#[serde(rename_all = "camelCase")] +pub struct ProviderUsageSnapshotEntry { + pub provider_id: String, + pub model_id: String, + pub last_used_at: String, + pub input_tokens: i32, + pub output_tokens: i32, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cache_read_input_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cache_write_input_tokens: Option, +} + +impl ExtensionState for ProviderUsageSnapshotState { + const EXTENSION_NAME: &'static str = "provider_usage"; + const VERSION: &'static str = "v0"; +} + +impl ProviderUsageSnapshotState { + pub fn add_usage(&mut self, provider_id: &str, model_id: &str, used_at: String, usage: Usage) { + let entry = self + .entries + .iter_mut() + .find(|entry| entry.provider_id == provider_id && entry.model_id == model_id); + let entry = match entry { + Some(entry) => entry, + None => { + self.entries.push(ProviderUsageSnapshotEntry { + provider_id: provider_id.to_string(), + model_id: model_id.to_string(), + last_used_at: used_at.clone(), + input_tokens: 0, + output_tokens: 0, + cache_read_input_tokens: None, + cache_write_input_tokens: None, + }); + self.entries.last_mut().expect("entry was just pushed") + } + }; + + entry.last_used_at = used_at; + let usage = Usage::new(Some(entry.input_tokens), Some(entry.output_tokens), None) + .with_cache_tokens( + entry.cache_read_input_tokens, + entry.cache_write_input_tokens, + ) + + usage; + entry.input_tokens = usage.input_tokens.unwrap_or(0); + entry.output_tokens = usage.output_tokens.unwrap_or(0); + entry.cache_read_input_tokens = usage.cache_read_input_tokens; + entry.cache_write_input_tokens = usage.cache_write_input_tokens; + } + + pub fn extension_data_with_usage( + extension_data: &ExtensionData, + provider_id: &str, + model_id: &str, + used_at: String, + usage: Usage, + ) -> Result { + let mut extension_data = extension_data.clone(); + let mut state = Self::from_extension_data(&extension_data).unwrap_or_default(); + state.add_usage(provider_id, model_id, used_at, usage); + state.to_extension_data(&mut extension_data)?; + Ok(extension_data) + } +} + #[cfg(test)] mod tests { use super::*; @@ -175,6 +251,37 @@ mod tests { data } + #[test] + fn provider_usage_snapshot_accumulates_by_provider_and_model() { + let mut state = ProviderUsageSnapshotState::default(); + + state.add_usage( + "test-provider", + "test-model", + "2026-07-02T10:00:00Z".to_string(), + Usage::new(Some(10), Some(2), Some(12)).with_cache_tokens(Some(3), None), + ); + state.add_usage( + "test-provider", + "test-model", + "2026-07-02T10:01:00Z".to_string(), + Usage::new(Some(5), Some(4), Some(9)).with_cache_tokens(None, Some(1)), + ); + + assert_eq!( + state.entries, + vec![ProviderUsageSnapshotEntry { + provider_id: "test-provider".to_string(), + model_id: "test-model".to_string(), + last_used_at: "2026-07-02T10:01:00Z".to_string(), + input_tokens: 15, + output_tokens: 6, + cache_read_input_tokens: Some(3), + cache_write_input_tokens: Some(1), + }] + ); + } + #[test_case( Some(extension_data_with(vec![test_extension()])), Some(vec![test_extension()]) diff --git a/crates/goose/src/session/mod.rs b/crates/goose/src/session/mod.rs index c7c5230e8013..3d607778f5e0 100644 --- a/crates/goose/src/session/mod.rs +++ b/crates/goose/src/session/mod.rs @@ -15,7 +15,9 @@ pub use diagnostics::{ DiagnosticsExtensions, DiagnosticsLevel, DiagnosticsLogs, DiagnosticsPrompt, DiagnosticsReport, DiagnosticsScheduledRecipe, DiagnosticsTextFile, SystemInfo, }; -pub use extension_data::{EnabledExtensionsState, ExtensionData, ExtensionState, TodoState}; +pub use extension_data::{ + EnabledExtensionsState, ExtensionData, ExtensionState, ProviderUsageSnapshotState, TodoState, +}; pub use session_manager::{ Session, SessionInsights, SessionManager, SessionNameUpdate, SessionType, SessionUpdateBuilder, }; diff --git a/ui/desktop/src/acp/__tests__/sessionNotificationAdapter.test.ts b/ui/desktop/src/acp/__tests__/sessionNotificationAdapter.test.ts index 22e054677652..e41aa81a3708 100644 --- a/ui/desktop/src/acp/__tests__/sessionNotificationAdapter.test.ts +++ b/ui/desktop/src/acp/__tests__/sessionNotificationAdapter.test.ts @@ -521,6 +521,50 @@ describe('createAcpSessionNotificationAdapter', () => { }); }); }); + + it('maps provider usage snapshot metadata from usage updates', () => { + const adapter = createAcpSessionNotificationAdapter(); + + expect( + adapter.apply( + acpUpdate({ + sessionUpdate: 'usage_update', + used: 42, + size: 200, + _meta: { + goose: { + providerUsage: { + entries: [ + { + providerId: 'test-provider', + modelId: 'test-model', + lastUsedAt: '2026-07-02T10:00:00Z', + inputTokens: 10, + outputTokens: 5, + }, + ], + }, + }, + }, + }) + ) + ).toEqual([ + { + type: 'tokenState', + tokenState: { + providerUsage: [ + { + providerId: 'test-provider', + modelId: 'test-model', + lastUsedAt: '2026-07-02T10:00:00Z', + inputTokens: 10, + outputTokens: 5, + }, + ], + }, + }, + ]); + }); }); describe('applyGoose', () => { diff --git a/ui/desktop/src/acp/adapter/shared.ts b/ui/desktop/src/acp/adapter/shared.ts index 073962701b85..2964fda66040 100644 --- a/ui/desktop/src/acp/adapter/shared.ts +++ b/ui/desktop/src/acp/adapter/shared.ts @@ -1,5 +1,5 @@ import type { ToolCall, ToolCallUpdate } from '@agentclientprotocol/sdk'; -import type { TokenState } from '../../types/chat'; +import type { ProviderUsageEntry, TokenState } from '../../types/chat'; import type { Message, NotificationEvent } from '../../types/message'; export type AcpChatStateChange = @@ -78,6 +78,26 @@ export function getGooseActiveRunId(update: { _meta?: unknown }): string | null : undefined; } +export function getGooseProviderUsage(update: { + _meta?: unknown; +}): ProviderUsageEntry[] | undefined { + if (!isRecord(update._meta)) { + return undefined; + } + + const goose = update._meta.goose; + if (!isRecord(goose)) { + return undefined; + } + + const providerUsage = goose.providerUsage; + if (!isRecord(providerUsage) || !Array.isArray(providerUsage.entries)) { + return undefined; + } + + return providerUsage.entries as ProviderUsageEntry[]; +} + export function rawInputToArguments(rawInput: unknown): Record { return isRecord(rawInput) ? rawInput : {}; } diff --git a/ui/desktop/src/acp/providers.ts b/ui/desktop/src/acp/providers.ts index e9e967f40564..785e8ff36a8b 100644 --- a/ui/desktop/src/acp/providers.ts +++ b/ui/desktop/src/acp/providers.ts @@ -155,6 +155,8 @@ export async function acpGetCanonicalModelInfo( model: string ): Promise { const client = await getAcpClient(); + // TODO: Rename this legacy API to acpGetModelInfo once the ACP method name changes. + // The server now resolves provider model metadata first and falls back to canonical data. const { modelInfo } = await client.goose.providersCanonicalModelInfo_unstable({ provider, model, diff --git a/ui/desktop/src/acp/sessionNotificationAdapter.ts b/ui/desktop/src/acp/sessionNotificationAdapter.ts index d881d5923b3c..e629429ad6d7 100644 --- a/ui/desktop/src/acp/sessionNotificationAdapter.ts +++ b/ui/desktop/src/acp/sessionNotificationAdapter.ts @@ -14,6 +14,7 @@ import { type AdapterState, cloneMessage, getGooseActiveRunId, + getGooseProviderUsage, } from './adapter/shared'; import { applyToolCall, applyToolCallUpdate } from './adapter/tools'; import type { AcpElicitationRequest } from './elicitationRequests'; @@ -91,8 +92,10 @@ function applyAcpSessionNotification( }, ]; } - case 'usage_update': - return []; + case 'usage_update': { + const providerUsage = getGooseProviderUsage(update); + return providerUsage ? [{ type: 'tokenState', tokenState: { providerUsage } }] : []; + } default: return []; } diff --git a/ui/desktop/src/components/BaseChat.tsx b/ui/desktop/src/components/BaseChat.tsx index e6e3bf8f383c..49b1e1a2a651 100644 --- a/ui/desktop/src/components/BaseChat.tsx +++ b/ui/desktop/src/components/BaseChat.tsx @@ -511,6 +511,7 @@ export default function BaseChat({ undefined } accumulatedCost={tokenState?.accumulatedCost ?? session?.accumulated_cost ?? undefined} + providerUsage={tokenState?.providerUsage} droppedFiles={droppedFiles} onFilesProcessed={() => setDroppedFiles([])} // Clear dropped files after processing messages={messages} diff --git a/ui/desktop/src/components/ChatInput.tsx b/ui/desktop/src/components/ChatInput.tsx index 73e0881d10f1..cf49ccdac9f6 100644 --- a/ui/desktop/src/components/ChatInput.tsx +++ b/ui/desktop/src/components/ChatInput.tsx @@ -28,6 +28,7 @@ import { MessageQueue, QueuedMessage } from './MessageQueue'; import { detectInterruption } from '../utils/interruptionDetector'; import { DiagnosticsModal } from './ui/Diagnostics'; import type { Message } from '../types/message'; +import type { ProviderUsageEntry } from '../types/chat'; import { getInitialWorkingDir } from '../utils/workingDir'; import { getPredefinedModelsFromEnv } from './settings/models/predefinedModelsUtils'; import { trackFileAttached, trackVoiceDictation, trackDiagnosticsOpened } from '../utils/analytics'; @@ -173,6 +174,7 @@ interface ChatInputProps { accumulatedInputTokens?: number; accumulatedOutputTokens?: number; accumulatedCost?: number | null; + providerUsage?: ProviderUsageEntry[]; messages?: Message[]; disableAnimation?: boolean; recipe?: Recipe | null; @@ -208,6 +210,7 @@ export default function ChatInput({ accumulatedInputTokens, accumulatedOutputTokens, accumulatedCost, + providerUsage, messages = [], disableAnimation = false, recipe: _recipe, @@ -1693,6 +1696,7 @@ export default function ChatInput({ inputTokens={accumulatedInputTokens} outputTokens={accumulatedOutputTokens} accumulatedCost={accumulatedCost} + providerUsage={providerUsage} model={effectiveModel} provider={effectiveProvider} /> diff --git a/ui/desktop/src/components/bottom_menu/CostTracker.test.tsx b/ui/desktop/src/components/bottom_menu/CostTracker.test.tsx new file mode 100644 index 000000000000..9fdd26872421 --- /dev/null +++ b/ui/desktop/src/components/bottom_menu/CostTracker.test.tsx @@ -0,0 +1,364 @@ +import { render, screen, type RenderOptions } from '@testing-library/react'; +import type React from 'react'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { IntlTestWrapper } from '../../i18n/test-utils'; +import { fetchCanonicalModelInfo } from '../../utils/canonical'; +import { CostTracker } from './CostTracker'; + +vi.mock('../../utils/canonical', () => ({ + fetchCanonicalModelInfo: vi.fn(), +})); + +vi.mock('../ui/Tooltip', () => ({ + Tooltip: ({ children }: { children: React.ReactNode }) => <>{children}, + TooltipTrigger: ({ children }: { children: React.ReactNode }) => <>{children}, + TooltipContent: ({ children }: { children: React.ReactNode }) =>
{children}
, +})); + +const renderWithIntl = (ui: React.ReactElement, options?: RenderOptions) => + render(ui, { wrapper: IntlTestWrapper, ...options }); + +const fetchCanonicalModelInfoMock = vi.mocked(fetchCanonicalModelInfo); + +describe('CostTracker', () => { + beforeEach(() => { + vi.clearAllMocks(); + fetchCanonicalModelInfoMock.mockResolvedValue({ + provider: 'test-provider', + model: 'gpt-5.5-2026-04-23', + contextLimit: 128000, + maxOutputTokens: null, + reasoning: true, + inputTokenCost: 500, + outputTokenCost: 3000, + cacheReadTokenCost: null, + cacheWriteTokenCost: null, + currency: 'RUB', + }); + }); + + it('shows input and output cost breakdown for provider usage snapshots', async () => { + renderWithIntl( + + ); + + expect(await screen.findByText('₽8.00')).toBeInTheDocument(); + expect( + screen.getByText(/Input: 10,000 tokens \(₽5.00\) \| Output: 1,000 tokens \(₽3.00\)/) + ).toBeInTheDocument(); + expect(screen.getByText(/Total session cost: ₽8.00/)).toBeInTheDocument(); + expect(fetchCanonicalModelInfoMock).toHaveBeenCalledWith( + 'test-provider', + 'gpt-5.5-2026-04-23' + ); + }); + + it('uses the same cost calculation for current session token usage', async () => { + renderWithIntl( + + ); + + expect(await screen.findByText('₽8.00')).toBeInTheDocument(); + expect( + screen.getByText(/Input: 10,000 tokens \(₽5.00\) \| Output: 1,000 tokens \(₽3.00\)/) + ).toBeInTheDocument(); + expect(fetchCanonicalModelInfoMock).toHaveBeenCalledWith('test-provider', 'gpt-5.5'); + }); + + it('prefers accumulated cost when provider usage snapshots are unavailable', async () => { + renderWithIntl( + + ); + + expect(await screen.findByText('0.03')).toBeInTheDocument(); + expect(screen.getByText(/Total session cost: 0.03/)).toBeInTheDocument(); + expect(fetchCanonicalModelInfoMock).not.toHaveBeenCalled(); + }); + + it('prefers accumulated cost when provider usage snapshots are partial', async () => { + renderWithIntl( + + ); + + expect(await screen.findByText('12.34')).toBeInTheDocument(); + expect(screen.getByText(/Total session cost: 12.34/)).toBeInTheDocument(); + expect(fetchCanonicalModelInfoMock).not.toHaveBeenCalled(); + }); + + it('does not price partial provider usage snapshots without accumulated cost', async () => { + renderWithIntl( + + ); + + expect(await screen.findByText('0.0000')).toBeInTheDocument(); + expect(screen.getByText(/Pricing data unavailable for gpt-5.5/)).toBeInTheDocument(); + expect(fetchCanonicalModelInfoMock).not.toHaveBeenCalled(); + }); + + it('prices cached input tokens with cache rates when available', async () => { + fetchCanonicalModelInfoMock.mockResolvedValue({ + provider: 'test-provider', + model: 'gpt-5.5-2026-04-23', + contextLimit: 128000, + maxOutputTokens: null, + reasoning: true, + inputTokenCost: 500, + outputTokenCost: 3000, + cacheReadTokenCost: 100, + cacheWriteTokenCost: 700, + currency: 'RUB', + }); + + renderWithIntl( + + ); + + expect(await screen.findByText('₽6.60')).toBeInTheDocument(); + expect( + screen.getByText(/Input: 10,000 tokens \(₽3.60\) \| Output: 1,000 tokens \(₽3.00\)/) + ).toBeInTheDocument(); + }); + + it('adds usage from different providers with the same currency using each provider price', async () => { + fetchCanonicalModelInfoMock.mockImplementation(async (provider) => ({ + provider, + model: 'shared-model', + contextLimit: 128000, + maxOutputTokens: null, + reasoning: true, + inputTokenCost: provider === 'cheap-provider' ? 100 : 500, + outputTokenCost: provider === 'cheap-provider' ? 200 : 1000, + cacheReadTokenCost: null, + cacheWriteTokenCost: null, + currency: 'USD', + })); + + renderWithIntl( + + ); + + expect(await screen.findByText('$7.20')).toBeInTheDocument(); + expect( + screen.getByText(/Input: 20,000 tokens \(\$6.00\) \| Output: 2,000 tokens \(\$1.20\)/) + ).toBeInTheDocument(); + expect(fetchCanonicalModelInfoMock).toHaveBeenCalledWith('cheap-provider', 'shared-model'); + expect(fetchCanonicalModelInfoMock).toHaveBeenCalledWith('expensive-provider', 'shared-model'); + }); + + it('keeps separate totals when provider usage contains different currencies', async () => { + fetchCanonicalModelInfoMock.mockImplementation(async (provider) => ({ + provider, + model: 'shared-model', + contextLimit: 128000, + maxOutputTokens: null, + reasoning: true, + inputTokenCost: provider === 'rub-provider' ? 500 : 1, + outputTokenCost: provider === 'rub-provider' ? 3000 : 2, + cacheReadTokenCost: null, + cacheWriteTokenCost: null, + currency: provider === 'rub-provider' ? 'RUB' : 'USD', + })); + + renderWithIntl( + + ); + + expect(await screen.findByText('₽8.00 + $3.00')).toBeInTheDocument(); + expect( + screen.getByText( + /Input: 1,010,000 tokens \(₽5.00 \+ \$1.00\) \| Output: 1,001,000 tokens \(₽3.00 \+ \$2.00\)/ + ) + ).toBeInTheDocument(); + }); + + it('treats free provider usage as zero cost when mixed with paid usage', async () => { + fetchCanonicalModelInfoMock.mockImplementation(async (provider) => { + if (provider === 'free-provider') { + return { + provider, + model: 'llama3', + contextLimit: 128000, + maxOutputTokens: null, + reasoning: false, + inputTokenCost: 0, + outputTokenCost: 0, + cacheReadTokenCost: null, + cacheWriteTokenCost: null, + currency: '$', + }; + } + + return { + provider, + model: 'gpt-5.5', + contextLimit: 128000, + maxOutputTokens: null, + reasoning: true, + inputTokenCost: 500, + outputTokenCost: 3000, + cacheReadTokenCost: null, + cacheWriteTokenCost: null, + currency: 'RUB', + }; + }); + + renderWithIntl( + + ); + + expect(await screen.findByText('₽8.00')).toBeInTheDocument(); + expect( + screen.getByText(/Input: 20,000 tokens \(₽5.00\) \| Output: 2,000 tokens \(₽3.00\)/) + ).toBeInTheDocument(); + }); + + it('does not infer free pricing from provider names', async () => { + fetchCanonicalModelInfoMock.mockResolvedValue({ + provider: 'ollama', + model: 'llama3', + contextLimit: 128000, + maxOutputTokens: null, + reasoning: false, + inputTokenCost: null, + outputTokenCost: null, + cacheReadTokenCost: null, + cacheWriteTokenCost: null, + currency: '$', + }); + + renderWithIntl( + + ); + + expect(await screen.findByText('0.0000')).toBeInTheDocument(); + expect(screen.getByText(/Pricing data unavailable for llama3/)).toBeInTheDocument(); + }); +}); diff --git a/ui/desktop/src/components/bottom_menu/CostTracker.tsx b/ui/desktop/src/components/bottom_menu/CostTracker.tsx index bbfd9ff235d4..392522d46775 100644 --- a/ui/desktop/src/components/bottom_menu/CostTracker.tsx +++ b/ui/desktop/src/components/bottom_menu/CostTracker.tsx @@ -3,6 +3,7 @@ import { CoinIcon } from '../icons'; import { Tooltip, TooltipContent, TooltipTrigger } from '../ui/Tooltip'; import { fetchCanonicalModelInfo, type CanonicalModelInfo } from '../../utils/canonical'; import { defineMessages, useIntl } from '../../i18n'; +import type { ProviderUsageEntry } from '../../types/chat'; const i18n = defineMessages({ pricingUnavailable: { @@ -27,6 +28,7 @@ interface CostTrackerProps { inputTokens?: number; outputTokens?: number; accumulatedCost?: number | null; + providerUsage?: ProviderUsageEntry[]; model: string | null; provider: string | null; } @@ -35,11 +37,12 @@ export function CostTracker({ inputTokens = 0, outputTokens = 0, accumulatedCost, + providerUsage, model: currentModel, provider: currentProvider, }: CostTrackerProps) { const intl = useIntl(); - const [costInfo, setCostInfo] = useState(null); + const [costEstimate, setCostEstimate] = useState(null); const [isLoading, setIsLoading] = useState(true); const [showPricing, setShowPricing] = useState(true); const [pricingFailed, setPricingFailed] = useState(false); @@ -64,41 +67,73 @@ export function CostTracker({ useEffect(() => { const loadCostInfo = async () => { if (!currentModel || !currentProvider) { + setCostEstimate(null); + setIsLoading(false); + return; + } + + const hasPartialProviderUsage = + providerUsage?.length && + !providerUsageCoversTokenTotals(providerUsage, inputTokens, outputTokens); + + if (accumulatedCost != null && (!providerUsage?.length || hasPartialProviderUsage)) { + setCostEstimate(null); + setPricingFailed(false); + setIsLoading(false); + return; + } + + if (hasPartialProviderUsage) { + setCostEstimate(null); + setPricingFailed(true); setIsLoading(false); return; } setIsLoading(true); try { - const costData = await fetchCanonicalModelInfo(currentProvider, currentModel); - if (costData) { - setCostInfo(costData); - setPricingFailed(false); - } else { - setPricingFailed(true); - setCostInfo(null); - } + const estimate = await calculateUsageCost( + providerUsage?.length + ? providerUsage + : [ + { + providerId: currentProvider, + modelId: currentModel, + lastUsedAt: '', + inputTokens, + outputTokens, + }, + ] + ); + setCostEstimate(estimate); + setPricingFailed(!estimate); } catch { setPricingFailed(true); - setCostInfo(null); + setCostEstimate(null); } finally { setIsLoading(false); } }; loadCostInfo(); - }, [currentModel, currentProvider]); + }, [currentModel, currentProvider, inputTokens, outputTokens, accumulatedCost, providerUsage]); // Return null early if pricing is disabled if (!showPricing) { return null; } - const calculateCost = (): number => { - return accumulatedCost ?? 0; - }; - - const formatCost = (cost: number): string => cost.toFixed(2); + const renderCost = (displayCost: string, tooltip: string) => ( + + +
+ + {displayCost} +
+
+ {tooltip} +
+ ); // Show loading state or when we don't have model/provider info if (!currentModel || !currentProvider) { @@ -114,86 +149,153 @@ export function CostTracker({ ); } - if ( - accumulatedCost == null && - (!costInfo || - (costInfo.inputTokenCost === undefined && costInfo.outputTokenCost === undefined)) - ) { - const freeProviders = ['ollama', 'local', 'localhost']; - if (freeProviders.includes(currentProvider.toLowerCase())) { - return ( -
- - {inputTokens.toLocaleString()}↑ {outputTokens.toLocaleString()}↓ - -
- ); - } - - // Otherwise show as unavailable - const getUnavailableTooltip = () => { - if (pricingFailed) { - return intl.formatMessage(i18n.pricingUnavailable, { model: currentModel }); - } - return intl.formatMessage(i18n.costUnavailable, { - model: currentModel, - inputTokens: inputTokens.toLocaleString(), - outputTokens: outputTokens.toLocaleString(), - }); - }; - - return ( - - -
- - 0.0000 -
-
- {getUnavailableTooltip()} -
+ if (accumulatedCost == null && !costEstimate) { + return renderCost( + '0.0000', + pricingFailed + ? intl.formatMessage(i18n.pricingUnavailable, { model: currentModel }) + : intl.formatMessage(i18n.costUnavailable, { + model: currentModel, + inputTokens: inputTokens.toLocaleString(), + outputTokens: outputTokens.toLocaleString(), + }) ); } - const totalCost = calculateCost(); + const totalCost = costEstimate ? formatCostTotals(costEstimate.total) : null; + const displayCost = + totalCost || (accumulatedCost == null ? '0.0000' : accumulatedCost.toFixed(2)); // Build tooltip content const getTooltipContent = (): string => { - if (pricingFailed) { + if (pricingFailed && accumulatedCost == null) { return intl.formatMessage(i18n.pricingUnavailable, { model: `${currentProvider}/${currentModel}` }); } - const currency = costInfo?.currency || '$'; - - if (accumulatedCost != null) { - return intl.formatMessage(i18n.totalSessionCost, { cost: `${currency}${totalCost.toFixed(4)}` }) + if (costEstimate) { + return intl.formatMessage(i18n.totalSessionCost, { cost: totalCost }) + `\n` + intl.formatMessage(i18n.inputOutputTooltip, { - inputTokens: inputTokens.toLocaleString(), - inputCost: `${currency}${((inputTokens * (costInfo?.inputTokenCost || 0)) / 1_000_000).toFixed(6)}`, - outputTokens: outputTokens.toLocaleString(), - outputCost: `${currency}${((outputTokens * (costInfo?.outputTokenCost || 0)) / 1_000_000).toFixed(6)}`, + inputTokens: costEstimate.inputTokens.toLocaleString(), + inputCost: formatCostTotals(costEstimate.input), + outputTokens: costEstimate.outputTokens.toLocaleString(), + outputCost: formatCostTotals(costEstimate.output), }); } - const inputCostStr = `${currency}${((inputTokens * (costInfo?.inputTokenCost || 0)) / 1_000_000).toFixed(6)}`; - const outputCostStr = `${currency}${((outputTokens * (costInfo?.outputTokenCost || 0)) / 1_000_000).toFixed(6)}`; - return intl.formatMessage(i18n.inputOutputTooltip, { - inputTokens: inputTokens.toLocaleString(), - inputCost: inputCostStr, - outputTokens: outputTokens.toLocaleString(), - outputCost: outputCostStr, - }); + return intl.formatMessage(i18n.totalSessionCost, { cost: displayCost }) + + `\n` + intl.formatMessage(i18n.inputOutputTooltip, { + inputTokens: inputTokens.toLocaleString(), + inputCost: 'unknown', + outputTokens: outputTokens.toLocaleString(), + outputCost: 'unknown', + }); + }; + + return renderCost(displayCost, getTooltipContent()); +} + +type CostEstimate = { + total: Record; + input: Record; + output: Record; + inputTokens: number; + outputTokens: number; +}; + +async function calculateUsageCost(entries: ProviderUsageEntry[]): Promise { + const estimate: CostEstimate = { + total: {}, + input: {}, + output: {}, + inputTokens: 0, + outputTokens: 0, }; - return ( - - -
- - {formatCost(totalCost)} -
-
- {getTooltipContent()} -
+ for (const entry of entries) { + const modelInfo = await fetchCanonicalModelInfo(entry.providerId, entry.modelId); + const cost = calculateEntryCost(entry, modelInfo); + if (!cost) { + return null; + } + estimate.inputTokens += entry.inputTokens; + estimate.outputTokens += entry.outputTokens; + const currency = modelInfo?.currency || '$'; + addCost(estimate.input, currency, cost.input); + addCost(estimate.output, currency, cost.output); + addCost(estimate.total, currency, cost.input + cost.output); + } + + return Object.keys(estimate.total).length ? estimate : null; +} + +function calculateEntryCost( + entry: ProviderUsageEntry, + modelInfo: CanonicalModelInfo | null +): { input: number; output: number } | null { + const inputTokenCost = modelInfo?.inputTokenCost; + const outputTokenCost = modelInfo?.outputTokenCost; + if (inputTokenCost == null || outputTokenCost == null) { + return null; + } + + const cacheReadTokens = entry.cacheReadInputTokens ?? 0; + const cacheWriteTokens = entry.cacheWriteInputTokens ?? 0; + const uncachedInputTokens = Math.max( + 0, + entry.inputTokens - cacheReadTokens - cacheWriteTokens ); + const cacheReadCost = modelInfo?.cacheReadTokenCost ?? inputTokenCost; + const cacheWriteCost = modelInfo?.cacheWriteTokenCost ?? inputTokenCost; + + const input = + (uncachedInputTokens * inputTokenCost + + cacheReadTokens * cacheReadCost + + cacheWriteTokens * cacheWriteCost) / + 1_000_000; + const output = (entry.outputTokens * outputTokenCost) / 1_000_000; + + return { input, output }; +} + +function providerUsageCoversTokenTotals( + entries: ProviderUsageEntry[], + inputTokens: number, + outputTokens: number +): boolean { + const totals = entries.reduce( + (totals, entry) => ({ + input: totals.input + entry.inputTokens, + output: totals.output + entry.outputTokens, + }), + { input: 0, output: 0 } + ); + + return totals.input >= inputTokens && totals.output >= outputTokens; +} + +function addCost(totals: Record, currency: string, amount: number): void { + totals[currency] = (totals[currency] ?? 0) + amount; +} + +function formatCostTotals(totals: Record): string { + const entries = Object.entries(totals); + const displayEntries = entries.filter(([, amount]) => amount > 0); + + return (displayEntries.length ? displayEntries : entries.slice(0, 1)) + .map(([currency, amount]) => formatMoney(amount, currency)) + .join(' + '); +} + +function formatMoney(amount: number, currency: string, digits = 2): string { + const unit = currency.trim().toUpperCase(); + const value = amount.toFixed(digits); + const symbol = + { + USD: '$', + EUR: '€', + GBP: '£', + RUB: '₽', + }[unit] ?? currency.trim(); + + return /^[A-Z]{3}$/.test(symbol) ? `${value} ${symbol}` : `${symbol}${value}`; } diff --git a/ui/desktop/src/types/chat.ts b/ui/desktop/src/types/chat.ts index 32e007c74bbf..25573195854a 100644 --- a/ui/desktop/src/types/chat.ts +++ b/ui/desktop/src/types/chat.ts @@ -13,6 +13,17 @@ export type TokenState = { inputTokens: number; outputTokens: number; totalTokens: number; + providerUsage?: ProviderUsageEntry[]; +}; + +export type ProviderUsageEntry = { + providerId: string; + modelId: string; + lastUsedAt: string; + inputTokens: number; + outputTokens: number; + cacheReadInputTokens?: number; + cacheWriteInputTokens?: number; }; export interface ChatType {