Skip to content

Commit 3e2e5eb

Browse files
committed
Keep validator cache updates on the event loop
1 parent dcc8afa commit 3e2e5eb

2 files changed

Lines changed: 50 additions & 6 deletions

File tree

‎src/mcp/client/session.py‎

Lines changed: 3 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1140,6 +1140,8 @@ async def validate_tool_result(self, name: str, result: types.CallToolResult) ->
11401140
if validator is None:
11411141
# First compilation lazily reads jsonschema's bundled schemas.
11421142
validator = await anyio.to_thread.run_sync(self._output_schema_validator, name, output_schema)
1143+
if _same_schema(self._tool_output_schemas.get(name), output_schema):
1144+
self._tool_output_validators[name] = validator
11431145

11441146
from jsonschema import exceptions as jsonschema_exceptions
11451147
from referencing.exceptions import Unresolvable
@@ -1174,18 +1176,13 @@ def _output_schema_validator(self, name: str, output_schema: dict[str, Any]) ->
11741176
from jsonschema.validators import validator_for
11751177
from referencing import Registry
11761178

1177-
if (validator := self._tool_output_validators.get(name)) is not None:
1178-
return validator
1179-
11801179
validator_cls = validator_for(output_schema)
11811180
try:
11821181
validator_cls.check_schema(output_schema)
11831182
except SchemaError as e:
11841183
raise RuntimeError(f"Invalid schema for tool {name}: {e}")
11851184
# An explicit empty registry: `$ref`s resolve within the schema document and the bundled metaschemas.
1186-
validator = validator_cls(output_schema, registry=Registry())
1187-
self._tool_output_validators[name] = validator
1188-
return validator
1185+
return validator_cls(output_schema, registry=Registry())
11891186

11901187
async def list_prompts(self, *, params: types.PaginatedRequestParams | None = None) -> types.ListPromptsResult:
11911188
"""Send a prompts/list request.

‎tests/client/test_session_promotions.py‎

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,11 @@
11
"""`dispatch_input_request` and `validate_tool_result` are public `ClientSession` API."""
22

3+
import anyio
4+
import anyio.from_thread
35
import mcp_types as types
46
import pytest
7+
from jsonschema.protocols import Validator
8+
from jsonschema.validators import validator_for
59
from mcp_types import (
610
CallToolResult,
711
ErrorData,
@@ -125,3 +129,46 @@ async def on_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParam
125129
await client.session.list_tools()
126130
with pytest.raises(RuntimeError, match="Invalid structured content returned by tool t"):
127131
await client.session.validate_tool_result("t", integer_result)
132+
133+
134+
@pytest.mark.anyio
135+
async def test_schema_change_during_compilation_does_not_cache_the_old_validator(
136+
monkeypatch: pytest.MonkeyPatch,
137+
) -> None:
138+
"""SDK-defined: a concurrent tool relisting cannot leave a stale compiled validator cached."""
139+
schemas = [
140+
{"type": "object", "properties": {"x": {"type": "integer"}}, "required": ["x"]},
141+
{"type": "object", "properties": {"x": {"type": "string"}}, "required": ["x"]},
142+
]
143+
144+
async def on_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams | None) -> ListToolsResult:
145+
return ListToolsResult(tools=[Tool(name="t", input_schema={"type": "object"}, output_schema=schemas.pop(0))])
146+
147+
compilation_started = anyio.Event()
148+
continue_compilation = anyio.Event()
149+
150+
def paused_validator_for(schema: dict[str, object], default: type[Validator] | None = None) -> type[Validator]:
151+
validator = validator_for(schema) if default is None else validator_for(schema, default=default)
152+
if not compilation_started.is_set():
153+
anyio.from_thread.run_sync(compilation_started.set)
154+
anyio.from_thread.run(continue_compilation.wait)
155+
return validator
156+
157+
monkeypatch.setattr("jsonschema.validators.validator_for", paused_validator_for)
158+
server = Server("test-server", on_list_tools=on_list_tools)
159+
async with Client(server) as client:
160+
await client.session.list_tools()
161+
with anyio.fail_after(5):
162+
async with anyio.create_task_group() as task_group:
163+
task_group.start_soon(
164+
client.session.validate_tool_result,
165+
"t",
166+
CallToolResult(content=[], structured_content={"x": 1}),
167+
)
168+
await compilation_started.wait()
169+
try:
170+
await client.session.list_tools()
171+
finally:
172+
continue_compilation.set()
173+
174+
await client.session.validate_tool_result("t", CallToolResult(content=[], structured_content={"x": "yes"}))

0 commit comments

Comments
 (0)