|
1 | 1 | """`dispatch_input_request` and `validate_tool_result` are public `ClientSession` API.""" |
2 | 2 |
|
| 3 | +import anyio |
| 4 | +import anyio.from_thread |
3 | 5 | import mcp_types as types |
4 | 6 | import pytest |
| 7 | +from jsonschema.protocols import Validator |
| 8 | +from jsonschema.validators import validator_for |
5 | 9 | from mcp_types import ( |
6 | 10 | CallToolResult, |
7 | 11 | ErrorData, |
@@ -125,3 +129,46 @@ async def on_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParam |
125 | 129 | await client.session.list_tools() |
126 | 130 | with pytest.raises(RuntimeError, match="Invalid structured content returned by tool t"): |
127 | 131 | 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