Skip to content

Commit 60b1acf

Browse files
author
Jianke LIN
committed
Validate full sampling tool result history
1 parent e942d00 commit 60b1acf

3 files changed

Lines changed: 222 additions & 22 deletions

File tree

src/mcp/server/validation.py

Lines changed: 25 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
This module provides validation logic for sampling and elicitation requests.
44
"""
55

6-
from mcp_types import INVALID_PARAMS, ClientCapabilities, SamplingMessage, Tool, ToolChoice
6+
from mcp_types import INVALID_PARAMS, ClientCapabilities, SamplingMessage, SamplingMessageContentBlock, Tool, ToolChoice
77

88
from mcp.shared.exceptions import MCPError
99

@@ -53,6 +53,7 @@ def validate_tool_use_result_messages(messages: list[SamplingMessage]) -> None:
5353
1. Messages with tool_result content contain ONLY tool_result content
5454
2. tool_result messages are preceded by a message with tool_use
5555
3. tool_result IDs match the tool_use IDs from the previous message
56+
4. Every tool_use message in the history is followed by matching tool_result content
5657
5758
See: https://github.com/modelcontextprotocol/modelcontextprotocol/issues/1577
5859
@@ -65,24 +66,26 @@ def validate_tool_use_result_messages(messages: list[SamplingMessage]) -> None:
6566
if not messages:
6667
return
6768

68-
last_content = messages[-1].content_as_list
69-
has_tool_results = any(c.type == "tool_result" for c in last_content)
70-
71-
previous_content = messages[-2].content_as_list if len(messages) >= 2 else None
72-
has_previous_tool_use = previous_content and any(c.type == "tool_use" for c in previous_content)
73-
74-
if has_tool_results:
75-
# Per spec: "SamplingMessage with tool result content blocks
76-
# MUST NOT contain other content types."
77-
if any(c.type != "tool_result" for c in last_content):
78-
raise ValueError("The last message must contain only tool_result content if any is present")
79-
if previous_content is None:
80-
raise ValueError("tool_result requires a previous message containing tool_use")
81-
if not has_previous_tool_use:
82-
raise ValueError("tool_result blocks do not match any tool_use in the previous message")
83-
84-
if has_previous_tool_use and previous_content:
85-
tool_use_ids = {c.id for c in previous_content if c.type == "tool_use"}
86-
tool_result_ids = {c.tool_use_id for c in last_content if c.type == "tool_result"}
87-
if tool_use_ids != tool_result_ids:
88-
raise ValueError("ids of tool_result blocks and tool_use blocks from previous message do not match")
69+
previous_content: list[SamplingMessageContentBlock] | None = None
70+
for content in (message.content_as_list for message in messages):
71+
has_tool_results = any(c.type == "tool_result" for c in content)
72+
previous_tool_use_ids: set[str] = set()
73+
if previous_content is not None:
74+
previous_tool_use_ids = {c.id for c in previous_content if c.type == "tool_use"}
75+
76+
if has_tool_results:
77+
# Per spec: "SamplingMessage with tool result content blocks
78+
# MUST NOT contain other content types."
79+
if any(c.type != "tool_result" for c in content):
80+
raise ValueError("A message must contain only tool_result content if any is present")
81+
if previous_content is None:
82+
raise ValueError("tool_result requires a previous message containing tool_use")
83+
if not previous_tool_use_ids:
84+
raise ValueError("tool_result blocks do not match any tool_use in the previous message")
85+
86+
if previous_tool_use_ids:
87+
tool_result_ids = {c.tool_use_id for c in content if c.type == "tool_result"}
88+
if previous_tool_use_ids != tool_result_ids:
89+
raise ValueError("ids of tool_result blocks and tool_use blocks from previous message do not match")
90+
91+
previous_content = content

tests/server/test_session.py

Lines changed: 112 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -197,6 +197,118 @@ async def test_send_request_skips_the_surface_gate_when_method_absent_at_version
197197
assert isinstance(result, types.EmptyResult)
198198

199199

200+
@pytest.mark.anyio
201+
async def test_create_message_tool_result_validation():
202+
"""Test tool_use/tool_result validation in create_message."""
203+
dispatcher = StubDispatcher(
204+
result={"role": "assistant", "content": [{"type": "text", "text": "ok"}], "model": "m"}
205+
)
206+
session = _make_session(
207+
dispatcher, capabilities=ClientCapabilities(sampling=SamplingCapability(tools=SamplingToolsCapability()))
208+
)
209+
tool = types.Tool(name="test_tool", input_schema={"type": "object"})
210+
text = types.TextContent(type="text", text="hello")
211+
tool_use = types.ToolUseContent(type="tool_use", id="call_1", name="test_tool", input={})
212+
tool_result = types.ToolResultContent(type="tool_result", tool_use_id="call_1", content=[])
213+
214+
# Case 1: tool_result mixed with other content
215+
with pytest.raises(ValueError, match="only tool_result content"):
216+
await session.create_message(
217+
messages=[
218+
types.SamplingMessage(role="user", content=text),
219+
types.SamplingMessage(role="assistant", content=tool_use),
220+
types.SamplingMessage(role="user", content=[tool_result, text]),
221+
],
222+
max_tokens=100,
223+
tools=[tool],
224+
)
225+
226+
# Case 2: tool_result without previous message
227+
with pytest.raises(ValueError, match="requires a previous message"):
228+
await session.create_message(
229+
messages=[types.SamplingMessage(role="user", content=tool_result)],
230+
max_tokens=100,
231+
tools=[tool],
232+
)
233+
234+
# Case 3: tool_result without previous tool_use
235+
with pytest.raises(ValueError, match="do not match any tool_use"):
236+
await session.create_message(
237+
messages=[
238+
types.SamplingMessage(role="user", content=text),
239+
types.SamplingMessage(role="user", content=tool_result),
240+
],
241+
max_tokens=100,
242+
tools=[tool],
243+
)
244+
245+
# Case 4: mismatched tool IDs
246+
with pytest.raises(ValueError, match="ids of tool_result blocks and tool_use blocks"):
247+
await session.create_message(
248+
messages=[
249+
types.SamplingMessage(role="user", content=text),
250+
types.SamplingMessage(role="assistant", content=tool_use),
251+
types.SamplingMessage(
252+
role="user",
253+
content=types.ToolResultContent(type="tool_result", tool_use_id="wrong_id", content=[]),
254+
),
255+
],
256+
max_tokens=100,
257+
tools=[tool],
258+
)
259+
260+
# Case 4b: earlier mismatched tool result with a later plain message
261+
with pytest.raises(ValueError, match="ids of tool_result blocks and tool_use blocks"):
262+
await session.create_message(
263+
messages=[
264+
types.SamplingMessage(role="assistant", content=tool_use),
265+
types.SamplingMessage(
266+
role="user",
267+
content=types.ToolResultContent(type="tool_result", tool_use_id="wrong_id", content=[]),
268+
),
269+
types.SamplingMessage(role="assistant", content=text),
270+
],
271+
max_tokens=100,
272+
tools=[tool],
273+
)
274+
275+
# Case 5: text-only message with tools (no tool_results) - passes validation
276+
await session.create_message(
277+
messages=[types.SamplingMessage(role="user", content=text)],
278+
max_tokens=100,
279+
tools=[tool],
280+
)
281+
282+
# Case 6: valid matching tool_result/tool_use IDs - passes validation
283+
await session.create_message(
284+
messages=[
285+
types.SamplingMessage(role="user", content=text),
286+
types.SamplingMessage(role="assistant", content=tool_use),
287+
types.SamplingMessage(role="user", content=tool_result),
288+
],
289+
max_tokens=100,
290+
tools=[tool],
291+
)
292+
293+
# Case 7: validation runs even without `tools` parameter
294+
# (tool loop continuation may omit tools while containing tool_result)
295+
with pytest.raises(ValueError, match="do not match any tool_use"):
296+
await session.create_message(
297+
messages=[
298+
types.SamplingMessage(role="user", content=text),
299+
types.SamplingMessage(role="user", content=tool_result),
300+
],
301+
max_tokens=100,
302+
)
303+
304+
# Case 8: empty messages list - skips validation entirely
305+
no_tools_session = _make_session(
306+
StubDispatcher(result={"role": "assistant", "content": {"type": "text", "text": "ok"}, "model": "m"}),
307+
capabilities=ClientCapabilities(sampling=SamplingCapability(tools=SamplingToolsCapability())),
308+
)
309+
await no_tools_session.create_message(messages=[], max_tokens=100)
310+
311+
200312
@pytest.mark.anyio
201313
async def test_send_request_validates_result_alias_only():
202314
"""Peer results validate alias-only; a snake_case key from the wire is

tests/server/test_validation.py

Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -108,6 +108,27 @@ def test_validate_tool_use_result_messages_raises_when_tool_result_mixed_with_ot
108108
validate_tool_use_result_messages(messages)
109109

110110

111+
def test_validate_tool_use_result_messages_raises_for_earlier_mixed_tool_result() -> None:
112+
"""Raises when an earlier message mixes tool_result with other content."""
113+
messages = [
114+
SamplingMessage(
115+
role="assistant",
116+
content=ToolUseContent(type="tool_use", id="tool-1", name="test", input={}),
117+
),
118+
SamplingMessage(
119+
role="user",
120+
content=[
121+
ToolResultContent(type="tool_result", tool_use_id="tool-1"),
122+
TextContent(type="text", text="also this"),
123+
],
124+
),
125+
SamplingMessage(role="assistant", content=TextContent(type="text", text="done")),
126+
]
127+
128+
with pytest.raises(ValueError, match="only tool_result content"):
129+
validate_tool_use_result_messages(messages)
130+
131+
111132
def test_validate_tool_use_result_messages_raises_when_tool_result_without_previous_tool_use() -> None:
112133
"""Raises when tool_result appears without preceding tool_use."""
113134
messages = [
@@ -146,6 +167,39 @@ def test_validate_tool_use_result_messages_raises_when_tool_result_ids_dont_matc
146167
validate_tool_use_result_messages(messages)
147168

148169

170+
def test_validate_tool_use_result_messages_raises_when_earlier_tool_result_ids_dont_match_tool_use() -> None:
171+
"""Raises when an earlier tool_result does not match the previous tool_use."""
172+
messages = [
173+
SamplingMessage(
174+
role="assistant",
175+
content=ToolUseContent(type="tool_use", id="tool-1", name="test", input={}),
176+
),
177+
SamplingMessage(
178+
role="user",
179+
content=ToolResultContent(type="tool_result", tool_use_id="tool-2"),
180+
),
181+
SamplingMessage(role="assistant", content=TextContent(type="text", text="done")),
182+
]
183+
184+
with pytest.raises(ValueError, match="do not match"):
185+
validate_tool_use_result_messages(messages)
186+
187+
188+
def test_validate_tool_use_result_messages_raises_when_tool_use_is_not_answered() -> None:
189+
"""Raises when a tool_use is followed by a non-tool_result message."""
190+
messages = [
191+
SamplingMessage(
192+
role="assistant",
193+
content=ToolUseContent(type="tool_use", id="tool-1", name="test", input={}),
194+
),
195+
SamplingMessage(role="user", content=TextContent(type="text", text="not a result")),
196+
SamplingMessage(role="assistant", content=TextContent(type="text", text="done")),
197+
]
198+
199+
with pytest.raises(ValueError, match="do not match"):
200+
validate_tool_use_result_messages(messages)
201+
202+
149203
def test_validate_tool_use_result_messages_no_error_when_tool_result_matches_tool_use() -> None:
150204
"""No error when tool_result IDs match tool_use IDs."""
151205
messages = [
@@ -159,3 +213,34 @@ def test_validate_tool_use_result_messages_no_error_when_tool_result_matches_too
159213
),
160214
]
161215
validate_tool_use_result_messages(messages) # Should not raise
216+
217+
218+
def test_validate_tool_use_result_messages_no_error_for_multiple_tool_pairs() -> None:
219+
"""No error when every tool_use in the history has a matching tool_result."""
220+
messages = [
221+
SamplingMessage(role="user", content=TextContent(type="text", text="first")),
222+
SamplingMessage(
223+
role="assistant",
224+
content=ToolUseContent(type="tool_use", id="tool-1", name="test", input={}),
225+
),
226+
SamplingMessage(
227+
role="user",
228+
content=ToolResultContent(type="tool_result", tool_use_id="tool-1"),
229+
),
230+
SamplingMessage(
231+
role="assistant",
232+
content=[
233+
ToolUseContent(type="tool_use", id="tool-2", name="test", input={}),
234+
ToolUseContent(type="tool_use", id="tool-3", name="test", input={}),
235+
],
236+
),
237+
SamplingMessage(
238+
role="user",
239+
content=[
240+
ToolResultContent(type="tool_result", tool_use_id="tool-3"),
241+
ToolResultContent(type="tool_result", tool_use_id="tool-2"),
242+
],
243+
),
244+
]
245+
246+
validate_tool_use_result_messages(messages)

0 commit comments

Comments
 (0)