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] <open-swe@users.noreply.github.com>
Co-authored-by: bracesproul <braceasproul@gmail.com>
This commit is contained in:
open-swe[bot] 2025-07-27 23:56:50 +00:00 • committed by GitHub
parent eb2311955a
commit fb4523539d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
25 changed files with 581 additions and 52 deletions

View file

@ -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 =

View file

@ -1,4 +1,5 @@
import {
getModelManager,
loadModel,
supportsParallelToolCallsParam,
Task,
@ -70,6 +71,8 @@ export async function generateAction(
config: GraphConfig,
): Promise<PlannerGraphUpdate> {
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),
};
}

View file

@ -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<PlannerGraphUpdate> {
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),
};
}

View file

@ -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<PlannerGraphUpdate> {
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<typeof condenseContextTool.schema>
).notes,
tokenData: trackCachePerformance(response),
tokenData: trackCachePerformance(response, modelName),
};
}

View file

@ -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<GraphUpdate> {
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),
};
}

View file

@ -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<GraphUpdate> {
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),
};
}

View file

@ -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 }),
};
}

View file

@ -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:
<user_request>
@ -148,6 +149,8 @@ export async function summarizeHistory(
config: GraphConfig,
): Promise<GraphUpdate> {
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),
};
}

View file

@ -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),
};
}

View file

@ -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),
};
}

View file

@ -1,4 +1,5 @@
import {
getModelManager,
loadModel,
supportsParallelToolCallsParam,
Task,
@ -112,6 +113,8 @@ export async function generateReviewActions(
config: GraphConfig,
): Promise<ReviewerGraphUpdate> {
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),
};
}

View file

@ -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<DiagnoseErrorInputs>;
@ -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),
};
}

View file

@ -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(

View file

@ -1,2 +1,3 @@
export * from "./load-model.js";
export * from "./constants.js";
export * from "./model-manager.js";

View file

@ -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
*/

View file

@ -341,12 +341,10 @@ export function ThreadView({
/>
)}
<TokenUsage
tokenData={
[
plannerStream.values.tokenData,
programmerStream.values.tokenData,
].filter(Boolean) as CacheMetrics[]
}
tokenData={[
...(plannerStream.values.tokenData ?? []),
...(programmerStream.values.tokenData ?? []),
].filter(Boolean)}
/>
</div>
</div>

View file

@ -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 (
<HoverCard>
<HoverCardTrigger asChild>
@ -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" : ""}`}
</Badge>
</button>
</HoverCardTrigger>
<HoverCardContent className="w-80">
<HoverCardContent className="w-96">
<div className="space-y-4">
<div className="flex items-center gap-2">
<ChartNoAxesColumnIncreasing className="h-4 w-4" />
@ -114,11 +268,14 @@ export function TokenUsage({ tokenData }: TokenUsageProps) {
<div className="flex items-center gap-1.5">
<Coins className="h-3 w-3 text-amber-500 dark:text-amber-400" />
<span className="text-muted-foreground text-xs font-medium">
Cost
{hasModelData ? "Estimated Cost" : "Cost"}
</span>
</div>
<span className="text-sm font-semibold">
${metrics.totalCost.toFixed(2)}
$
{(hasModelData ? totalModelCost : metrics.totalCost).toFixed(
2,
)}
</span>
</div>
@ -150,6 +307,112 @@ export function TokenUsage({ tokenData }: TokenUsageProps) {
)}
</div>
</div>
{hasModelData && modelTokenData.length > 1 && (
<>
<Separator />
<Collapsible
open={isExpanded}
onOpenChange={setIsExpanded}
>
<CollapsibleTrigger className="hover:text-foreground flex w-full items-center justify-between text-sm font-medium">
<span>Per-Model Breakdown</span>
{isExpanded ? (
<ChevronDown className="h-4 w-4" />
) : (
<ChevronRight className="h-4 w-4" />
)}
</CollapsibleTrigger>
<CollapsibleContent className="space-y-3 pt-3">
{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 (
<div
key={index}
className="space-y-2 rounded-lg border p-3"
>
<div className="flex items-center justify-between">
<span className="text-muted-foreground truncate text-xs font-medium">
{model.model.replace(/^(anthropic|openai):/, "")}
</span>
<Badge
variant="outline"
className="text-xs"
>
${modelCost.toFixed(3)}
</Badge>
</div>
<div className="space-y-2 text-xs">
<div className="flex justify-between">
<div className="flex items-center gap-1.5">
<Zap className="h-3 w-3 text-blue-500 dark:text-blue-400" />
<span className="text-muted-foreground font-medium">
Input
</span>
</div>
<span className="font-semibold">
{(
model.inputTokens +
model.cacheCreationInputTokens +
model.cacheReadInputTokens
).toLocaleString()}
</span>
</div>
<div className="flex justify-between">
<div className="flex items-center gap-1.5">
<TrendingUp className="h-3 w-3 text-green-500 dark:text-green-400" />
<span className="text-muted-foreground font-medium">
Output
</span>
</div>
<span className="font-semibold">
{model.outputTokens.toLocaleString()}
</span>
</div>
</div>
{modelCachedTokens > 0 && (
<div className="flex justify-between text-xs">
<span className="text-xs font-medium text-blue-600 dark:text-blue-400">
Cache Percentage
</span>
<Badge
variant="outline"
className="border-blue-200 text-blue-600 dark:border-blue-800 dark:text-blue-400"
>
{modelCachePercentage}%
</Badge>
</div>
)}
</div>
);
})}
<div className="text-muted-foreground border-t pt-2 text-xs">
* Estimated costs. Please review all token usage and pricing
to ensure accuracy.
</div>
</CollapsibleContent>
</Collapsible>
</>
)}
</div>
</HoverCardContent>
</HoverCard>

View file

@ -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: ["<rootDir>/src/**/*.test.ts"],
};

View file

@ -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"
},

View file

@ -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
});
});
});

View file

@ -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<string, ModelTokenData>();
// 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());
}

View file

@ -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<CacheMetrics>().optional(), {
tokenData: withLangGraph(z.custom<ModelTokenData[]>().optional(), {
reducer: {
schema: z.custom<CacheMetrics>().optional(),
schema: z.custom<ModelTokenData[]>().optional(),
fn: tokenDataReducer,
},
}),

View file

@ -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<CacheMetrics>().optional(), {
tokenData: withLangGraph(z.custom<ModelTokenData[]>().optional(), {
reducer: {
schema: z.custom<CacheMetrics>().optional(),
schema: z.custom<ModelTokenData[]>().optional(),
fn: tokenDataReducer,
},
}),

View file

@ -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<CacheMetrics>().optional(), {
tokenData: withLangGraph(z.custom<ModelTokenData[]>().optional(), {
reducer: {
schema: z.custom<CacheMetrics>().optional(),
schema: z.custom<ModelTokenData[]>().optional(),
fn: tokenDataReducer,
},
}),

View file

@ -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