From fb4523539d09cdb21ed2a2a92c4c4f0c5386c21d Mon Sep 17 00:00:00 2001 From: "open-swe[bot]" <215916821+open-swe[bot]@users.noreply.github.com> Date: Sun, 27 Jul 2025 23:56:50 +0000 Subject: [PATCH] fix: update tokenData structure to support multiple LLMs (#555) * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * Apply patch * fix backend * fix frontend * cr --------- Co-authored-by: open-swe[bot] Co-authored-by: bracesproul --- .../planner/nodes/determine-needs-context.ts | 5 +- .../planner/nodes/generate-message/index.ts | 5 +- .../planner/nodes/generate-plan/index.ts | 5 +- .../src/graphs/planner/nodes/notetaker.ts | 5 +- .../programmer/nodes/generate-conclusion.ts | 5 +- .../nodes/generate-message/index.ts | 5 +- .../src/graphs/programmer/nodes/open-pr.ts | 5 +- .../programmer/nodes/summarize-history.ts | 5 +- .../graphs/programmer/nodes/update-plan.ts | 5 +- .../src/graphs/reviewer/nodes/final-review.ts | 5 +- .../nodes/generate-review-actions/index.ts | 5 +- .../src/graphs/shared/diagnose-error.ts | 9 +- apps/open-swe/src/utils/caching.ts | 15 +- apps/open-swe/src/utils/llms/index.ts | 1 + apps/open-swe/src/utils/llms/model-manager.ts | 8 + apps/web/src/components/v2/thread-view.tsx | 10 +- apps/web/src/components/v2/token-usage.tsx | 279 +++++++++++++++++- packages/shared/jest.config.js | 19 ++ packages/shared/package.json | 9 +- packages/shared/src/__tests__/caching.test.ts | 153 ++++++++++ packages/shared/src/caching.ts | 47 ++- packages/shared/src/open-swe/planner/types.ts | 6 +- .../shared/src/open-swe/reviewer/types.ts | 6 +- packages/shared/src/open-swe/types.ts | 12 +- yarn.lock | 4 + 25 files changed, 581 insertions(+), 52 deletions(-) create mode 100644 packages/shared/jest.config.js create mode 100644 packages/shared/src/__tests__/caching.test.ts diff --git a/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts b/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts index 24cc2296..c9f6fdc8 100644 --- a/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts +++ b/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts @@ -17,6 +17,7 @@ import { getMessageContentString } from "@open-swe/shared/messages"; import { filterHiddenMessages } from "../../../utils/message/filter-hidden.js"; import { createLogger, LogLevel } from "../../../utils/logger.js"; import { trackCachePerformance } from "../../../utils/caching.js"; +import { getModelManager } from "../../../utils/llms/model-manager.js"; const logger = createLogger(LogLevel.INFO, "DetermineNeedsContext"); @@ -120,6 +121,8 @@ export async function determineNeedsContext( getMissingMessages(state, config), loadModel(config, Task.ROUTER), ]); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.ROUTER); if (!missingMessages.length) { throw new Error( "Can not determine if more context is needed if there are no missing messages.", @@ -155,7 +158,7 @@ export async function determineNeedsContext( const commandUpdate: PlannerGraphUpdate = { messages: missingMessages, - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; const shouldGatherContext = diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts b/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts index 58a76023..c46ed52b 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts @@ -1,4 +1,5 @@ import { + getModelManager, loadModel, supportsParallelToolCallsParam, Task, @@ -70,6 +71,8 @@ export async function generateAction( config: GraphConfig, ): Promise { const model = await loadModel(config, Task.PROGRAMMER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.PROGRAMMER, @@ -145,6 +148,6 @@ export async function generateAction( return { messages: [...missingMessages, response], ...(latestTaskPlan && { taskPlan: latestTaskPlan }), - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts b/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts index 14a33888..ee9c81e0 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts @@ -23,6 +23,7 @@ import { getScratchpad } from "../../utils/scratchpad-notes.js"; import { SCRATCHPAD_PROMPT, SYSTEM_PROMPT } from "./prompt.js"; import { DO_NOT_RENDER_ID_PREFIX } from "@open-swe/shared/constants"; import { filterMessagesWithoutContent } from "../../../../utils/message/content.js"; +import { getModelManager } from "../../../../utils/llms/model-manager.js"; import { trackCachePerformance } from "../../../../utils/caching.js"; function formatSystemPrompt(state: PlannerGraphState): string { @@ -54,6 +55,8 @@ export async function generatePlan( config: GraphConfig, ): Promise { const model = await loadModel(config, Task.PROGRAMMER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.SUMMARIZER, @@ -125,6 +128,6 @@ export async function generatePlan( proposedPlanTitle: proposedPlanArgs.title, proposedPlan: proposedPlanArgs.plan, ...(newSessionId && { sandboxSessionId: newSessionId }), - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/planner/nodes/notetaker.ts b/apps/open-swe/src/graphs/planner/nodes/notetaker.ts index c89993ae..cd6f5711 100644 --- a/apps/open-swe/src/graphs/planner/nodes/notetaker.ts +++ b/apps/open-swe/src/graphs/planner/nodes/notetaker.ts @@ -18,6 +18,7 @@ import { ToolMessage } from "@langchain/core/messages"; import { DO_NOT_RENDER_ID_PREFIX } from "@open-swe/shared/constants"; import { createWriteTechnicalNotesToolFields } from "@open-swe/shared/open-swe/tools"; import { trackCachePerformance } from "../../../utils/caching.js"; +import { getModelManager } from "../../../utils/llms/model-manager.js"; const SCRATCHPAD_PROMPT = `You've also wrote technical notes to a scratchpad throughout the context gathering process. Ensure you include/incorporate these notes, or the highest quality parts of these notes in your conclusion notes. @@ -102,6 +103,8 @@ export async function notetaker( config: GraphConfig, ): Promise { const model = await loadModel(config, Task.SUMMARIZER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.SUMMARIZER, @@ -146,6 +149,6 @@ ${state.messages.map(getMessageString).join("\n")}`; contextGatheringNotes: ( toolCall.args as z.infer ).notes, - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/programmer/nodes/generate-conclusion.ts b/apps/open-swe/src/graphs/programmer/nodes/generate-conclusion.ts index 24a75c08..b2c01b2c 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/generate-conclusion.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/generate-conclusion.ts @@ -16,6 +16,7 @@ import { } from "@open-swe/shared/open-swe/tasks"; import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js"; import { trackCachePerformance } from "../../../utils/caching.js"; +import { getModelManager } from "../../../utils/llms/model-manager.js"; const logger = createLogger(LogLevel.INFO, "GenerateConclusionNode"); @@ -40,6 +41,8 @@ export async function generateConclusion( config: GraphConfig, ): Promise { const model = await loadModel(config, Task.SUMMARIZER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER); const userRequestPrompt = formatUserRequestPrompt(state.messages); const userMessage = `${userRequestPrompt} @@ -83,6 +86,6 @@ Given all of this, please respond with the concise conclusion. Do not include an messages: [response], internalMessages: [response], taskPlan: updatedTaskPlan, - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts b/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts index 442d8ffe..5f2aa788 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts @@ -5,6 +5,7 @@ import { GraphUpdate, } from "@open-swe/shared/open-swe/types"; import { + getModelManager, loadModel, supportsParallelToolCallsParam, Task, @@ -141,6 +142,8 @@ export async function generateAction( config: GraphConfig, ): Promise { const model = await loadModel(config, Task.PROGRAMMER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.PROGRAMMER, @@ -246,6 +249,6 @@ export async function generateAction( internalMessages: newMessagesList, ...(newSandboxSessionId && { sandboxSessionId: newSandboxSessionId }), ...(latestTaskPlan && { taskPlan: latestTaskPlan }), - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts index 10f0b6b6..9005af2c 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts @@ -36,6 +36,7 @@ import { import { getRepoAbsolutePath } from "@open-swe/shared/git"; import { createOpenPrToolFields } from "@open-swe/shared/open-swe/tools"; import { trackCachePerformance } from "../../../utils/caching.js"; +import { getModelManager } from "../../../utils/llms/model-manager.js"; import { GitHubPullRequest, GitHubPullRequestList, @@ -116,6 +117,8 @@ export async function openPullRequest( const openPrTool = createOpenPrToolFields(); // use the router model since this is a simple task that doesn't need an advanced model const model = await loadModel(config, Task.ROUTER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.ROUTER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.ROUTER, @@ -219,7 +222,7 @@ export async function openPullRequest( }), ...(codebaseTree && { codebaseTree }), ...(dependenciesInstalled !== null && { dependenciesInstalled }), - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), ...(updatedTaskPlan && { taskPlan: updatedTaskPlan }), }; } diff --git a/apps/open-swe/src/graphs/programmer/nodes/summarize-history.ts b/apps/open-swe/src/graphs/programmer/nodes/summarize-history.ts index ad6fb489..713ae33d 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/summarize-history.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/summarize-history.ts @@ -21,6 +21,7 @@ import { createConversationHistorySummaryToolFields } from "@open-swe/shared/ope import { formatUserRequestPrompt } from "../../../utils/user-request.js"; import { getMessagesSinceLastSummary } from "../../../utils/tokens.js"; import { trackCachePerformance } from "../../../utils/caching.js"; +import { getModelManager } from "../../../utils/llms/model-manager.js"; const SINGLE_USER_REQUEST_PROMPT = `Here is the user's request: @@ -148,6 +149,8 @@ export async function summarizeHistory( config: GraphConfig, ): Promise { const model = await loadModel(config, Task.SUMMARIZER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER); const plan = getActivePlanItems(state.taskPlan); const conversationHistoryToSummarize = await getMessagesSinceLastSummary( @@ -190,6 +193,6 @@ export async function summarizeHistory( return { messages: summaryMessages, internalMessages: newInternalMessages, - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts b/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts index 4016ea00..1a9685b4 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts @@ -27,6 +27,7 @@ import { createLogger, LogLevel } from "../../../utils/logger.js"; import { createUpdatePlanToolFields } from "@open-swe/shared/open-swe/tools"; import { formatCustomRulesPrompt } from "../../../utils/custom-rules.js"; import { trackCachePerformance } from "../../../utils/caching.js"; +import { getModelManager } from "../../../utils/llms/model-manager.js"; const logger = createLogger(LogLevel.INFO, "UpdatePlanNode"); @@ -124,6 +125,8 @@ export async function updatePlan( }); const model = await loadModel(config, Task.PROGRAMMER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.PROGRAMMER, @@ -204,6 +207,6 @@ export async function updatePlan( messages: [toolMessage], internalMessages: [toolMessage], taskPlan: newTaskPlan, - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts b/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts index 67952446..dc414d4c 100644 --- a/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts +++ b/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts @@ -30,6 +30,7 @@ import { ToolMessage, } from "@langchain/core/messages"; import { trackCachePerformance } from "../../../utils/caching.js"; +import { getModelManager } from "../../../utils/llms/model-manager.js"; import { createScratchpadTool } from "../../../tools/scratchpad.js"; const SYSTEM_PROMPT = `You are a code reviewer for a software engineer working on a large codebase. @@ -117,6 +118,8 @@ export async function finalReview( const incompleteTool = createCodeReviewMarkTaskNotCompleteFields(); const tools = [completedTool, incompleteTool]; const model = await loadModel(config, Task.PROGRAMMER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.PROGRAMMER, @@ -206,6 +209,6 @@ export async function finalReview( messages: messagesUpdate, internalMessages: messagesUpdate, reviewsCount: (state.reviewsCount || 0) + 1, - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts b/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts index 8524f8cb..496836cd 100644 --- a/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts +++ b/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts @@ -1,4 +1,5 @@ import { + getModelManager, loadModel, supportsParallelToolCallsParam, Task, @@ -112,6 +113,8 @@ export async function generateReviewActions( config: GraphConfig, ): Promise { const model = await loadModel(config, Task.PROGRAMMER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.PROGRAMMER, @@ -166,6 +169,6 @@ export async function generateReviewActions( return { messages: [response], reviewerMessages: [response], - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/graphs/shared/diagnose-error.ts b/apps/open-swe/src/graphs/shared/diagnose-error.ts index cf99d832..352a26a5 100644 --- a/apps/open-swe/src/graphs/shared/diagnose-error.ts +++ b/apps/open-swe/src/graphs/shared/diagnose-error.ts @@ -7,7 +7,7 @@ import { import { createDiagnoseErrorToolFields } from "@open-swe/shared/open-swe/tools"; import { z } from "zod"; -import { CacheMetrics, GraphConfig } from "@open-swe/shared/open-swe/types"; +import { ModelTokenData, GraphConfig } from "@open-swe/shared/open-swe/types"; import { createLogger, LogLevel } from "../../utils/logger.js"; import { getAllLastFailedActions } from "../../utils/tool-message-error.js"; import { getMessageString } from "../../utils/message/content.js"; @@ -17,6 +17,7 @@ import { Task, } from "../../utils/llms/index.js"; import { trackCachePerformance } from "../../utils/caching.js"; +import { getModelManager } from "../../utils/llms/model-manager.js"; const logger = createLogger(LogLevel.INFO, "SharedDiagnoseError"); @@ -76,7 +77,7 @@ const formatUserPrompt = (messages: BaseMessage[]): string => { interface DiagnoseErrorInputs { messages: BaseMessage[]; codebaseTree: string; - tokenData?: CacheMetrics; + tokenData?: ModelTokenData[]; } type DiagnoseErrorUpdate = Partial; @@ -95,6 +96,8 @@ export async function diagnoseError( logger.info("The last few tool calls resulted in errors. Diagnosing error."); const model = await loadModel(config, Task.SUMMARIZER); + const modelManager = getModelManager(); + const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER); const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( config, Task.SUMMARIZER, @@ -142,6 +145,6 @@ export async function diagnoseError( return { messages: [response, toolMessage], - tokenData: trackCachePerformance(response), + tokenData: trackCachePerformance(response, modelName), }; } diff --git a/apps/open-swe/src/utils/caching.ts b/apps/open-swe/src/utils/caching.ts index 46ff7de8..8a818a0f 100644 --- a/apps/open-swe/src/utils/caching.ts +++ b/apps/open-swe/src/utils/caching.ts @@ -9,7 +9,7 @@ import { MessageContent, ToolMessage, } from "@langchain/core/messages"; -import { CacheMetrics } from "@open-swe/shared/open-swe/types"; +import { CacheMetrics, ModelTokenData } from "@open-swe/shared/open-swe/types"; import { createLogger, LogLevel } from "./logger.js"; import { calculateCostSavings } from "@open-swe/shared/caching"; @@ -21,7 +21,10 @@ export interface CacheablePromptSegment { cache_control?: { type: "ephemeral" }; } -export function trackCachePerformance(response: AIMessageChunk): CacheMetrics { +export function trackCachePerformance( + response: AIMessageChunk, + model: string, +): ModelTokenData[] { const metrics: CacheMetrics = { cacheCreationInputTokens: response.usage_metadata?.input_token_details?.cache_creation || 0, @@ -41,12 +44,18 @@ export function trackCachePerformance(response: AIMessageChunk): CacheMetrics { const costSavings = calculateCostSavings(metrics).totalSavings; logger.info("Cache Performance", { + model, cacheHitRate: `${(cacheHitRate * 100).toFixed(2)}%`, costSavings: `$${costSavings.toFixed(4)}`, ...metrics, }); - return metrics; + return [ + { + ...metrics, + model, + }, + ]; } function addCacheControlToMessageContent( diff --git a/apps/open-swe/src/utils/llms/index.ts b/apps/open-swe/src/utils/llms/index.ts index 02da7f36..d6143dfd 100644 --- a/apps/open-swe/src/utils/llms/index.ts +++ b/apps/open-swe/src/utils/llms/index.ts @@ -1,2 +1,3 @@ export * from "./load-model.js"; export * from "./constants.js"; +export * from "./model-manager.js"; diff --git a/apps/open-swe/src/utils/llms/model-manager.ts b/apps/open-swe/src/utils/llms/model-manager.ts index 076db70a..cd638673 100644 --- a/apps/open-swe/src/utils/llms/model-manager.ts +++ b/apps/open-swe/src/utils/llms/model-manager.ts @@ -232,6 +232,14 @@ export class ModelManager { return configs; } + /** + * Get the model name for a task from GraphConfig + */ + public getModelNameForTask(config: GraphConfig, task: Task): string { + const baseConfig = this.getBaseConfigForTask(config, task); + return baseConfig.modelName; + } + /** * Get base configuration for a task from GraphConfig */ diff --git a/apps/web/src/components/v2/thread-view.tsx b/apps/web/src/components/v2/thread-view.tsx index 04119a14..40342fe4 100644 --- a/apps/web/src/components/v2/thread-view.tsx +++ b/apps/web/src/components/v2/thread-view.tsx @@ -341,12 +341,10 @@ export function ThreadView({ /> )} diff --git a/apps/web/src/components/v2/token-usage.tsx b/apps/web/src/components/v2/token-usage.tsx index 935d15ca..8a983e6f 100644 --- a/apps/web/src/components/v2/token-usage.tsx +++ b/apps/web/src/components/v2/token-usage.tsx @@ -1,5 +1,8 @@ -import { CacheMetrics } from "@open-swe/shared/open-swe/types"; -import { calculateCostSavings } from "@open-swe/shared/caching"; +import { CacheMetrics, ModelTokenData } from "@open-swe/shared/open-swe/types"; +import { + calculateCostSavings, + tokenDataReducer, +} from "@open-swe/shared/caching"; import { Badge } from "../ui/badge"; import { Separator } from "../ui/separator"; import { @@ -7,18 +10,34 @@ import { HoverCardContent, HoverCardTrigger, } from "../ui/hover-card"; +import { + Collapsible, + CollapsibleContent, + CollapsibleTrigger, +} from "../ui/collapsible"; import { ChartNoAxesColumnIncreasing, + ChevronDown, + ChevronRight, Coins, TrendingUp, Zap, } from "lucide-react"; +import { useState } from "react"; interface TokenUsageProps { - tokenData?: CacheMetrics[]; + tokenData?: ModelTokenData[] | CacheMetrics[]; } -function mergeTokenData(tokenDataArray: CacheMetrics[]): CacheMetrics { +function isModelTokenData( + data: ModelTokenData[] | CacheMetrics[], +): data is ModelTokenData[] { + return data.length > 0 && "model" in data[0]; +} + +function mergeTokenData( + tokenDataArray: ModelTokenData[] | CacheMetrics[], +): CacheMetrics { return tokenDataArray.reduce( (merged, current) => ({ cacheCreationInputTokens: @@ -37,7 +56,130 @@ function mergeTokenData(tokenDataArray: CacheMetrics[]): CacheMetrics { ); } +function mergeModelTokenData(tokenData: ModelTokenData[]): ModelTokenData[] { + if (tokenData.length <= 1) { + return tokenData; + } + const [firstTokenData, ...restTokenData] = tokenData; + return tokenDataReducer([firstTokenData], restTokenData); +} + +function getModelPricingPlaceholder(model: string): { + inputPrice: number; + outputPrice: number; + cachePrice: number; +} { + // Actual Claude pricing per 1M tokens based on official pricing table + const pricingMap: Record< + string, + { inputPrice: number; outputPrice: number; cachePrice: number } + > = { + // Claude 4 models + "anthropic:claude-4-opus": { + inputPrice: 15.0, + outputPrice: 75.0, + cachePrice: 1.5, + }, + "anthropic:claude-4-sonnet": { + inputPrice: 3.0, + outputPrice: 15.0, + cachePrice: 0.3, + }, + + // Claude 3.7 models + "anthropic:claude-3-7-sonnet": { + inputPrice: 3.0, + outputPrice: 15.0, + cachePrice: 0.3, + }, + + // Claude 3.5 models + "anthropic:claude-3-5-sonnet-20241022": { + inputPrice: 3.0, + outputPrice: 15.0, + cachePrice: 0.3, + }, + "anthropic:claude-3-5-sonnet-20240620": { + inputPrice: 3.0, + outputPrice: 15.0, + cachePrice: 0.3, + }, + "anthropic:claude-3-5-haiku-20241022": { + inputPrice: 0.8, + outputPrice: 4.0, + cachePrice: 0.08, + }, + + // Claude 3 models + "anthropic:claude-3-opus-20240229": { + inputPrice: 15.0, + outputPrice: 75.0, + cachePrice: 1.5, + }, + "anthropic:claude-3-haiku-20240307": { + inputPrice: 0.25, + outputPrice: 1.25, + cachePrice: 0.03, + }, + + // OpenAI models (actual pricing - no caching support) + "openai:o4": { inputPrice: 1.1, outputPrice: 4.4, cachePrice: 1.1 }, + "openai:o4-mini": { inputPrice: 1.1, outputPrice: 4.4, cachePrice: 1.1 }, + "openai:o3": { inputPrice: 2.0, outputPrice: 8.0, cachePrice: 2.0 }, + "openai:o3-mini": { inputPrice: 1.1, outputPrice: 4.4, cachePrice: 1.1 }, + "openai:gpt-4o": { inputPrice: 2.5, outputPrice: 10.0, cachePrice: 2.5 }, + "openai:gpt-4o-mini": { + inputPrice: 0.15, + outputPrice: 0.6, + cachePrice: 0.15, + }, + "openai:gpt-4.1": { inputPrice: 2.0, outputPrice: 8.0, cachePrice: 2.0 }, + "openai:gpt-4.1-mini": { + inputPrice: 0.4, + outputPrice: 1.6, + cachePrice: 0.4, + }, + // Legacy OpenAI models + "openai:o1-preview": { + inputPrice: 15.0, + outputPrice: 60.0, + cachePrice: 15.0, + }, + "openai:o1-mini": { inputPrice: 3.0, outputPrice: 12.0, cachePrice: 3.0 }, + + // Google Gemini models (actual pricing - no caching support) + "google-genai:gemini-2.5-pro": { + inputPrice: 1.5, + outputPrice: 10.0, + cachePrice: 1.5, + }, + "google-genai:gemini-2.5-flash": { + inputPrice: 0.3, + outputPrice: 2.5, + cachePrice: 0.3, + }, + }; + + return ( + pricingMap[model] || { inputPrice: 1.0, outputPrice: 5.0, cachePrice: 0.5 } + ); // Default fallback +} + +function calculateModelCost(modelData: ModelTokenData): number { + const pricing = getModelPricingPlaceholder(modelData.model); + const baseInputCost = + (modelData.inputTokens * pricing.inputPrice) / 1_000_000; + const cacheCreationCost = + (modelData.cacheCreationInputTokens * pricing.inputPrice) / 1_000_000; // Cache creation uses base input price + const cacheHitCost = + (modelData.cacheReadInputTokens * pricing.cachePrice) / 1_000_000; // Cache hits use discounted price + const outputCost = (modelData.outputTokens * pricing.outputPrice) / 1_000_000; + return baseInputCost + cacheCreationCost + cacheHitCost + outputCost; +} + export function TokenUsage({ tokenData }: TokenUsageProps) { + const [isExpanded, setIsExpanded] = useState(false); + if (!tokenData || tokenData.length === 0) return null; const mergedTokenData = mergeTokenData(tokenData); @@ -52,6 +194,16 @@ export function TokenUsage({ tokenData }: TokenUsageProps) { ).toFixed(2); const metrics = calculateCostSavings(mergedTokenData); + const hasModelData = isModelTokenData(tokenData); + const modelTokenData = hasModelData + ? mergeModelTokenData(tokenData as ModelTokenData[]) + : []; + + // Calculate total cost using model-specific pricing if available + const totalModelCost = hasModelData + ? modelTokenData.reduce((sum, model) => sum + calculateModelCost(model), 0) + : metrics.totalCost; + return ( @@ -61,11 +213,13 @@ export function TokenUsage({ tokenData }: TokenUsageProps) { variant="secondary" className="text-xs" > - {tokenData.length} agent{tokenData.length !== 1 ? "s" : ""} + {hasModelData + ? `${modelTokenData.length} model${modelTokenData.length !== 1 ? "s" : ""}` + : `${tokenData.length} agent${tokenData.length !== 1 ? "s" : ""}`} - +
@@ -114,11 +268,14 @@ export function TokenUsage({ tokenData }: TokenUsageProps) {
- Cost + {hasModelData ? "Estimated Cost" : "Cost"}
- ${metrics.totalCost.toFixed(2)} + $ + {(hasModelData ? totalModelCost : metrics.totalCost).toFixed( + 2, + )}
@@ -150,6 +307,112 @@ export function TokenUsage({ tokenData }: TokenUsageProps) { )}
+ + {hasModelData && modelTokenData.length > 1 && ( + <> + + + + Per-Model Breakdown + {isExpanded ? ( + + ) : ( + + )} + + + {modelTokenData.map((model, index) => { + const modelCost = calculateModelCost(model); + const modelTotalTokens = + model.inputTokens + + model.outputTokens + + model.cacheCreationInputTokens + + model.cacheReadInputTokens; + const modelCachedTokens = + model.cacheCreationInputTokens + + model.cacheReadInputTokens; + const modelCachePercentage = + modelTotalTokens > 0 + ? ( + (modelCachedTokens / modelTotalTokens) * + 100 + ).toFixed(1) + : "0"; + + return ( +
+
+ + {model.model.replace(/^(anthropic|openai):/, "")} + + + ${modelCost.toFixed(3)} + +
+ +
+
+
+ + + Input + +
+ + {( + model.inputTokens + + model.cacheCreationInputTokens + + model.cacheReadInputTokens + ).toLocaleString()} + +
+
+
+ + + Output + +
+ + {model.outputTokens.toLocaleString()} + +
+
+ + {modelCachedTokens > 0 && ( +
+ + Cache Percentage + + + {modelCachePercentage}% + +
+ )} +
+ ); + })} + +
+ * Estimated costs. Please review all token usage and pricing + to ensure accuracy. +
+
+
+ + )}
diff --git a/packages/shared/jest.config.js b/packages/shared/jest.config.js new file mode 100644 index 00000000..507052db --- /dev/null +++ b/packages/shared/jest.config.js @@ -0,0 +1,19 @@ +export default { + preset: "ts-jest/presets/default-esm", + moduleNameMapper: { + "^(\\.{1,2}/.*)\\.js$": "$1", + }, + transform: { + "^.+\\.tsx?$": [ + "ts-jest", + { + useESM: true, + }, + ], + }, + extensionsToTreatAsEsm: [".ts"], + setupFiles: ["dotenv/config"], + passWithNoTests: true, + testTimeout: 20_000, + testMatch: ["/src/**/*.test.ts"], +}; diff --git a/packages/shared/package.json b/packages/shared/package.json index 9e40838d..a6b9b359 100644 --- a/packages/shared/package.json +++ b/packages/shared/package.json @@ -14,7 +14,10 @@ "lint": "eslint .", "lint:fix": "eslint . --fix", "format": "prettier --write .", - "format:check": "prettier --check ." + "format:check": "prettier --check .", + "test": "NODE_OPTIONS=--experimental-vm-modules yarn run jest --config jest.config.js --testPathIgnorePatterns=int.test.ts", + "test:int": "node --experimental-vm-modules node_modules/jest/bin/jest.js --config jest.config.js --testPathPattern=int.test.ts", + "test:single": "NODE_OPTIONS=--experimental-vm-modules yarn run jest --config jest.config.js --testTimeout 100000" }, "dependencies": { "@langchain/core": "^0.3.65", @@ -27,8 +30,10 @@ "devDependencies": { "@eslint/eslintrc": "^3.1.0", "@eslint/js": "^9.19.0", + "@jest/globals": "^29.7.0", "@octokit/types": "^12.0.0", "@tsconfig/recommended": "^1.0.8", + "@types/jest": "^29.5.0", "@types/node": "^22.13.5", "dotenv": "^16.4.7", "eslint": "^9.19.0", @@ -36,7 +41,9 @@ "eslint-plugin-import": "^2.27.5", "eslint-plugin-no-instanceof": "^1.0.1", "eslint-plugin-prettier": "^4.2.1", + "jest": "^29.7.0", "prettier": "^3.5.2", + "ts-jest": "^29.1.0", "typescript": "~5.7.2", "typescript-eslint": "^8.22.0" }, diff --git a/packages/shared/src/__tests__/caching.test.ts b/packages/shared/src/__tests__/caching.test.ts new file mode 100644 index 00000000..fd15e4f9 --- /dev/null +++ b/packages/shared/src/__tests__/caching.test.ts @@ -0,0 +1,153 @@ +import { tokenDataReducer } from "../caching.js"; +import { ModelTokenData } from "../open-swe/types.js"; + +describe("tokenDataReducer", () => { + it("should merge objects with the same model string", () => { + const state: ModelTokenData[] = [ + { + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 100, + cacheReadInputTokens: 50, + inputTokens: 200, + outputTokens: 150, + }, + { + model: "openai:gpt-4.1-mini", + cacheCreationInputTokens: 80, + cacheReadInputTokens: 30, + inputTokens: 120, + outputTokens: 90, + }, + ]; + + const update: ModelTokenData[] = [ + { + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 25, + cacheReadInputTokens: 15, + inputTokens: 75, + outputTokens: 60, + }, + { + model: "openai:gpt-3.5-turbo", + cacheCreationInputTokens: 40, + cacheReadInputTokens: 20, + inputTokens: 100, + outputTokens: 80, + }, + ]; + + const result = tokenDataReducer(state, update); + + // Should have 3 models total (2 from state, 1 merged, 1 new) + expect(result).toHaveLength(3); + + // Find the merged anthropic model + const mergedAnthropic = result.find( + (data) => data.model === "anthropic:claude-sonnet-4-0", + ); + expect(mergedAnthropic).toEqual({ + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 125, // 100 + 25 + cacheReadInputTokens: 65, // 50 + 15 + inputTokens: 275, // 200 + 75 + outputTokens: 210, // 150 + 60 + }); + + // Find the unchanged openai gpt-4.1-mini model + const unchangedOpenAI = result.find( + (data) => data.model === "openai:gpt-4.1-mini", + ); + expect(unchangedOpenAI).toEqual({ + model: "openai:gpt-4.1-mini", + cacheCreationInputTokens: 80, + cacheReadInputTokens: 30, + inputTokens: 120, + outputTokens: 90, + }); + + // Find the new openai gpt-3.5-turbo model + const newOpenAI = result.find( + (data) => data.model === "openai:gpt-3.5-turbo", + ); + expect(newOpenAI).toEqual({ + model: "openai:gpt-3.5-turbo", + cacheCreationInputTokens: 40, + cacheReadInputTokens: 20, + inputTokens: 100, + outputTokens: 80, + }); + }); + + it("should return update array when state is undefined", () => { + const update: ModelTokenData[] = [ + { + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 100, + cacheReadInputTokens: 50, + inputTokens: 200, + outputTokens: 150, + }, + ]; + + const result = tokenDataReducer(undefined, update); + + expect(result).toEqual(update); + }); + + it("should handle empty update array", () => { + const state: ModelTokenData[] = [ + { + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 100, + cacheReadInputTokens: 50, + inputTokens: 200, + outputTokens: 150, + }, + ]; + + const result = tokenDataReducer(state, []); + + expect(result).toEqual(state); + }); + + it("should handle multiple updates for the same model", () => { + const state: ModelTokenData[] = [ + { + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 100, + cacheReadInputTokens: 50, + inputTokens: 200, + outputTokens: 150, + }, + ]; + + const update: ModelTokenData[] = [ + { + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 25, + cacheReadInputTokens: 15, + inputTokens: 75, + outputTokens: 60, + }, + { + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 10, + cacheReadInputTokens: 5, + inputTokens: 30, + outputTokens: 20, + }, + ]; + + const result = tokenDataReducer(state, update); + + expect(result).toHaveLength(1); + expect(result[0]).toEqual({ + model: "anthropic:claude-sonnet-4-0", + cacheCreationInputTokens: 135, // 100 + 25 + 10 + cacheReadInputTokens: 70, // 50 + 15 + 5 + inputTokens: 305, // 200 + 75 + 30 + outputTokens: 230, // 150 + 60 + 20 + }); + }); +}); diff --git a/packages/shared/src/caching.ts b/packages/shared/src/caching.ts index ca57e3af..64b6dec2 100644 --- a/packages/shared/src/caching.ts +++ b/packages/shared/src/caching.ts @@ -1,4 +1,4 @@ -import { CacheMetrics } from "./open-swe/types.js"; +import { CacheMetrics, ModelTokenData } from "./open-swe/types.js"; export function calculateCostSavings(metrics: CacheMetrics): { totalSavings: number; @@ -48,18 +48,41 @@ export function calculateCostSavings(metrics: CacheMetrics): { } export function tokenDataReducer( - state: CacheMetrics | undefined, - update: CacheMetrics, -): CacheMetrics { + state: ModelTokenData[] | undefined, + update: ModelTokenData[], +): ModelTokenData[] { if (!state) { return update; } - return { - cacheCreationInputTokens: - state.cacheCreationInputTokens + update.cacheCreationInputTokens, - cacheReadInputTokens: - state.cacheReadInputTokens + update.cacheReadInputTokens, - inputTokens: state.inputTokens + update.inputTokens, - outputTokens: state.outputTokens + update.outputTokens, - }; + + // Create a map to merge data by model + const modelMap = new Map(); + + // Add existing state data to the map + for (const data of state) { + modelMap.set(data.model, { ...data }); + } + + // Merge update data with existing data + for (const data of update) { + const existing = modelMap.get(data.model); + if (existing) { + // Merge the metrics for the same model + modelMap.set(data.model, { + model: data.model, + cacheCreationInputTokens: + existing.cacheCreationInputTokens + data.cacheCreationInputTokens, + cacheReadInputTokens: + existing.cacheReadInputTokens + data.cacheReadInputTokens, + inputTokens: existing.inputTokens + data.inputTokens, + outputTokens: existing.outputTokens + data.outputTokens, + }); + } else { + // Add new model data + modelMap.set(data.model, { ...data }); + } + } + + // Convert map back to array + return Array.from(modelMap.values()); } diff --git a/packages/shared/src/open-swe/planner/types.ts b/packages/shared/src/open-swe/planner/types.ts index 964bfb08..2a0c041a 100644 --- a/packages/shared/src/open-swe/planner/types.ts +++ b/packages/shared/src/open-swe/planner/types.ts @@ -3,8 +3,8 @@ import { z } from "zod"; import { MessagesZodState } from "@langchain/langgraph"; import { AgentSession, - CacheMetrics, CustomRules, + ModelTokenData, TargetRepository, TaskPlan, } from "../types.js"; @@ -103,9 +103,9 @@ export const PlannerGraphStateObj = MessagesZodState.extend({ fn: (_state, update) => update, }, }), - tokenData: withLangGraph(z.custom().optional(), { + tokenData: withLangGraph(z.custom().optional(), { reducer: { - schema: z.custom().optional(), + schema: z.custom().optional(), fn: tokenDataReducer, }, }), diff --git a/packages/shared/src/open-swe/reviewer/types.ts b/packages/shared/src/open-swe/reviewer/types.ts index 9a47d5c8..4475eb2e 100644 --- a/packages/shared/src/open-swe/reviewer/types.ts +++ b/packages/shared/src/open-swe/reviewer/types.ts @@ -6,8 +6,8 @@ import { MessagesZodState, } from "@langchain/langgraph"; import { - CacheMetrics, CustomRules, + ModelTokenData, TargetRepository, TaskPlan, } from "../types.js"; @@ -116,9 +116,9 @@ export const ReviewerGraphStateObj = MessagesZodState.extend({ }, default: () => 0, }), - tokenData: withLangGraph(z.custom().optional(), { + tokenData: withLangGraph(z.custom().optional(), { reducer: { - schema: z.custom().optional(), + schema: z.custom().optional(), fn: tokenDataReducer, }, }), diff --git a/packages/shared/src/open-swe/types.ts b/packages/shared/src/open-swe/types.ts index 129d4711..fc9a239f 100644 --- a/packages/shared/src/open-swe/types.ts +++ b/packages/shared/src/open-swe/types.ts @@ -34,6 +34,14 @@ export interface CacheMetrics { outputTokens: number; } +export interface ModelTokenData extends CacheMetrics { + /** + * The model name that generated this token usage data + * e.g., "anthropic:claude-sonnet-4-0", "openai:gpt-4.1-mini" + */ + model: string; +} + export type PlanItem = { /** * The index of the plan item. This is the order in which @@ -274,9 +282,9 @@ export const GraphAnnotation = MessagesZodState.extend({ default: () => 0, }), - tokenData: withLangGraph(z.custom().optional(), { + tokenData: withLangGraph(z.custom().optional(), { reducer: { - schema: z.custom().optional(), + schema: z.custom().optional(), fn: tokenDataReducer, }, }), diff --git a/yarn.lock b/yarn.lock index 390176c9..ecdd2d0d 100644 --- a/yarn.lock +++ b/yarn.lock @@ -4181,12 +4181,14 @@ __metadata: dependencies: "@eslint/eslintrc": ^3.1.0 "@eslint/js": ^9.19.0 + "@jest/globals": ^29.7.0 "@langchain/core": ^0.3.65 "@langchain/langgraph": ^0.3.8 "@langchain/langgraph-sdk": ^0.0.95 "@octokit/rest": ^22.0.0 "@octokit/types": ^12.0.0 "@tsconfig/recommended": ^1.0.8 + "@types/jest": ^29.5.0 "@types/node": ^22.13.5 dotenv: ^16.4.7 eslint: ^9.19.0 @@ -4194,8 +4196,10 @@ __metadata: eslint-plugin-import: ^2.27.5 eslint-plugin-no-instanceof: ^1.0.1 eslint-plugin-prettier: ^4.2.1 + jest: ^29.7.0 jsonwebtoken: ^9.0.2 prettier: ^3.5.2 + ts-jest: ^29.1.0 typescript: ~5.7.2 typescript-eslint: ^8.22.0 zod: ^3.25.32