diff --git a/apps/open-swe/src/graphs/manager/nodes/classify-message/index.ts b/apps/open-swe/src/graphs/manager/nodes/classify-message/index.ts index 200ac96a..e155f61e 100644 --- a/apps/open-swe/src/graphs/manager/nodes/classify-message/index.ts +++ b/apps/open-swe/src/graphs/manager/nodes/classify-message/index.ts @@ -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, diff --git a/apps/open-swe/src/graphs/manager/utils/generate-issue-fields.ts b/apps/open-swe/src/graphs/manager/utils/generate-issue-fields.ts index 2ad73da3..b2bb62fd 100644 --- a/apps/open-swe/src/graphs/manager/utils/generate-issue-fields.ts +++ b/apps/open-swe/src/graphs/manager/utils/generate-issue-fields.ts @@ -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], { diff --git a/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts b/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts index c9f6fdc8..00a6bc8c 100644 --- a/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts +++ b/apps/open-swe/src/graphs/planner/nodes/determine-needs-context.ts @@ -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 { 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, diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts b/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts index e8ef3afd..11cab739 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-message/index.ts @@ -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 { - 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); diff --git a/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts b/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts index 8a9263b1..08df4555 100644 --- a/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts +++ b/apps/open-swe/src/graphs/planner/nodes/generate-plan/index.ts @@ -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 { - 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], { diff --git a/apps/open-swe/src/graphs/planner/nodes/notetaker.ts b/apps/open-swe/src/graphs/planner/nodes/notetaker.ts index cd6f5711..5e4230b5 100644 --- a/apps/open-swe/src/graphs/planner/nodes/notetaker.ts +++ b/apps/open-swe/src/graphs/planner/nodes/notetaker.ts @@ -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 { - 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, diff --git a/apps/open-swe/src/graphs/planner/nodes/rewrite-plan.ts b/apps/open-swe/src/graphs/planner/nodes/rewrite-plan.ts index d9ae0f4a..38d24e99 100644 --- a/apps/open-swe/src/graphs/planner/nodes/rewrite-plan.ts +++ b/apps/open-swe/src/graphs/planner/nodes/rewrite-plan.ts @@ -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, diff --git a/apps/open-swe/src/graphs/programmer/nodes/diagnose-error.ts b/apps/open-swe/src/graphs/programmer/nodes/diagnose-error.ts index f8da05a0..5b1d0b29 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/diagnose-error.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/diagnose-error.ts @@ -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, diff --git a/apps/open-swe/src/graphs/programmer/nodes/generate-conclusion.ts b/apps/open-swe/src/graphs/programmer/nodes/generate-conclusion.ts index b2c01b2c..758260cb 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/generate-conclusion.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/generate-conclusion.ts @@ -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 { - 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} diff --git a/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts b/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts index fe3693f7..f37e04ff 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/generate-message/index.ts @@ -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 { 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, }); diff --git a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts index 89af0f0c..f66042ae 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/open-pr.ts @@ -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, diff --git a/apps/open-swe/src/graphs/programmer/nodes/summarize-history.ts b/apps/open-swe/src/graphs/programmer/nodes/summarize-history.ts index 713ae33d..eb325ff6 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/summarize-history.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/summarize-history.ts @@ -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 { - 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( diff --git a/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts b/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts index 260d1f1b..d2167c0b 100644 --- a/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts +++ b/apps/open-swe/src/graphs/programmer/nodes/update-plan.ts @@ -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, diff --git a/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts b/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts index 2b73867e..0453ca96 100644 --- a/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts +++ b/apps/open-swe/src/graphs/reviewer/nodes/final-review.ts @@ -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", diff --git a/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts b/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts index d4a19570..c3dc7211 100644 --- a/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts +++ b/apps/open-swe/src/graphs/reviewer/nodes/generate-review-actions/index.ts @@ -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 { 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, }); diff --git a/apps/open-swe/src/graphs/shared/diagnose-error.ts b/apps/open-swe/src/graphs/shared/diagnose-error.ts index 352a26a5..459d8843 100644 --- a/apps/open-swe/src/graphs/shared/diagnose-error.ts +++ b/apps/open-swe/src/graphs/shared/diagnose-error.ts @@ -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, diff --git a/apps/open-swe/src/tools/search-documents-for/index.ts b/apps/open-swe/src/tools/search-documents-for/index.ts index dab09533..008afbf2 100644 --- a/apps/open-swe/src/tools/search-documents-for/index.ts +++ b/apps/open-swe/src/tools/search-documents-for/index.ts @@ -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}", diff --git a/apps/open-swe/src/utils/llms/index.ts b/apps/open-swe/src/utils/llms/index.ts index d6143dfd..ea4be988 100644 --- a/apps/open-swe/src/utils/llms/index.ts +++ b/apps/open-swe/src/utils/llms/index.ts @@ -1,3 +1,2 @@ export * from "./load-model.js"; -export * from "./constants.js"; export * from "./model-manager.js"; diff --git a/apps/open-swe/src/utils/llms/load-model.ts b/apps/open-swe/src/utils/llms/load-model.ts index a167be3f..7ff37256 100644 --- a/apps/open-swe/src/utils/llms/load-model.ts +++ b/apps/open-swe/src/utils/llms/load-model.ts @@ -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; providerMessages?: Record; @@ -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`] ?? diff --git a/apps/open-swe/src/utils/llms/model-manager.ts b/apps/open-swe/src/utils/llms/model-manager.ts index 8fbf1340..28d28232 100644 --- a/apps/open-swe/src/utils/llms/model-manager.ts +++ b/apps/open-swe/src/utils/llms/model-manager.ts @@ -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> = { + const defaultModels: Record> = { 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", }, }; diff --git a/apps/open-swe/src/utils/mcp-output/index.ts b/apps/open-swe/src/utils/mcp-output/index.ts index db126df3..7fdd2f36 100644 --- a/apps/open-swe/src/utils/mcp-output/index.ts +++ b/apps/open-swe/src/utils/mcp-output/index.ts @@ -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}", diff --git a/apps/open-swe/src/utils/runtime-fallback.ts b/apps/open-swe/src/utils/runtime-fallback.ts index 5f15c5e2..f03febc7 100644 --- a/apps/open-swe/src/utils/runtime-fallback.ts +++ b/apps/open-swe/src/utils/runtime-fallback.ts @@ -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 { private primaryRunnable: any; private config: GraphConfig; - private task: Task; + private task: LLMTask; private modelManager: ModelManager; private providerTools?: Record; private providerMessages?: Record; @@ -53,7 +53,7 @@ export class FallbackRunnable< constructor( primaryRunnable: any, config: GraphConfig, - task: Task, + task: LLMTask, modelManager: ModelManager, options?: { providerTools?: Record; diff --git a/apps/web/src/components/api-key-banner.tsx b/apps/web/src/components/api-key-banner.tsx index 56b432bc..19706540 100644 --- a/apps/web/src/components/api-key-banner.tsx +++ b/apps/web/src/components/api-key-banner.tsx @@ -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; } diff --git a/apps/web/src/components/v2/terminal-input.tsx b/apps/web/src/components/v2/terminal-input.tsx index 9e8cf3cf..b5ef154b 100644 --- a/apps/web/src/components/v2/terminal-input.tsx +++ b/apps/web/src/components/v2/terminal-input.tsx @@ -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 = ( +

+ {API_KEY_REQUIRED_MESSAGE} Please add your API key(s) in{" "} + + settings + +

+); + +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({ 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( -

- {API_KEY_REQUIRED_MESSAGE} Please add your API key(s) in{" "} - - settings - -

, - { - 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({