feat: per-message usage/cost tracking with derived session totals (#10172)

Co-authored-by: Douwe M Osinga <douwe@sidewalklabs.com>
This commit is contained in:
filip
2026-07-05 13:39:10 -07:00
committed by GitHub
parent 1df5f60735
commit 9b837f1a46
21 changed files with 1025 additions and 120 deletions
@@ -43,6 +43,10 @@ impl Conversation {
&self.0
}
pub fn messages_mut(&mut self) -> &mut Vec<Message> {
&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
@@ -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<String>,
}
#[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<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output_tokens: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub total_tokens: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read_tokens: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_write_tokens: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cost: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cost_source: Option<CostSource>,
#[serde(skip_serializing_if = "Option::is_none")]
pub elapsed_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub time_to_first_token_ms: Option<u64>,
#[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<Box<MessageUsage>>,
}
impl Default for MessageMetadata {
@@ -690,6 +736,7 @@ impl Default for MessageMetadata {
agent_visible: true,
inference: None,
steer: false,
usage: None,
}
}
}
@@ -9,6 +9,17 @@ pub struct ProviderUsage {
pub usage: Usage,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stats: Option<ProviderStats>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cost: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cost_source: Option<CostSource>,
}
#[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
}
}
@@ -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<Usage> {
}
}
fn provider_usage_with_cost(
model: String,
usage: Usage,
data: &Value,
fallback_cost: Option<f64>,
) -> 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
@@ -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<f64> {
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())
}
+13 -5
View File
@@ -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(),
@@ -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(),
+6 -3
View File
@@ -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,
+17 -12
View File
@@ -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(
+3 -1
View File
@@ -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<Client>,
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)?;
}
+40 -8
View File
@@ -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<UsageUpdates> {
pub(super) fn build_usage_updates(
session: &Session,
totals: &SessionUsageTotals,
) -> Option<UsageUpdates> {
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<UsageUpdates> {
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<Client>,
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]
+1 -5
View File
@@ -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)
}
}
+1 -1
View File
@@ -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 {
+1 -5
View File
@@ -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)
}
+22 -4
View File
@@ -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<String> = None;
let mut pending_turn_usage: Option<ProviderUsage> = 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?;
}
@@ -483,6 +483,36 @@ impl SummonClient {
})
}
async fn create_subagent_session(
&self,
task_config: &TaskConfig,
name: String,
) -> Result<crate::session::Session, String> {
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<ServerNotification>,
subscribers: Arc<Mutex<Vec<mpsc::Sender<ServerNotification>>>>,
@@ -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::<ServerNotification>();
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();
+30 -26
View File
@@ -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<String>,
usage: &ProviderUsage,
is_compaction_usage: bool,
) -> Result<()> {
) -> Result<ProviderUsage> {
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<f64>,
usage: &ProviderUsage,
provider_name: &str,
) -> Option<f64> {
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<f64>, Option<CostSource>) {
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),
}
}
}
+1 -1
View File
@@ -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};
+1
View File
@@ -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)?;
+629 -6
View File
@@ -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<String>,
#[serde(default)]
pub parent_session_id: Option<String>,
#[serde(default)]
pub last_message_snippet: Option<String>,
}
@@ -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<Option<DateTime<Utc>>>,
project_id: Option<Option<String>>,
parent_session_id: Option<Option<String>>,
}
#[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<f64>,
}
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<String>) -> 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<SessionUsageTotals> {
self.storage.get_session_usage_totals(id).await
}
pub async fn record_usage_metrics(
&self,
session_id: &str,
schedule_id: Option<String>,
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<String> {
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<Sqlite> {
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<String>,
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<SessionUsageTotals> {
let pool = self.pool().await?;
let rows = sqlx::query_as::<
_,
(
Option<i64>,
Option<i64>,
Option<i64>,
Option<i64>,
Option<i64>,
Option<f64>,
Option<i64>,
Option<i64>,
Option<i64>,
Option<i64>,
Option<i64>,
Option<f64>,
),
>(
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<f64> = None;
let larger =
|acc: Option<i64>, ledger: Option<i64>| 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<String> {
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);
}
}
+82
View File
@@ -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