|
| 1 | +"""Regression test for #3122. |
| 2 | +
|
| 3 | +``Server._handle_message`` used to log recorded warnings while still inside the |
| 4 | +``warnings.catch_warnings(record=True)`` block. Because ``record=True`` forces |
| 5 | +an "always" filter, a warning emitted by a logging handler during that logging |
| 6 | +step was appended to the very list being iterated, so the loop never |
| 7 | +terminated. The loop is synchronous, so task cancellation could not interrupt |
| 8 | +it. |
| 9 | +
|
| 10 | +The fix snapshots the recorded warnings and logs them after leaving the |
| 11 | +``catch_warnings`` block, so handler-emitted warnings can no longer extend the |
| 12 | +iteration. This test drives ``_handle_message`` with a request that records one |
| 13 | +warning and a logging handler that itself warns on every record; the handler |
| 14 | +must be invoked exactly once. The handler stops re-warning after a cap so that |
| 15 | +the pre-fix infinite loop terminates and fails the assertion instead of hanging |
| 16 | +the test session. |
| 17 | +""" |
| 18 | + |
| 19 | +import logging |
| 20 | +import warnings |
| 21 | +from unittest.mock import AsyncMock, Mock |
| 22 | + |
| 23 | +import pytest |
| 24 | + |
| 25 | +import mcp.types as types |
| 26 | +from mcp.server.lowlevel.server import Server |
| 27 | +from mcp.server.session import ServerSession |
| 28 | +from mcp.shared.session import RequestResponder |
| 29 | + |
| 30 | + |
| 31 | +class _WarningEmittingHandler(logging.Handler): |
| 32 | + """A logging handler whose emit() raises a warning, mimicking e.g. a |
| 33 | + timestamp formatter that calls a deprecated API on every record. It stops |
| 34 | + after ``cap`` records so a regressed (looping) build still terminates.""" |
| 35 | + |
| 36 | + def __init__(self, cap: int = 100) -> None: |
| 37 | + super().__init__() |
| 38 | + self.emit_count = 0 |
| 39 | + self._cap = cap |
| 40 | + |
| 41 | + def emit(self, record: logging.LogRecord) -> None: |
| 42 | + self.emit_count += 1 |
| 43 | + if self.emit_count <= self._cap: |
| 44 | + warnings.warn("warning raised while logging", stacklevel=1) |
| 45 | + |
| 46 | + |
| 47 | +@pytest.mark.anyio |
| 48 | +async def test_handle_message_logs_each_warning_once_when_handler_warns(): |
| 49 | + server = Server("test-server") |
| 50 | + |
| 51 | + session = Mock(spec=ServerSession) |
| 52 | + session.send_log_message = AsyncMock() |
| 53 | + |
| 54 | + async def _handle_request_that_warns(*args: object, **kwargs: object) -> None: |
| 55 | + warnings.warn("warning raised while handling the request", stacklevel=1) |
| 56 | + |
| 57 | + server._handle_request = _handle_request_that_warns # type: ignore[assignment] |
| 58 | + |
| 59 | + responder = Mock(spec=RequestResponder) |
| 60 | + responder.request = types.ClientRequest(root=types.PingRequest(method="ping")) |
| 61 | + responder.__enter__ = Mock(return_value=responder) |
| 62 | + responder.__exit__ = Mock(return_value=None) |
| 63 | + |
| 64 | + server_logger = logging.getLogger("mcp.server.lowlevel.server") |
| 65 | + handler = _WarningEmittingHandler() |
| 66 | + server_logger.addHandler(handler) |
| 67 | + previous_level = server_logger.level |
| 68 | + server_logger.setLevel(logging.INFO) |
| 69 | + |
| 70 | + try: |
| 71 | + with warnings.catch_warnings(): |
| 72 | + warnings.simplefilter("always") |
| 73 | + await server._handle_message(responder, session, {}, raise_exceptions=False) |
| 74 | + finally: |
| 75 | + server_logger.removeHandler(handler) |
| 76 | + server_logger.setLevel(previous_level) |
| 77 | + |
| 78 | + assert handler.emit_count == 1 |
0 commit comments