Skip to content
Merged
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
874 changes: 423 additions & 451 deletions docs-site/package-lock.json

Large diffs are not rendered by default.

2 changes: 1 addition & 1 deletion docs-site/package.json
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,6 @@
"devDependencies": {
"@cloudflare/workers-types": "^4.20251004.0",
"@typescript/native-preview": "^7.0.0-dev.20251006.1",
"wrangler": "^4.42.0"
"wrangler": "^4.59.1"
}
}
1 change: 1 addition & 0 deletions docs/mcp/server.md
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,7 @@ async def sampling_callback(
annotations=None,
meta=None,
),
meta=None,
)
]
"""
Expand Down
1 change: 0 additions & 1 deletion examples/pydantic_ai_examples/weather_agent_gradio.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,7 +99,6 @@ def select_data(message: gr.SelectData) -> str:
past_messages = gr.State([])
chatbot = gr.Chatbot(
label='Packing Assistant',
type='messages',
avatar_images=(None, 'https://ai.pydantic.dev/img/logo-white.svg'),
examples=[
{'text': 'What is the weather like in Miami?'},
Expand Down
4 changes: 2 additions & 2 deletions examples/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -56,8 +56,8 @@ dependencies = [
"rich>=13.9.2",
"uvicorn>=0.32.0",
"devtools>=0.12.2",
"gradio>=5.9.0",
"mcp[cli]>=1.4.1",
"gradio>=5.31.0",
"mcp[cli]>=1.25.0",
"modal>=1.0.4",
"duckdb>=1.3.2",
"datasets>=4.0.0",
Expand Down
17 changes: 14 additions & 3 deletions pydantic_ai_slim/pydantic_ai/_mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
from collections.abc import Sequence
from typing import Literal

from typing_extensions import assert_never

from . import exceptions, messages

try:
Expand Down Expand Up @@ -33,11 +35,16 @@ def map_from_mcp_params(params: mcp_types.CreateMessageRequestParams) -> list[me
# TODO(Marcelo): We can reuse the `_map_tool_result_part` from the mcp module here.
if isinstance(content, mcp_types.TextContent):
user_part_content: str | Sequence[messages.UserContent] = content.text
else:
# image content
elif isinstance(content, (mcp_types.ImageContent, mcp_types.AudioContent)):
user_part_content = [
messages.BinaryContent(data=base64.b64decode(content.data), media_type=content.mimeType)
]
elif isinstance(content, list):
raise NotImplementedError('list content type is not yet supported')
elif isinstance(content, (mcp_types.ToolUseContent, mcp_types.ToolResultContent)):
raise NotImplementedError(f'{type(content).__name__} cannot be used as user content')
else:
assert_never(content)

request_parts.append(messages.UserPromptPart(content=user_part_content))
else:
Expand All @@ -47,7 +54,10 @@ def map_from_mcp_params(params: mcp_types.CreateMessageRequestParams) -> list[me
pai_messages.append(messages.ModelRequest(parts=request_parts))
request_parts = []

response_parts.append(map_from_sampling_content(content))
if isinstance(content, (mcp_types.TextContent, mcp_types.ImageContent, mcp_types.AudioContent)):
response_parts.append(map_from_sampling_content(content))
Comment thread
DouweM marked this conversation as resolved.
else:
raise NotImplementedError(f'Unsupported assistant content type: {type(content).__name__}')

if response_parts:
pai_messages.append(messages.ModelResponse(parts=response_parts))
Expand Down Expand Up @@ -123,4 +133,5 @@ def map_from_sampling_content(
if isinstance(content, mcp_types.TextContent): # pragma: no branch
return messages.TextPart(content=content.text)
else:
# TODO: Add support for Image/Audio using FilePart.
raise NotImplementedError('Image and Audio responses in sampling are not yet supported')
136 changes: 66 additions & 70 deletions pydantic_ai_slim/pydantic_ai/mcp.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,13 @@
from __future__ import annotations

import base64
import functools
import os
import re
import warnings
from abc import ABC, abstractmethod
from asyncio import Lock
from collections.abc import AsyncIterator, Awaitable, Callable, Sequence
from contextlib import AbstractAsyncContextManager, AsyncExitStack, asynccontextmanager
from contextlib import AsyncExitStack, asynccontextmanager
from dataclasses import dataclass, field, replace
from datetime import timedelta
from pathlib import Path
Expand All @@ -32,7 +31,7 @@
from mcp.client.session import ClientSession, ElicitationFnT, LoggingFnT
from mcp.client.sse import sse_client
from mcp.client.stdio import StdioServerParameters, stdio_client
from mcp.client.streamable_http import GetSessionIdCallback, streamablehttp_client
from mcp.client.streamable_http import streamable_http_client
from mcp.shared import exceptions as mcp_exceptions
from mcp.shared.context import RequestContext
from mcp.shared.message import SessionMessage
Expand Down Expand Up @@ -1113,67 +1112,6 @@ def __init__(
client_info=client_info,
)

@property
@abstractmethod
def _transport_client(
self,
) -> Callable[
...,
AbstractAsyncContextManager[
tuple[
MemoryObjectReceiveStream[SessionMessage | Exception],
MemoryObjectSendStream[SessionMessage],
GetSessionIdCallback,
],
]
| AbstractAsyncContextManager[
tuple[
MemoryObjectReceiveStream[SessionMessage | Exception],
MemoryObjectSendStream[SessionMessage],
]
],
]: ...

@asynccontextmanager
async def client_streams(
self,
) -> AsyncIterator[
tuple[
MemoryObjectReceiveStream[SessionMessage | Exception],
MemoryObjectSendStream[SessionMessage],
]
]:
if self.http_client and self.headers:
raise ValueError('`http_client` is mutually exclusive with `headers`.') # pragma: no cover

transport_client_partial = functools.partial(
self._transport_client,
url=self.url,
timeout=self.timeout,
sse_read_timeout=self.read_timeout,
)

if self.http_client is not None:
# TODO: Clean up once https://github.com/modelcontextprotocol/python-sdk/pull/1177 lands.
@asynccontextmanager
async def httpx_client_factory(
headers: dict[str, str] | None = None,
timeout: httpx.Timeout | None = None,
auth: httpx.Auth | None = None,
) -> AsyncIterator[httpx.AsyncClient]:
assert self.http_client is not None
yield self.http_client

async with transport_client_partial(httpx_client_factory=httpx_client_factory) as (
read_stream,
write_stream,
*_,
):
yield read_stream, write_stream
else:
async with transport_client_partial(headers=self.headers) as (read_stream, write_stream, *_):
yield read_stream, write_stream

def __repr__(self) -> str: # pragma: no cover
repr_args = [
f'url={self.url!r}',
Expand Down Expand Up @@ -1217,9 +1155,47 @@ def __get_pydantic_core_schema__(cls, _: Any, __: Any) -> CoreSchema:
),
)

@property
def _transport_client(self):
return sse_client # pragma: no cover
# sse_client has a hang bug (https://github.com/modelcontextprotocol/python-sdk/issues/1811)
# that prevents testing SSE transport in CI.
# TODO: Remove pragma and add a test
# once https://github.com/modelcontextprotocol/python-sdk/pull/1838 is released.
@asynccontextmanager
async def client_streams( # pragma: no cover
self,
) -> AsyncIterator[
tuple[
MemoryObjectReceiveStream[SessionMessage | Exception],
MemoryObjectSendStream[SessionMessage],
]
]:
if self.http_client and self.headers:
raise ValueError('`http_client` is mutually exclusive with `headers`.')

if self.http_client is not None:

def httpx_client_factory(
headers: dict[str, str] | None = None,
timeout: httpx.Timeout | None = None,
auth: httpx.Auth | None = None,
) -> httpx.AsyncClient:
assert self.http_client is not None
return self.http_client

async with sse_client(
url=self.url,
timeout=self.timeout,
sse_read_timeout=self.read_timeout,
httpx_client_factory=httpx_client_factory,
) as (read_stream, write_stream, *_):
yield read_stream, write_stream
else:
async with sse_client(
url=self.url,
timeout=self.timeout,
sse_read_timeout=self.read_timeout,
headers=self.headers,
) as (read_stream, write_stream, *_):
yield read_stream, write_stream

def __eq__(self, value: object, /) -> bool:
return super().__eq__(value) and isinstance(value, MCPServerSSE) and self.url == value.url
Expand Down Expand Up @@ -1281,9 +1257,29 @@ def __get_pydantic_core_schema__(cls, _: Any, __: Any) -> CoreSchema:
),
)

@property
def _transport_client(self):
return streamablehttp_client
@asynccontextmanager
async def client_streams(
self,
) -> AsyncIterator[
tuple[
MemoryObjectReceiveStream[SessionMessage | Exception],
MemoryObjectSendStream[SessionMessage],
]
]:
if self.http_client and self.headers:
raise ValueError('`http_client` is mutually exclusive with `headers`.')

aexit_stack = AsyncExitStack()
http_client = self.http_client or await aexit_stack.enter_async_context(
httpx.AsyncClient(timeout=httpx.Timeout(self.timeout, read=self.read_timeout), headers=self.headers)
Comment thread
dsfaccini marked this conversation as resolved.
)
read_stream, write_stream, *_ = await aexit_stack.enter_async_context(
streamable_http_client(self.url, http_client=http_client)
)
try:
yield read_stream, write_stream
finally:
await aexit_stack.aclose()

def __eq__(self, value: object, /) -> bool:
return super().__eq__(value) and isinstance(value, MCPServerStreamableHTTP) and self.url == value.url
Expand Down
4 changes: 2 additions & 2 deletions pydantic_ai_slim/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -105,9 +105,9 @@ cli = [
"pyperclip>=1.9.0",
]
# MCP
mcp = ["mcp>=1.18.0"]
mcp = ["mcp>=1.25.0"]
# FastMCP
fastmcp = ["fastmcp>=2.12.0"]
fastmcp = ["fastmcp>=2.14.0"]
# Evals
evals = ["pydantic-evals=={{ version }}"]
# UI
Expand Down
33 changes: 32 additions & 1 deletion tests/test_mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@
ImageContent,
Implementation,
TextContent,
ToolUseContent,
)

from pydantic_ai._mcp import map_from_mcp_params, map_from_model_response, map_from_pai_messages
Expand Down Expand Up @@ -1659,6 +1660,20 @@ def test_map_from_mcp_params_model_response():
)


def test_map_from_mcp_params_unsupported_user_content():
params = CreateMessageRequestParams(
messages=[
SamplingMessage(
role='user',
content=ToolUseContent(type='tool_use', id='123', name='tool', input={}),
),
],
maxTokens=8,
)
with pytest.raises(NotImplementedError, match='ToolUseContent cannot be used as user content'):
map_from_mcp_params(params)


def test_map_from_pai_messages_with_binary_content():
"""Test that map_from_pai_messages correctly converts image and audio content to MCP format.

Expand All @@ -1678,7 +1693,11 @@ def test_map_from_pai_messages_with_binary_content():
assert system_prompt == ''
assert [m.model_dump(by_alias=True) for m in sampling_msgs] == snapshot(
[
{'role': 'user', 'content': {'type': 'text', 'text': 'text message', 'annotations': None, '_meta': None}},
{
'role': 'user',
'content': {'type': 'text', 'text': 'text message', 'annotations': None, '_meta': None},
'_meta': None,
},
{
'role': 'user',
'content': {
Expand All @@ -1688,6 +1707,7 @@ def test_map_from_pai_messages_with_binary_content():
'annotations': None,
'_meta': None,
},
'_meta': None,
},
]
)
Expand Down Expand Up @@ -2234,6 +2254,17 @@ async def test_custom_http_client_not_closed():
assert not custom_http_client.is_closed


async def test_http_client_mutually_exclusive_with_headers():
server = MCPServerStreamableHTTP(
url='https://example.com/mcp',
http_client=cached_async_http_client(),
headers={'Authorization': 'Bearer token'},
)
with pytest.raises(ValueError, match='`http_client` is mutually exclusive with `headers`'):
async with server:
pass


# ============================================================================
# Tool and Resource Caching Tests
# ============================================================================
Expand Down
Loading