import asyncio import unittest from ipython_shell.jupyter.transport import ( InputRequest, JupyterTransport, KernelMessage, ) class TransportTests(unittest.TestCase): def test_kernel_message_keeps_content_and_buffers(self): message = KernelMessage( 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"]) 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_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_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()