mirror of
https://github.com/rohitg00/ai-engineering-from-scratch.git
synced 2026-10-02 01:54:39 +08:00
fix(mcp-08): reuse strict response decoding
This commit is contained in:
@@ -97,7 +97,13 @@ def decode_rpc_response(
|
||||
response: dict[str, Any],
|
||||
expected_id: int | str,
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
if response.get("jsonrpc") != "2.0" or response.get("id") != expected_id:
|
||||
response_id = response.get("id")
|
||||
if (
|
||||
response.get("jsonrpc") != "2.0"
|
||||
or type(response_id) not in (int, str)
|
||||
or type(response_id) is not type(expected_id)
|
||||
or response_id != expected_id
|
||||
):
|
||||
raise RuntimeError("invalid JSON-RPC response envelope")
|
||||
has_result = "result" in response
|
||||
has_error = "error" in response
|
||||
@@ -478,11 +484,12 @@ class MultiServerClient:
|
||||
else:
|
||||
raise RuntimeError(f"{peer.name}: protocol era not selected")
|
||||
response = self._send(peer, message)
|
||||
if response is None or response.get("id") != request_id:
|
||||
raise RuntimeError(f"{peer.name}: missing or mismatched response")
|
||||
if "error" in response:
|
||||
raise RuntimeError(f"{peer.name}: RPC error {response['error']}")
|
||||
result = dict(response["result"])
|
||||
if not isinstance(response, dict):
|
||||
raise RuntimeError(f"{peer.name}: missing response")
|
||||
kind, payload = decode_rpc_response(response, request_id)
|
||||
if kind != "result":
|
||||
raise RuntimeError(f"{peer.name}: RPC error {payload}")
|
||||
result = dict(payload)
|
||||
if peer.era == "modern" and "resultType" not in result:
|
||||
raise RuntimeError(f"{peer.name}: modern result omitted resultType")
|
||||
if peer.era == "legacy":
|
||||
|
||||
@@ -54,9 +54,19 @@ class McpClientTests(unittest.TestCase):
|
||||
def modern_error(
|
||||
request: dict,
|
||||
timeout_ms: int | None = None,
|
||||
*,
|
||||
sink: list[dict] = received,
|
||||
error_code: int = code,
|
||||
error_message: str = message,
|
||||
error_data: dict | None = data,
|
||||
) -> dict:
|
||||
received.append(request)
|
||||
return main.rpc_error(request.get("id"), code, message, data)
|
||||
sink.append(request)
|
||||
return main.rpc_error(
|
||||
request.get("id"),
|
||||
error_code,
|
||||
error_message,
|
||||
error_data,
|
||||
)
|
||||
|
||||
client = main.MultiServerClient()
|
||||
client.add_server("broken-modern", modern_error, allow_legacy=True)
|
||||
@@ -98,9 +108,15 @@ class McpClientTests(unittest.TestCase):
|
||||
with self.subTest(signal=signal):
|
||||
received = []
|
||||
|
||||
def unavailable(message: dict, timeout_ms: int | None = None) -> dict | None:
|
||||
received.append(message)
|
||||
if signal == "closed":
|
||||
def unavailable(
|
||||
message: dict,
|
||||
timeout_ms: int | None = None,
|
||||
*,
|
||||
sink: list[dict] = received,
|
||||
current_signal: str = signal,
|
||||
) -> dict | None:
|
||||
sink.append(message)
|
||||
if current_signal == "closed":
|
||||
raise ConnectionError("transport closed")
|
||||
return None
|
||||
|
||||
@@ -189,6 +205,39 @@ class McpClientTests(unittest.TestCase):
|
||||
["server/discover"],
|
||||
)
|
||||
|
||||
def test_request_rejects_a_response_without_result_or_error(self) -> None:
|
||||
server = main.ModernFakeServer("modern", [make_tool("search")])
|
||||
client = main.MultiServerClient()
|
||||
client.add_server("modern", server)
|
||||
client.connect_all()
|
||||
|
||||
def malformed_response(
|
||||
message: dict,
|
||||
timeout_ms: int | None = None,
|
||||
) -> dict:
|
||||
return {"jsonrpc": "2.0", "id": message["id"]}
|
||||
|
||||
client.peers["modern"].transport = malformed_response
|
||||
with self.assertRaisesRegex(RuntimeError, "exactly one of result or error"):
|
||||
client.discover_tools()
|
||||
|
||||
def test_discovery_rejects_a_boolean_response_id_for_an_integer_request(self) -> None:
|
||||
server = main.ModernFakeServer("modern", [make_tool("search")])
|
||||
|
||||
def boolean_id_response(
|
||||
message: dict,
|
||||
timeout_ms: int | None = None,
|
||||
) -> dict | None:
|
||||
response = server(message, timeout_ms)
|
||||
if response is not None:
|
||||
response["id"] = True
|
||||
return response
|
||||
|
||||
client = main.MultiServerClient()
|
||||
client.add_server("modern", boolean_id_response)
|
||||
with self.assertRaisesRegex(RuntimeError, "invalid JSON-RPC response envelope"):
|
||||
client.connect_all()
|
||||
|
||||
def test_merge_is_deterministic_and_prefixes_collisions(self) -> None:
|
||||
alpha = main.ModernFakeServer("alpha", [make_tool("search"), make_tool("write")])
|
||||
beta = main.ModernFakeServer("beta", [make_tool("read"), make_tool("search")])
|
||||
|
||||
Reference in New Issue
Block a user