Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ dependencies = [
"anyio>=4.0.0",
"sniffio>=1.0.0",
"typing_extensions>=4.0.0; python_version<'3.11'",
"mcp>=1.23.0,<2.0.0",
"mcp>=1.23.0,<3.0.0",
]

[project.optional-dependencies]
Expand Down
55 changes: 45 additions & 10 deletions src/claude_agent_sdk/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -384,9 +384,12 @@ def create_sdk_mcp_server(
from mcp.server import Server
from mcp.types import (
AudioContent,
CallToolRequestParams,
CallToolResult,
EmbeddedResource,
ImageContent,
ListToolsResult,
PaginatedRequestParams,
ResourceLink,
TextContent,
Tool,
Expand Down Expand Up @@ -446,14 +449,10 @@ def _build_meta(tool_def: "SdkMcpTool[Any]") -> dict[str, Any] | None:
for tool_def in tools
]

# Register list_tools handler to expose available tools
@server.list_tools() # type: ignore[no-untyped-call,untyped-decorator]
async def list_tools() -> list[Tool]:
"""Return the list of available tools."""
return cached_tool_list

# Register call_tool handler to execute tools
@server.call_tool() # type: ignore[untyped-decorator]
async def call_tool(name: str, arguments: dict[str, Any]) -> Any:
"""Execute a tool by name with given arguments."""
if name not in tool_map:
Expand All @@ -477,11 +476,15 @@ async def call_tool(name: str, arguments: dict[str, Any]) -> Any:
if item_type == "text":
content.append(TextContent(type="text", text=item["text"]))
elif item_type == "image":
# Built from the wire names: MCP SDK v2 renamed the
# field to mime_type and kept mimeType as its alias.
content.append(
ImageContent(
type="image",
data=item["data"],
mimeType=item["mimeType"],
ImageContent.model_validate(
{
"type": "image",
"data": item["data"],
"mimeType": item["mimeType"],
}
)
)
elif item_type == "resource_link":
Expand Down Expand Up @@ -517,8 +520,40 @@ async def call_tool(name: str, arguments: dict[str, Any]) -> Any:
item_type,
)

return CallToolResult(
content=content, isError=result.get("is_error", False)
return CallToolResult.model_validate(
{"content": content, "isError": result.get("is_error", False)}
)

# Register the handlers. MCP SDK v1 uses decorators; v2 registers by
# method name and passes a request context plus parsed params. Only one
# of the two APIs exists on any given install, so go through Any.
registry: Any = server
if hasattr(server, "list_tools"):
registry.list_tools()(list_tools)
registry.call_tool()(call_tool)
else:

async def on_list_tools(ctx: Any, params: Any) -> Any:
return ListToolsResult(tools=await list_tools())

async def on_call_tool(ctx: Any, params: Any) -> Any:
# v1's call_tool decorator reports handler exceptions as an
# error result rather than a protocol error; keep that here.
try:
return await call_tool(params.name, params.arguments or {})
except Exception as e:
return CallToolResult.model_validate(
{
"content": [TextContent(type="text", text=str(e))],
"isError": True,
}
)

registry.add_request_handler(
"tools/list", PaginatedRequestParams, on_list_tools
)
registry.add_request_handler(
"tools/call", CallToolRequestParams, on_call_tool
)

# Return SDK server configuration
Expand Down
70 changes: 36 additions & 34 deletions src/claude_agent_sdk/_internal/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,23 @@
DEFERRING_TASK_TYPES = frozenset({"local_agent", "local_workflow"})


async def _call_sdk_mcp_handler(server: Any, request: Any) -> Any:
"""Invoke an SDK MCP server's handler for a request, or None if unhandled.

MCP SDK v1 keys handlers by request type and wraps the result in a
ServerResult; v2 keys them by method name and returns the result directly.
"""
handlers = getattr(server, "request_handlers", None)
if handlers is not None:
handler = handlers.get(type(request))
return (await handler(request)).root if handler else None
entry = server.get_request_handler(request.method)
if entry is None:
return None
# Handlers registered by create_sdk_mcp_server() ignore the context.
return await entry.handler(None, request.params)


def _convert_hook_output_for_cli(hook_output: dict[str, Any]) -> dict[str, Any]:
"""Convert Python-safe field names to CLI-expected field names.

Expand Down Expand Up @@ -646,30 +663,15 @@ async def _handle_sdk_mcp_request(

elif method == "tools/list":
request = ListToolsRequest(method=method)
handler = server.request_handlers.get(ListToolsRequest)
if handler:
result = await handler(request)
# Convert MCP result to JSONRPC response
tools_data = []
for tool in result.root.tools: # type: ignore[union-attr]
tool_data: dict[str, Any] = {
"name": tool.name,
"description": tool.description,
"inputSchema": (
tool.inputSchema.model_dump()
if hasattr(tool.inputSchema, "model_dump")
else tool.inputSchema
)
if tool.inputSchema
else {},
}
if tool.annotations:
tool_data["annotations"] = tool.annotations.model_dump(
exclude_none=True
)
if tool.meta:
tool_data["_meta"] = tool.meta
tools_data.append(tool_data)
result = await _call_sdk_mcp_handler(server, request)
if result is not None:
# Convert MCP result to JSONRPC response. Dumping by alias
# yields the wire names (inputSchema, _meta, ...) on both
# MCP SDK v1 and v2, which renamed the fields to snake_case.
tools_data = [
tool.model_dump(by_alias=True, exclude_none=True, mode="json")
for tool in result.tools
]
return {
"jsonrpc": "2.0",
"id": message.get("id"),
Expand All @@ -683,24 +685,21 @@ async def _handle_sdk_mcp_request(
name=params.get("name"), arguments=params.get("arguments", {})
),
)
handler = server.request_handlers.get(CallToolRequest)
if handler:
result = await handler(call_request)
result = await _call_sdk_mcp_handler(server, call_request)
if result is not None:
# Convert MCP result to JSONRPC response
content = []
for item in result.root.content: # type: ignore[union-attr]
for item in result.content:
item_type = getattr(item, "type", None)
if item_type == "text":
content.append(
{"type": "text", "text": getattr(item, "text", "")}
)
elif item_type == "image":
content.append(
{
"type": "image",
"data": getattr(item, "data", ""),
"mimeType": getattr(item, "mimeType", ""),
}
item.model_dump(
by_alias=True, exclude_none=True, mode="json"
)
)
elif item_type == "resource_link":
parts = []
Expand Down Expand Up @@ -736,7 +735,10 @@ async def _handle_sdk_mcp_request(
)

response_data = {"content": content}
if hasattr(result.root, "isError") and result.root.isError:
# MCP SDK v2 renamed isError to is_error.
if getattr(result, "isError", None) or getattr(
result, "is_error", None
):
response_data["isError"] = True # type: ignore[assignment]

return {
Expand Down
Loading