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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -16,8 +16,8 @@ import {
import {
loadModel,
supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
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.
@ -265,10 +265,10 @@ export async function rewritePlan(
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(
config,
Task.PLANNER,
LLMTask.PLANNER,
);
const tasksToModify = await identifyTasksToModify(
state,

View file

@ -17,8 +17,8 @@ import { getMessageContentString } from "@open-swe/shared/messages";
import {
loadModel,
supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { z } from "zod";
import { createLogger, LogLevel } from "../../../utils/logger.js";
import {
@ -108,10 +108,10 @@ export async function diagnoseError(
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(
config,
Task.SUMMARIZER,
LLMTask.SUMMARIZER,
);
const modelWithTools = model.bindTools([diagnoseErrorTool], {
tool_choice: diagnoseErrorTool.name,

View file

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

View file

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

View file

@ -21,8 +21,8 @@ import { z } from "zod";
import {
loadModel,
supportsParallelToolCallsParam,
Task,
} from "../../../utils/llms/index.js";
import { LLMTask } from "@open-swe/shared/open-swe/llm-task";
import { formatPlanPromptWithSummaries } from "../../../utils/plan-prompt.js";
import { formatUserRequestPrompt } from "../../../utils/user-request.js";
import { AIMessage, BaseMessage, ToolMessage } from "@langchain/core/messages";
@ -138,12 +138,12 @@ export async function openPullRequest(
const openPrTool = createOpenPrToolFields();
// use the router model since this is a simple task that doesn't need an advanced model
const model = await loadModel(config, Task.ROUTER);
const model = await loadModel(config, LLMTask.ROUTER);
const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.ROUTER);
const modelName = modelManager.getModelNameForTask(config, LLMTask.ROUTER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config,
Task.ROUTER,
LLMTask.ROUTER,
);
const modelWithTool = model.bindTools([openPrTool], {
tool_choice: openPrTool.name,

View file

@ -5,7 +5,8 @@ import {
GraphUpdate,
PlanItem,
} 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 {
AIMessage,
BaseMessage,
@ -148,9 +149,12 @@ export async function summarizeHistory(
state: GraphState,
config: GraphConfig,
): Promise<GraphUpdate> {
const model = await loadModel(config, Task.SUMMARIZER);
const model = await loadModel(config, LLMTask.SUMMARIZER);
const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.SUMMARIZER);
const modelName = modelManager.getModelNameForTask(
config,
LLMTask.SUMMARIZER,
);
const plan = getActivePlanItems(state.taskPlan);
const conversationHistoryToSummarize = await getMessagesSinceLastSummary(

View file

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

View file

@ -17,8 +17,8 @@ import {
import {
loadModel,
supportsParallelToolCallsParam,
Task,
} 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 { z } from "zod";
import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js";
@ -117,12 +117,12 @@ export async function finalReview(
const completedTool = createCodeReviewMarkTaskCompletedFields();
const incompleteTool = createCodeReviewMarkTaskNotCompleteFields();
const tools = [completedTool, incompleteTool];
const model = await loadModel(config, Task.REVIEWER);
const model = await loadModel(config, LLMTask.REVIEWER);
const modelManager = getModelManager();
const modelName = modelManager.getModelNameForTask(config, Task.REVIEWER);
const modelName = modelManager.getModelNameForTask(config, LLMTask.REVIEWER);
const modelSupportsParallelToolCallsParam = supportsParallelToolCallsParam(
config,
Task.REVIEWER,
LLMTask.REVIEWER,
);
const modelWithTools = model.bindTools(tools, {
tool_choice: "any",

View file

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

View file

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

View file

@ -2,7 +2,8 @@ import { tool } from "@langchain/core/tools";
import { createLogger, LogLevel } from "../../utils/logger.js";
import { createSearchDocumentForToolFields } from "@open-swe/shared/open-swe/tools";
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 { getMessageContentString } from "@open-swe/shared/messages";
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(
"{DOCUMENT_PAGE_CONTENT}",

View file

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

View file

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

View file

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

View file

@ -1,5 +1,6 @@
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 { DOCUMENT_TOC_GENERATION_PROMPT } from "./prompt.js";
import { getMessageContentString } from "@open-swe/shared/messages";
@ -29,7 +30,7 @@ export async function handleMcpDocumentationOutput(
});
try {
const model = await loadModel(config, Task.SUMMARIZER);
const model = await loadModel(config, LLMTask.SUMMARIZER);
const systemPrompt = DOCUMENT_TOC_GENERATION_PROMPT.replace(
"{DOCUMENT_PAGE_CONTENT}",

View file

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

View file

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

View file

@ -20,6 +20,9 @@ import {
} from "@open-swe/shared/constants";
import { ManagerGraphUpdate } from "@open-swe/shared/open-swe/manager/types";
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 {
placeholder?: string;
@ -36,6 +39,24 @@ interface TerminalInputProps {
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({
placeholder = "Enter your command...",
disabled = false,
@ -55,6 +76,7 @@ export function TerminalInput({
const { getConfig } = useConfigStore();
const { selectedRepository } = useGitHubAppProvider();
const [loading, setLoading] = useState(false);
const { user, isLoading: isUserLoading } = useUser();
const stream = useStream<GraphState>({
apiUrl,
@ -71,8 +93,28 @@ export function TerminalInput({
});
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);
const trimmedMessage = message.trim();
if (trimmedMessage.length > 0 || contentBlocks.length > 0) {
const newHumanMessage = new HumanMessage({
id: uuidv4(),
@ -91,6 +133,7 @@ export function TerminalInput({
targetRepository: selectedRepository,
autoAcceptPlan,
};
const run = await stream.client.runs.create(
newThreadId,
MANAGER_GRAPH_ID,
@ -99,7 +142,7 @@ export function TerminalInput({
config: {
recursion_limit: 400,
configurable: {
...getConfig(DEFAULT_CONFIG_KEY),
...defaultConfig,
},
},
ifNotExists: "create",
@ -144,20 +187,8 @@ export function TerminalInput({
e.message.includes(API_KEY_REQUIRED_MESSAGE)
) {
toast.error(
<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>,
{
richColors: true,
duration: 30_000,
closeButton: true,
},
MISSING_API_KEYS_TOAST_CONTENT,
MISSING_API_KEYS_TOAST_OPTIONS,
);
}
} finally {
@ -205,7 +236,9 @@ export function TerminalInput({
<Button
onClick={handleSend}
disabled={disabled || !message.trim() || !selectedRepository}
disabled={
disabled || !message.trim() || !selectedRepository || isUserLoading
}
size="icon"
variant="brand"
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,
* generating plans, taking context gathering actions, etc.
@ -28,23 +28,23 @@ export enum Task {
}
export const TASK_TO_CONFIG_DEFAULTS_MAP = {
[Task.PLANNER]: {
[LLMTask.PLANNER]: {
modelName: "anthropic:claude-sonnet-4-0",
temperature: 0,
},
[Task.PROGRAMMER]: {
[LLMTask.PROGRAMMER]: {
modelName: "anthropic:claude-sonnet-4-0",
temperature: 0,
},
[Task.REVIEWER]: {
[LLMTask.REVIEWER]: {
modelName: "anthropic:claude-sonnet-4-0",
temperature: 0,
},
[Task.ROUTER]: {
[LLMTask.ROUTER]: {
modelName: "anthropic:claude-3-5-haiku-latest",
temperature: 0,
},
[Task.SUMMARIZER]: {
[LLMTask.SUMMARIZER]: {
modelName: "anthropic:claude-3-5-haiku-latest",
temperature: 0,
},