mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-05 06:02:15 +00:00
[open-swe] feat: integrate MCPDoc Server into planner and programmer (#349)
* feat: added LangGraph's MCP server * chore: code cleaning * fix: packages fix * feat: mcp servers user configurable * chore: code cleaning * chore: formatting fixes Co-authored-by: Brace Sproul <braceasproul@gmail.com> * fix: types for http streamable mcp * chore: yarn file updated * chore: linting errors * cr * Update apps/open-swe/src/utils/mcp-client.ts --------- Co-authored-by: Brace Sproul <braceasproul@gmail.com>
This commit is contained in:
parent
9082b86799
commit
2a05332693
11 changed files with 1390 additions and 1748 deletions
|
|
@ -29,6 +29,7 @@
|
||||||
"@langchain/google-genai": "^0.2.9",
|
"@langchain/google-genai": "^0.2.9",
|
||||||
"@langchain/langgraph": "^0.3.3",
|
"@langchain/langgraph": "^0.3.3",
|
||||||
"@langchain/langgraph-sdk": "^0.0.85",
|
"@langchain/langgraph-sdk": "^0.0.85",
|
||||||
|
"@langchain/mcp-adapters": "^0.5.2",
|
||||||
"@langchain/openai": "^0.5.10",
|
"@langchain/openai": "^0.5.10",
|
||||||
"@mendable/firecrawl-js": "^1.29.1",
|
"@mendable/firecrawl-js": "^1.29.1",
|
||||||
"@octokit/app": "^16.0.1",
|
"@octokit/app": "^16.0.1",
|
||||||
|
|
|
||||||
|
|
@ -22,6 +22,7 @@ import { getTaskPlanFromIssue } from "../../../../utils/github/issue-task.js";
|
||||||
import { createRgTool } from "../../../../tools/rg.js";
|
import { createRgTool } from "../../../../tools/rg.js";
|
||||||
import { formatCustomRulesPrompt } from "../../../../utils/custom-rules.js";
|
import { formatCustomRulesPrompt } from "../../../../utils/custom-rules.js";
|
||||||
import { createPlannerNotesTool } from "../../../../tools/planner-notes.js";
|
import { createPlannerNotesTool } from "../../../../tools/planner-notes.js";
|
||||||
|
import { getMcpTools } from "../../../../utils/mcp-client.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "GeneratePlanningMessageNode");
|
const logger = createLogger(LogLevel.INFO, "GeneratePlanningMessageNode");
|
||||||
|
|
||||||
|
|
@ -50,12 +51,19 @@ export async function generateAction(
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<PlannerGraphUpdate> {
|
): Promise<PlannerGraphUpdate> {
|
||||||
const model = await loadModel(config, Task.ACTION_GENERATOR);
|
const model = await loadModel(config, Task.ACTION_GENERATOR);
|
||||||
|
const mcpTools = await getMcpTools(config);
|
||||||
|
|
||||||
const tools = [
|
const tools = [
|
||||||
createRgTool(state),
|
createRgTool(state),
|
||||||
createShellTool(state),
|
createShellTool(state),
|
||||||
createPlannerNotesTool(),
|
createPlannerNotesTool(),
|
||||||
createGetURLContentTool(),
|
createGetURLContentTool(),
|
||||||
|
...mcpTools,
|
||||||
];
|
];
|
||||||
|
logger.info(
|
||||||
|
`MCP tools added to Planner: ${mcpTools.map((t) => t.name).join(", ")}`,
|
||||||
|
);
|
||||||
|
|
||||||
const modelWithTools = model.bindTools(tools, {
|
const modelWithTools = model.bindTools(tools, {
|
||||||
tool_choice: "auto",
|
tool_choice: "auto",
|
||||||
parallel_tool_calls: true,
|
parallel_tool_calls: true,
|
||||||
|
|
|
||||||
|
|
@ -9,8 +9,10 @@ import {
|
||||||
PlannerGraphUpdate,
|
PlannerGraphUpdate,
|
||||||
} from "@open-swe/shared/open-swe/planner/types";
|
} from "@open-swe/shared/open-swe/planner/types";
|
||||||
import { createLogger, LogLevel } from "../../../utils/logger.js";
|
import { createLogger, LogLevel } from "../../../utils/logger.js";
|
||||||
import { zodSchemaToString } from "../../../utils/zod-to-string.js";
|
import {
|
||||||
import { formatBadArgsError } from "../../../utils/zod-to-string.js";
|
safeSchemaToString,
|
||||||
|
safeBadArgsError,
|
||||||
|
} from "../../../utils/zod-to-string.js";
|
||||||
import { truncateOutput } from "../../../utils/truncate-outputs.js";
|
import { truncateOutput } from "../../../utils/truncate-outputs.js";
|
||||||
import { createRgTool } from "../../../tools/rg.js";
|
import { createRgTool } from "../../../tools/rg.js";
|
||||||
import {
|
import {
|
||||||
|
|
@ -20,12 +22,13 @@ import {
|
||||||
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { daytonaClient } from "../../../utils/sandbox.js";
|
import { daytonaClient } from "../../../utils/sandbox.js";
|
||||||
import { createPlannerNotesTool } from "../../../tools/planner-notes.js";
|
import { createPlannerNotesTool } from "../../../tools/planner-notes.js";
|
||||||
|
import { getMcpTools } from "../../../utils/mcp-client.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
||||||
|
|
||||||
export async function takeActions(
|
export async function takeActions(
|
||||||
state: PlannerGraphState,
|
state: PlannerGraphState,
|
||||||
_config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<PlannerGraphUpdate> {
|
): Promise<PlannerGraphUpdate> {
|
||||||
const { messages } = state;
|
const { messages } = state;
|
||||||
const lastMessage = messages[messages.length - 1];
|
const lastMessage = messages[messages.length - 1];
|
||||||
|
|
@ -38,12 +41,19 @@ export async function takeActions(
|
||||||
const rgTool = createRgTool(state);
|
const rgTool = createRgTool(state);
|
||||||
const plannerNotesTool = createPlannerNotesTool();
|
const plannerNotesTool = createPlannerNotesTool();
|
||||||
const getURLContentTool = createGetURLContentTool();
|
const getURLContentTool = createGetURLContentTool();
|
||||||
const toolsMap = {
|
|
||||||
[shellTool.name]: shellTool,
|
const mcpTools = await getMcpTools(config);
|
||||||
[rgTool.name]: rgTool,
|
|
||||||
[plannerNotesTool.name]: plannerNotesTool,
|
const allTools = [
|
||||||
[getURLContentTool.name]: getURLContentTool,
|
shellTool,
|
||||||
};
|
rgTool,
|
||||||
|
plannerNotesTool,
|
||||||
|
getURLContentTool,
|
||||||
|
...mcpTools,
|
||||||
|
];
|
||||||
|
const toolsMap = Object.fromEntries(
|
||||||
|
allTools.map((tool) => [tool.name, tool]),
|
||||||
|
);
|
||||||
|
|
||||||
const toolCalls = lastMessage.tool_calls;
|
const toolCalls = lastMessage.tool_calls;
|
||||||
if (!toolCalls?.length) {
|
if (!toolCalls?.length) {
|
||||||
|
|
@ -77,8 +87,13 @@ export async function takeActions(
|
||||||
result: string;
|
result: string;
|
||||||
status: "success" | "error";
|
status: "success" | "error";
|
||||||
};
|
};
|
||||||
result = toolResult.result;
|
if (typeof toolResult === "string") {
|
||||||
toolCallStatus = toolResult.status;
|
result = toolResult;
|
||||||
|
toolCallStatus = "success";
|
||||||
|
} else {
|
||||||
|
result = toolResult.result;
|
||||||
|
toolCallStatus = toolResult.status;
|
||||||
|
}
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
toolCallStatus = "error";
|
toolCallStatus = "error";
|
||||||
if (
|
if (
|
||||||
|
|
@ -87,9 +102,9 @@ export async function takeActions(
|
||||||
) {
|
) {
|
||||||
logger.error("Received tool input did not match expected schema", {
|
logger.error("Received tool input did not match expected schema", {
|
||||||
toolCall,
|
toolCall,
|
||||||
expectedSchema: zodSchemaToString(tool.schema),
|
expectedSchema: safeSchemaToString(tool.schema),
|
||||||
});
|
});
|
||||||
result = formatBadArgsError(tool.schema, toolCall.args);
|
result = safeBadArgsError(tool.schema, toolCall.args, toolCall.name);
|
||||||
} else {
|
} else {
|
||||||
logger.error("Failed to call tool", {
|
logger.error("Failed to call tool", {
|
||||||
...(e instanceof Error
|
...(e instanceof Error
|
||||||
|
|
|
||||||
|
|
@ -28,6 +28,7 @@ import { getTaskPlanFromIssue } from "../../../../utils/github/issue-task.js";
|
||||||
import { createRgTool } from "../../../../tools/rg.js";
|
import { createRgTool } from "../../../../tools/rg.js";
|
||||||
import { createInstallDependenciesTool } from "../../../../tools/install-dependencies.js";
|
import { createInstallDependenciesTool } from "../../../../tools/install-dependencies.js";
|
||||||
import { formatCustomRulesPrompt } from "../../../../utils/custom-rules.js";
|
import { formatCustomRulesPrompt } from "../../../../utils/custom-rules.js";
|
||||||
|
import { getMcpTools } from "../../../../utils/mcp-client.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "GenerateMessageNode");
|
const logger = createLogger(LogLevel.INFO, "GenerateMessageNode");
|
||||||
|
|
||||||
|
|
@ -72,6 +73,8 @@ export async function generateAction(
|
||||||
config: GraphConfig,
|
config: GraphConfig,
|
||||||
): Promise<GraphUpdate> {
|
): Promise<GraphUpdate> {
|
||||||
const model = await loadModel(config, Task.ACTION_GENERATOR);
|
const model = await loadModel(config, Task.ACTION_GENERATOR);
|
||||||
|
const mcpTools = await getMcpTools(config);
|
||||||
|
|
||||||
const tools = [
|
const tools = [
|
||||||
createRgTool(state),
|
createRgTool(state),
|
||||||
createShellTool(state),
|
createShellTool(state),
|
||||||
|
|
@ -79,11 +82,16 @@ export async function generateAction(
|
||||||
createRequestHumanHelpToolFields(),
|
createRequestHumanHelpToolFields(),
|
||||||
createUpdatePlanToolFields(),
|
createUpdatePlanToolFields(),
|
||||||
createGetURLContentTool(),
|
createGetURLContentTool(),
|
||||||
|
...mcpTools,
|
||||||
// Only provide the dependencies installed tool if they're not already installed.
|
// Only provide the dependencies installed tool if they're not already installed.
|
||||||
...(state.dependenciesInstalled
|
...(state.dependenciesInstalled
|
||||||
? []
|
? []
|
||||||
: [createInstallDependenciesTool(state)]),
|
: [createInstallDependenciesTool(state)]),
|
||||||
];
|
];
|
||||||
|
logger.info(
|
||||||
|
`MCP tools added to Programmer: ${mcpTools.map((t) => t.name).join(", ")}`,
|
||||||
|
);
|
||||||
|
|
||||||
const modelWithTools = model.bindTools(tools, {
|
const modelWithTools = model.bindTools(tools, {
|
||||||
tool_choice: "auto",
|
tool_choice: "auto",
|
||||||
parallel_tool_calls: true,
|
parallel_tool_calls: true,
|
||||||
|
|
|
||||||
|
|
@ -15,8 +15,8 @@ import {
|
||||||
getChangedFilesStatus,
|
getChangedFilesStatus,
|
||||||
} from "../../../utils/github/git.js";
|
} from "../../../utils/github/git.js";
|
||||||
import {
|
import {
|
||||||
formatBadArgsError,
|
safeSchemaToString,
|
||||||
zodSchemaToString,
|
safeBadArgsError,
|
||||||
} from "../../../utils/zod-to-string.js";
|
} from "../../../utils/zod-to-string.js";
|
||||||
import { Command } from "@langchain/langgraph";
|
import { Command } from "@langchain/langgraph";
|
||||||
import { truncateOutput } from "../../../utils/truncate-outputs.js";
|
import { truncateOutput } from "../../../utils/truncate-outputs.js";
|
||||||
|
|
@ -26,6 +26,7 @@ import { getRepoAbsolutePath } from "@open-swe/shared/git";
|
||||||
import { shouldDiagnoseError } from "../utils/tool-message-error.js";
|
import { shouldDiagnoseError } from "../utils/tool-message-error.js";
|
||||||
import { createInstallDependenciesTool } from "../../../tools/install-dependencies.js";
|
import { createInstallDependenciesTool } from "../../../tools/install-dependencies.js";
|
||||||
import { createRgTool } from "../../../tools/rg.js";
|
import { createRgTool } from "../../../tools/rg.js";
|
||||||
|
import { getMcpTools } from "../../../utils/mcp-client.js";
|
||||||
|
|
||||||
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
const logger = createLogger(LogLevel.INFO, "TakeAction");
|
||||||
|
|
||||||
|
|
@ -50,13 +51,20 @@ export async function takeAction(
|
||||||
const rgTool = createRgTool(state);
|
const rgTool = createRgTool(state);
|
||||||
const installDependenciesTool = createInstallDependenciesTool(state);
|
const installDependenciesTool = createInstallDependenciesTool(state);
|
||||||
const getURLContentTool = createGetURLContentTool();
|
const getURLContentTool = createGetURLContentTool();
|
||||||
const toolsMap = {
|
|
||||||
[applyPatchTool.name]: applyPatchTool,
|
const mcpTools = await getMcpTools(config);
|
||||||
[shellTool.name]: shellTool,
|
|
||||||
[rgTool.name]: rgTool,
|
const allTools = [
|
||||||
[installDependenciesTool.name]: installDependenciesTool,
|
shellTool,
|
||||||
[getURLContentTool.name]: getURLContentTool,
|
rgTool,
|
||||||
};
|
installDependenciesTool,
|
||||||
|
applyPatchTool,
|
||||||
|
getURLContentTool,
|
||||||
|
...mcpTools,
|
||||||
|
];
|
||||||
|
const toolsMap = Object.fromEntries(
|
||||||
|
allTools.map((tool) => [tool.name, tool]),
|
||||||
|
);
|
||||||
|
|
||||||
const toolCalls = lastMessage.tool_calls;
|
const toolCalls = lastMessage.tool_calls;
|
||||||
if (!toolCalls?.length) {
|
if (!toolCalls?.length) {
|
||||||
|
|
@ -83,8 +91,13 @@ export async function takeAction(
|
||||||
const toolResult: { result: string; status: "success" | "error" } =
|
const toolResult: { result: string; status: "success" | "error" } =
|
||||||
// @ts-expect-error tool.invoke types are weird here...
|
// @ts-expect-error tool.invoke types are weird here...
|
||||||
await tool.invoke(toolCall.args);
|
await tool.invoke(toolCall.args);
|
||||||
result = toolResult.result;
|
if (typeof toolResult === "string") {
|
||||||
toolCallStatus = toolResult.status;
|
result = toolResult;
|
||||||
|
toolCallStatus = "success";
|
||||||
|
} else {
|
||||||
|
result = toolResult.result;
|
||||||
|
toolCallStatus = toolResult.status;
|
||||||
|
}
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
toolCallStatus = "error";
|
toolCallStatus = "error";
|
||||||
if (
|
if (
|
||||||
|
|
@ -93,9 +106,9 @@ export async function takeAction(
|
||||||
) {
|
) {
|
||||||
logger.error("Received tool input did not match expected schema", {
|
logger.error("Received tool input did not match expected schema", {
|
||||||
toolCall,
|
toolCall,
|
||||||
expectedSchema: zodSchemaToString(tool.schema),
|
expectedSchema: safeSchemaToString(tool.schema),
|
||||||
});
|
});
|
||||||
result = formatBadArgsError(tool.schema, toolCall.args);
|
result = safeBadArgsError(tool.schema, toolCall.args, toolCall.name);
|
||||||
} else {
|
} else {
|
||||||
logger.error("Failed to call tool", {
|
logger.error("Failed to call tool", {
|
||||||
...(e instanceof Error
|
...(e instanceof Error
|
||||||
|
|
|
||||||
82
apps/open-swe/src/utils/mcp-client.ts
Normal file
82
apps/open-swe/src/utils/mcp-client.ts
Normal file
|
|
@ -0,0 +1,82 @@
|
||||||
|
import { MultiServerMCPClient } from "@langchain/mcp-adapters";
|
||||||
|
import type { StructuredToolInterface } from "@langchain/core/tools";
|
||||||
|
import { GraphConfig } from "@open-swe/shared/open-swe/types";
|
||||||
|
import {
|
||||||
|
McpServerConfigSchema,
|
||||||
|
McpServers,
|
||||||
|
} from "@open-swe/shared/open-swe/mcp";
|
||||||
|
import { createLogger, LogLevel } from "./logger.js";
|
||||||
|
import { DEFAULT_MCP_SERVERS } from "@open-swe/shared/constants";
|
||||||
|
|
||||||
|
const logger = createLogger(LogLevel.INFO, "MCP Client");
|
||||||
|
|
||||||
|
// Singleton instance of the MCP client
|
||||||
|
let mcpClientInstance: MultiServerMCPClient | null = null;
|
||||||
|
let lastConfigHash: string | null = null;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Returns a shared MCP client instance
|
||||||
|
*/
|
||||||
|
export function mcpClient(mcpServers: McpServers): MultiServerMCPClient {
|
||||||
|
const serversToUse = mcpServers;
|
||||||
|
const configHash = JSON.stringify(serversToUse);
|
||||||
|
|
||||||
|
// Recreate client if configuration changed
|
||||||
|
if (!mcpClientInstance || lastConfigHash !== configHash) {
|
||||||
|
mcpClientInstance = new MultiServerMCPClient({
|
||||||
|
additionalToolNamePrefix: "",
|
||||||
|
mcpServers: serversToUse,
|
||||||
|
});
|
||||||
|
lastConfigHash = configHash;
|
||||||
|
logger.info(
|
||||||
|
`MCP client initialized with ${Object.keys(serversToUse).length} servers: ${Object.keys(serversToUse).join(", ")}`,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
return mcpClientInstance;
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Gets MCP tools with configurable servers
|
||||||
|
* @param config GraphConfig containing optional MCP servers configuration
|
||||||
|
* @returns Array of MCP tools, empty array if error occurs
|
||||||
|
*/
|
||||||
|
export async function getMcpTools(
|
||||||
|
config: GraphConfig,
|
||||||
|
): Promise<StructuredToolInterface[]> {
|
||||||
|
try {
|
||||||
|
// TODO: Remove default MCP servers obj once UI is implemented
|
||||||
|
const mergedServers: McpServers = { ...DEFAULT_MCP_SERVERS };
|
||||||
|
|
||||||
|
const mcpServersConfig = config?.configurable?.["mcpServers"];
|
||||||
|
if (mcpServersConfig) {
|
||||||
|
try {
|
||||||
|
const userServers: McpServers = JSON.parse(mcpServersConfig);
|
||||||
|
for (const serverName in userServers) {
|
||||||
|
const serverConfig = userServers[serverName];
|
||||||
|
if (!serverConfig) continue;
|
||||||
|
try {
|
||||||
|
McpServerConfigSchema.parse(serverConfig);
|
||||||
|
mergedServers[serverName] = serverConfig;
|
||||||
|
} catch (error) {
|
||||||
|
logger.warn(
|
||||||
|
`Failed to parse MCP server configuration for ${serverName}: ${error}. Skipping.`,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
logger.warn(
|
||||||
|
`Failed to parse user MCP servers configuration: ${error}. Using defaults only.`,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!mergedServers) return [];
|
||||||
|
|
||||||
|
const client = mcpClient(mergedServers);
|
||||||
|
const tools = await client.getTools();
|
||||||
|
return tools;
|
||||||
|
} catch (error) {
|
||||||
|
logger.error(`Error getting MCP tools: ${error}`);
|
||||||
|
return [];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
import { z } from "zod";
|
import { z } from "zod";
|
||||||
|
import { truncateOutput } from "./truncate-outputs.js";
|
||||||
|
|
||||||
export function getMissingKeysFromObjectSchema(
|
export function getMissingKeysFromObjectSchema(
|
||||||
schema: z.ZodTypeAny,
|
schema: z.ZodTypeAny,
|
||||||
|
|
@ -68,3 +69,37 @@ export function formatBadArgsError(schema: z.ZodTypeAny, args: any) {
|
||||||
"\n - ",
|
"\n - ",
|
||||||
)}\n`;
|
)}\n`;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export function safeSchemaToString(schema: unknown): string {
|
||||||
|
if (schema instanceof z.ZodType) {
|
||||||
|
try {
|
||||||
|
const result = zodSchemaToString(schema);
|
||||||
|
return truncateOutput(result);
|
||||||
|
} catch {
|
||||||
|
const result = JSON.stringify(schema); // fallback to JSON.stringify
|
||||||
|
return truncateOutput(result);
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
const result = JSON.stringify(schema);
|
||||||
|
return truncateOutput(result);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
export function safeBadArgsError(
|
||||||
|
schema: unknown,
|
||||||
|
args: any,
|
||||||
|
toolName: string,
|
||||||
|
): string {
|
||||||
|
if (schema instanceof z.ZodType) {
|
||||||
|
try {
|
||||||
|
const result = formatBadArgsError(schema, args);
|
||||||
|
return truncateOutput(result);
|
||||||
|
} catch {
|
||||||
|
const schemaString = truncateOutput(JSON.stringify(schema));
|
||||||
|
return `Invalid arguments for tool "${toolName}". Expected schema: ${schemaString}`;
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
const schemaString = truncateOutput(JSON.stringify(schema));
|
||||||
|
return `Invalid arguments for tool "${toolName}". Expected schema: ${schemaString}`;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -20,3 +20,19 @@ export const PROGRAMMER_GRAPH_ID = "programmer";
|
||||||
|
|
||||||
export const GITHUB_USER_ID_HEADER = "x-github-user-id";
|
export const GITHUB_USER_ID_HEADER = "x-github-user-id";
|
||||||
export const GITHUB_USER_LOGIN_HEADER = "x-github-user-login";
|
export const GITHUB_USER_LOGIN_HEADER = "x-github-user-login";
|
||||||
|
|
||||||
|
export const DEFAULT_MCP_SERVERS = {
|
||||||
|
"langgraph-docs-mcp": {
|
||||||
|
command: "uvx",
|
||||||
|
args: [
|
||||||
|
"--from",
|
||||||
|
"mcpdoc",
|
||||||
|
"mcpdoc",
|
||||||
|
"--urls",
|
||||||
|
"LangGraphPY:https://langchain-ai.github.io/langgraph/llms.txt LangGraphJS:https://langchain-ai.github.io/langgraphjs/llms.txt",
|
||||||
|
"--transport",
|
||||||
|
"stdio",
|
||||||
|
],
|
||||||
|
stderr: "inherit" as const,
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
|
||||||
212
packages/shared/src/open-swe/mcp.ts
Normal file
212
packages/shared/src/open-swe/mcp.ts
Normal file
|
|
@ -0,0 +1,212 @@
|
||||||
|
import { z } from "zod";
|
||||||
|
import type { OAuthClientProvider } from "@modelcontextprotocol/sdk/client/auth.js";
|
||||||
|
|
||||||
|
export const oAuthClientProviderSchema = z.custom<OAuthClientProvider>(
|
||||||
|
(val) => {
|
||||||
|
if (!val || typeof val !== "object") return false;
|
||||||
|
|
||||||
|
// Check required properties and methods exist
|
||||||
|
const requiredMethods = [
|
||||||
|
"redirectUrl",
|
||||||
|
"clientMetadata",
|
||||||
|
"clientInformation",
|
||||||
|
"tokens",
|
||||||
|
"saveTokens",
|
||||||
|
];
|
||||||
|
|
||||||
|
// redirectUrl can be a string, URL, or getter returning string/URL
|
||||||
|
if (!("redirectUrl" in val)) return false;
|
||||||
|
|
||||||
|
// clientMetadata can be an object or getter returning an object
|
||||||
|
if (!("clientMetadata" in val)) return false;
|
||||||
|
|
||||||
|
// Check that required methods exist (they can be functions or getters)
|
||||||
|
for (const method of requiredMethods) {
|
||||||
|
if (!(method in val)) return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
return true;
|
||||||
|
},
|
||||||
|
{
|
||||||
|
message:
|
||||||
|
"Must be a valid OAuthClientProvider implementation with required properties: redirectUrl, clientMetadata, clientInformation, tokens, saveTokens",
|
||||||
|
},
|
||||||
|
);
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Stdio transport restart configuration
|
||||||
|
*/
|
||||||
|
export const stdioRestartSchema = z
|
||||||
|
.object({
|
||||||
|
/**
|
||||||
|
* Whether to automatically restart the process if it exits
|
||||||
|
*/
|
||||||
|
enabled: z
|
||||||
|
.boolean()
|
||||||
|
.describe("Whether to automatically restart the process if it exits")
|
||||||
|
.optional(),
|
||||||
|
/**
|
||||||
|
* Maximum number of restart attempts
|
||||||
|
*/
|
||||||
|
maxAttempts: z
|
||||||
|
.number()
|
||||||
|
.describe("The maximum number of restart attempts")
|
||||||
|
.optional(),
|
||||||
|
/**
|
||||||
|
* Delay in milliseconds between restart attempts
|
||||||
|
*/
|
||||||
|
delayMs: z
|
||||||
|
.number()
|
||||||
|
.describe("The delay in milliseconds between restart attempts")
|
||||||
|
.optional(),
|
||||||
|
})
|
||||||
|
.describe("Configuration for stdio transport restart");
|
||||||
|
|
||||||
|
export const stdioConnectionSchema = z
|
||||||
|
.object({
|
||||||
|
/**
|
||||||
|
* Optional transport type, inferred from the structure of the config if not provided. Included
|
||||||
|
* for compatibility with common MCP client config file formats.
|
||||||
|
*/
|
||||||
|
transport: z.literal("stdio").optional(),
|
||||||
|
/**
|
||||||
|
* Optional transport type, inferred from the structure of the config if not provided. Included
|
||||||
|
* for compatibility with common MCP client config file formats.
|
||||||
|
*/
|
||||||
|
type: z.literal("stdio").optional(),
|
||||||
|
/**
|
||||||
|
* The executable to run the server (e.g. `node`, `npx`, etc)
|
||||||
|
*/
|
||||||
|
command: z.string().describe("The executable to run the server"),
|
||||||
|
/**
|
||||||
|
* Array of command line arguments to pass to the executable
|
||||||
|
*/
|
||||||
|
args: z
|
||||||
|
.array(z.string())
|
||||||
|
.describe("Command line arguments to pass to the executable"),
|
||||||
|
/**
|
||||||
|
* Environment variables to set when spawning the process.
|
||||||
|
*/
|
||||||
|
env: z
|
||||||
|
.record(z.string())
|
||||||
|
.describe("The environment to use when spawning the process")
|
||||||
|
.optional(),
|
||||||
|
/**
|
||||||
|
* The encoding to use when reading from the process
|
||||||
|
*/
|
||||||
|
encoding: z
|
||||||
|
.string()
|
||||||
|
.describe("The encoding to use when reading from the process")
|
||||||
|
.optional(),
|
||||||
|
/**
|
||||||
|
* How to handle stderr of the child process. This matches the semantics of Node's `child_process.spawn`
|
||||||
|
*
|
||||||
|
* The default is "inherit", meaning messages to stderr will be printed to the parent process's stderr.
|
||||||
|
*
|
||||||
|
* @default "inherit"
|
||||||
|
*/
|
||||||
|
stderr: z
|
||||||
|
.union([
|
||||||
|
z.literal("overlapped"),
|
||||||
|
z.literal("pipe"),
|
||||||
|
z.literal("ignore"),
|
||||||
|
z.literal("inherit"),
|
||||||
|
])
|
||||||
|
.optional()
|
||||||
|
.default("inherit"),
|
||||||
|
/**
|
||||||
|
* The working directory to use when spawning the process.
|
||||||
|
*/
|
||||||
|
cwd: z
|
||||||
|
.string()
|
||||||
|
.describe("The working directory to use when spawning the process")
|
||||||
|
.optional(),
|
||||||
|
/**
|
||||||
|
* Additional restart settings
|
||||||
|
*/
|
||||||
|
restart: stdioRestartSchema.optional(),
|
||||||
|
})
|
||||||
|
.describe("Configuration for stdio transport connection");
|
||||||
|
|
||||||
|
export const streamableHttpReconnectSchema = z
|
||||||
|
.object({
|
||||||
|
/**
|
||||||
|
* Whether to automatically reconnect if the connection is lost
|
||||||
|
*/
|
||||||
|
enabled: z
|
||||||
|
.boolean()
|
||||||
|
.describe("Whether to automatically reconnect if the connection is lost")
|
||||||
|
.optional(),
|
||||||
|
/**
|
||||||
|
* Maximum number of reconnection attempts
|
||||||
|
*/
|
||||||
|
maxAttempts: z
|
||||||
|
.number()
|
||||||
|
.describe("The maximum number of reconnection attempts")
|
||||||
|
.optional(),
|
||||||
|
/**
|
||||||
|
* Delay in milliseconds between reconnection attempts
|
||||||
|
*/
|
||||||
|
delayMs: z
|
||||||
|
.number()
|
||||||
|
.describe("The delay in milliseconds between reconnection attempts")
|
||||||
|
.optional(),
|
||||||
|
})
|
||||||
|
.describe("Configuration for streamable HTTP transport reconnection");
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Streamable HTTP transport connection
|
||||||
|
*/
|
||||||
|
export const streamableHttpConnectionSchema = z
|
||||||
|
.object({
|
||||||
|
/**
|
||||||
|
* Optional transport type, inferred from the structure of the config. If "sse", will not attempt
|
||||||
|
* to connect using streamable HTTP.
|
||||||
|
*/
|
||||||
|
transport: z.union([z.literal("http"), z.literal("sse")]).optional(),
|
||||||
|
/**
|
||||||
|
* Optional transport type, inferred from the structure of the config. If "sse", will not attempt
|
||||||
|
* to connect using streamable HTTP.
|
||||||
|
*/
|
||||||
|
type: z.union([z.literal("http"), z.literal("sse")]).optional(),
|
||||||
|
/**
|
||||||
|
* The URL to connect to
|
||||||
|
*/
|
||||||
|
url: z.string().url(),
|
||||||
|
/**
|
||||||
|
* Additional headers to send with the request, useful for authentication
|
||||||
|
*/
|
||||||
|
headers: z.record(z.string()).optional(),
|
||||||
|
/**
|
||||||
|
* OAuth client provider for automatic authentication handling.
|
||||||
|
* When provided, the transport will automatically handle token refresh,
|
||||||
|
* 401 error retries, and OAuth 2.0 flows according to RFC 6750.
|
||||||
|
* This is the recommended approach for authentication instead of manual headers.
|
||||||
|
*/
|
||||||
|
authProvider: oAuthClientProviderSchema.optional(),
|
||||||
|
/**
|
||||||
|
* Additional reconnection settings.
|
||||||
|
*/
|
||||||
|
reconnect: streamableHttpReconnectSchema.optional(),
|
||||||
|
/**
|
||||||
|
* Whether to automatically fallback to SSE if Streamable HTTP is not available or not supported
|
||||||
|
*
|
||||||
|
* @default true
|
||||||
|
*/
|
||||||
|
automaticSSEFallback: z.boolean().optional().default(true),
|
||||||
|
})
|
||||||
|
.describe("Configuration for streamable HTTP transport connection");
|
||||||
|
|
||||||
|
export const McpServerConfigSchema = z.union([
|
||||||
|
stdioConnectionSchema,
|
||||||
|
streamableHttpConnectionSchema,
|
||||||
|
]);
|
||||||
|
|
||||||
|
export type McpServerConfig = z.infer<typeof McpServerConfigSchema>;
|
||||||
|
|
||||||
|
export type McpServers = {
|
||||||
|
/**
|
||||||
|
* Map of server names to their configurations
|
||||||
|
*/
|
||||||
|
[serverName: string]: McpServerConfig;
|
||||||
|
};
|
||||||
|
|
@ -19,6 +19,7 @@ import {
|
||||||
GITHUB_TOKEN_COOKIE,
|
GITHUB_TOKEN_COOKIE,
|
||||||
GITHUB_USER_ID_HEADER,
|
GITHUB_USER_ID_HEADER,
|
||||||
GITHUB_USER_LOGIN_HEADER,
|
GITHUB_USER_LOGIN_HEADER,
|
||||||
|
DEFAULT_MCP_SERVERS,
|
||||||
} from "../constants.js";
|
} from "../constants.js";
|
||||||
import { withLangGraph } from "@langchain/langgraph/zod";
|
import { withLangGraph } from "@langchain/langgraph/zod";
|
||||||
import { BaseMessage } from "@langchain/core/messages";
|
import { BaseMessage } from "@langchain/core/messages";
|
||||||
|
|
@ -413,6 +414,14 @@ export const GraphConfigurationMetadata: {
|
||||||
type: "hidden",
|
type: "hidden",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
mcpServers: {
|
||||||
|
x_open_swe_ui_config: {
|
||||||
|
type: "json",
|
||||||
|
default: JSON.stringify(DEFAULT_MCP_SERVERS),
|
||||||
|
description:
|
||||||
|
"JSON configuration for custom MCP servers. LangGraph docs server is set by default.",
|
||||||
|
},
|
||||||
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
export const GraphConfiguration = z.object({
|
export const GraphConfiguration = z.object({
|
||||||
|
|
@ -589,6 +598,15 @@ export const GraphConfiguration = z.object({
|
||||||
.string()
|
.string()
|
||||||
.optional()
|
.optional()
|
||||||
.langgraph.metadata(GraphConfigurationMetadata[GITHUB_INSTALLATION_NAME]),
|
.langgraph.metadata(GraphConfigurationMetadata[GITHUB_INSTALLATION_NAME]),
|
||||||
|
/**
|
||||||
|
* Custom MCP servers configuration as JSON string. Merges with default servers.
|
||||||
|
* @default Default LangGraph docs MCP server
|
||||||
|
*/
|
||||||
|
mcpServers: z
|
||||||
|
.string()
|
||||||
|
.optional()
|
||||||
|
.default(JSON.stringify(DEFAULT_MCP_SERVERS))
|
||||||
|
.langgraph.metadata(GraphConfigurationMetadata.mcpServers),
|
||||||
});
|
});
|
||||||
|
|
||||||
export type GraphConfig = LangGraphRunnableConfig<
|
export type GraphConfig = LangGraphRunnableConfig<
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue