fix: Rename Task to LLMTask and improve api key required message (#652)

* fix: Rename task to llmtask and improve api key required message

* cr

* cr
This commit is contained in:
Brace Sproul 2025-08-03 13:02:03 -07:00 • committed by GitHub
parent 98b62bdbc0
commit f5e7af5db0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 212 additions and 141 deletions

View file

@ -14,8 +14,8 @@ import { z } from "zod";
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../../utils/llms/index.js"; } from "../../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { Command, END } from "@langchain/langgraph"; import { Command, END } from "@langchain/langgraph";
import { getMessageContentString } from "@open-swe/shared/messages"; import { getMessageContentString } from "@open-swe/shared/messages";
import { import {
@ -108,10 +108,10 @@ export async function classifyMessage(
description: "Respond to the user's message and determine how to route it.", description: "Respond to the user's message and determine how to route it.",
schema, schema,
}; };
const model = await loadModel(config, Task.ROUTER); const model = await loadModel(config, LLMTask.ROUTER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.ROUTER, LLMTask.ROUTER,
); );
const modelWithTools = model.bindTools([respondAndRouteTool], { const modelWithTools = model.bindTools([respondAndRouteTool], {
tool_choice: respondAndRouteTool.name, tool_choice: respondAndRouteTool.name,

View file

@ -4,15 +4,15 @@ import { z } from "zod";
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js"; } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { getMessageString } from "../../../utils/message/content.js"; import { getMessageString } from "../../../utils/message/content.js";
export async function createIssueFieldsFromMessages( export async function createIssueFieldsFromMessages(
messages: BaseMessage[], messages: BaseMessage[],
configurable: GraphConfig["configurable"], configurable: GraphConfig["configurable"],
): Promise<{ title: string; body: string }> { ): Promise<{ title: string; body: string }> {
const model = await loadModel({ configurable }, Task.ROUTER); const model = await loadModel({ configurable }, LLMTask.ROUTER);
const githubIssueTool = { const githubIssueTool = {
name: "create_github_issue", name: "create_github_issue",
description: "Create a new GitHub issue with the given title and body.", description: "Create a new GitHub issue with the given title and body.",
@ -31,7 +31,7 @@ export async function createIssueFieldsFromMessages(
}; };
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
{ configurable }, { configurable },
Task.ROUTER, LLMTask.ROUTER,
); );
const modelWithTools = model const modelWithTools = model
.bindTools([githubIssueTool], { .bindTools([githubIssueTool], {

View file

@ -8,8 +8,8 @@ import { z } from "zod";
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js"; } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { getMissingMessages } from "../../../utils/github/issue-messages.js"; import { getMissingMessages } from "../../../utils/github/issue-messages.js";
import { getMessageString } from "../../../utils/message/content.js"; import { getMessageString } from "../../../utils/message/content.js";
import { isHumanMessage } from "@langchain/core/messages"; import { isHumanMessage } from "@langchain/core/messages";
@ -119,10 +119,10 @@ export async function determineNeedsContext(
): Promise<Command> { ): Promise<Command> {
const [missingMessages, model] = await Promise.all([ const [missingMessages, model] = await Promise.all([
getMissingMessages(state, config), getMissingMessages(state, config),
loadModel(config, Task.ROUTER), loadModel(config, LLMTask.ROUTER),
]); ]);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.ROUTER); const modelName = modelManager.getModelNameForTask(config, LLMTask.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.",
@ -130,7 +130,7 @@ export async function determineNeedsContext(
} }
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.ROUTER, LLMTask.ROUTER,
); );
const modelWithTools = model.bindTools([determineContextTool], { const modelWithTools = model.bindTools([determineContextTool], {
tool_choice: determineContextTool.name, tool_choice: determineContextTool.name,

View file

@ -2,8 +2,8 @@ import {
getModelManager, getModelManager,
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../../utils/llms/index.js"; } from "../../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { import {
createGetURLContentTool, createGetURLContentTool,
createShellTool, createShellTool,
@ -86,12 +86,12 @@ export async function generateAction(
state: PlannerGraphState, state: PlannerGraphState,
config: GraphConfig, config: GraphConfig,
): Promise<PlannerGraphUpdate> { ): Promise<PlannerGraphUpdate> {
const model = await loadModel(config, Task.PLANNER); const model = await loadModel(config, LLMTask.PLANNER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.PLANNER); const modelName = modelManager.getModelNameForTask(config, LLMTask.PLANNER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.PLANNER, LLMTask.PLANNER,
); );
const mcpTools = await getMcpTools(config); const mcpTools = await getMcpTools(config);

View file

@ -5,8 +5,8 @@ import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../../utils/llms/index.js"; } from "../../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { import {
PlannerGraphState, PlannerGraphState,
PlannerGraphUpdate, PlannerGraphUpdate,
@ -55,12 +55,12 @@ export async function generatePlan(
state: PlannerGraphState, state: PlannerGraphState,
config: GraphConfig, config: GraphConfig,
): Promise<PlannerGraphUpdate> { ): Promise<PlannerGraphUpdate> {
const model = await loadModel(config, Task.PLANNER); const model = await loadModel(config, LLMTask.PLANNER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.PLANNER); const modelName = modelManager.getModelNameForTask(config, LLMTask.PLANNER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.PLANNER, LLMTask.PLANNER,
); );
const sessionPlanTool = createSessionPlanToolFields(); const sessionPlanTool = createSessionPlanToolFields();
const modelWithTools = model.bindTools([sessionPlanTool], { const modelWithTools = model.bindTools([sessionPlanTool], {

View file

@ -8,8 +8,8 @@ import {
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js"; } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { getMessageString } from "../../../utils/message/content.js"; import { getMessageString } from "../../../utils/message/content.js";
import { formatUserRequestPrompt } from "../../../utils/user-request.js"; import { formatUserRequestPrompt } from "../../../utils/user-request.js";
import { formatCustomRulesPrompt } from "../../../utils/custom-rules.js"; import { formatCustomRulesPrompt } from "../../../utils/custom-rules.js";
@ -102,12 +102,15 @@ export async function notetaker(
state: PlannerGraphState, state: PlannerGraphState,
config: GraphConfig, config: GraphConfig,
): Promise<PlannerGraphUpdate> { ): Promise<PlannerGraphUpdate> {
const model = await loadModel(config, Task.SUMMARIZER); const model = await loadModel(config, LLMTask.SUMMARIZER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER); const modelName = modelManager.getModelNameForTask(
config,
LLMTask.SUMMARIZER,
);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.SUMMARIZER, LLMTask.SUMMARIZER,
); );
const modelWithTools = model.bindTools([condenseContextTool], { const modelWithTools = model.bindTools([condenseContextTool], {
tool_choice: condenseContextTool.name, tool_choice: condenseContextTool.name,

View file

@ -16,8 +16,8 @@ import {
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js"; } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { FallbackRunnable } from "../../../utils/runtime-fallback.js"; import { FallbackRunnable } from "../../../utils/runtime-fallback.js";
const systemPromptIdentifyChanges = `You are operating as an agentic coding assistant built by LangChain. You've previously been given a task to generate a plan of action for, to address the user's initial request. const systemPromptIdentifyChanges = `You are operating as an agentic coding assistant built by LangChain. You've previously been given a task to generate a plan of action for, to address the user's initial request.
@ -265,10 +265,10 @@ export async function rewritePlan(
throw new Error("No plan change request found."); throw new Error("No plan change request found.");
} }
const model = await loadModel(config, Task.PLANNER); const model = await loadModel(config, LLMTask.PLANNER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.PLANNER, LLMTask.PLANNER,
); );
const tasksToModify = await identifyTasksToModify( const tasksToModify = await identifyTasksToModify(
state, state,

View file

@ -17,8 +17,8 @@ import { getMessageContentString } from "@open-swe/shared/messages";
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js"; } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { z } from "zod"; import { z } from "zod";
import { createLogger, LogLevel } from "../../../utils/logger.js"; import { createLogger, LogLevel } from "../../../utils/logger.js";
import { import {
@ -108,10 +108,10 @@ export async function diagnoseError(
logger.info("The last two tool calls resulted in errors. Diagnosing error."); logger.info("The last two tool calls resulted in errors. Diagnosing error.");
const model = await loadModel(config, Task.SUMMARIZER); const model = await loadModel(config, LLMTask.SUMMARIZER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.SUMMARIZER, LLMTask.SUMMARIZER,
); );
const modelWithTools = model.bindTools([diagnoseErrorTool], { const modelWithTools = model.bindTools([diagnoseErrorTool], {
tool_choice: diagnoseErrorTool.name, tool_choice: diagnoseErrorTool.name,

View file

@ -4,7 +4,8 @@ import {
GraphUpdate, GraphUpdate,
PlanItem, PlanItem,
} from "@open-swe/shared/open-swe/types"; } from "@open-swe/shared/open-swe/types";
import { loadModel, Task } from "../../../utils/llms/index.js"; import { loadModel } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { getMessageContentString } from "@open-swe/shared/messages"; import { getMessageContentString } from "@open-swe/shared/messages";
import { getMessageString } from "../../../utils/message/content.js"; import { getMessageString } from "../../../utils/message/content.js";
import { createLogger, LogLevel } from "../../../utils/logger.js"; import { createLogger, LogLevel } from "../../../utils/logger.js";
@ -40,9 +41,12 @@ export async function generateConclusion(
state: GraphState, state: GraphState,
config: GraphConfig, config: GraphConfig,
): Promise<GraphUpdate> { ): Promise<GraphUpdate> {
const model = await loadModel(config, Task.SUMMARIZER); const model = await loadModel(config, LLMTask.SUMMARIZER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER); const modelName = modelManager.getModelNameForTask(
config,
LLMTask.SUMMARIZER,
);
const userRequestPrompt = formatUserRequestPrompt(state.messages); const userRequestPrompt = formatUserRequestPrompt(state.messages);
const userMessage = `${userRequestPrompt} const userMessage = `${userRequestPrompt}

View file

@ -10,8 +10,8 @@ import {
loadModel, loadModel,
Provider, Provider,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../../utils/llms/index.js"; } from "../../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { import {
createShellTool, createShellTool,
createApplyPatchTool, createApplyPatchTool,
@ -264,10 +264,13 @@ export async function generateAction(
config: GraphConfig, config: GraphConfig,
): Promise<GraphUpdate> { ): Promise<GraphUpdate> {
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); const modelName = modelManager.getModelNameForTask(
config,
LLMTask.PROGRAMMER,
);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.PROGRAMMER, LLMTask.PROGRAMMER,
); );
const markTaskCompletedTool = createMarkTaskCompletedToolFields(); const markTaskCompletedTool = createMarkTaskCompletedToolFields();
const isAnthropicModel = modelName.includes("claude-"); const isAnthropicModel = modelName.includes("claude-");
@ -286,7 +289,7 @@ export async function generateAction(
}, },
); );
const model = await loadModel(config, Task.PROGRAMMER, { const model = await loadModel(config, LLMTask.PROGRAMMER, {
providerTools: providerTools, providerTools: providerTools,
providerMessages: providerMessages, providerMessages: providerMessages,
}); });

View file

@ -21,8 +21,8 @@ import { z } from "zod";
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js"; } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { formatPlanPromptWithSummaries } from "../../../utils/plan-prompt.js"; import { formatPlanPromptWithSummaries } from "../../../utils/plan-prompt.js";
import { formatUserRequestPrompt } from "../../../utils/user-request.js"; import { formatUserRequestPrompt } from "../../../utils/user-request.js";
import { AIMessage, BaseMessage, ToolMessage } from "@langchain/core/messages"; import { AIMessage, BaseMessage, ToolMessage } from "@langchain/core/messages";
@ -138,12 +138,12 @@ 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, LLMTask.ROUTER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.ROUTER); const modelName = modelManager.getModelNameForTask(config, LLMTask.ROUTER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.ROUTER, LLMTask.ROUTER,
); );
const modelWithTool = model.bindTools([openPrTool], { const modelWithTool = model.bindTools([openPrTool], {
tool_choice: openPrTool.name, tool_choice: openPrTool.name,

View file

@ -5,7 +5,8 @@ import {
GraphUpdate, GraphUpdate,
PlanItem, PlanItem,
} from "@open-swe/shared/open-swe/types"; } from "@open-swe/shared/open-swe/types";
import { loadModel, Task } from "../../../utils/llms/index.js"; import { loadModel } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { import {
AIMessage, AIMessage,
BaseMessage, BaseMessage,
@ -148,9 +149,12 @@ export async function summarizeHistory(
state: GraphState, state: GraphState,
config: GraphConfig, config: GraphConfig,
): Promise<GraphUpdate> { ): Promise<GraphUpdate> {
const model = await loadModel(config, Task.SUMMARIZER); const model = await loadModel(config, LLMTask.SUMMARIZER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER); const modelName = modelManager.getModelNameForTask(
config,
LLMTask.SUMMARIZER,
);
const plan = getActivePlanItems(state.taskPlan); const plan = getActivePlanItems(state.taskPlan);
const conversationHistoryToSummarize = await getMessagesSinceLastSummary( const conversationHistoryToSummarize = await getMessagesSinceLastSummary(

View file

@ -9,8 +9,8 @@ import {
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js"; } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { z } from "zod"; import { z } from "zod";
import { import {
getActiveTask, getActiveTask,
@ -148,12 +148,15 @@ export async function updatePlan(
...updatePlanToolCall, ...updatePlanToolCall,
}); });
const model = await loadModel(config, Task.PROGRAMMER); const model = await loadModel(config, LLMTask.PROGRAMMER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.PROGRAMMER); const modelName = modelManager.getModelNameForTask(
config,
LLMTask.PROGRAMMER,
);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.PROGRAMMER, LLMTask.PROGRAMMER,
); );
const modelWithTools = model.bindTools([updatePlanTool], { const modelWithTools = model.bindTools([updatePlanTool], {
tool_choice: updatePlanTool.name, tool_choice: updatePlanTool.name,

View file

@ -17,8 +17,8 @@ import {
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js"; } from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { GraphConfig, PlanItem } from "@open-swe/shared/open-swe/types"; import { GraphConfig, PlanItem } from "@open-swe/shared/open-swe/types";
import { z } from "zod"; import { z } from "zod";
import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js"; import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js";
@ -117,12 +117,12 @@ export async function finalReview(
const completedTool = createCodeReviewMarkTaskCompletedFields(); const completedTool = createCodeReviewMarkTaskCompletedFields();
const incompleteTool = createCodeReviewMarkTaskNotCompleteFields(); const incompleteTool = createCodeReviewMarkTaskNotCompleteFields();
const tools = [completedTool, incompleteTool]; const tools = [completedTool, incompleteTool];
const model = await loadModel(config, Task.REVIEWER); const model = await loadModel(config, LLMTask.REVIEWER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.REVIEWER); const modelName = modelManager.getModelNameForTask(config, LLMTask.REVIEWER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.REVIEWER, LLMTask.REVIEWER,
); );
const modelWithTools = model.bindTools(tools, { const modelWithTools = model.bindTools(tools, {
tool_choice: "any", tool_choice: "any",

View file

@ -3,8 +3,8 @@ import {
loadModel, loadModel,
Provider, Provider,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../../../utils/llms/index.js"; } from "../../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { import {
ReviewerGraphState, ReviewerGraphState,
ReviewerGraphUpdate, ReviewerGraphUpdate,
@ -189,10 +189,10 @@ export async function generateReviewActions(
config: GraphConfig, config: GraphConfig,
): Promise<ReviewerGraphUpdate> { ): Promise<ReviewerGraphUpdate> {
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.REVIEWER); const modelName = modelManager.getModelNameForTask(config, LLMTask.REVIEWER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.REVIEWER, LLMTask.REVIEWER,
); );
const isAnthropicModel = modelName.includes("claude-"); const isAnthropicModel = modelName.includes("claude-");
@ -201,7 +201,7 @@ export async function generateReviewActions(
config, config,
); );
const model = await loadModel(config, Task.REVIEWER, { const model = await loadModel(config, LLMTask.REVIEWER, {
providerTools, providerTools,
providerMessages, providerMessages,
}); });

View file

@ -14,8 +14,8 @@ import { getMessageString } from "../../utils/message/content.js";
import { import {
loadModel, loadModel,
supportsParallelToolCallsParam, supportsParallelToolCallsParam,
Task,
} from "../../utils/llms/index.js"; } from "../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { trackCachePerformance } from "../../utils/caching.js"; import { trackCachePerformance } from "../../utils/caching.js";
import { getModelManager } from "../../utils/llms/model-manager.js"; import { getModelManager } from "../../utils/llms/model-manager.js";
@ -95,12 +95,15 @@ 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, LLMTask.SUMMARIZER);
const modelManager = getModelManager(); const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER); const modelName = modelManager.getModelNameForTask(
config,
LLMTask.SUMMARIZER,
);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam( const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config, config,
Task.SUMMARIZER, LLMTask.SUMMARIZER,
); );
const modelWithTools = model.bindTools([diagnoseErrorTool], { const modelWithTools = model.bindTools([diagnoseErrorTool], {
tool_choice: diagnoseErrorTool.name, tool_choice: diagnoseErrorTool.name,

View file

@ -2,7 +2,8 @@ import { tool } from "@langchain/core/tools";
import { createLogger, LogLevel } from "../../utils/logger.js"; import { createLogger, LogLevel } from "../../utils/logger.js";
import { createSearchDocumentForToolFields } from "@open-swe/shared/open-swe/tools"; import { createSearchDocumentForToolFields } from "@open-swe/shared/open-swe/tools";
import { FireCrawlLoader } from "@langchain/community/document_loaders/web/firecrawl"; import { FireCrawlLoader } from "@langchain/community/document_loaders/web/firecrawl";
import { loadModel, Task } from "../../utils/llms/index.js"; import { loadModel } from "../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { GraphConfig, GraphState } from "@open-swe/shared/open-swe/types"; import { GraphConfig, GraphState } from "@open-swe/shared/open-swe/types";
import { getMessageContentString } from "@open-swe/shared/messages"; import { getMessageContentString } from "@open-swe/shared/messages";
import { DOCUMENT_SEARCH_PROMPT } from "./prompt.js"; import { DOCUMENT_SEARCH_PROMPT } from "./prompt.js";
@ -76,7 +77,7 @@ export function createSearchDocumentForTool(
}; };
} }
const model = await loadModel(config, Task.SUMMARIZER); const model = await loadModel(config, LLMTask.SUMMARIZER);
const searchPrompt = DOCUMENT_SEARCH_PROMPT.replace( const searchPrompt = DOCUMENT_SEARCH_PROMPT.replace(
"{DOCUMENT_PAGE_CONTENT}", "{DOCUMENT_PAGE_CONTENT}",

View file

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

View file

@ -1,13 +1,16 @@
import { GraphConfig } from "@open-swe/shared/open-swe/types"; import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { getModelManager, Provider } from "./model-manager.js"; import { getModelManager, Provider } from "./model-manager.js";
import { FallbackRunnable } from "../runtime-fallback.js"; import { FallbackRunnable } from "../runtime-fallback.js";
import { Task, TASK_TO_CONFIG_DEFAULTS_MAP } from "./constants.js";
import { BindToolsInput } from "@langchain/core/language_models/chat_models"; import { BindToolsInput } from "@langchain/core/language_models/chat_models";
import { BaseMessageLike } from "@langchain/core/messages"; import { BaseMessageLike } from "@langchain/core/messages";
import {
LLMTask,
TASK_TO_CONFIG_DEFAULTS_MAP,
} from "@open-swe/shared/open-swe/llm-task";
export async function loadModel( export async function loadModel(
config: GraphConfig, config: GraphConfig,
task: Task, task: LLMTask,
options?: { options?: {
providerTools?: Record<Provider, BindToolsInput[]>; providerTools?: Record<Provider, BindToolsInput[]>;
providerMessages?: Record<Provider, BaseMessageLike[]>; providerMessages?: Record<Provider, BaseMessageLike[]>;
@ -33,7 +36,7 @@ export const MODELS_NO_PARALLEL_TOOL_CALLING = ["openai:o3", "openai:o3-mini"];
export function supportsParallelToolCallsParam( export function supportsParallelToolCallsParam(
config: GraphConfig, config: GraphConfig,
task: Task, task: LLMTask,
): boolean { ): boolean {
const modelStr = const modelStr =
config.configurable?.[`${task}ModelName`] ?? config.configurable?.[`${task}ModelName`] ??

View file

@ -4,10 +4,12 @@ import {
} from "langchain/chat_models/universal"; } from "langchain/chat_models/universal";
import { GraphConfig } from "@open-swe/shared/open-swe/types"; import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { createLogger, LogLevel } from "../logger.js"; import { createLogger, LogLevel } from "../logger.js";
import { Task } from "./constants.js"; import {
LLMTask,
TASK_TO_CONFIG_DEFAULTS_MAP,
} from "@open-swe/shared/open-swe/llm-task";
import { isAllowedUser } from "@open-swe/shared/github/allowed-users"; import { isAllowedUser } from "@open-swe/shared/github/allowed-users";
import { decryptSecret } from "@open-swe/shared/crypto"; import { decryptSecret } from "@open-swe/shared/crypto";
import { TASK_TO_CONFIG_DEFAULTS_MAP } from "./constants.js";
import { API_KEY_REQUIRED_MESSAGE } from "@open-swe/shared/constants"; import { API_KEY_REQUIRED_MESSAGE } from "@open-swe/shared/constants";
const logger = createLogger(LogLevel.INFO, "ModelManager"); const logger = createLogger(LogLevel.INFO, "ModelManager");
@ -101,7 +103,7 @@ export class ModelManager {
/** /**
* Load a single model (no fallback during loading) * Load a single model (no fallback during loading)
*/ */
async loadModel(graphConfig: GraphConfig, task: Task) { async loadModel(graphConfig: GraphConfig, task: LLMTask) {
const baseConfig = this.getBaseConfigForTask(graphConfig, task); const baseConfig = this.getBaseConfigForTask(graphConfig, task);
const model = await this.initializeModel(baseConfig, graphConfig); const model = await this.initializeModel(baseConfig, graphConfig);
return model; return model;
@ -199,7 +201,7 @@ export class ModelManager {
public getModelConfigs( public getModelConfigs(
config: GraphConfig, config: GraphConfig,
task: Task, task: LLMTask,
selectedModel: ConfigurableModel, selectedModel: ConfigurableModel,
) { ) {
const configs: ModelLoadConfig[] = []; const configs: ModelLoadConfig[] = [];
@ -264,7 +266,7 @@ export class ModelManager {
/** /**
* Get the model name for a task from GraphConfig * Get the model name for a task from GraphConfig
*/ */
public getModelNameForTask(config: GraphConfig, task: Task): string { public getModelNameForTask(config: GraphConfig, task: LLMTask): string {
const baseConfig = this.getBaseConfigForTask(config, task); const baseConfig = this.getBaseConfigForTask(config, task);
return baseConfig.modelName; return baseConfig.modelName;
} }
@ -274,34 +276,34 @@ export class ModelManager {
*/ */
private getBaseConfigForTask( private getBaseConfigForTask(
config: GraphConfig, config: GraphConfig,
task: Task, task: LLMTask,
): ModelLoadConfig { ): ModelLoadConfig {
const taskMap = { const taskMap = {
[Task.PLANNER]: { [LLMTask.PLANNER]: {
modelName: modelName:
config.configurable?.[`${task}ModelName`] ?? config.configurable?.[`${task}ModelName`] ??
TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName, TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName,
temperature: config.configurable?.[`${task}Temperature`] ?? 0, temperature: config.configurable?.[`${task}Temperature`] ?? 0,
}, },
[Task.PROGRAMMER]: { [LLMTask.PROGRAMMER]: {
modelName: modelName:
config.configurable?.[`${task}ModelName`] ?? config.configurable?.[`${task}ModelName`] ??
TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName, TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName,
temperature: config.configurable?.[`${task}Temperature`] ?? 0, temperature: config.configurable?.[`${task}Temperature`] ?? 0,
}, },
[Task.REVIEWER]: { [LLMTask.REVIEWER]: {
modelName: modelName:
config.configurable?.[`${task}ModelName`] ?? config.configurable?.[`${task}ModelName`] ??
TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName, TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName,
temperature: config.configurable?.[`${task}Temperature`] ?? 0, temperature: config.configurable?.[`${task}Temperature`] ?? 0,
}, },
[Task.ROUTER]: { [LLMTask.ROUTER]: {
modelName: modelName:
config.configurable?.[`${task}ModelName`] ?? config.configurable?.[`${task}ModelName`] ??
TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName, TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName,
temperature: config.configurable?.[`${task}Temperature`] ?? 0, temperature: config.configurable?.[`${task}Temperature`] ?? 0,
}, },
[Task.SUMMARIZER]: { [LLMTask.SUMMARIZER]: {
modelName: modelName:
config.configurable?.[`${task}ModelName`] ?? config.configurable?.[`${task}ModelName`] ??
TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName, TASK_TO_CONFIG_DEFAULTS_MAP[task].modelName,
@ -341,29 +343,29 @@ export class ModelManager {
*/ */
private getDefaultModelForProvider( private getDefaultModelForProvider(
provider: Provider, provider: Provider,
task: Task, task: LLMTask,
): ModelLoadConfig | null { ): ModelLoadConfig | null {
const defaultModels: Record<Provider, Record<Task, string>> = { const defaultModels: Record<Provider, Record<LLMTask, string>> = {
anthropic: { anthropic: {
[Task.PLANNER]: "claude-sonnet-4-0", [LLMTask.PLANNER]: "claude-sonnet-4-0",
[Task.PROGRAMMER]: "claude-sonnet-4-0", [LLMTask.PROGRAMMER]: "claude-sonnet-4-0",
[Task.REVIEWER]: "claude-sonnet-4-0", [LLMTask.REVIEWER]: "claude-sonnet-4-0",
[Task.ROUTER]: "claude-3-5-haiku-latest", [LLMTask.ROUTER]: "claude-3-5-haiku-latest",
[Task.SUMMARIZER]: "claude-sonnet-4-0", [LLMTask.SUMMARIZER]: "claude-sonnet-4-0",
}, },
"google-genai": { "google-genai": {
[Task.PLANNER]: "gemini-2.5-flash", [LLMTask.PLANNER]: "gemini-2.5-flash",
[Task.PROGRAMMER]: "gemini-2.5-pro", [LLMTask.PROGRAMMER]: "gemini-2.5-pro",
[Task.REVIEWER]: "gemini-2.5-flash", [LLMTask.REVIEWER]: "gemini-2.5-flash",
[Task.ROUTER]: "gemini-2.5-flash", [LLMTask.ROUTER]: "gemini-2.5-flash",
[Task.SUMMARIZER]: "gemini-2.5-pro", [LLMTask.SUMMARIZER]: "gemini-2.5-pro",
}, },
openai: { openai: {
[Task.PLANNER]: "o3", [LLMTask.PLANNER]: "o3",
[Task.PROGRAMMER]: "gpt-4.1", [LLMTask.PROGRAMMER]: "gpt-4.1",
[Task.REVIEWER]: "o3", [LLMTask.REVIEWER]: "o3",
[Task.ROUTER]: "gpt-4o-mini", [LLMTask.ROUTER]: "gpt-4o-mini",
[Task.SUMMARIZER]: "gpt-4.1-mini", [LLMTask.SUMMARIZER]: "gpt-4.1-mini",
}, },
}; };

View file

@ -1,5 +1,6 @@
import { GraphConfig } from "@open-swe/shared/open-swe/types"; import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { loadModel, Task } from "../llms/index.js"; import { loadModel } from "../llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { createLogger, LogLevel } from "../logger.js"; import { createLogger, LogLevel } from "../logger.js";
import { DOCUMENT_TOC_GENERATION_PROMPT } from "./prompt.js"; import { DOCUMENT_TOC_GENERATION_PROMPT } from "./prompt.js";
import { getMessageContentString } from "@open-swe/shared/messages"; import { getMessageContentString } from "@open-swe/shared/messages";
@ -29,7 +30,7 @@ export async function handleMcpDocumentationOutput(
}); });
try { try {
const model = await loadModel(config, Task.SUMMARIZER); const model = await loadModel(config, LLMTask.SUMMARIZER);
const systemPrompt = DOCUMENT_TOC_GENERATION_PROMPT.replace( const systemPrompt = DOCUMENT_TOC_GENERATION_PROMPT.replace(
"{DOCUMENT_PAGE_CONTENT}", "{DOCUMENT_PAGE_CONTENT}",

View file

@ -1,5 +1,5 @@
import { GraphConfig } from "@open-swe/shared/open-swe/types"; import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { Task } from "./llms/index.js"; import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { ModelManager, Provider } from "./llms/model-manager.js"; import { ModelManager, Provider } from "./llms/model-manager.js";
import { createLogger, LogLevel } from "./logger.js"; import { createLogger, LogLevel } from "./logger.js";
import { Runnable, RunnableConfig } from "@langchain/core/runnables"; import { Runnable, RunnableConfig } from "@langchain/core/runnables";
@ -45,7 +45,7 @@ export class FallbackRunnable<
> extends ConfigurableModel<RunInput, CallOptions> { > extends ConfigurableModel<RunInput, CallOptions> {
private primaryRunnable: any; private primaryRunnable: any;
private config: GraphConfig; private config: GraphConfig;
private task: Task; private task: LLMTask;
private modelManager: ModelManager; private modelManager: ModelManager;
private providerTools?: Record<Provider, BindToolsInput[]>; private providerTools?: Record<Provider, BindToolsInput[]>;
private providerMessages?: Record<Provider, BaseMessageLike[]>; private providerMessages?: Record<Provider, BaseMessageLike[]>;
@ -53,7 +53,7 @@ export class FallbackRunnable<
constructor( constructor(
primaryRunnable: any, primaryRunnable: any,
config: GraphConfig, config: GraphConfig,
task: Task, task: LLMTask,
modelManager: ModelManager, modelManager: ModelManager,
options?: { options?: {
providerTools?: Record<Provider, BindToolsInput[]>; providerTools?: Record<Provider, BindToolsInput[]>;

View file

@ -8,6 +8,7 @@ import { useUser } from "@/hooks/useUser";
import { useConfigStore, DEFAULT_CONFIG_KEY } from "@/hooks/useConfigStore"; import { useConfigStore, DEFAULT_CONFIG_KEY } from "@/hooks/useConfigStore";
import { isAllowedUser } from "@open-swe/shared/github/allowed-users"; import { isAllowedUser } from "@open-swe/shared/github/allowed-users";
import Link from "next/link"; import Link from "next/link";
import { hasApiKeySet } from "@/lib/api-keys";
const API_KEY_BANNER_DISMISSED_KEY = "api_key_banner_dismissed"; const API_KEY_BANNER_DISMISSED_KEY = "api_key_banner_dismissed";
@ -28,29 +29,15 @@ export function ApiKeyBanner() {
} }
}, []); }, []);
// Don't show banner if: const userIsAllowed = user && isAllowedUser(user.login);
// - Still loading user data
// - User is not authenticated
// - User has dismissed the banner
if (isLoading || !user || dismissed) {
return null;
}
// Check if user is in the allowed list if (
const userIsAllowed = isAllowedUser(user.login); isLoading ||
!user ||
// If user is allowed, they don't need API keys dismissed ||
if (userIsAllowed) { userIsAllowed ||
return null; hasApiKeySet(config)
} ) {
// Check if user has any API keys configured
const apiKeys = config.apiKeys || {};
const hasApiKeys =
apiKeys.anthropicApiKey || apiKeys.openaiApiKey || apiKeys.googleApiKey;
// If user has API keys, don't show banner
if (hasApiKeys) {
return null; return null;
} }

View file

@ -20,6 +20,9 @@ import {
} from "@open-swe/shared/constants"; } from "@open-swe/shared/constants";
import { ManagerGraphUpdate } from "@open-swe/shared/open-swe/manager/types"; import { ManagerGraphUpdate } from "@open-swe/shared/open-swe/manager/types";
import { useDraftStorage } from "@/hooks/useDraftStorage"; import { useDraftStorage } from "@/hooks/useDraftStorage";
import { hasApiKeySet } from "@/lib/api-keys";
import { useUser } from "@/hooks/useUser";
import { isAllowedUser } from "@open-swe/shared/github/allowed-users";
interface TerminalInputProps { interface TerminalInputProps {
placeholder?: string; placeholder?: string;
@ -36,6 +39,24 @@ interface TerminalInputProps {
draftToLoad?: string; draftToLoad?: string;
} }
const MISSING_API_KEYS_TOAST_CONTENT = (
<p>
{API_KEY_REQUIRED_MESSAGE} Please add your API key(s) in{" "}
<a
className="text-blue-500 underline underline-offset-1 dark:text-blue-400"
href="/settings?tab=api-keys"
>
settings
</a>
</p>
);
const MISSING_API_KEYS_TOAST_OPTIONS = {
richColors: true,
duration: 30_000,
closeButton: true,
};
export function TerminalInput({ export function TerminalInput({
placeholder = "Enter your command...", placeholder = "Enter your command...",
disabled = false, disabled = false,
@ -55,6 +76,7 @@ export function TerminalInput({
const { getConfig } = useConfigStore(); const { getConfig } = useConfigStore();
const { selectedRepository } = useGitHubAppProvider(); const { selectedRepository } = useGitHubAppProvider();
const [loading, setLoading] = useState(false); const [loading, setLoading] = useState(false);
const { user, isLoading: isUserLoading } = useUser();
const stream = useStream<GraphState>({ const stream = useStream<GraphState>({
apiUrl, apiUrl,
@ -71,8 +93,28 @@ export function TerminalInput({
}); });
return; return;
} }
if (!user) {
toast.error("User not found. Please sign in first", {
richColors: true,
closeButton: true,
});
return;
}
const defaultConfig = getConfig(DEFAULT_CONFIG_KEY);
if (!isAllowedUser(user.login) && !hasApiKeySet(defaultConfig)) {
toast.error(
MISSING_API_KEYS_TOAST_CONTENT,
MISSING_API_KEYS_TOAST_OPTIONS,
);
}
setLoading(true); setLoading(true);
const trimmedMessage = message.trim(); const trimmedMessage = message.trim();
if (trimmedMessage.length > 0 || contentBlocks.length > 0) { if (trimmedMessage.length > 0 || contentBlocks.length > 0) {
const newHumanMessage = new HumanMessage({ const newHumanMessage = new HumanMessage({
id: uuidv4(), id: uuidv4(),
@ -91,6 +133,7 @@ export function TerminalInput({
targetRepository: selectedRepository, targetRepository: selectedRepository,
autoAcceptPlan, autoAcceptPlan,
}; };
const run = await stream.client.runs.create( const run = await stream.client.runs.create(
newThreadId, newThreadId,
MANAGER_GRAPH_ID, MANAGER_GRAPH_ID,
@ -99,7 +142,7 @@ export function TerminalInput({
config: { config: {
recursion_limit: 400, recursion_limit: 400,
configurable: { configurable: {
...getConfig(DEFAULT_CONFIG_KEY), ...defaultConfig,
}, },
}, },
ifNotExists: "create", ifNotExists: "create",
@ -144,20 +187,8 @@ export function TerminalInput({
e.message.includes(API_KEY_REQUIRED_MESSAGE) e.message.includes(API_KEY_REQUIRED_MESSAGE)
) { ) {
toast.error( toast.error(
<p> MISSING_API_KEYS_TOAST_CONTENT,
{API_KEY_REQUIRED_MESSAGE} Please add your API key(s) in{" "} MISSING_API_KEYS_TOAST_OPTIONS,
<a
className="text-blue-500 underline underline-offset-1 dark:text-blue-400"
href="/settings?tab=api-keys"
>
settings
</a>
</p>,
{
richColors: true,
duration: 30_000,
closeButton: true,
},
); );
} }
} finally { } finally {
@ -205,7 +236,9 @@ export function TerminalInput({
<Button <Button
onClick={handleSend} onClick={handleSend}
disabled={disabled || !message.trim() || !selectedRepository} disabled={
disabled || !message.trim() || !selectedRepository || isUserLoading
}
size="icon" size="icon"
variant="brand" variant="brand"
className="ml-auto size-8 rounded-full border border-white/20 transition-all duration-200 hover:border-white/30 disabled:border-transparent" className="ml-auto size-8 rounded-full border border-white/20 transition-all duration-200 hover:border-white/30 disabled:border-transparent"

View file

@ -0,0 +1,25 @@
export function hasApiKeySet(config: Record<string, any>) {
const modelNameKeys = Object.keys(config).filter((key) =>
key.endsWith("ModelName"),
);
const enabledProviders = modelNameKeys
.map((key) => config[key])
.map((p) => p.split(":")[0]);
const apiKeys = config.apiKeys || {};
// No providers enabled means user is using default model: anthropic
if (enabledProviders.length === 0 && !apiKeys.anthropicApiKey) {
return false;
}
if (
(enabledProviders.includes("anthropic") && !apiKeys.anthropicApiKey) ||
(enabledProviders.includes("openai") && !apiKeys.openaiApiKey) ||
(enabledProviders.includes("google-genai") && !apiKeys.googleApiKey)
) {
return false;
}
return true;
}

View file

@ -1,4 +1,4 @@
export enum Task { export enum LLMTask {
/** /**
* Used for programmer tasks. This includes: writing code, * Used for programmer tasks. This includes: writing code,
* generating plans, taking context gathering actions, etc. * generating plans, taking context gathering actions, etc.
@ -28,23 +28,23 @@ export enum Task {
} }
export const TASK_TO_CONFIG_DEFAULTS_MAP = { export const TASK_TO_CONFIG_DEFAULTS_MAP = {
[Task.PLANNER]: { [LLMTask.PLANNER]: {
modelName: "anthropic:claude-sonnet-4-0", modelName: "anthropic:claude-sonnet-4-0",
temperature: 0, temperature: 0,
}, },
[Task.PROGRAMMER]: { [LLMTask.PROGRAMMER]: {
modelName: "anthropic:claude-sonnet-4-0", modelName: "anthropic:claude-sonnet-4-0",
temperature: 0, temperature: 0,
}, },
[Task.REVIEWER]: { [LLMTask.REVIEWER]: {
modelName: "anthropic:claude-sonnet-4-0", modelName: "anthropic:claude-sonnet-4-0",
temperature: 0, temperature: 0,
}, },
[Task.ROUTER]: { [LLMTask.ROUTER]: {
modelName: "anthropic:claude-3-5-haiku-latest", modelName: "anthropic:claude-3-5-haiku-latest",
temperature: 0, temperature: 0,
}, },
[Task.SUMMARIZER]: { [LLMTask.SUMMARIZER]: {
modelName: "anthropic:claude-3-5-haiku-latest", modelName: "anthropic:claude-3-5-haiku-latest",
temperature: 0, temperature: 0,
}, },