diff --git a/crates/goose-server/src/routes/reply.rs b/crates/goose-server/src/routes/reply.rs index 09fc5b8b..68d0b4ca 100644 --- a/crates/goose-server/src/routes/reply.rs +++ b/crates/goose-server/src/routes/reply.rs @@ -164,6 +164,7 @@ pub async fn get_token_state(session_manager: &SessionManager, session_id: &str) accumulated_input_tokens: session.accumulated_input_tokens.unwrap_or(0), accumulated_output_tokens: session.accumulated_output_tokens.unwrap_or(0), accumulated_total_tokens: session.accumulated_total_tokens.unwrap_or(0), + accumulated_cost: session.accumulated_cost, }) .inspect_err(|e| { tracing::warn!( diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 6cda8fed..5597d52f 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -4293,6 +4293,7 @@ print(\"hello, world\") accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens, + accumulated_cost: None, schedule_id: None, recipe: None, user_recipe_values: None, diff --git a/crates/goose/src/agents/reply_parts.rs b/crates/goose/src/agents/reply_parts.rs index f4939bc9..89e7535a 100644 --- a/crates/goose/src/agents/reply_parts.rs +++ b/crates/goose/src/agents/reply_parts.rs @@ -527,6 +527,12 @@ impl Agent { let accumulated_output = accumulate(session.accumulated_output_tokens, usage.usage.output_tokens); + let accumulated_cost = session + .provider_name + .as_deref() + .and_then(|pn| self.accumulate_cost(session.accumulated_cost, usage, pn)) + .or(session.accumulated_cost); + let (current_total, current_input, current_output) = if is_compaction_usage { // After compaction: summary output becomes new input context let new_input = usage.usage.output_tokens; @@ -548,11 +554,32 @@ impl Agent { .accumulated_total_tokens(accumulated_total) .accumulated_input_tokens(accumulated_input) .accumulated_output_tokens(accumulated_output) + .accumulated_cost(accumulated_cost) .apply() .await?; Ok(()) } + + fn accumulate_cost( + &self, + existing: Option, + usage: &ProviderUsage, + provider_name: &str, + ) -> Option { + let canonical = + crate::providers::canonical::maybe_get_canonical_model(provider_name, &usage.model)?; + + let input_price = canonical.cost.input?; + let output_price = canonical.cost.output?; + + let input_tokens = usage.usage.input_tokens.unwrap_or(0) as f64; + let output_tokens = usage.usage.output_tokens.unwrap_or(0) as f64; + + let chunk_cost = (input_tokens * input_price + output_tokens * output_price) / 1_000_000.0; + + Some(existing.unwrap_or(0.0) + chunk_cost) + } } /// Check whether a tool should be callable by an app based on MCP Apps visibility metadata. diff --git a/crates/goose/src/conversation/message.rs b/crates/goose/src/conversation/message.rs index 9f86d3af..180012f6 100644 --- a/crates/goose/src/conversation/message.rs +++ b/crates/goose/src/conversation/message.rs @@ -1026,6 +1026,7 @@ pub struct TokenState { pub accumulated_input_tokens: i32, pub accumulated_output_tokens: i32, pub accumulated_total_tokens: i32, + pub accumulated_cost: Option, } #[cfg(test)] diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index 37b27266..954a561f 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -19,7 +19,7 @@ use std::sync::{Arc, LazyLock}; use tracing::{info, warn}; use utoipa::ToSchema; -pub const CURRENT_SCHEMA_VERSION: i32 = 12; +pub const CURRENT_SCHEMA_VERSION: i32 = 13; pub const SESSIONS_FOLDER: &str = "sessions"; pub const DB_NAME: &str = "sessions.db"; @@ -72,6 +72,7 @@ pub struct Session { pub accumulated_total_tokens: Option, pub accumulated_input_tokens: Option, pub accumulated_output_tokens: Option, + pub accumulated_cost: Option, pub schedule_id: Option, pub recipe: Option, pub user_recipe_values: Option>, @@ -101,6 +102,7 @@ pub struct SessionUpdateBuilder<'a> { accumulated_total_tokens: Option>, accumulated_input_tokens: Option>, accumulated_output_tokens: Option>, + accumulated_cost: Option>, schedule_id: Option>, recipe: Option>, user_recipe_values: Option>>, @@ -135,6 +137,7 @@ impl<'a> SessionUpdateBuilder<'a> { accumulated_total_tokens: None, accumulated_input_tokens: None, accumulated_output_tokens: None, + accumulated_cost: None, schedule_id: None, recipe: None, user_recipe_values: None, @@ -213,6 +216,11 @@ impl<'a> SessionUpdateBuilder<'a> { self } + pub fn accumulated_cost(mut self, cost: Option) -> Self { + self.accumulated_cost = Some(cost); + self + } + pub fn schedule_id(mut self, schedule_id: Option) -> Self { self.schedule_id = Some(schedule_id); self @@ -490,6 +498,7 @@ impl Default for Session { accumulated_total_tokens: None, accumulated_input_tokens: None, accumulated_output_tokens: None, + accumulated_cost: None, schedule_id: None, recipe: None, user_recipe_values: None, @@ -557,6 +566,7 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session { accumulated_total_tokens: row.try_get("accumulated_total_tokens")?, accumulated_input_tokens: row.try_get("accumulated_input_tokens")?, accumulated_output_tokens: row.try_get("accumulated_output_tokens")?, + accumulated_cost: row.try_get("accumulated_cost").ok().flatten(), schedule_id: row.try_get("schedule_id")?, recipe, user_recipe_values, @@ -666,6 +676,7 @@ impl SessionStorage { accumulated_total_tokens INTEGER, accumulated_input_tokens INTEGER, accumulated_output_tokens INTEGER, + accumulated_cost REAL, schedule_id TEXT, recipe_json TEXT, user_recipe_values_json TEXT, @@ -786,9 +797,10 @@ impl SessionStorage { id, name, user_set_name, session_type, working_dir, created_at, updated_at, extension_data, total_tokens, input_tokens, output_tokens, accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens, + accumulated_cost, schedule_id, recipe_json, user_recipe_values_json, provider_name, model_config_json, goose_mode - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) "#, ) .bind(&session.id) @@ -805,6 +817,7 @@ impl SessionStorage { .bind(session.accumulated_total_tokens) .bind(session.accumulated_input_tokens) .bind(session.accumulated_output_tokens) + .bind(session.accumulated_cost) .bind(&session.schedule_id) .bind(recipe_json) .bind(user_recipe_values_json) @@ -1087,6 +1100,19 @@ impl SessionStorage { .await?; } } + 13 => { + let has_accumulated_cost = sqlx::query_scalar::<_, i32>( + "SELECT COUNT(*) FROM pragma_table_info('sessions') WHERE name = 'accumulated_cost'", + ) + .fetch_one(&mut **tx) + .await? + > 0; + if !has_accumulated_cost { + sqlx::query("ALTER TABLE sessions ADD COLUMN accumulated_cost REAL") + .execute(&mut **tx) + .await?; + } + } _ => { anyhow::bail!("Unknown migration version: {}", version); } @@ -1147,6 +1173,7 @@ impl SessionStorage { SELECT id, working_dir, name, description, user_set_name, session_type, created_at, updated_at, extension_data, total_tokens, input_tokens, output_tokens, accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens, + accumulated_cost, schedule_id, recipe_json, user_recipe_values_json, provider_name, model_config_json, goose_mode, archived_at, project_id @@ -1207,6 +1234,7 @@ impl SessionStorage { builder.accumulated_output_tokens, "accumulated_output_tokens" ); + add_update!(builder.accumulated_cost, "accumulated_cost"); add_update!(builder.schedule_id, "schedule_id"); add_update!(builder.recipe, "recipe_json"); add_update!(builder.user_recipe_values, "user_recipe_values_json"); @@ -1259,6 +1287,9 @@ impl SessionStorage { if let Some(aot) = builder.accumulated_output_tokens { q = q.bind(aot); } + if let Some(ac) = builder.accumulated_cost { + q = q.bind(ac); + } if let Some(sid) = builder.schedule_id { q = q.bind(sid); } @@ -1445,6 +1476,7 @@ impl SessionStorage { SELECT s.id, s.working_dir, s.name, s.description, s.user_set_name, s.session_type, s.created_at, s.updated_at, s.extension_data, s.total_tokens, s.input_tokens, s.output_tokens, s.accumulated_total_tokens, s.accumulated_input_tokens, s.accumulated_output_tokens, + 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, @@ -1564,6 +1596,7 @@ impl SessionStorage { .accumulated_total_tokens(import.accumulated_total_tokens) .accumulated_input_tokens(import.accumulated_input_tokens) .accumulated_output_tokens(import.accumulated_output_tokens) + .accumulated_cost(import.accumulated_cost) .schedule_id(import.schedule_id) .recipe(import.recipe) .user_recipe_values(import.user_recipe_values); diff --git a/ui/desktop/openapi.json b/ui/desktop/openapi.json index 3f0a16d6..9823f833 100644 --- a/ui/desktop/openapi.json +++ b/ui/desktop/openapi.json @@ -7868,6 +7868,11 @@ "message_count" ], "properties": { + "accumulated_cost": { + "type": "number", + "format": "double", + "nullable": true + }, "accumulated_input_tokens": { "type": "integer", "format": "int32", @@ -8567,6 +8572,11 @@ "accumulatedTotalTokens" ], "properties": { + "accumulatedCost": { + "type": "number", + "format": "double", + "nullable": true + }, "accumulatedInputTokens": { "type": "integer", "format": "int32" diff --git a/ui/desktop/src/api/types.gen.ts b/ui/desktop/src/api/types.gen.ts index 5f410551..691d9e67 100644 --- a/ui/desktop/src/api/types.gen.ts +++ b/ui/desktop/src/api/types.gen.ts @@ -1272,6 +1272,7 @@ export type ScheduledJob = { }; export type Session = { + accumulated_cost?: number | null; accumulated_input_tokens?: number | null; accumulated_output_tokens?: number | null; accumulated_total_tokens?: number | null; @@ -1480,6 +1481,7 @@ export type ThinkingContent = { }; export type TokenState = { + accumulatedCost?: number | null; accumulatedInputTokens: number; accumulatedOutputTokens: number; accumulatedTotalTokens: number; diff --git a/ui/desktop/src/components/BaseChat.tsx b/ui/desktop/src/components/BaseChat.tsx index 23e80ba8..a72f5984 100644 --- a/ui/desktop/src/components/BaseChat.tsx +++ b/ui/desktop/src/components/BaseChat.tsx @@ -28,7 +28,6 @@ import { RecipeHeader } from './RecipeHeader'; import { RecipeWarningModal } from './ui/RecipeWarningModal'; import { scanRecipe } from '../recipe'; import { UserInput } from '../types/message'; -import { useCostTracking } from '../hooks/useCostTracking'; import RecipeActivities from './recipes/RecipeActivities'; import { useToolCount } from './alerts/useToolCount'; import { getThinkingMessage, getTextAndImageContent } from '../types/message'; @@ -196,14 +195,6 @@ export default function BaseChat({ handleSubmit(input); }; - const { sessionCosts } = useCostTracking({ - sessionInputTokens: session?.accumulated_input_tokens || 0, - sessionOutputTokens: session?.accumulated_output_tokens || 0, - localInputTokens: 0, - localOutputTokens: 0, - session, - }); - const sessionModel = session?.model_config?.model_name ?? null; const sessionProvider = session?.provider_name ?? null; const sessionLoaded = session !== undefined; @@ -511,11 +502,14 @@ export default function BaseChat({ accumulatedOutputTokens={ tokenState?.accumulatedOutputTokens ?? session?.accumulated_output_tokens ?? undefined } + accumulatedCost={ + tokenState?.accumulatedCost ?? session?.accumulated_cost ?? undefined + } droppedFiles={droppedFiles} onFilesProcessed={() => setDroppedFiles([])} // Clear dropped files after processing messages={messages} disableAnimation={disableAnimation} - sessionCosts={sessionCosts} + recipe={recipe} recipeAccepted={!hasNotAcceptedRecipe} initialPrompt={initialPrompt} diff --git a/ui/desktop/src/components/ChatInput.tsx b/ui/desktop/src/components/ChatInput.tsx index 2878cb05..1c450db9 100644 --- a/ui/desktop/src/components/ChatInput.tsx +++ b/ui/desktop/src/components/ChatInput.tsx @@ -167,14 +167,8 @@ interface ChatInputProps { totalTokens?: number; accumulatedInputTokens?: number; accumulatedOutputTokens?: number; + accumulatedCost?: number | null; messages?: Message[]; - sessionCosts?: { - [key: string]: { - inputTokens: number; - outputTokens: number; - totalCost: number; - }; - }; disableAnimation?: boolean; recipe?: Recipe | null; recipeId?: string | null; @@ -203,9 +197,9 @@ export default function ChatInput({ totalTokens, accumulatedInputTokens, accumulatedOutputTokens, + accumulatedCost, messages = [], disableAnimation = false, - sessionCosts, recipe, recipeId, recipeAccepted, @@ -1690,7 +1684,7 @@ export default function ChatInput({ diff --git a/ui/desktop/src/components/Hub.tsx b/ui/desktop/src/components/Hub.tsx index b7e971d1..3cdeb21a 100644 --- a/ui/desktop/src/components/Hub.tsx +++ b/ui/desktop/src/components/Hub.tsx @@ -107,7 +107,6 @@ export default function Hub({ onFilesProcessed={() => {}} messages={[]} disableAnimation={false} - sessionCosts={undefined} toolCount={0} onWorkingDirChange={setWorkingDir} inputRef={inputRef} diff --git a/ui/desktop/src/components/bottom_menu/CostTracker.tsx b/ui/desktop/src/components/bottom_menu/CostTracker.tsx index e64ab96f..1ade649f 100644 --- a/ui/desktop/src/components/bottom_menu/CostTracker.tsx +++ b/ui/desktop/src/components/bottom_menu/CostTracker.tsx @@ -14,10 +14,6 @@ const i18n = defineMessages({ id: 'costTracker.costUnavailable', defaultMessage: 'Cost data not available for {model} ({inputTokens} input, {outputTokens} output tokens)', }, - sessionCostBreakdown: { - id: 'costTracker.sessionCostBreakdown', - defaultMessage: 'Session cost breakdown:', - }, totalSessionCost: { id: 'costTracker.totalSessionCost', defaultMessage: 'Total session cost: {cost}', @@ -31,13 +27,7 @@ const i18n = defineMessages({ interface CostTrackerProps { inputTokens?: number; outputTokens?: number; - sessionCosts?: { - [key: string]: { - inputTokens: number; - outputTokens: number; - totalCost: number; - }; - }; + accumulatedCost?: number | null; model: string | null; provider: string | null; } @@ -45,7 +35,7 @@ interface CostTrackerProps { export function CostTracker({ inputTokens = 0, outputTokens = 0, - sessionCosts, + accumulatedCost, model: currentModel, provider: currentProvider, }: CostTrackerProps) { @@ -106,41 +96,7 @@ export function CostTracker({ } const calculateCost = (): number => { - // If we have session costs, calculate the total across all models - if (sessionCosts) { - let totalCost = 0; - - // Add up all historical costs from different models - Object.values(sessionCosts).forEach((modelCost) => { - totalCost += modelCost.totalCost; - }); - - // Add current model cost if we have pricing info - if ( - costInfo && - (costInfo.input_token_cost !== undefined || costInfo.output_token_cost !== undefined) - ) { - const currentInputCost = (inputTokens * (costInfo.input_token_cost || 0)) / 1_000_000; - const currentOutputCost = (outputTokens * (costInfo.output_token_cost || 0)) / 1_000_000; - totalCost += currentInputCost + currentOutputCost; - } - - return totalCost; - } - - // Fallback to simple calculation for current model only - if ( - !costInfo || - (costInfo.input_token_cost === undefined && costInfo.output_token_cost === undefined) - ) { - return 0; - } - - const inputCost = (inputTokens * (costInfo.input_token_cost || 0)) / 1_000_000; - const outputCost = (outputTokens * (costInfo.output_token_cost || 0)) / 1_000_000; - const total = inputCost + outputCost; - - return total; + return accumulatedCost ?? 0; }; const formatCost = (cost: number): string => { @@ -165,10 +121,10 @@ export function CostTracker({ ); } - // If no cost info found, try to return a default if ( - !costInfo || - (costInfo.input_token_cost === undefined && costInfo.output_token_cost === undefined) + accumulatedCost == null && + (!costInfo || + (costInfo.input_token_cost === undefined && costInfo.output_token_cost === undefined)) ) { const freeProviders = ['ollama', 'local', 'localhost']; if (freeProviders.includes(currentProvider.toLowerCase())) { @@ -216,38 +172,22 @@ export function CostTracker({ // Build tooltip content const getTooltipContent = (): string => { - // Handle error states first if (pricingFailed) { return intl.formatMessage(i18n.pricingUnavailable, { model: `${currentProvider}/${currentModel}` }); } - // Handle session costs - if (sessionCosts && Object.keys(sessionCosts).length > 0) { - // Show session breakdown - let tooltip = intl.formatMessage(i18n.sessionCostBreakdown) + '\n'; + const currency = costInfo?.currency || '$'; - Object.entries(sessionCosts).forEach(([modelKey, cost]) => { - const costStr = `${costInfo?.currency || '$'}${cost.totalCost.toFixed(6)}`; - tooltip += `${modelKey}: ${costStr} (${cost.inputTokens.toLocaleString()} in, ${cost.outputTokens.toLocaleString()} out)\n`; - }); - - // Add current model if it has costs - if (costInfo && (inputTokens > 0 || outputTokens > 0)) { - const currentCost = - (inputTokens * (costInfo.input_token_cost || 0) + - outputTokens * (costInfo.output_token_cost || 0)) / - 1_000_000; - if (currentCost > 0) { - tooltip += `${currentProvider}/${currentModel} (current): ${costInfo.currency || '$'}${currentCost.toFixed(6)} (${inputTokens.toLocaleString()} in, ${outputTokens.toLocaleString()} out)\n`; - } - } - - tooltip += '\n' + intl.formatMessage(i18n.totalSessionCost, { cost: `${costInfo?.currency || '$'}${totalCost.toFixed(6)}` }); - return tooltip; + if (accumulatedCost != null) { + return intl.formatMessage(i18n.totalSessionCost, { cost: `${currency}${totalCost.toFixed(4)}` }) + + `\n` + intl.formatMessage(i18n.inputOutputTooltip, { + inputTokens: inputTokens.toLocaleString(), + inputCost: `${currency}${((inputTokens * (costInfo?.input_token_cost || 0)) / 1_000_000).toFixed(6)}`, + outputTokens: outputTokens.toLocaleString(), + outputCost: `${currency}${((outputTokens * (costInfo?.output_token_cost || 0)) / 1_000_000).toFixed(6)}`, + }); } - // Default tooltip for single model - const currency = costInfo?.currency || '$'; const inputCostStr = `${currency}${((inputTokens * (costInfo?.input_token_cost || 0)) / 1_000_000).toFixed(6)}`; const outputCostStr = `${currency}${((outputTokens * (costInfo?.output_token_cost || 0)) / 1_000_000).toFixed(6)}`; return intl.formatMessage(i18n.inputOutputTooltip, { diff --git a/ui/desktop/src/hooks/useCostTracking.ts b/ui/desktop/src/hooks/useCostTracking.ts deleted file mode 100644 index 8b6ad058..00000000 --- a/ui/desktop/src/hooks/useCostTracking.ts +++ /dev/null @@ -1,92 +0,0 @@ -import { useEffect, useRef, useState } from 'react'; -import { fetchCanonicalModelInfo } from '../utils/canonical'; -import { Session } from '../api'; - -interface UseCostTrackingProps { - sessionInputTokens: number; - sessionOutputTokens: number; - localInputTokens: number; - localOutputTokens: number; - session?: Session | null; -} - -export const useCostTracking = ({ - sessionInputTokens, - sessionOutputTokens, - localInputTokens, - localOutputTokens, - session, -}: UseCostTrackingProps) => { - const [sessionCosts, setSessionCosts] = useState<{ - [key: string]: { - inputTokens: number; - outputTokens: number; - totalCost: number; - }; - }>({}); - - const currentModel = session?.model_config?.model_name ?? undefined; - const currentProvider = session?.provider_name ?? undefined; - const prevModelRef = useRef(undefined); - const prevProviderRef = useRef(undefined); - - // Handle model changes and accumulate costs - useEffect(() => { - if (!currentModel || !currentProvider) return; - - const handleModelChange = async () => { - if ( - prevModelRef.current !== undefined && - prevProviderRef.current !== undefined && - (prevModelRef.current !== currentModel || prevProviderRef.current !== currentProvider) - ) { - // Model/provider has changed, save the costs for the previous model - const prevKey = `${prevProviderRef.current}/${prevModelRef.current}`; - - // Get pricing info for the previous model - const prevCostInfo = await fetchCanonicalModelInfo( - prevProviderRef.current, - prevModelRef.current - ); - - if (prevCostInfo) { - const prevInputCost = - ((sessionInputTokens || localInputTokens) * (prevCostInfo.input_token_cost || 0)) / - 1_000_000; - const prevOutputCost = - ((sessionOutputTokens || localOutputTokens) * (prevCostInfo.output_token_cost || 0)) / - 1_000_000; - const prevTotalCost = prevInputCost + prevOutputCost; - - // Save the accumulated costs for this model - setSessionCosts((prev) => ({ - ...prev, - [prevKey]: { - inputTokens: sessionInputTokens || localInputTokens, - outputTokens: sessionOutputTokens || localOutputTokens, - totalCost: prevTotalCost, - }, - })); - } - - } - - prevModelRef.current = currentModel || undefined; - prevProviderRef.current = currentProvider || undefined; - }; - - handleModelChange(); - }, [ - currentModel, - currentProvider, - sessionInputTokens, - sessionOutputTokens, - localInputTokens, - localOutputTokens, - session, - ]); - - return { - sessionCosts, - }; -}; diff --git a/ui/desktop/src/i18n/messages/en.json b/ui/desktop/src/i18n/messages/en.json index faada7a3..c01ae221 100644 --- a/ui/desktop/src/i18n/messages/en.json +++ b/ui/desktop/src/i18n/messages/en.json @@ -314,9 +314,7 @@ "costTracker.pricingUnavailable": { "defaultMessage": "Pricing data unavailable for {model}" }, - "costTracker.sessionCostBreakdown": { - "defaultMessage": "Session cost breakdown:" - }, + "costTracker.totalSessionCost": { "defaultMessage": "Total session cost: {cost}" },