From 9b837f1a463f87fbe1bf06e1bb32bce204999916 Mon Sep 17 00:00:00 2001 From: filip <44206832+filipkujawa@users.noreply.github.com> Date: Sun, 5 Jul 2026 13:39:10 -0700 Subject: [PATCH] feat: per-message usage/cost tracking with derived session totals (#10172) Co-authored-by: Douwe M Osinga --- .../goose-provider-types/src/conversation.rs | 4 + .../src/conversation/message.rs | 47 ++ .../src/conversation/token_usage.rs | 25 +- .../src/formats/anthropic.rs | 58 +- .../src/formats/openai.rs | 21 +- crates/goose-providers/src/openai.rs | 18 +- .../goose-providers/src/openai_compatible.rs | 12 +- crates/goose-server/src/openapi.rs | 9 +- crates/goose-server/src/routes/reply.rs | 29 +- crates/goose/src/acp/response_builder.rs | 4 +- crates/goose/src/acp/server.rs | 48 +- crates/goose/src/acp/server/fork_session.rs | 6 +- crates/goose/src/acp/server/load_session.rs | 2 +- crates/goose/src/acp/server/new_session.rs | 6 +- crates/goose/src/agents/agent.rs | 26 +- .../src/agents/platform_extensions/summon.rs | 54 +- crates/goose/src/agents/reply_parts.rs | 56 +- crates/goose/src/providers/base.rs | 2 +- crates/goose/src/providers/openrouter.rs | 1 + crates/goose/src/session/session_manager.rs | 635 +++++++++++++++++- ui/desktop/openapi.json | 82 +++ 21 files changed, 1025 insertions(+), 120 deletions(-) diff --git a/crates/goose-provider-types/src/conversation.rs b/crates/goose-provider-types/src/conversation.rs index 457a6b329..0541cf485 100644 --- a/crates/goose-provider-types/src/conversation.rs +++ b/crates/goose-provider-types/src/conversation.rs @@ -43,6 +43,10 @@ impl Conversation { &self.0 } + pub fn messages_mut(&mut self) -> &mut Vec { + &mut self.0 + } + pub fn push(&mut self, message: Message) { if message.content.is_empty() && message.metadata.inference.is_some() { if let Some(existing) = self diff --git a/crates/goose-provider-types/src/conversation/message.rs b/crates/goose-provider-types/src/conversation/message.rs index 38c3a5006..092ce0dd2 100644 --- a/crates/goose-provider-types/src/conversation/message.rs +++ b/crates/goose-provider-types/src/conversation/message.rs @@ -1,3 +1,4 @@ +use crate::conversation::token_usage::{CostSource, ProviderUsage}; use crate::conversation::tool_result_serde; use crate::mcp_utils::extract_text_from_resource; use crate::utils::sanitize_unicode_tags; @@ -666,6 +667,49 @@ pub struct InferenceMetadata { pub resolved_model: Option, } +#[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize, Debug, Default)] +#[serde(rename_all = "camelCase")] +pub struct MessageUsage { + #[serde(skip_serializing_if = "Option::is_none")] + pub input_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub output_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub total_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_write_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cost: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cost_source: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub elapsed_ms: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub time_to_first_token_ms: Option, + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + pub is_compaction: bool, +} + +impl MessageUsage { + pub fn from_provider_usage(usage: &ProviderUsage, is_compaction: bool) -> Self { + let stats = usage.stats.as_ref(); + MessageUsage { + input_tokens: usage.usage.input_tokens, + output_tokens: usage.usage.output_tokens, + total_tokens: usage.usage.total_tokens, + cache_read_tokens: usage.usage.cache_read_input_tokens, + cache_write_tokens: usage.usage.cache_write_input_tokens, + cost: usage.cost, + cost_source: usage.cost_source, + elapsed_ms: stats.and_then(|s| s.elapsed_ms), + time_to_first_token_ms: stats.and_then(|s| s.time_to_first_token_ms), + is_compaction, + } + } +} + #[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize, Debug)] /// Metadata for message visibility and model inference details #[serde(rename_all = "camelCase")] @@ -681,6 +725,8 @@ pub struct MessageMetadata { /// without matching user-visible text. Never sent to providers. #[serde(default, skip_serializing_if = "std::ops::Not::not")] pub steer: bool, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub usage: Option>, } impl Default for MessageMetadata { @@ -690,6 +736,7 @@ impl Default for MessageMetadata { agent_visible: true, inference: None, steer: false, + usage: None, } } } diff --git a/crates/goose-provider-types/src/conversation/token_usage.rs b/crates/goose-provider-types/src/conversation/token_usage.rs index d62a4d370..925228a3d 100644 --- a/crates/goose-provider-types/src/conversation/token_usage.rs +++ b/crates/goose-provider-types/src/conversation/token_usage.rs @@ -9,6 +9,17 @@ pub struct ProviderUsage { pub usage: Usage, #[serde(default, skip_serializing_if = "Option::is_none")] pub stats: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cost: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cost_source: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, ToSchema)] +#[serde(rename_all = "snake_case")] +pub enum CostSource { + ProviderReported, + Estimated, } #[derive(Debug, Clone, Serialize, Deserialize, Default)] @@ -35,6 +46,8 @@ impl ProviderUsage { model, usage, stats: None, + cost: None, + cost_source: None, } } @@ -43,14 +56,10 @@ impl ProviderUsage { self } - /// Combine this ProviderUsage with another, adding their token counts - /// Uses the model from this ProviderUsage - pub fn combine_with(&self, other: &ProviderUsage) -> ProviderUsage { - ProviderUsage { - model: self.model.clone(), - usage: self.usage + other.usage, - stats: self.stats.clone().or_else(|| other.stats.clone()), - } + pub fn with_cost(mut self, cost: f64, source: CostSource) -> Self { + self.cost = Some(cost); + self.cost_source = Some(source); + self } } diff --git a/crates/goose-provider-types/src/formats/anthropic.rs b/crates/goose-provider-types/src/formats/anthropic.rs index 8b3feaeb2..4c22def87 100644 --- a/crates/goose-provider-types/src/formats/anthropic.rs +++ b/crates/goose-provider-types/src/formats/anthropic.rs @@ -1,7 +1,7 @@ use crate::canonical::maybe_get_canonical_model; use crate::canonical::ThinkingMode; use crate::conversation::message::{Message, MessageContent}; -use crate::conversation::token_usage::{ProviderUsage, Usage}; +use crate::conversation::token_usage::{CostSource, ProviderUsage, Usage}; use crate::errors::ProviderError; use crate::images::{convert_image, ImageFormat}; use crate::mcp_utils::extract_text_from_resource; @@ -580,6 +580,19 @@ pub fn get_usage(data: &Value) -> Result { } } +fn provider_usage_with_cost( + model: String, + usage: Usage, + data: &Value, + fallback_cost: Option, +) -> ProviderUsage { + let provider_usage = ProviderUsage::new(model, usage); + match super::openai::get_cost(data).or(fallback_cost) { + Some(cost) => provider_usage.with_cost(cost, CostSource::ProviderReported), + None => provider_usage, + } +} + pub fn thinking_effort(model_config: &ModelConfig) -> ThinkingEffort { model_config .thinking_effort() @@ -810,7 +823,7 @@ where .and_then(|v| v.as_str()) .unwrap_or("unknown") .to_string(); - final_usage = Some(ProviderUsage::new(model, usage)); + final_usage = Some(provider_usage_with_cost(model, usage, usage_data, None)); } } continue; @@ -944,13 +957,18 @@ where if let Some(existing_usage) = &final_usage { let merged_usage = merge_delta_usage(&existing_usage.usage, &delta_usage, usage_data); - final_usage = Some(ProviderUsage::new(existing_usage.model.clone(), merged_usage)); + final_usage = Some(provider_usage_with_cost( + existing_usage.model.clone(), + merged_usage, + usage_data, + existing_usage.cost, + )); } else { let model = event.data.get("model") .and_then(|v| v.as_str()) .unwrap_or("unknown") .to_string(); - final_usage = Some(ProviderUsage::new(model, delta_usage)); + final_usage = Some(provider_usage_with_cost(model, delta_usage, usage_data, None)); } } if let Some(delta) = event.data.get("delta") { @@ -994,7 +1012,8 @@ where .and_then(|v| v.as_str()) .unwrap_or("unknown") .to_string(); - final_usage = Some(ProviderUsage::new(model, usage)); + let fallback_cost = final_usage.as_ref().and_then(|u| u.cost); + final_usage = Some(provider_usage_with_cost(model, usage, usage_data, fallback_cost)); } break; } @@ -2021,6 +2040,35 @@ mod tests { assert_eq!(usage.usage.cache_write_input_tokens, Some(10000)); } + #[tokio::test] + async fn test_streaming_preserves_provider_cost_from_delta() { + let events = concat!( + r#"data: {"type":"message_start","message":{"id":"m1","role":"assistant","content":[],"model":"glm-4.7","usage":{"input_tokens":100,"output_tokens":0}}}"#, + "\n", + r#"data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}"#, + "\n", + r#"data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}"#, + "\n", + r#"data: {"type":"content_block_stop","index":0}"#, + "\n", + r#"data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":50,"cost":0.0123}}"#, + "\n", + r#"data: {"type":"message_stop"}"#, + ); + + let usage = collect_stream_results(events) + .await + .into_iter() + .filter_map(|r| r.ok().and_then(|(_, usage)| usage)) + .next_back() + .expect("stream should yield usage"); + + assert_eq!(usage.cost, Some(0.0123)); + assert_eq!(usage.cost_source, Some(CostSource::ProviderReported)); + assert_eq!(usage.usage.input_tokens, Some(100)); + assert_eq!(usage.usage.output_tokens, Some(50)); + } + #[tokio::test] async fn test_streaming_delta_usage_is_cumulative_and_wins() { // Server tool use grows input during the turn: the final diff --git a/crates/goose-provider-types/src/formats/openai.rs b/crates/goose-provider-types/src/formats/openai.rs index 496d3da54..191a5f3a8 100644 --- a/crates/goose-provider-types/src/formats/openai.rs +++ b/crates/goose-provider-types/src/formats/openai.rs @@ -1,5 +1,5 @@ use crate::conversation::message::{Message, MessageContent, ProviderMetadata}; -use crate::conversation::token_usage::{ProviderUsage, Usage}; +use crate::conversation::token_usage::{CostSource, ProviderUsage, Usage}; use crate::errors::ProviderError; use crate::images::{convert_image, detect_image_path, load_image_file, ImageFormat}; use crate::json::{parse_tool_arguments, truncation_error_message}; @@ -836,6 +836,13 @@ pub fn get_usage(usage: &Value) -> Usage { .with_cache_tokens(cache_read_input_tokens, cache_write_input_tokens) } +pub fn get_cost(usage: &Value) -> Option { + usage + .get("cost") + .and_then(|v| v.as_f64()) + .filter(|c| c.is_finite() && *c >= 0.0) +} + fn extract_usage_with_output_tokens( chunk: &StreamingChunk, fallback_model: Option<&str>, @@ -844,11 +851,13 @@ fn extract_usage_with_output_tokens( .usage .as_ref() .and_then(|u| { - chunk - .model - .as_deref() - .or(fallback_model) - .map(|model| ProviderUsage::new(model.to_string(), get_usage(u))) + chunk.model.as_deref().or(fallback_model).map(|model| { + let usage = ProviderUsage::new(model.to_string(), get_usage(u)); + match get_cost(u) { + Some(cost) => usage.with_cost(cost, CostSource::ProviderReported), + None => usage, + } + }) }) .filter(|u| u.usage.output_tokens.is_some()) } diff --git a/crates/goose-providers/src/openai.rs b/crates/goose-providers/src/openai.rs index a5e6e8630..11218218c 100644 --- a/crates/goose-providers/src/openai.rs +++ b/crates/goose-providers/src/openai.rs @@ -3,12 +3,12 @@ use super::base::{ConfigKey, ModelInfo, Provider, ProviderMetadata}; use super::retry::ProviderRetry; use crate::api_client::{AuthMethod, TlsConfig}; use crate::conversation::message::Message; -use crate::conversation::token_usage::ProviderUsage; +use crate::conversation::token_usage::{CostSource, ProviderUsage}; use crate::declarative::{DeclarativeProviderConfig, KeyResolver}; use crate::errors::ProviderError; use crate::formats::openai::is_openai_responses_model; use crate::formats::openai::{ - create_request_with_options, get_usage, response_to_message, OpenAiFormatOptions, + create_request_with_options, get_cost, get_usage, response_to_message, OpenAiFormatOptions, }; use crate::formats::openai_responses::{ create_responses_request, get_responses_usage, responses_api_to_message, ResponsesApiResponse, @@ -641,7 +641,11 @@ impl Provider for OpenAiProvider { let message = responses_api_to_message(&responses_api_response)?; let usage_data = get_responses_usage(&responses_api_response); - let usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); + let usage_json = json.get("usage").unwrap_or(&serde_json::Value::Null); + let mut usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); + if let Some(cost) = get_cost(usage_json) { + usage = usage.with_cost(cost, CostSource::ProviderReported); + } log.write( &serde_json::to_value(&message).unwrap_or_default(), @@ -689,8 +693,12 @@ impl Provider for OpenAiProvider { ProviderError::RequestFailed(format!("Failed to parse message: {}", e)) })?; - let usage_data = get_usage(json.get("usage").unwrap_or(&serde_json::Value::Null)); - let usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); + let usage_json = json.get("usage").unwrap_or(&serde_json::Value::Null); + let usage_data = get_usage(usage_json); + let mut usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); + if let Some(cost) = get_cost(usage_json) { + usage = usage.with_cost(cost, CostSource::ProviderReported); + } log.write( &serde_json::to_value(&message).unwrap_or_default(), diff --git a/crates/goose-providers/src/openai_compatible.rs b/crates/goose-providers/src/openai_compatible.rs index 5ce39c817..78f96aed9 100644 --- a/crates/goose-providers/src/openai_compatible.rs +++ b/crates/goose-providers/src/openai_compatible.rs @@ -1,4 +1,4 @@ -use crate::conversation::token_usage::ProviderUsage; +use crate::conversation::token_usage::{CostSource, ProviderUsage}; use crate::images::ImageFormat; use anyhow::Error; use async_stream::try_stream; @@ -18,7 +18,7 @@ use super::retry::ProviderRetry; use crate::conversation::message::Message; use crate::errors::ProviderError; use crate::formats::openai::{ - create_request, get_usage, response_to_message, response_to_streaming_message, + create_request, get_cost, get_usage, response_to_message, response_to_streaming_message, }; use crate::formats::openai_responses::responses_api_to_streaming_message; use crate::model::ModelConfig; @@ -143,8 +143,12 @@ impl Provider for OpenAiCompatibleProvider { ProviderError::RequestFailed(format!("Failed to parse message: {}", e)) })?; - let usage_data = get_usage(json.get("usage").unwrap_or(&serde_json::Value::Null)); - let usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); + let usage_json = json.get("usage").unwrap_or(&serde_json::Value::Null); + let usage_data = get_usage(usage_json); + let mut usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); + if let Some(cost) = get_cost(usage_json) { + usage = usage.with_cost(cost, CostSource::ProviderReported); + } log.write( &serde_json::to_value(&message).unwrap_or_default(), diff --git a/crates/goose-server/src/openapi.rs b/crates/goose-server/src/openapi.rs index ec679a7dc..fa43e6f75 100644 --- a/crates/goose-server/src/openapi.rs +++ b/crates/goose-server/src/openapi.rs @@ -28,10 +28,11 @@ use goose::config::declarative_providers::{ }; use goose::conversation::message::{ ActionRequired, ActionRequiredData, FrontendToolRequest, InferenceMetadata, Message, - MessageContent, MessageMetadata, RedactedThinkingContent, SystemNotificationContent, - SystemNotificationType, ThinkingContent, TokenState, ToolConfirmationRequest, ToolRequest, - ToolResponse, + MessageContent, MessageMetadata, MessageUsage, RedactedThinkingContent, + SystemNotificationContent, SystemNotificationType, ThinkingContent, TokenState, + ToolConfirmationRequest, ToolRequest, ToolResponse, }; +use goose::providers::base::CostSource; use crate::routes::recipe_utils::RecipeManifest; use crate::routes::reply::MessageEvent; @@ -512,6 +513,8 @@ derive_utoipa!(IconTheme as IconThemeSchema); MessageContent, MessageMetadata, InferenceMetadata, + MessageUsage, + CostSource, TokenState, Usage, ContentSchema, diff --git a/crates/goose-server/src/routes/reply.rs b/crates/goose-server/src/routes/reply.rs index dbc6ed5ab..e06162849 100644 --- a/crates/goose-server/src/routes/reply.rs +++ b/crates/goose-server/src/routes/reply.rs @@ -154,18 +154,23 @@ pub enum MessageEvent { } pub async fn get_token_state(session_manager: &SessionManager, session_id: &str) -> TokenState { - session_manager - .get_session(session_id, false) - .await - .map(|session| TokenState::from(&session)) - .inspect_err(|e| { - tracing::warn!( - "Failed to fetch session token state for {}: {}", - session_id, - e - ); - }) - .unwrap_or_default() + let session = match session_manager.get_session(session_id, false).await { + Ok(session) => session, + Err(e) => { + tracing::warn!("Failed to fetch session token state for {session_id}: {e}"); + return TokenState::default(); + } + }; + + match session_manager.get_session_usage_totals(session_id).await { + Ok(totals) => { + goose::session::session_manager::token_state_from_session_and_totals(&session, &totals) + } + Err(e) => { + tracing::warn!("Failed to aggregate usage for {session_id}: {e}"); + TokenState::from(&session) + } + } } async fn stream_event( diff --git a/crates/goose/src/acp/response_builder.rs b/crates/goose/src/acp/response_builder.rs index babb0fd25..dfb314bd4 100644 --- a/crates/goose/src/acp/response_builder.rs +++ b/crates/goose/src/acp/response_builder.rs @@ -1,6 +1,7 @@ use crate::agents::ExtensionLoadResult; use crate::config::{Config, GooseMode}; use crate::providers::inventory::{ProviderInventoryEntry, ProviderInventoryService}; +use crate::session::session_manager::SessionUsageTotals; use crate::session::Session; use crate::slash_commands::types::{SlashCommandEntry, SlashCommandSource}; use agent_client_protocol::schema::v1::{ @@ -401,10 +402,11 @@ fn available_commands_update(working_dir: &std::path::Path) -> AvailableCommands pub(super) fn send_session_setup_notifications( cx: &ConnectionTo, session: &Session, + totals: &SessionUsageTotals, supports_goose_custom_notifications: bool, ) -> Result<(), agent_client_protocol::Error> { let session_id = SessionId::new(session.id.clone()); - if let Some(updates) = build_usage_updates(session) { + if let Some(updates) = build_usage_updates(session, totals) { if supports_goose_custom_notifications { cx.send_notification(updates.custom)?; } diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 5365f1b9b..103126c29 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -34,6 +34,7 @@ use crate::providers::inventory::{ RefreshSkipReason, }; use crate::scheduler_trait::SchedulerTrait; +use crate::session::session_manager::SessionUsageTotals; use crate::session::{ EnabledExtensionsState, ExtensionData, ExtensionState, Session, SessionManager, SessionType, }; @@ -836,13 +837,16 @@ pub(super) struct UsageUpdates { pub(super) standard: UsageUpdate, } -pub(super) fn build_usage_updates(session: &Session) -> Option { +pub(super) fn build_usage_updates( + session: &Session, + totals: &SessionUsageTotals, +) -> 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; let accumulated_input_tokens = - to_nonnegative_u64(session.accumulated_usage.input_tokens).unwrap_or(0); + to_nonnegative_u64(totals.accumulated_usage.input_tokens).unwrap_or(0); let accumulated_output_tokens = - to_nonnegative_u64(session.accumulated_usage.output_tokens).unwrap_or(0); + to_nonnegative_u64(totals.accumulated_usage.output_tokens).unwrap_or(0); Some(UsageUpdates { custom: GooseSessionNotification { session_id: session.id.clone(), @@ -851,12 +855,12 @@ pub(super) fn build_usage_updates(session: &Session) -> Option { context_limit: ctx_limit, accumulated_input_tokens, accumulated_output_tokens, - accumulated_cost: session.accumulated_cost, + accumulated_cost: totals.accumulated_cost, }), }, standard: { let mut standard = UsageUpdate::new(used, ctx_limit); - if let Some(amount) = session.accumulated_cost { + if let Some(amount) = totals.accumulated_cost { standard = standard.cost(Cost::new(amount, "USD")); } standard @@ -890,6 +894,24 @@ impl GooseAcpAgent { .unwrap_or(false) } + pub(super) async fn notify_session_setup( + &self, + cx: &ConnectionTo, + session: &Session, + ) -> Result<(), agent_client_protocol::Error> { + let totals = self + .session_manager + .get_session_usage_totals(&session.id) + .await + .unwrap_or_default(); + send_session_setup_notifications( + cx, + session, + &totals, + self.supports_goose_custom_notifications(), + ) + } + pub(super) fn supports_recipe_param_requests(&self) -> bool { self.client_supports_recipe_param_requests .get() @@ -2670,7 +2692,12 @@ impl GooseAcpAgent { .get_session(&session_id, false) .await .internal_err_ctx("Failed to load session")?; - if let Some(updates) = build_usage_updates(&session) { + let totals = self + .session_manager + .get_session_usage_totals(&session_id) + .await + .unwrap_or_default(); + if let Some(updates) = build_usage_updates(&session, &totals) { if self.supports_goose_custom_notifications() { cx.send_notification(updates.custom)?; } @@ -3876,7 +3903,12 @@ print(\"hello, world\") goose_providers::model::ModelConfig::new("test-model") .with_context_limit(Some(258_000)), ); - let updates = build_usage_updates(&session).expect("usage updates should be present"); + let totals = SessionUsageTotals { + accumulated_usage: session.accumulated_usage, + accumulated_cost: session.accumulated_cost, + }; + let updates = + build_usage_updates(&session, &totals).expect("usage updates should be present"); assert_eq!(updates.custom.session_id, "session-1"); let usage = match updates.custom.update { GooseSessionUpdate::UsageUpdate(usage) => usage, @@ -3894,7 +3926,7 @@ print(\"hello, world\") TokenUsage::new(Some(80), Some(40), Some(120)), TokenUsage::default(), ); - assert!(build_usage_updates(&session).is_none()); + assert!(build_usage_updates(&session, &SessionUsageTotals::default()).is_none()); } #[test] diff --git a/crates/goose/src/acp/server/fork_session.rs b/crates/goose/src/acp/server/fork_session.rs index 2f08e3477..4812f8cfd 100644 --- a/crates/goose/src/acp/server/fork_session.rs +++ b/crates/goose/src/acp/server/fork_session.rs @@ -72,11 +72,7 @@ impl GooseAcpAgent { if let Some(co) = config_options { response = response.config_options(co); } - send_session_setup_notifications( - cx, - &goose_session, - self.supports_goose_custom_notifications(), - )?; + self.notify_session_setup(cx, &goose_session).await?; Ok(response) } } diff --git a/crates/goose/src/acp/server/load_session.rs b/crates/goose/src/acp/server/load_session.rs index fa5538d08..dc8ea3db5 100644 --- a/crates/goose/src/acp/server/load_session.rs +++ b/crates/goose/src/acp/server/load_session.rs @@ -209,7 +209,7 @@ impl GooseAcpAgent { let (mode_state, config_options) = build_session_setup_config(&self.provider_inventory, &session).await?; - send_session_setup_notifications(cx, &session, self.supports_goose_custom_notifications())?; + self.notify_session_setup(cx, &session).await?; let mut response = LoadSessionResponse::new().modes(mode_state); if let Some(co) = config_options { diff --git a/crates/goose/src/acp/server/new_session.rs b/crates/goose/src/acp/server/new_session.rs index 68934047e..fe7c28f03 100644 --- a/crates/goose/src/acp/server/new_session.rs +++ b/crates/goose/src/acp/server/new_session.rs @@ -84,11 +84,7 @@ impl GooseAcpAgent { let response = self .build_new_session_response(&reloaded_session, &extension_results) .await?; - super::send_session_setup_notifications( - cx, - &reloaded_session, - self.supports_goose_custom_notifications(), - )?; + self.notify_session_setup(cx, &reloaded_session).await?; Ok(response) } diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index d4a8f592e..e14aa3516 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -34,7 +34,7 @@ use crate::context_mgmt::{ check_if_compaction_needed, compact_messages, DEFAULT_COMPACTION_THRESHOLD, }; use crate::conversation::message::{ - ActionRequiredData, InferenceMetadata, Message, MessageContent, ProviderMetadata, + ActionRequiredData, InferenceMetadata, Message, MessageContent, MessageUsage, ProviderMetadata, SystemNotificationType, ToolRequest, }; use crate::conversation::{debug_conversation_fix, fix_conversation, Conversation}; @@ -53,6 +53,7 @@ use crate::session::{Session, SessionManager, SessionNameUpdate}; use crate::tool_inspection::ToolInspectionManager; use crate::tool_monitor::RepetitionInspector; use crate::utils::is_token_cancelled; +use goose_providers::conversation::token_usage::ProviderUsage; use goose_providers::errors::ProviderError; use goose_providers::thinking::ThinkingEffort; use regex::Regex; @@ -268,6 +269,17 @@ pub enum AgentEvent { HistoryReplaced(Conversation), } +fn attach_turn_usage(messages: &mut Conversation, usage: &ProviderUsage) { + if let Some(message) = messages + .messages_mut() + .iter_mut() + .rev() + .find(|m| m.role == rmcp::model::Role::Assistant) + { + message.metadata.usage = Some(Box::new(MessageUsage::from_provider_usage(usage, false))); + } +} + impl Default for Agent { fn default() -> Self { Self::new() @@ -2018,6 +2030,7 @@ impl Agent { let mut did_recovery_compact_this_iteration = false; let mut exit_chat = false; let mut pending_final_output: Option = None; + let mut pending_turn_usage: Option = None; // Track whether this provider turn has already emitted visible // thinking so a later tool-call chunk can suppress replayed @@ -2034,8 +2047,9 @@ impl Agent { 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()); + let enriched = self.update_session_metrics(&session_config.id, session_config.schedule_id.clone(), usage, false).await?; + yield AgentEvent::Usage(enriched.clone()); + pending_turn_usage = Some(enriched); } if let Some(response) = response { @@ -2655,7 +2669,7 @@ impl Agent { yield AgentEvent::Message(message); } - let messages_to_add = if let Some(ref inference) = inference { + let mut messages_to_add = if let Some(ref inference) = inference { Conversation::new_unvalidated( messages_to_add .into_iter() @@ -2665,6 +2679,10 @@ impl Agent { messages_to_add }; + if let Some(usage) = pending_turn_usage.take() { + attach_turn_usage(&mut messages_to_add, &usage); + } + for msg in &messages_to_add { session_manager.add_message(&session_config.id, msg).await?; } diff --git a/crates/goose/src/agents/platform_extensions/summon.rs b/crates/goose/src/agents/platform_extensions/summon.rs index 7e3dfcc29..255241983 100644 --- a/crates/goose/src/agents/platform_extensions/summon.rs +++ b/crates/goose/src/agents/platform_extensions/summon.rs @@ -483,6 +483,36 @@ impl SummonClient { }) } + async fn create_subagent_session( + &self, + task_config: &TaskConfig, + name: String, + ) -> Result { + let session = self + .context + .session_manager + .create_session( + task_config.parent_working_dir.clone(), + name, + SessionType::SubAgent, + GooseMode::Auto, + ) + .await + .map_err(|e| format!("Failed to create subagent session: {}", e))?; + + if !task_config.parent_session_id.is_empty() { + self.context + .session_manager + .update(&session.id) + .parent_session_id(Some(task_config.parent_session_id.clone())) + .apply() + .await + .map_err(|e| format!("Failed to link subagent to parent session: {}", e))?; + } + + Ok(session) + } + fn spawn_notification_bridge( mut notif_rx: tokio::sync::mpsc::UnboundedReceiver, subscribers: Arc>>>, @@ -1255,16 +1285,8 @@ impl SummonClient { .with_use_login_shell_path(self.context.use_login_shell_path); let subagent_session = self - .context - .session_manager - .create_session( - task_config.parent_working_dir.clone(), - "Delegated task".to_string(), - SessionType::SubAgent, - GooseMode::Auto, - ) - .await - .map_err(|e| format!("Failed to create subagent session: {}", e))?; + .create_subagent_session(&task_config, "Delegated task".to_string()) + .await?; let (notif_tx, notif_rx) = tokio::sync::mpsc::unbounded_channel::(); Self::spawn_notification_bridge( @@ -1805,16 +1827,8 @@ impl SummonClient { .with_use_login_shell_path(self.context.use_login_shell_path); let subagent_session = self - .context - .session_manager - .create_session( - task_config.parent_working_dir.clone(), - description.clone(), - SessionType::SubAgent, - GooseMode::Auto, - ) - .await - .map_err(|e| format!("Failed to create subagent session: {}", e))?; + .create_subagent_session(&task_config, description.clone()) + .await?; let task_id = subagent_session.id.clone(); diff --git a/crates/goose/src/agents/reply_parts.rs b/crates/goose/src/agents/reply_parts.rs index c1c366888..5eadd8ab4 100644 --- a/crates/goose/src/agents/reply_parts.rs +++ b/crates/goose/src/agents/reply_parts.rs @@ -12,7 +12,7 @@ use super::super::agents::Agent; #[cfg(feature = "code-mode")] use crate::agents::platform_extensions::code_execution; use crate::config::Config; -use crate::conversation::message::{Message, MessageContent, ToolRequest}; +use crate::conversation::message::{Message, MessageContent, MessageUsage, ToolRequest}; use crate::conversation::Conversation; #[cfg(test)] use crate::providers::base::stream_from_single_message; @@ -21,7 +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 goose_providers::conversation::token_usage::{ProviderUsage, Usage}; +use goose_providers::conversation::token_usage::{CostSource, ProviderUsage, Usage}; use goose_providers::model::ModelConfig; use rmcp::model::Tool; use tracing::warn; @@ -523,17 +523,17 @@ impl Agent { schedule_id: Option, usage: &ProviderUsage, is_compaction_usage: bool, - ) -> Result<()> { + ) -> Result { let manager = self.config.session_manager.clone(); let session = manager.get_session(session_id, false).await?; - let accumulated_usage = session.accumulated_usage + usage.usage; + let (chunk_cost, cost_source) = + self.resolve_chunk_cost(usage, session.provider_name.as_deref()); - let accumulated_cost = session - .provider_name - .as_deref() - .and_then(|pn| self.accumulate_cost(session.accumulated_cost, usage, pn)) - .or(session.accumulated_cost); + let mut enriched = usage.clone(); + enriched.cost = chunk_cost; + enriched.cost_source = cost_source; + let ledger = MessageUsage::from_provider_usage(&enriched, is_compaction_usage); let current_usage = if is_compaction_usage { // After compaction: summary output becomes new input context @@ -544,29 +544,33 @@ impl Agent { }; manager - .update(session_id) - .schedule_id(schedule_id) - .usage(current_usage) - .accumulated_usage(accumulated_usage) - .accumulated_cost(accumulated_cost) - .apply() + .record_usage_metrics( + session_id, + schedule_id, + current_usage, + &usage.model, + &ledger, + ) .await?; - Ok(()) + Ok(enriched) } - fn accumulate_cost( + fn resolve_chunk_cost( &self, - existing: Option, usage: &ProviderUsage, - provider_name: &str, - ) -> Option { - let canonical = - crate::providers::canonical::maybe_get_canonical_model(provider_name, &usage.model)?; - - let chunk_cost = canonical.cost.estimate_cost(&usage.usage)?; - - Some(existing.unwrap_or(0.0) + chunk_cost) + provider_name: Option<&str>, + ) -> (Option, Option) { + if let Some(cost) = usage.cost { + return (Some(cost), Some(CostSource::ProviderReported)); + } + match provider_name + .and_then(|pn| crate::providers::canonical::maybe_get_canonical_model(pn, &usage.model)) + .and_then(|canonical| canonical.cost.estimate_cost(&usage.usage)) + { + Some(cost) => (Some(cost), Some(CostSource::Estimated)), + None => (None, None), + } } } diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index 4f3c096c0..c601c7901 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -2,7 +2,7 @@ use super::api_client::TlsConfig; use anyhow::Result; use futures::future::BoxFuture; pub use goose_providers::conversation::token_usage::{ - DraftStats, ProviderStats, ProviderUsage, Usage, + CostSource, DraftStats, ProviderStats, ProviderUsage, Usage, }; use serde::{Deserialize, Serialize}; diff --git a/crates/goose/src/providers/openrouter.rs b/crates/goose/src/providers/openrouter.rs index 85c3b0d88..92249bd10 100644 --- a/crates/goose/src/providers/openrouter.rs +++ b/crates/goose/src/providers/openrouter.rs @@ -326,6 +326,7 @@ impl Provider for OpenRouterProvider { if let Some(obj) = payload.as_object_mut() { obj.insert("transforms".to_string(), json!(["middle-out"])); + obj.insert("usage".to_string(), json!({ "include": true })); } let mut log = start_log(model_config, &payload)?; diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index d9989bbab..4dc6160fb 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -1,7 +1,8 @@ use crate::config::paths::Paths; use crate::config::GooseMode; -use crate::conversation::message::{Message, TokenState}; +use crate::conversation::message::{Message, MessageUsage, TokenState}; use crate::conversation::Conversation; +use crate::providers::base::CostSource; use crate::providers::base::Provider; use crate::recipe::Recipe; use crate::session::extension_data::ExtensionData; @@ -23,7 +24,7 @@ use std::sync::{Arc, LazyLock}; use tracing::{info, warn}; use utoipa::ToSchema; -pub const CURRENT_SCHEMA_VERSION: i32 = 14; +pub const CURRENT_SCHEMA_VERSION: i32 = 15; pub const SESSIONS_FOLDER: &str = "sessions"; pub const DB_NAME: &str = "sessions.db"; const MILLISECOND_TIMESTAMP_THRESHOLD: i64 = 10_000_000_000; @@ -92,6 +93,8 @@ pub struct Session { #[serde(default)] pub project_id: Option, #[serde(default)] + pub parent_session_id: Option, + #[serde(default)] pub last_message_snippet: Option, } @@ -119,6 +122,31 @@ impl From<&Session> for TokenState { } } +pub fn token_state_from_session_and_totals( + session: &Session, + totals: &SessionUsageTotals, +) -> TokenState { + TokenState { + input_tokens: session.usage.input_tokens.unwrap_or(0), + output_tokens: session.usage.output_tokens.unwrap_or(0), + total_tokens: session.usage.total_tokens.unwrap_or(0), + cache_read_tokens: session.usage.cache_read_input_tokens.unwrap_or(0), + cache_write_tokens: session.usage.cache_write_input_tokens.unwrap_or(0), + accumulated_input_tokens: totals.accumulated_usage.input_tokens.unwrap_or(0), + accumulated_output_tokens: totals.accumulated_usage.output_tokens.unwrap_or(0), + accumulated_total_tokens: totals.accumulated_usage.total_tokens.unwrap_or(0), + accumulated_cache_read_tokens: totals + .accumulated_usage + .cache_read_input_tokens + .unwrap_or(0), + accumulated_cache_write_tokens: totals + .accumulated_usage + .cache_write_input_tokens + .unwrap_or(0), + accumulated_cost: totals.accumulated_cost, + } +} + pub struct SessionUpdateBuilder<'a> { session_manager: &'a SessionManager, session_id: String, @@ -139,6 +167,7 @@ pub struct SessionUpdateBuilder<'a> { archived_at: Option>>, project_id: Option>, + parent_session_id: Option>, } #[derive(Serialize, ToSchema, Debug)] @@ -148,6 +177,12 @@ pub struct SessionInsights { pub total_tokens: i64, } +#[derive(Debug, Clone, Default)] +pub struct SessionUsageTotals { + pub accumulated_usage: Usage, + pub accumulated_cost: Option, +} + impl<'a> SessionUpdateBuilder<'a> { fn new(session_manager: &'a SessionManager, session_id: String) -> Self { Self { @@ -169,6 +204,7 @@ impl<'a> SessionUpdateBuilder<'a> { goose_mode: None, archived_at: None, project_id: None, + parent_session_id: None, } } @@ -271,6 +307,11 @@ impl<'a> SessionUpdateBuilder<'a> { self.project_id = Some(project_id); self } + + pub fn parent_session_id(mut self, parent_session_id: Option) -> Self { + self.parent_session_id = Some(parent_session_id); + self + } } pub struct SessionManager { @@ -430,6 +471,23 @@ impl SessionManager { .await } + pub async fn get_session_usage_totals(&self, id: &str) -> Result { + self.storage.get_session_usage_totals(id).await + } + + pub async fn record_usage_metrics( + &self, + session_id: &str, + schedule_id: Option, + current_usage: Usage, + model: &str, + ledger: &MessageUsage, + ) -> Result<()> { + self.storage + .record_usage_metrics(session_id, schedule_id, current_usage, model, ledger) + .await + } + pub async fn export_session(&self, id: &str) -> Result { self.storage.export_session(id).await } @@ -648,6 +706,7 @@ impl Default for Session { goose_mode: GooseMode::default(), archived_at: None, project_id: None, + parent_session_id: None, last_message_snippet: None, } } @@ -742,11 +801,49 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session { .unwrap_or_default(), archived_at: row.try_get("archived_at").ok(), project_id: row.try_get("project_id").ok().flatten(), + parent_session_id: row.try_get("parent_session_id").ok().flatten(), last_message_snippet: None, }) } } +async fn insert_usage_ledger_row( + tx: &mut sqlx::Transaction<'_, Sqlite>, + session_id: &str, + model: Option<&str>, + usage: &MessageUsage, +) -> Result<()> { + let cost_source = usage.cost_source.map(|cs| match cs { + CostSource::ProviderReported => "provider_reported", + CostSource::Estimated => "estimated", + }); + + sqlx::query( + r#" + INSERT INTO usage_ledger ( + session_id, created_timestamp, model, + input_tokens, output_tokens, total_tokens, + cache_read_tokens, cache_write_tokens, + cost, cost_source, is_compaction + ) + VALUES (?, strftime('%s','now'), ?, ?, ?, ?, ?, ?, ?, ?, ?) + "#, + ) + .bind(session_id) + .bind(model) + .bind(usage.input_tokens) + .bind(usage.output_tokens) + .bind(usage.total_tokens) + .bind(usage.cache_read_tokens) + .bind(usage.cache_write_tokens) + .bind(usage.cost) + .bind(cost_source) + .bind(usage.is_compaction as i64) + .execute(&mut **tx) + .await?; + Ok(()) +} + impl SessionStorage { fn create_pool(path: &Path) -> Pool { if let Some(parent) = path.parent() { @@ -864,7 +961,8 @@ impl SessionStorage { model_config_json TEXT, goose_mode TEXT NOT NULL DEFAULT 'auto', archived_at TIMESTAMP, - project_id TEXT + project_id TEXT, + parent_session_id TEXT ) "#, ) @@ -889,6 +987,27 @@ impl SessionStorage { .execute(&mut *tx) .await?; + sqlx::query( + r#" + CREATE TABLE IF NOT EXISTS usage_ledger ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE, + created_timestamp INTEGER NOT NULL, + model TEXT, + input_tokens INTEGER, + output_tokens INTEGER, + total_tokens INTEGER, + cache_read_tokens INTEGER, + cache_write_tokens INTEGER, + cost REAL, + cost_source TEXT, + is_compaction INTEGER DEFAULT 0 + ) + "#, + ) + .execute(&mut *tx) + .await?; + sqlx::query("CREATE INDEX IF NOT EXISTS idx_messages_session ON messages(session_id)") .execute(&mut *tx) .await?; @@ -904,6 +1023,16 @@ impl SessionStorage { sqlx::query("CREATE INDEX IF NOT EXISTS idx_sessions_type ON sessions(session_type)") .execute(&mut *tx) .await?; + sqlx::query( + "CREATE INDEX IF NOT EXISTS idx_sessions_parent ON sessions(parent_session_id)", + ) + .execute(&mut *tx) + .await?; + sqlx::query( + "CREATE INDEX IF NOT EXISTS idx_usage_ledger_session ON usage_ledger(session_id)", + ) + .execute(&mut *tx) + .await?; tx.commit().await?; @@ -1326,6 +1455,50 @@ impl SessionStorage { } } } + 15 => { + let has_parent = sqlx::query_scalar::<_, i32>( + "SELECT COUNT(*) FROM pragma_table_info('sessions') WHERE name = 'parent_session_id'", + ) + .fetch_one(&mut **tx) + .await? + > 0; + if !has_parent { + sqlx::query("ALTER TABLE sessions ADD COLUMN parent_session_id TEXT") + .execute(&mut **tx) + .await?; + sqlx::query( + "CREATE INDEX IF NOT EXISTS idx_sessions_parent ON sessions(parent_session_id)", + ) + .execute(&mut **tx) + .await?; + } + + sqlx::query( + r#" + CREATE TABLE IF NOT EXISTS usage_ledger ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE, + created_timestamp INTEGER NOT NULL, + model TEXT, + input_tokens INTEGER, + output_tokens INTEGER, + total_tokens INTEGER, + cache_read_tokens INTEGER, + cache_write_tokens INTEGER, + cost REAL, + cost_source TEXT, + is_compaction INTEGER DEFAULT 0 + ) + "#, + ) + .execute(&mut **tx) + .await?; + sqlx::query( + "CREATE INDEX IF NOT EXISTS idx_usage_ledger_session ON usage_ledger(session_id)", + ) + .execute(&mut **tx) + .await?; + } _ => { anyhow::bail!("Unknown migration version: {}", version); } @@ -1391,7 +1564,7 @@ impl SessionStorage { accumulated_cost, schedule_id, recipe_json, user_recipe_values_json, provider_name, model_config_json, goose_mode, - archived_at, project_id + archived_at, project_id, parent_session_id FROM sessions WHERE id = ? "#, @@ -1470,6 +1643,7 @@ impl SessionStorage { add_update!(builder.archived_at, "archived_at"); add_update!(builder.project_id, "project_id"); + add_update!(builder.parent_session_id, "parent_session_id"); if updates.is_empty() { return Ok(()); @@ -1546,6 +1720,9 @@ impl SessionStorage { if let Some(ref project_id) = builder.project_id { q = q.bind(project_id.as_ref()); } + if let Some(ref parent_session_id) = builder.parent_session_id { + q = q.bind(parent_session_id.as_ref()); + } let pool = self.pool().await?; let mut tx = pool.begin_with("BEGIN IMMEDIATE").await?; @@ -1738,7 +1915,7 @@ impl SessionStorage { s.accumulated_cost, s.schedule_id, s.recipe_json, s.user_recipe_values_json, s.provider_name, s.model_config_json, s.goose_mode, - s.archived_at, s.project_id, + s.archived_at, s.project_id, s.parent_session_id, COUNT(m.id) as message_count, MAX({}) as last_message_timestamp, {} as sort_timestamp @@ -1864,6 +2041,11 @@ impl SessionStorage { .execute(&mut *tx) .await?; + sqlx::query("DELETE FROM usage_ledger WHERE session_id = ?") + .bind(session_id) + .execute(&mut *tx) + .await?; + sqlx::query("DELETE FROM sessions WHERE id = ?") .bind(session_id) .execute(&mut *tx) @@ -1906,6 +2088,180 @@ impl SessionStorage { }) } + async fn record_usage_metrics( + &self, + session_id: &str, + schedule_id: Option, + current_usage: Usage, + model: &str, + ledger: &MessageUsage, + ) -> Result<()> { + let pool = self.pool().await?; + let mut tx = pool.begin_with("BEGIN IMMEDIATE").await?; + + sqlx::query( + r#" + INSERT INTO usage_ledger ( + session_id, created_timestamp, + input_tokens, output_tokens, total_tokens, + cache_read_tokens, cache_write_tokens, + cost, cost_source + ) + SELECT s.id, strftime('%s','now'), + MAX(COALESCE(s.accumulated_input_tokens, 0) - l.input_sum, 0), + MAX(COALESCE(s.accumulated_output_tokens, 0) - l.output_sum, 0), + MAX(COALESCE(s.accumulated_total_tokens, 0) - l.total_sum, 0), + MAX(COALESCE(s.accumulated_cache_read_tokens, 0) - l.cache_read_sum, 0), + MAX(COALESCE(s.accumulated_cache_write_tokens, 0) - l.cache_write_sum, 0), + CASE WHEN s.accumulated_cost IS NULL OR s.accumulated_cost <= l.cost_sum THEN NULL + ELSE s.accumulated_cost - l.cost_sum END, + 'carried_forward' + FROM sessions s, + (SELECT COALESCE(SUM(input_tokens), 0) AS input_sum, + COALESCE(SUM(output_tokens), 0) AS output_sum, + COALESCE(SUM(total_tokens), 0) AS total_sum, + COALESCE(SUM(cache_read_tokens), 0) AS cache_read_sum, + COALESCE(SUM(cache_write_tokens), 0) AS cache_write_sum, + COALESCE(SUM(cost), 0.0) AS cost_sum + FROM usage_ledger WHERE session_id = ?) l + WHERE s.id = ? + AND (COALESCE(s.accumulated_input_tokens, 0) > l.input_sum + OR COALESCE(s.accumulated_output_tokens, 0) > l.output_sum + OR COALESCE(s.accumulated_total_tokens, 0) > l.total_sum + OR COALESCE(s.accumulated_cost, 0.0) > l.cost_sum + 1e-9) + "#, + ) + .bind(session_id) + .bind(session_id) + .execute(&mut *tx) + .await?; + + sqlx::query( + r#" + UPDATE sessions SET + schedule_id = ?, + total_tokens = ?, input_tokens = ?, output_tokens = ?, + cache_read_tokens = ?, cache_write_tokens = ?, + accumulated_total_tokens = COALESCE(accumulated_total_tokens, 0) + ?, + accumulated_input_tokens = COALESCE(accumulated_input_tokens, 0) + ?, + accumulated_output_tokens = COALESCE(accumulated_output_tokens, 0) + ?, + accumulated_cache_read_tokens = COALESCE(accumulated_cache_read_tokens, 0) + ?, + accumulated_cache_write_tokens = COALESCE(accumulated_cache_write_tokens, 0) + ?, + accumulated_cost = CASE + WHEN ? IS NULL THEN accumulated_cost + ELSE COALESCE(accumulated_cost, 0) + ? + END, + updated_at = datetime('now') + WHERE id = ? + "#, + ) + .bind(schedule_id) + .bind(current_usage.total_tokens) + .bind(current_usage.input_tokens) + .bind(current_usage.output_tokens) + .bind(current_usage.cache_read_input_tokens) + .bind(current_usage.cache_write_input_tokens) + .bind(ledger.total_tokens.unwrap_or(0)) + .bind(ledger.input_tokens.unwrap_or(0)) + .bind(ledger.output_tokens.unwrap_or(0)) + .bind(ledger.cache_read_tokens.unwrap_or(0)) + .bind(ledger.cache_write_tokens.unwrap_or(0)) + .bind(ledger.cost) + .bind(ledger.cost) + .bind(session_id) + .execute(&mut *tx) + .await?; + + insert_usage_ledger_row(&mut tx, session_id, Some(model), ledger).await?; + + tx.commit().await?; + Ok(()) + } + + async fn get_session_usage_totals(&self, session_id: &str) -> Result { + let pool = self.pool().await?; + let rows = sqlx::query_as::< + _, + ( + Option, + Option, + Option, + Option, + Option, + Option, + Option, + Option, + Option, + Option, + Option, + Option, + ), + >( + r#" + WITH RECURSIVE tree(id) AS ( + SELECT id FROM sessions WHERE id = ? + UNION + SELECT s.id FROM sessions s JOIN tree ON s.parent_session_id = tree.id + ) + SELECT + s.accumulated_input_tokens, s.accumulated_output_tokens, s.accumulated_total_tokens, + s.accumulated_cache_read_tokens, s.accumulated_cache_write_tokens, s.accumulated_cost, + SUM(u.input_tokens), SUM(u.output_tokens), SUM(u.total_tokens), + SUM(u.cache_read_tokens), SUM(u.cache_write_tokens), SUM(u.cost) + FROM sessions s + LEFT JOIN usage_ledger u ON u.session_id = s.id + WHERE s.id IN (SELECT id FROM tree) + GROUP BY s.id + "#, + ) + .bind(session_id) + .fetch_all(pool) + .await?; + + let mut input = 0i64; + let mut output = 0i64; + let mut total = 0i64; + let mut cache_read = 0i64; + let mut cache_write = 0i64; + let mut cost: Option = None; + + let larger = + |acc: Option, ledger: Option| acc.unwrap_or(0).max(ledger.unwrap_or(0)); + + for row in rows { + let ( + acc_in, + acc_out, + acc_total, + acc_cr, + acc_cw, + acc_cost, + l_in, + l_out, + l_total, + l_cr, + l_cw, + l_cost, + ) = row; + input += larger(acc_in, l_in); + output += larger(acc_out, l_out); + total += larger(acc_total, l_total); + cache_read += larger(acc_cr, l_cr); + cache_write += larger(acc_cw, l_cw); + if acc_cost.is_some() || l_cost.is_some() { + let c = acc_cost.unwrap_or(0.0).max(l_cost.unwrap_or(0.0)); + cost = Some(cost.unwrap_or(0.0) + c); + } + } + + let opt = |v: i64| Some(i32::try_from(v).unwrap_or(i32::MAX)); + Ok(SessionUsageTotals { + accumulated_usage: Usage::new(opt(input), opt(output), opt(total)) + .with_cache_tokens(opt(cache_read), opt(cache_write)), + accumulated_cost: cost, + }) + } + async fn export_session(&self, id: &str) -> Result { let session = self.get_session(id, true).await?; serde_json::to_string_pretty(&session).map_err(Into::into) @@ -2189,7 +2545,7 @@ mod tests { use super::*; use crate::conversation::message::{Message, MessageContent}; use crate::providers::base::MessageStream; - use goose_providers::conversation::token_usage::ProviderUsage; + use goose_providers::conversation::token_usage::{CostSource, ProviderUsage}; use goose_providers::errors::ProviderError; use rmcp::model::Tool; use tempfile::TempDir; @@ -3619,4 +3975,271 @@ mod tests { assert_eq!(loaded.usage, usage); assert_eq!(loaded.accumulated_usage, accumulated_usage); } + + fn message_usage(input: i32, output: i32, cost: f64, is_compaction: bool) -> MessageUsage { + MessageUsage { + input_tokens: Some(input), + output_tokens: Some(output), + total_tokens: Some(input + output), + cost: Some(cost), + cost_source: Some(CostSource::Estimated), + is_compaction, + ..Default::default() + } + } + + async fn new_session(sm: &SessionManager) -> String { + sm.create_session( + PathBuf::from("/tmp"), + "s".to_string(), + SessionType::User, + GooseMode::default(), + ) + .await + .unwrap() + .id + } + + async fn seed_ledger( + sm: &SessionManager, + session_id: &str, + usage: &MessageUsage, + ) -> Result<()> { + let pool = sm.storage().pool().await?; + let mut tx = pool.begin().await?; + insert_usage_ledger_row(&mut tx, session_id, None, usage).await?; + tx.commit().await?; + Ok(()) + } + + #[tokio::test] + async fn test_usage_totals_include_subagent_tree() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let parent = new_session(&sm).await; + let child = new_session(&sm).await; + sm.update(&child) + .parent_session_id(Some(parent.clone())) + .apply() + .await + .unwrap(); + + seed_ledger(&sm, &parent, &message_usage(100, 20, 0.10, false)) + .await + .unwrap(); + seed_ledger(&sm, &child, &message_usage(40, 8, 0.04, false)) + .await + .unwrap(); + + let parent_totals = sm.get_session_usage_totals(&parent).await.unwrap(); + assert_eq!(parent_totals.accumulated_usage.input_tokens, Some(140)); + assert!((parent_totals.accumulated_cost.unwrap() - 0.14).abs() < 1e-9); + + let child_totals = sm.get_session_usage_totals(&child).await.unwrap(); + assert_eq!(child_totals.accumulated_usage.input_tokens, Some(40)); + assert!((child_totals.accumulated_cost.unwrap() - 0.04).abs() < 1e-9); + } + + #[tokio::test] + async fn test_ledger_reconciles_spend_recorded_on_pre_v15_builds() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let id = new_session(&sm).await; + + sm.update(&id) + .accumulated_usage(Usage::new(Some(5000), Some(1000), Some(6000))) + .accumulated_cost(Some(5.0)) + .apply() + .await + .unwrap(); + + sm.record_usage_metrics( + &id, + None, + Usage::new(Some(100), Some(20), Some(120)), + "test-model", + &message_usage(100, 20, 0.01, false), + ) + .await + .unwrap(); + + let totals = sm.get_session_usage_totals(&id).await.unwrap(); + assert_eq!(totals.accumulated_usage.total_tokens, Some(6120)); + assert!((totals.accumulated_cost.unwrap() - 5.01).abs() < 1e-9); + + let session = sm.get_session(&id, false).await.unwrap(); + sm.update(&id) + .accumulated_usage( + session.accumulated_usage + Usage::new(Some(500), Some(50), Some(550)), + ) + .accumulated_cost(Some(session.accumulated_cost.unwrap() + 0.50)) + .apply() + .await + .unwrap(); + + sm.record_usage_metrics( + &id, + None, + Usage::new(Some(30), Some(5), Some(35)), + "test-model", + &message_usage(30, 5, 0.03, false), + ) + .await + .unwrap(); + + let totals = sm.get_session_usage_totals(&id).await.unwrap(); + assert_eq!(totals.accumulated_usage.input_tokens, Some(5630)); + assert_eq!(totals.accumulated_usage.output_tokens, Some(1075)); + assert_eq!(totals.accumulated_usage.total_tokens, Some(6705)); + assert!((totals.accumulated_cost.unwrap() - 5.54).abs() < 1e-9); + + let session = sm.get_session(&id, false).await.unwrap(); + assert_eq!(session.accumulated_usage, totals.accumulated_usage); + assert!((session.accumulated_cost.unwrap() - 5.54).abs() < 1e-9); + } + + #[tokio::test] + async fn test_usage_totals_read_through_unreconciled_drift() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let id = new_session(&sm).await; + + sm.record_usage_metrics( + &id, + None, + Usage::new(Some(100), Some(20), Some(120)), + "test-model", + &message_usage(100, 20, 0.10, false), + ) + .await + .unwrap(); + + let session = sm.get_session(&id, false).await.unwrap(); + sm.update(&id) + .accumulated_usage( + session.accumulated_usage + Usage::new(Some(500), Some(50), Some(550)), + ) + .accumulated_cost(Some(session.accumulated_cost.unwrap() + 0.50)) + .apply() + .await + .unwrap(); + + let totals = sm.get_session_usage_totals(&id).await.unwrap(); + assert_eq!(totals.accumulated_usage.input_tokens, Some(600)); + assert_eq!(totals.accumulated_usage.total_tokens, Some(670)); + assert!((totals.accumulated_cost.unwrap() - 0.60).abs() < 1e-9); + } + + #[tokio::test] + async fn test_usage_totals_fall_back_to_accumulated_for_legacy_sessions() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let id = new_session(&sm).await; + + sm.update(&id) + .accumulated_usage(Usage::new(Some(500), Some(100), Some(600))) + .accumulated_cost(Some(0.42)) + .apply() + .await + .unwrap(); + + let totals = sm.get_session_usage_totals(&id).await.unwrap(); + assert_eq!(totals.accumulated_usage.input_tokens, Some(500)); + assert_eq!(totals.accumulated_cost, Some(0.42)); + } + + #[tokio::test] + async fn test_usage_ledger_survives_conversation_replace() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let id = new_session(&sm).await; + + seed_ledger(&sm, &id, &message_usage(1000, 200, 1.0, false)) + .await + .unwrap(); + seed_ledger(&sm, &id, &message_usage(50, 10, 0.05, true)) + .await + .unwrap(); + + sm.replace_conversation(&id, &Conversation::default()) + .await + .unwrap(); + + let totals = sm.get_session_usage_totals(&id).await.unwrap(); + assert_eq!(totals.accumulated_usage.total_tokens, Some(1260)); + assert!((totals.accumulated_cost.unwrap() - 1.05).abs() < 1e-9); + } + + #[tokio::test] + async fn test_usage_totals_mixed_legacy_and_ledger_tree() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let parent = new_session(&sm).await; + let child = new_session(&sm).await; + sm.update(&child) + .parent_session_id(Some(parent.clone())) + .apply() + .await + .unwrap(); + + seed_ledger(&sm, &parent, &message_usage(100, 20, 0.10, false)) + .await + .unwrap(); + sm.update(&child) + .accumulated_usage(Usage::new(Some(300), Some(60), Some(360))) + .accumulated_cost(Some(0.25)) + .apply() + .await + .unwrap(); + + let totals = sm.get_session_usage_totals(&parent).await.unwrap(); + assert_eq!(totals.accumulated_usage.input_tokens, Some(400)); + assert_eq!(totals.accumulated_usage.output_tokens, Some(80)); + assert!((totals.accumulated_cost.unwrap() - 0.35).abs() < 1e-9); + } + + #[tokio::test] + async fn test_delete_session_with_ledger_rows() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let id = new_session(&sm).await; + + seed_ledger(&sm, &id, &message_usage(100, 20, 0.10, false)) + .await + .unwrap(); + + sm.delete_session(&id).await.unwrap(); + assert!(sm.get_session(&id, false).await.is_err()); + } + + #[tokio::test] + async fn test_pre_v15_delete_cascades_ledger_rows() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let id = new_session(&sm).await; + + seed_ledger(&sm, &id, &message_usage(100, 20, 0.10, false)) + .await + .unwrap(); + + let pool = sm.storage().pool().await.unwrap(); + sqlx::query("DELETE FROM messages WHERE session_id = ?") + .bind(&id) + .execute(pool) + .await + .unwrap(); + sqlx::query("DELETE FROM sessions WHERE id = ?") + .bind(&id) + .execute(pool) + .await + .unwrap(); + + let remaining: i64 = + sqlx::query_scalar("SELECT COUNT(*) FROM usage_ledger WHERE session_id = ?") + .bind(&id) + .fetch_one(pool) + .await + .unwrap(); + assert_eq!(remaining, 0); + } } diff --git a/ui/desktop/openapi.json b/ui/desktop/openapi.json index 943efbbab..413878159 100644 --- a/ui/desktop/openapi.json +++ b/ui/desktop/openapi.json @@ -3361,6 +3361,14 @@ "$ref": "#/components/schemas/Message" } }, + "CostSource": { + "type": "string", + "description": "How the `cost` on a usage record was determined.", + "enum": [ + "provider_reported", + "estimated" + ] + }, "CreateCustomProviderResponse": { "type": "object", "required": [ @@ -5116,12 +5124,81 @@ "type": "boolean", "description": "Whether this message is a steer injected into an active run. UI-only:\nsurfaced as `_meta.goose.steer` so clients can mark the steer boundary\nwithout matching user-visible text. Never sent to providers." }, + "usage": { + "allOf": [ + { + "$ref": "#/components/schemas/MessageUsage" + } + ], + "nullable": true + }, "userVisible": { "type": "boolean", "description": "Whether the message should be visible to the user in the UI" } } }, + "MessageUsage": { + "type": "object", + "description": "Token usage and cost of a single provider call, attached to the turn's\nassistant message for display and recorded to the session's usage ledger.", + "properties": { + "cacheReadTokens": { + "type": "integer", + "format": "int32", + "nullable": true + }, + "cacheWriteTokens": { + "type": "integer", + "format": "int32", + "nullable": true + }, + "cost": { + "type": "number", + "format": "double", + "nullable": true + }, + "costSource": { + "allOf": [ + { + "$ref": "#/components/schemas/CostSource" + } + ], + "nullable": true + }, + "elapsedMs": { + "type": "integer", + "format": "int64", + "description": "Wall-clock generation time, used by the client for a tokens/sec readout.", + "nullable": true, + "minimum": 0 + }, + "inputTokens": { + "type": "integer", + "format": "int32", + "nullable": true + }, + "isCompaction": { + "type": "boolean", + "description": "Usage from a compaction/summarization call rather than a normal turn.\nAggregation counts it; the client can badge it." + }, + "outputTokens": { + "type": "integer", + "format": "int32", + "nullable": true + }, + "timeToFirstTokenMs": { + "type": "integer", + "format": "int64", + "nullable": true, + "minimum": 0 + }, + "totalTokens": { + "type": "integer", + "format": "int32", + "nullable": true + } + } + }, "ModelCapabilities": { "type": "object", "required": [ @@ -6488,6 +6565,11 @@ "name": { "type": "string" }, + "parent_session_id": { + "type": "string", + "description": "For sub-agent sessions, the session that spawned them.", + "nullable": true + }, "project_id": { "type": "string", "nullable": true