mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-04 09:02:11 +00:00
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:
parent
eb2311955a
commit
fb4523539d
25 changed files with 581 additions and 52 deletions
|
|
@ -17,6 +17,7 @@ import { getMessageContentString } from "@open-swe/shared/messages";
|
||||||
import { filterHiddenMessages } from "../../../utils/message/filter-hidden.js";
|
import { filterHiddenMessages } from "../../../utils/message/filter-hidden.js";
|
||||||
import { createLogger, LogLevel } from "../../../utils/logger.js";
|
import { createLogger, LogLevel } from "../../../utils/logger.js";
|
||||||
import { trackCachePerformance } from "../../../utils/caching.js";
|
import { trackCachePerformance } from "../../../utils/caching.js";
|
||||||
|
import { getModelManager } from "../../../utils/llms/model-manager.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "DetermineNeedsContext");
|
const logger = createLogger(LogLevel.INFO, "DetermineNeedsContext");
|
||||||
|
|
||||||
|
|
@ -120,6 +121,8 @@ export async function determineNeedsContext(
|
||||||
getMissingMessages(state, config),
|
getMissingMessages(state, config),
|
||||||
loadModel(config, Task.ROUTER),
|
loadModel(config, Task.ROUTER),
|
||||||
]);
|
]);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.ROUTER);
|
||||||
if (!missingMessages.length) {
|
if (!missingMessages.length) {
|
||||||
throw new Error(
|
throw new Error(
|
||||||
"Can not determine if more context is needed if there are no missing messages.",
|
"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 = {
|
const commandUpdate: PlannerGraphUpdate = {
|
||||||
messages: missingMessages,
|
messages: missingMessages,
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
|
|
||||||
const shouldGatherContext =
|
const shouldGatherContext =
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
import {
|
import {
|
||||||
|
getModelManager,
|
||||||
loadModel,
|
loadModel,
|
||||||
supportsParallelToolCallsParam,
|
supportsParallelToolCallsParam,
|
||||||
Task,
|
Task,
|
||||||
|
|
@ -70,6 +71,8 @@ export async function generateAction(
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<PlannerGraphUpdate> {
|
): Promise<PlannerGraphUpdate> {
|
||||||
const model = await loadModel(config, Task.PROGRAMMER);
|
const model = await loadModel(config, Task.PROGRAMMER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.PROGRAMMER,
|
Task.PROGRAMMER,
|
||||||
|
|
@ -145,6 +148,6 @@ export async function generateAction(
|
||||||
return {
|
return {
|
||||||
messages: [...missingMessages, response],
|
messages: [...missingMessages, response],
|
||||||
...(latestTaskPlan && { taskPlan: latestTaskPlan }),
|
...(latestTaskPlan && { taskPlan: latestTaskPlan }),
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -23,6 +23,7 @@ import { getScratchpad } from "../../utils/scratchpad-notes.js";
|
||||||
import { SCRATCHPAD_PROMPT, SYSTEM_PROMPT } from "./prompt.js";
|
import { SCRATCHPAD_PROMPT, SYSTEM_PROMPT } from "./prompt.js";
|
||||||
import { DO_NOT_RENDER_ID_PREFIX } from "@open-swe/shared/constants";
|
import { DO_NOT_RENDER_ID_PREFIX } from "@open-swe/shared/constants";
|
||||||
import { filterMessagesWithoutContent } from "../../../../utils/message/content.js";
|
import { filterMessagesWithoutContent } from "../../../../utils/message/content.js";
|
||||||
|
import { getModelManager } from "../../../../utils/llms/model-manager.js";
|
||||||
import { trackCachePerformance } from "../../../../utils/caching.js";
|
import { trackCachePerformance } from "../../../../utils/caching.js";
|
||||||
|
|
||||||
function formatSystemPrompt(state: PlannerGraphState): string {
|
function formatSystemPrompt(state: PlannerGraphState): string {
|
||||||
|
|
@ -54,6 +55,8 @@ export async function generatePlan(
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<PlannerGraphUpdate> {
|
): Promise<PlannerGraphUpdate> {
|
||||||
const model = await loadModel(config, Task.PROGRAMMER);
|
const model = await loadModel(config, Task.PROGRAMMER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.SUMMARIZER,
|
Task.SUMMARIZER,
|
||||||
|
|
@ -125,6 +128,6 @@ export async function generatePlan(
|
||||||
proposedPlanTitle: proposedPlanArgs.title,
|
proposedPlanTitle: proposedPlanArgs.title,
|
||||||
proposedPlan: proposedPlanArgs.plan,
|
proposedPlan: proposedPlanArgs.plan,
|
||||||
...(newSessionId && { sandboxSessionId: newSessionId }),
|
...(newSessionId && { sandboxSessionId: newSessionId }),
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -18,6 +18,7 @@ import { ToolMessage } from "@langchain/core/messages";
|
||||||
import { DO_NOT_RENDER_ID_PREFIX } from "@open-swe/shared/constants";
|
import { DO_NOT_RENDER_ID_PREFIX } from "@open-swe/shared/constants";
|
||||||
import { createWriteTechnicalNotesToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createWriteTechnicalNotesToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
import { trackCachePerformance } from "../../../utils/caching.js";
|
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.
|
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,
|
config: GraphConfig,
|
||||||
): Promise<PlannerGraphUpdate> {
|
): Promise<PlannerGraphUpdate> {
|
||||||
const model = await loadModel(config, Task.SUMMARIZER);
|
const model = await loadModel(config, Task.SUMMARIZER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.SUMMARIZER,
|
Task.SUMMARIZER,
|
||||||
|
|
@ -146,6 +149,6 @@ ${state.messages.map(getMessageString).join("\n")}`;
|
||||||
contextGatheringNotes: (
|
contextGatheringNotes: (
|
||||||
toolCall.args as z.infer<typeof condenseContextTool.schema>
|
toolCall.args as z.infer<typeof condenseContextTool.schema>
|
||||||
).notes,
|
).notes,
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -16,6 +16,7 @@ import {
|
||||||
} from "@open-swe/shared/open-swe/tasks";
|
} from "@open-swe/shared/open-swe/tasks";
|
||||||
import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js";
|
import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js";
|
||||||
import { trackCachePerformance } from "../../../utils/caching.js";
|
import { trackCachePerformance } from "../../../utils/caching.js";
|
||||||
|
import { getModelManager } from "../../../utils/llms/model-manager.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "GenerateConclusionNode");
|
const logger = createLogger(LogLevel.INFO, "GenerateConclusionNode");
|
||||||
|
|
||||||
|
|
@ -40,6 +41,8 @@ export async function generateConclusion(
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<GraphUpdate> {
|
): Promise<GraphUpdate> {
|
||||||
const model = await loadModel(config, Task.SUMMARIZER);
|
const model = await loadModel(config, Task.SUMMARIZER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER);
|
||||||
|
|
||||||
const userRequestPrompt = formatUserRequestPrompt(state.messages);
|
const userRequestPrompt = formatUserRequestPrompt(state.messages);
|
||||||
const userMessage = `${userRequestPrompt}
|
const userMessage = `${userRequestPrompt}
|
||||||
|
|
@ -83,6 +86,6 @@ Given all of this, please respond with the concise conclusion. Do not include an
|
||||||
messages: [response],
|
messages: [response],
|
||||||
internalMessages: [response],
|
internalMessages: [response],
|
||||||
taskPlan: updatedTaskPlan,
|
taskPlan: updatedTaskPlan,
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import {
|
||||||
GraphUpdate,
|
GraphUpdate,
|
||||||
} from "@open-swe/shared/open-swe/types";
|
} from "@open-swe/shared/open-swe/types";
|
||||||
import {
|
import {
|
||||||
|
getModelManager,
|
||||||
loadModel,
|
loadModel,
|
||||||
supportsParallelToolCallsParam,
|
supportsParallelToolCallsParam,
|
||||||
Task,
|
Task,
|
||||||
|
|
@ -141,6 +142,8 @@ export async function generateAction(
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<GraphUpdate> {
|
): Promise<GraphUpdate> {
|
||||||
const model = await loadModel(config, Task.PROGRAMMER);
|
const model = await loadModel(config, Task.PROGRAMMER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.PROGRAMMER,
|
Task.PROGRAMMER,
|
||||||
|
|
@ -246,6 +249,6 @@ export async function generateAction(
|
||||||
internalMessages: newMessagesList,
|
internalMessages: newMessagesList,
|
||||||
...(newSandboxSessionId && { sandboxSessionId: newSandboxSessionId }),
|
...(newSandboxSessionId && { sandboxSessionId: newSandboxSessionId }),
|
||||||
...(latestTaskPlan && { taskPlan: latestTaskPlan }),
|
...(latestTaskPlan && { taskPlan: latestTaskPlan }),
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -36,6 +36,7 @@ import {
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { createOpenPrToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createOpenPrToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
import { trackCachePerformance } from "../../../utils/caching.js";
|
import { trackCachePerformance } from "../../../utils/caching.js";
|
||||||
|
import { getModelManager } from "../../../utils/llms/model-manager.js";
|
||||||
import {
|
import {
|
||||||
GitHubPullRequest,
|
GitHubPullRequest,
|
||||||
GitHubPullRequestList,
|
GitHubPullRequestList,
|
||||||
|
|
@ -116,6 +117,8 @@ export async function openPullRequest(
|
||||||
const openPrTool = createOpenPrToolFields();
|
const openPrTool = createOpenPrToolFields();
|
||||||
// use the router model since this is a simple task that doesn't need an advanced model
|
// 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 model = await loadModel(config, Task.ROUTER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.ROUTER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.ROUTER,
|
Task.ROUTER,
|
||||||
|
|
@ -219,7 +222,7 @@ export async function openPullRequest(
|
||||||
}),
|
}),
|
||||||
...(codebaseTree && { codebaseTree }),
|
...(codebaseTree && { codebaseTree }),
|
||||||
...(dependenciesInstalled !== null && { dependenciesInstalled }),
|
...(dependenciesInstalled !== null && { dependenciesInstalled }),
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
...(updatedTaskPlan && { taskPlan: updatedTaskPlan }),
|
...(updatedTaskPlan && { taskPlan: updatedTaskPlan }),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -21,6 +21,7 @@ import { createConversationHistorySummaryToolFields } from "@open-swe/shared/ope
|
||||||
import { formatUserRequestPrompt } from "../../../utils/user-request.js";
|
import { formatUserRequestPrompt } from "../../../utils/user-request.js";
|
||||||
import { getMessagesSinceLastSummary } from "../../../utils/tokens.js";
|
import { getMessagesSinceLastSummary } from "../../../utils/tokens.js";
|
||||||
import { trackCachePerformance } from "../../../utils/caching.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:
|
const SINGLE_USER_REQUEST_PROMPT = `Here is the user's request:
|
||||||
<user_request>
|
<user_request>
|
||||||
|
|
@ -148,6 +149,8 @@ export async function summarizeHistory(
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<GraphUpdate> {
|
): Promise<GraphUpdate> {
|
||||||
const model = await loadModel(config, Task.SUMMARIZER);
|
const model = await loadModel(config, Task.SUMMARIZER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER);
|
||||||
|
|
||||||
const plan = getActivePlanItems(state.taskPlan);
|
const plan = getActivePlanItems(state.taskPlan);
|
||||||
const conversationHistoryToSummarize = await getMessagesSinceLastSummary(
|
const conversationHistoryToSummarize = await getMessagesSinceLastSummary(
|
||||||
|
|
@ -190,6 +193,6 @@ export async function summarizeHistory(
|
||||||
return {
|
return {
|
||||||
messages: summaryMessages,
|
messages: summaryMessages,
|
||||||
internalMessages: newInternalMessages,
|
internalMessages: newInternalMessages,
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -27,6 +27,7 @@ import { createLogger, LogLevel } from "../../../utils/logger.js";
|
||||||
import { createUpdatePlanToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createUpdatePlanToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
import { formatCustomRulesPrompt } from "../../../utils/custom-rules.js";
|
import { formatCustomRulesPrompt } from "../../../utils/custom-rules.js";
|
||||||
import { trackCachePerformance } from "../../../utils/caching.js";
|
import { trackCachePerformance } from "../../../utils/caching.js";
|
||||||
|
import { getModelManager } from "../../../utils/llms/model-manager.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "UpdatePlanNode");
|
const logger = createLogger(LogLevel.INFO, "UpdatePlanNode");
|
||||||
|
|
||||||
|
|
@ -124,6 +125,8 @@ export async function updatePlan(
|
||||||
});
|
});
|
||||||
|
|
||||||
const model = await loadModel(config, Task.PROGRAMMER);
|
const model = await loadModel(config, Task.PROGRAMMER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.PROGRAMMER,
|
Task.PROGRAMMER,
|
||||||
|
|
@ -204,6 +207,6 @@ export async function updatePlan(
|
||||||
messages: [toolMessage],
|
messages: [toolMessage],
|
||||||
internalMessages: [toolMessage],
|
internalMessages: [toolMessage],
|
||||||
taskPlan: newTaskPlan,
|
taskPlan: newTaskPlan,
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -30,6 +30,7 @@ import {
|
||||||
ToolMessage,
|
ToolMessage,
|
||||||
} from "@langchain/core/messages";
|
} from "@langchain/core/messages";
|
||||||
import { trackCachePerformance } from "../../../utils/caching.js";
|
import { trackCachePerformance } from "../../../utils/caching.js";
|
||||||
|
import { getModelManager } from "../../../utils/llms/model-manager.js";
|
||||||
import { createScratchpadTool } from "../../../tools/scratchpad.js";
|
import { createScratchpadTool } from "../../../tools/scratchpad.js";
|
||||||
|
|
||||||
const SYSTEM_PROMPT = `You are a code reviewer for a software engineer working on a large codebase.
|
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 incompleteTool = createCodeReviewMarkTaskNotCompleteFields();
|
||||||
const tools = [completedTool, incompleteTool];
|
const tools = [completedTool, incompleteTool];
|
||||||
const model = await loadModel(config, Task.PROGRAMMER);
|
const model = await loadModel(config, Task.PROGRAMMER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.PROGRAMMER,
|
Task.PROGRAMMER,
|
||||||
|
|
@ -206,6 +209,6 @@ export async function finalReview(
|
||||||
messages: messagesUpdate,
|
messages: messagesUpdate,
|
||||||
internalMessages: messagesUpdate,
|
internalMessages: messagesUpdate,
|
||||||
reviewsCount: (state.reviewsCount || 0) + 1,
|
reviewsCount: (state.reviewsCount || 0) + 1,
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
import {
|
import {
|
||||||
|
getModelManager,
|
||||||
loadModel,
|
loadModel,
|
||||||
supportsParallelToolCallsParam,
|
supportsParallelToolCallsParam,
|
||||||
Task,
|
Task,
|
||||||
|
|
@ -112,6 +113,8 @@ export async function generateReviewActions(
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<ReviewerGraphUpdate> {
|
): Promise<ReviewerGraphUpdate> {
|
||||||
const model = await loadModel(config, Task.PROGRAMMER);
|
const model = await loadModel(config, Task.PROGRAMMER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.PROGRAMMER,
|
Task.PROGRAMMER,
|
||||||
|
|
@ -166,6 +169,6 @@ export async function generateReviewActions(
|
||||||
return {
|
return {
|
||||||
messages: [response],
|
messages: [response],
|
||||||
reviewerMessages: [response],
|
reviewerMessages: [response],
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,7 @@ import {
|
||||||
import { createDiagnoseErrorToolFields } from "@open-swe/shared/open-swe/tools";
|
import { createDiagnoseErrorToolFields } from "@open-swe/shared/open-swe/tools";
|
||||||
|
|
||||||
import { z } from "zod";
|
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 { createLogger, LogLevel } from "../../utils/logger.js";
|
||||||
import { getAllLastFailedActions } from "../../utils/tool-message-error.js";
|
import { getAllLastFailedActions } from "../../utils/tool-message-error.js";
|
||||||
import { getMessageString } from "../../utils/message/content.js";
|
import { getMessageString } from "../../utils/message/content.js";
|
||||||
|
|
@ -17,6 +17,7 @@ import {
|
||||||
Task,
|
Task,
|
||||||
} from "../../utils/llms/index.js";
|
} from "../../utils/llms/index.js";
|
||||||
import { trackCachePerformance } from "../../utils/caching.js";
|
import { trackCachePerformance } from "../../utils/caching.js";
|
||||||
|
import { getModelManager } from "../../utils/llms/model-manager.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "SharedDiagnoseError");
|
const logger = createLogger(LogLevel.INFO, "SharedDiagnoseError");
|
||||||
|
|
||||||
|
|
@ -76,7 +77,7 @@ const formatUserPrompt = (messages: BaseMessage[]): string => {
|
||||||
interface DiagnoseErrorInputs {
|
interface DiagnoseErrorInputs {
|
||||||
messages: BaseMessage[];
|
messages: BaseMessage[];
|
||||||
codebaseTree: string;
|
codebaseTree: string;
|
||||||
tokenData?: CacheMetrics;
|
tokenData?: ModelTokenData[];
|
||||||
}
|
}
|
||||||
|
|
||||||
type DiagnoseErrorUpdate = Partial<DiagnoseErrorInputs>;
|
type DiagnoseErrorUpdate = Partial<DiagnoseErrorInputs>;
|
||||||
|
|
@ -95,6 +96,8 @@ export async function diagnoseError(
|
||||||
logger.info("The last few tool calls resulted in errors. Diagnosing error.");
|
logger.info("The last few tool calls resulted in errors. Diagnosing error.");
|
||||||
|
|
||||||
const model = await loadModel(config, Task.SUMMARIZER);
|
const model = await loadModel(config, Task.SUMMARIZER);
|
||||||
|
const modelManager = getModelManager();
|
||||||
|
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER);
|
||||||
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
|
||||||
config,
|
config,
|
||||||
Task.SUMMARIZER,
|
Task.SUMMARIZER,
|
||||||
|
|
@ -142,6 +145,6 @@ export async function diagnoseError(
|
||||||
|
|
||||||
return {
|
return {
|
||||||
messages: [response, toolMessage],
|
messages: [response, toolMessage],
|
||||||
tokenData: trackCachePerformance(response),
|
tokenData: trackCachePerformance(response, modelName),
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -9,7 +9,7 @@ import {
|
||||||
MessageContent,
|
MessageContent,
|
||||||
ToolMessage,
|
ToolMessage,
|
||||||
} from "@langchain/core/messages";
|
} 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 { createLogger, LogLevel } from "./logger.js";
|
||||||
import { calculateCostSavings } from "@open-swe/shared/caching";
|
import { calculateCostSavings } from "@open-swe/shared/caching";
|
||||||
|
|
||||||
|
|
@ -21,7 +21,10 @@ export interface CacheablePromptSegment {
|
||||||
cache_control?: { type: "ephemeral" };
|
cache_control?: { type: "ephemeral" };
|
||||||
}
|
}
|
||||||
|
|
||||||
export function trackCachePerformance(response: AIMessageChunk): CacheMetrics {
|
export function trackCachePerformance(
|
||||||
|
response: AIMessageChunk,
|
||||||
|
model: string,
|
||||||
|
): ModelTokenData[] {
|
||||||
const metrics: CacheMetrics = {
|
const metrics: CacheMetrics = {
|
||||||
cacheCreationInputTokens:
|
cacheCreationInputTokens:
|
||||||
response.usage_metadata?.input_token_details?.cache_creation || 0,
|
response.usage_metadata?.input_token_details?.cache_creation || 0,
|
||||||
|
|
@ -41,12 +44,18 @@ export function trackCachePerformance(response: AIMessageChunk): CacheMetrics {
|
||||||
const costSavings = calculateCostSavings(metrics).totalSavings;
|
const costSavings = calculateCostSavings(metrics).totalSavings;
|
||||||
|
|
||||||
logger.info("Cache Performance", {
|
logger.info("Cache Performance", {
|
||||||
|
model,
|
||||||
cacheHitRate: `${(cacheHitRate * 100).toFixed(2)}%`,
|
cacheHitRate: `${(cacheHitRate * 100).toFixed(2)}%`,
|
||||||
costSavings: `$${costSavings.toFixed(4)}`,
|
costSavings: `$${costSavings.toFixed(4)}`,
|
||||||
...metrics,
|
...metrics,
|
||||||
});
|
});
|
||||||
|
|
||||||
return metrics;
|
return [
|
||||||
|
{
|
||||||
|
...metrics,
|
||||||
|
model,
|
||||||
|
},
|
||||||
|
];
|
||||||
}
|
}
|
||||||
|
|
||||||
function addCacheControlToMessageContent(
|
function addCacheControlToMessageContent(
|
||||||
|
|
|
||||||
|
|
@ -1,2 +1,3 @@
|
||||||
export * from "./load-model.js";
|
export * from "./load-model.js";
|
||||||
export * from "./constants.js";
|
export * from "./constants.js";
|
||||||
|
export * from "./model-manager.js";
|
||||||
|
|
|
||||||
|
|
@ -232,6 +232,14 @@ export class ModelManager {
|
||||||
return configs;
|
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
|
* Get base configuration for a task from GraphConfig
|
||||||
*/
|
*/
|
||||||
|
|
|
||||||
|
|
@ -341,12 +341,10 @@ export function ThreadView({
|
||||||
/>
|
/>
|
||||||
)}
|
)}
|
||||||
<TokenUsage
|
<TokenUsage
|
||||||
tokenData={
|
tokenData={[
|
||||||
[
|
...(plannerStream.values.tokenData ?? []),
|
||||||
plannerStream.values.tokenData,
|
...(programmerStream.values.tokenData ?? []),
|
||||||
programmerStream.values.tokenData,
|
].filter(Boolean)}
|
||||||
].filter(Boolean) as CacheMetrics[]
|
|
||||||
}
|
|
||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,8 @@
|
||||||
import { CacheMetrics } from "@open-swe/shared/open-swe/types";
|
import { CacheMetrics, ModelTokenData } from "@open-swe/shared/open-swe/types";
|
||||||
import { calculateCostSavings } from "@open-swe/shared/caching";
|
import {
|
||||||
|
calculateCostSavings,
|
||||||
|
tokenDataReducer,
|
||||||
|
} from "@open-swe/shared/caching";
|
||||||
import { Badge } from "../ui/badge";
|
import { Badge } from "../ui/badge";
|
||||||
import { Separator } from "../ui/separator";
|
import { Separator } from "../ui/separator";
|
||||||
import {
|
import {
|
||||||
|
|
@ -7,18 +10,34 @@ import {
|
||||||
HoverCardContent,
|
HoverCardContent,
|
||||||
HoverCardTrigger,
|
HoverCardTrigger,
|
||||||
} from "../ui/hover-card";
|
} from "../ui/hover-card";
|
||||||
|
import {
|
||||||
|
Collapsible,
|
||||||
|
CollapsibleContent,
|
||||||
|
CollapsibleTrigger,
|
||||||
|
} from "../ui/collapsible";
|
||||||
import {
|
import {
|
||||||
ChartNoAxesColumnIncreasing,
|
ChartNoAxesColumnIncreasing,
|
||||||
|
ChevronDown,
|
||||||
|
ChevronRight,
|
||||||
Coins,
|
Coins,
|
||||||
TrendingUp,
|
TrendingUp,
|
||||||
Zap,
|
Zap,
|
||||||
} from "lucide-react";
|
} from "lucide-react";
|
||||||
|
import { useState } from "react";
|
||||||
|
|
||||||
interface TokenUsageProps {
|
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(
|
return tokenDataArray.reduce(
|
||||||
(merged, current) => ({
|
(merged, current) => ({
|
||||||
cacheCreationInputTokens:
|
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) {
|
export function TokenUsage({ tokenData }: TokenUsageProps) {
|
||||||
|
const [isExpanded, setIsExpanded] = useState(false);
|
||||||
|
|
||||||
if (!tokenData || tokenData.length === 0) return null;
|
if (!tokenData || tokenData.length === 0) return null;
|
||||||
|
|
||||||
const mergedTokenData = mergeTokenData(tokenData);
|
const mergedTokenData = mergeTokenData(tokenData);
|
||||||
|
|
@ -52,6 +194,16 @@ export function TokenUsage({ tokenData }: TokenUsageProps) {
|
||||||
).toFixed(2);
|
).toFixed(2);
|
||||||
const metrics = calculateCostSavings(mergedTokenData);
|
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 (
|
return (
|
||||||
<HoverCard>
|
<HoverCard>
|
||||||
<HoverCardTrigger asChild>
|
<HoverCardTrigger asChild>
|
||||||
|
|
@ -61,11 +213,13 @@ export function TokenUsage({ tokenData }: TokenUsageProps) {
|
||||||
variant="secondary"
|
variant="secondary"
|
||||||
className="text-xs"
|
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>
|
</Badge>
|
||||||
</button>
|
</button>
|
||||||
</HoverCardTrigger>
|
</HoverCardTrigger>
|
||||||
<HoverCardContent className="w-80">
|
<HoverCardContent className="w-96">
|
||||||
<div className="space-y-4">
|
<div className="space-y-4">
|
||||||
<div className="flex items-center gap-2">
|
<div className="flex items-center gap-2">
|
||||||
<ChartNoAxesColumnIncreasing className="h-4 w-4" />
|
<ChartNoAxesColumnIncreasing className="h-4 w-4" />
|
||||||
|
|
@ -114,11 +268,14 @@ export function TokenUsage({ tokenData }: TokenUsageProps) {
|
||||||
<div className="flex items-center gap-1.5">
|
<div className="flex items-center gap-1.5">
|
||||||
<Coins className="h-3 w-3 text-amber-500 dark:text-amber-400" />
|
<Coins className="h-3 w-3 text-amber-500 dark:text-amber-400" />
|
||||||
<span className="text-muted-foreground text-xs font-medium">
|
<span className="text-muted-foreground text-xs font-medium">
|
||||||
Cost
|
{hasModelData ? "Estimated Cost" : "Cost"}
|
||||||
</span>
|
</span>
|
||||||
</div>
|
</div>
|
||||||
<span className="text-sm font-semibold">
|
<span className="text-sm font-semibold">
|
||||||
${metrics.totalCost.toFixed(2)}
|
$
|
||||||
|
{(hasModelData ? totalModelCost : metrics.totalCost).toFixed(
|
||||||
|
2,
|
||||||
|
)}
|
||||||
</span>
|
</span>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
|
@ -150,6 +307,112 @@ export function TokenUsage({ tokenData }: TokenUsageProps) {
|
||||||
)}
|
)}
|
||||||
</div>
|
</div>
|
||||||
</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>
|
</div>
|
||||||
</HoverCardContent>
|
</HoverCardContent>
|
||||||
</HoverCard>
|
</HoverCard>
|
||||||
|
|
|
||||||
19
packages/shared/jest.config.js
Normal file
19
packages/shared/jest.config.js
Normal 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"],
|
||||||
|
};
|
||||||
|
|
@ -14,7 +14,10 @@
|
||||||
"lint": "eslint .",
|
"lint": "eslint .",
|
||||||
"lint:fix": "eslint . --fix",
|
"lint:fix": "eslint . --fix",
|
||||||
"format": "prettier --write .",
|
"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": {
|
"dependencies": {
|
||||||
"@langchain/core": "^0.3.65",
|
"@langchain/core": "^0.3.65",
|
||||||
|
|
@ -27,8 +30,10 @@
|
||||||
"devDependencies": {
|
"devDependencies": {
|
||||||
"@eslint/eslintrc": "^3.1.0",
|
"@eslint/eslintrc": "^3.1.0",
|
||||||
"@eslint/js": "^9.19.0",
|
"@eslint/js": "^9.19.0",
|
||||||
|
"@jest/globals": "^29.7.0",
|
||||||
"@octokit/types": "^12.0.0",
|
"@octokit/types": "^12.0.0",
|
||||||
"@tsconfig/recommended": "^1.0.8",
|
"@tsconfig/recommended": "^1.0.8",
|
||||||
|
"@types/jest": "^29.5.0",
|
||||||
"@types/node": "^22.13.5",
|
"@types/node": "^22.13.5",
|
||||||
"dotenv": "^16.4.7",
|
"dotenv": "^16.4.7",
|
||||||
"eslint": "^9.19.0",
|
"eslint": "^9.19.0",
|
||||||
|
|
@ -36,7 +41,9 @@
|
||||||
"eslint-plugin-import": "^2.27.5",
|
"eslint-plugin-import": "^2.27.5",
|
||||||
"eslint-plugin-no-instanceof": "^1.0.1",
|
"eslint-plugin-no-instanceof": "^1.0.1",
|
||||||
"eslint-plugin-prettier": "^4.2.1",
|
"eslint-plugin-prettier": "^4.2.1",
|
||||||
|
"jest": "^29.7.0",
|
||||||
"prettier": "^3.5.2",
|
"prettier": "^3.5.2",
|
||||||
|
"ts-jest": "^29.1.0",
|
||||||
"typescript": "~5.7.2",
|
"typescript": "~5.7.2",
|
||||||
"typescript-eslint": "^8.22.0"
|
"typescript-eslint": "^8.22.0"
|
||||||
},
|
},
|
||||||
|
|
|
||||||
153
packages/shared/src/__tests__/caching.test.ts
Normal file
153
packages/shared/src/__tests__/caching.test.ts
Normal 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
|
||||||
|
});
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
import { CacheMetrics } from "./open-swe/types.js";
|
import { CacheMetrics, ModelTokenData } from "./open-swe/types.js";
|
||||||
|
|
||||||
export function calculateCostSavings(metrics: CacheMetrics): {
|
export function calculateCostSavings(metrics: CacheMetrics): {
|
||||||
totalSavings: number;
|
totalSavings: number;
|
||||||
|
|
@ -48,18 +48,41 @@ export function calculateCostSavings(metrics: CacheMetrics): {
|
||||||
}
|
}
|
||||||
|
|
||||||
export function tokenDataReducer(
|
export function tokenDataReducer(
|
||||||
state: CacheMetrics | undefined,
|
state: ModelTokenData[] | undefined,
|
||||||
update: CacheMetrics,
|
update: ModelTokenData[],
|
||||||
): CacheMetrics {
|
): ModelTokenData[] {
|
||||||
if (!state) {
|
if (!state) {
|
||||||
return update;
|
return update;
|
||||||
}
|
}
|
||||||
return {
|
|
||||||
cacheCreationInputTokens:
|
// Create a map to merge data by model
|
||||||
state.cacheCreationInputTokens + update.cacheCreationInputTokens,
|
const modelMap = new Map<string, ModelTokenData>();
|
||||||
cacheReadInputTokens:
|
|
||||||
state.cacheReadInputTokens + update.cacheReadInputTokens,
|
// Add existing state data to the map
|
||||||
inputTokens: state.inputTokens + update.inputTokens,
|
for (const data of state) {
|
||||||
outputTokens: state.outputTokens + update.outputTokens,
|
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());
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -3,8 +3,8 @@ import { z } from "zod";
|
||||||
import { MessagesZodState } from "@langchain/langgraph";
|
import { MessagesZodState } from "@langchain/langgraph";
|
||||||
import {
|
import {
|
||||||
AgentSession,
|
AgentSession,
|
||||||
CacheMetrics,
|
|
||||||
CustomRules,
|
CustomRules,
|
||||||
|
ModelTokenData,
|
||||||
TargetRepository,
|
TargetRepository,
|
||||||
TaskPlan,
|
TaskPlan,
|
||||||
} from "../types.js";
|
} from "../types.js";
|
||||||
|
|
@ -103,9 +103,9 @@ export const PlannerGraphStateObj = MessagesZodState.extend({
|
||||||
fn: (_state, update) => update,
|
fn: (_state, update) => update,
|
||||||
},
|
},
|
||||||
}),
|
}),
|
||||||
tokenData: withLangGraph(z.custom<CacheMetrics>().optional(), {
|
tokenData: withLangGraph(z.custom<ModelTokenData[]>().optional(), {
|
||||||
reducer: {
|
reducer: {
|
||||||
schema: z.custom<CacheMetrics>().optional(),
|
schema: z.custom<ModelTokenData[]>().optional(),
|
||||||
fn: tokenDataReducer,
|
fn: tokenDataReducer,
|
||||||
},
|
},
|
||||||
}),
|
}),
|
||||||
|
|
|
||||||
|
|
@ -6,8 +6,8 @@ import {
|
||||||
MessagesZodState,
|
MessagesZodState,
|
||||||
} from "@langchain/langgraph";
|
} from "@langchain/langgraph";
|
||||||
import {
|
import {
|
||||||
CacheMetrics,
|
|
||||||
CustomRules,
|
CustomRules,
|
||||||
|
ModelTokenData,
|
||||||
TargetRepository,
|
TargetRepository,
|
||||||
TaskPlan,
|
TaskPlan,
|
||||||
} from "../types.js";
|
} from "../types.js";
|
||||||
|
|
@ -116,9 +116,9 @@ export const ReviewerGraphStateObj = MessagesZodState.extend({
|
||||||
},
|
},
|
||||||
default: () => 0,
|
default: () => 0,
|
||||||
}),
|
}),
|
||||||
tokenData: withLangGraph(z.custom<CacheMetrics>().optional(), {
|
tokenData: withLangGraph(z.custom<ModelTokenData[]>().optional(), {
|
||||||
reducer: {
|
reducer: {
|
||||||
schema: z.custom<CacheMetrics>().optional(),
|
schema: z.custom<ModelTokenData[]>().optional(),
|
||||||
fn: tokenDataReducer,
|
fn: tokenDataReducer,
|
||||||
},
|
},
|
||||||
}),
|
}),
|
||||||
|
|
|
||||||
|
|
@ -34,6 +34,14 @@ export interface CacheMetrics {
|
||||||
outputTokens: number;
|
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 = {
|
export type PlanItem = {
|
||||||
/**
|
/**
|
||||||
* The index of the plan item. This is the order in which
|
* The index of the plan item. This is the order in which
|
||||||
|
|
@ -274,9 +282,9 @@ export const GraphAnnotation = MessagesZodState.extend({
|
||||||
default: () => 0,
|
default: () => 0,
|
||||||
}),
|
}),
|
||||||
|
|
||||||
tokenData: withLangGraph(z.custom<CacheMetrics>().optional(), {
|
tokenData: withLangGraph(z.custom<ModelTokenData[]>().optional(), {
|
||||||
reducer: {
|
reducer: {
|
||||||
schema: z.custom<CacheMetrics>().optional(),
|
schema: z.custom<ModelTokenData[]>().optional(),
|
||||||
fn: tokenDataReducer,
|
fn: tokenDataReducer,
|
||||||
},
|
},
|
||||||
}),
|
}),
|
||||||
|
|
|
||||||
|
|
@ -4181,12 +4181,14 @@ __metadata:
|
||||||
dependencies:
|
dependencies:
|
||||||
"@eslint/eslintrc": ^3.1.0
|
"@eslint/eslintrc": ^3.1.0
|
||||||
"@eslint/js": ^9.19.0
|
"@eslint/js": ^9.19.0
|
||||||
|
"@jest/globals": ^29.7.0
|
||||||
"@langchain/core": ^0.3.65
|
"@langchain/core": ^0.3.65
|
||||||
"@langchain/langgraph": ^0.3.8
|
"@langchain/langgraph": ^0.3.8
|
||||||
"@langchain/langgraph-sdk": ^0.0.95
|
"@langchain/langgraph-sdk": ^0.0.95
|
||||||
"@octokit/rest": ^22.0.0
|
"@octokit/rest": ^22.0.0
|
||||||
"@octokit/types": ^12.0.0
|
"@octokit/types": ^12.0.0
|
||||||
"@tsconfig/recommended": ^1.0.8
|
"@tsconfig/recommended": ^1.0.8
|
||||||
|
"@types/jest": ^29.5.0
|
||||||
"@types/node": ^22.13.5
|
"@types/node": ^22.13.5
|
||||||
dotenv: ^16.4.7
|
dotenv: ^16.4.7
|
||||||
eslint: ^9.19.0
|
eslint: ^9.19.0
|
||||||
|
|
@ -4194,8 +4196,10 @@ __metadata:
|
||||||
eslint-plugin-import: ^2.27.5
|
eslint-plugin-import: ^2.27.5
|
||||||
eslint-plugin-no-instanceof: ^1.0.1
|
eslint-plugin-no-instanceof: ^1.0.1
|
||||||
eslint-plugin-prettier: ^4.2.1
|
eslint-plugin-prettier: ^4.2.1
|
||||||
|
jest: ^29.7.0
|
||||||
jsonwebtoken: ^9.0.2
|
jsonwebtoken: ^9.0.2
|
||||||
prettier: ^3.5.2
|
prettier: ^3.5.2
|
||||||
|
ts-jest: ^29.1.0
|
||||||
typescript: ~5.7.2
|
typescript: ~5.7.2
|
||||||
typescript-eslint: ^8.22.0
|
typescript-eslint: ^8.22.0
|
||||||
zod: ^3.25.32
|
zod: ^3.25.32
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue