mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 10:23:14 +00:00
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:
parent
98b62bdbc0
commit
f5e7af5db0
26 changed files with 212 additions and 141 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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], {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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], {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -1,3 +1,2 @@
|
|||
export * from "./load-model.js";
|
||||
export * from "./constants.js";
|
||||
export * from "./model-manager.js";
|
||||
|
|
|
|||
|
|
@ -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`] ??
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
},
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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[]>;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
25
apps/web/src/lib/api-keys.ts
Normal file
25
apps/web/src/lib/api-keys.ts
Normal 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;
|
||||
}
|
||||
|
|
@ -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,
|
||||
},
|
||||
Loading…
Add table
Reference in a new issue