253 lines
8.9 KiB
Python
253 lines
8.9 KiB
Python
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()
|