From 329a0987b4a4475d54293f65a0ada6cee0df6365 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Wed, 10 Dec 2025 15:40:16 -0500 Subject: [PATCH] test: use McpError assertions now that exception propagation is fixed --- .../middleware/test_initialization_middleware.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/tests/server/middleware/test_initialization_middleware.py b/tests/server/middleware/test_initialization_middleware.py index a34920a4a..5c2266268 100644 --- a/tests/server/middleware/test_initialization_middleware.py +++ b/tests/server/middleware/test_initialization_middleware.py @@ -292,7 +292,7 @@ async def test_middleware_can_access_initialize_result(): async def test_middleware_mcp_error_during_initialization(): - """Test that McpError raised in middleware during initialization is sent to responder.""" + """Test that McpError raised in middleware during initialization is sent to client.""" server = FastMCP("TestServer") class ErrorThrowingMiddleware(Middleware): @@ -309,11 +309,12 @@ async def test_middleware_mcp_error_during_initialization(): server.add_middleware(ErrorThrowingMiddleware()) - with pytest.raises(Exception) as exc_info: + with pytest.raises(McpError) as exc_info: async with Client(server): pass - assert "Invalid initialization parameters" in str(exc_info.value) + assert exc_info.value.error.message == "Invalid initialization parameters" + assert exc_info.value.error.code == mt.INVALID_PARAMS async def test_middleware_mcp_error_before_call_next(): @@ -332,11 +333,12 @@ async def test_middleware_mcp_error_before_call_next(): server.add_middleware(EarlyErrorMiddleware()) - with pytest.raises(Exception) as exc_info: + with pytest.raises(McpError) as exc_info: async with Client(server): pass - assert "Request validation failed" in str(exc_info.value) + assert exc_info.value.error.message == "Request validation failed" + assert exc_info.value.error.code == mt.INVALID_REQUEST async def test_middleware_mcp_error_after_call_next():