import asyncio import os import unittest from unittest.mock import patch from ipython_shell.jupyter.messages import ( DisplayDataMessage, JupyterMessage, UnknownJupyterMessage, parse_jupyter_message, ) from ipython_shell.jupyter.transport import InputRequest, JupyterTransport class TransportTests(unittest.TestCase): def test_kernel_message_keeps_content_and_buffers(self): message = JupyterMessage( msg_type="display_data", parent_id="call-1", content={"data": {"text/plain": "42"}}, buffers=[b"binary"], ) self.assertEqual(message.parent_id, "call-1") self.assertEqual(message.content["data"], {"text/plain": "42"}) self.assertEqual(message.buffers, [b"binary"]) def test_parser_types_known_message_content(self): message = parse_jupyter_message( { "msg_id": "message-1", "msg_type": "display_data", "parent_header": {"msg_id": "call-1"}, "metadata": {}, "content": { "data": {"text/plain": "42"}, "metadata": {}, }, "buffers": [b"binary"], }, channel="iopub", ) self.assertIsInstance(message, DisplayDataMessage) assert isinstance(message, DisplayDataMessage) self.assertEqual(message.content.data["text/plain"], "42") self.assertEqual(message.parent_id, "call-1") self.assertEqual(message.buffers, [b"binary"]) def test_parser_preserves_unknown_message(self): message = parse_jupyter_message( { "msg_id": "message-2", "msg_type": "future_extension", "parent_header": {"msg_id": "call-1"}, "metadata": {"vendor": "example"}, "content": {"answer": 42}, "buffers": [], }, channel="iopub", ) self.assertIsInstance(message, UnknownJupyterMessage) self.assertEqual(message.msg_type, "future_extension") self.assertEqual(message.content["answer"], 42) self.assertEqual(message.metadata["vendor"], "example") class AsyncTransportTests(unittest.IsolatedAsyncioTestCase): async def test_transport_keeps_channel_readers_alive_during_a_call(self): transport = JupyterTransport() await transport.start() try: readers = set(transport._reader_tasks) call_id = await transport.execute("2 + 2") messages = [message async for message in transport.messages_for(call_id)] self.assertEqual(len(readers), 3) self.assertEqual(transport._reader_tasks, readers) self.assertTrue(all(not reader.done() for reader in readers)) finally: await transport.shutdown() self.assertTrue( any(message.msg_type == "execute_result" for message in messages) ) async def test_transport_executes_and_returns_execute_reply(self): transport = JupyterTransport() await transport.start() try: call_id = await transport.execute("2 + 2") messages = [message async for message in transport.messages_for(call_id)] finally: await transport.shutdown() self.assertTrue( any(message.msg_type == "execute_result" for message in messages) ) self.assertEqual(messages[-1].msg_type, "execute_reply") self.assertEqual(messages[-1].content.status, "ok") async def test_transport_preserves_mime_and_errors(self): transport = JupyterTransport() await transport.start() try: display_id = await transport.execute( "from IPython.display import display\n" "display({'text/plain': 'hello'}, raw=True)" ) display_messages = [ message async for message in transport.messages_for(display_id) ] error_id = await transport.execute("raise ValueError('boom')") error_messages = [ message async for message in transport.messages_for(error_id) ] finally: await transport.shutdown() display = next( message for message in display_messages if message.msg_type == "display_data" ) self.assertEqual(display.content.data["text/plain"], "hello") error = next( message for message in error_messages if message.msg_type == "error" ) self.assertEqual(error.content.ename, "ValueError") async def test_transport_waits_for_iopub_output_after_idle_and_reply(self): transport = JupyterTransport() transport.client = object() transport._active_call = "call-1" transport._message_queue = asyncio.Queue() transport._message_queue.put_nowait( ( "shell", { "msg_id": "reply-1", "msg_type": "execute_reply", "parent_header": {"msg_id": "call-1"}, "metadata": {}, "content": {"status": "ok", "execution_count": 1}, "buffers": [], }, ) ) transport._message_queue.put_nowait( ( "iopub", { "msg_id": "status-1", "msg_type": "status", "parent_header": {"msg_id": "call-1"}, "metadata": {}, "content": {"execution_state": "idle"}, "buffers": [], }, ) ) async def enqueue_late_display() -> None: await asyncio.sleep(0.15) await transport._message_queue.put( ( "iopub", { "msg_id": "display-1", "msg_type": "display_data", "parent_header": {"msg_id": "call-1"}, "metadata": {}, "content": { "data": {"image/png": "ZmFrZQ=="}, "metadata": {}, }, "buffers": [], }, ) ) asyncio.create_task(enqueue_late_display()) messages = [message async for message in transport.messages_for("call-1")] self.assertTrue(any(message.msg_type == "display_data" for message in messages)) async def test_transport_preserves_matplotlib_mime_output(self): transport = JupyterTransport() await transport.start() try: call_id = await transport.execute( "import matplotlib.pyplot as plt\nplt.plot([1, 2, 3], [4, 5, 6])" ) messages = [message async for message in transport.messages_for(call_id)] finally: await transport.shutdown() display = next( message for message in messages if message.msg_type == "display_data" ) self.assertIn("image/png", display.content.data) async def test_transport_keeps_mime_output_when_parent_forces_agg(self): # Streamlit sets MPLBACKEND=Agg in its own process. The kernel must # still use an inline backend so rich output survives the transport. with patch.dict(os.environ, {"MPLBACKEND": "Agg"}): transport = JupyterTransport() await transport.start() try: call_id = await transport.execute( "import matplotlib.pyplot as plt\nplt.plot([1, 2, 3], [4, 5, 6])" ) messages = [ message async for message in transport.messages_for(call_id) ] finally: await transport.shutdown() display = next( message for message in messages if message.msg_type == "display_data" ) self.assertIn("image/png", display.content.data) async def test_transport_routes_input_reply_without_parent_stdin(self): transport = JupyterTransport() await transport.start() try: call_id = await transport.execute("answer = input('name? '); answer") stream = transport.messages_for(call_id) while True: event = await asyncio.wait_for(anext(stream), timeout=5) if isinstance(event, InputRequest): input_request = event break self.assertEqual(input_request.prompt, "name? ") await transport.reply_to_input("Ada") messages = [message async for message in stream] finally: await transport.shutdown() self.assertTrue( any(message.msg_type == "execute_result" for message in messages) ) if __name__ == "__main__": unittest.main()