feat: preserve Jupyter MIME output in Streamlit
This commit is contained in:
+12
-16
@@ -4,7 +4,8 @@ import unittest
|
||||
|
||||
from ipython_shell.jupyter.app import JupyterApp
|
||||
from ipython_shell.jupyter.shell import JupyterShell
|
||||
from ipython_shell.jupyter.transport import InputRequest, KernelMessage
|
||||
from ipython_shell.jupyter.messages import ExecuteResultMessage
|
||||
from ipython_shell.jupyter.transport import InputRequest
|
||||
|
||||
|
||||
class JupyterShellTests(unittest.IsolatedAsyncioTestCase):
|
||||
@@ -30,19 +31,17 @@ class JupyterShellTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
self.assertEqual(
|
||||
next(
|
||||
message.content["data"]["text/plain"]
|
||||
message.content.data["text/plain"]
|
||||
for message in first
|
||||
if isinstance(message, KernelMessage)
|
||||
and message.msg_type == "execute_result"
|
||||
if isinstance(message, ExecuteResultMessage)
|
||||
),
|
||||
"41",
|
||||
)
|
||||
self.assertEqual(
|
||||
next(
|
||||
message.content["data"]["text/plain"]
|
||||
message.content.data["text/plain"]
|
||||
for message in second
|
||||
if isinstance(message, KernelMessage)
|
||||
and message.msg_type == "execute_result"
|
||||
if isinstance(message, ExecuteResultMessage)
|
||||
),
|
||||
"42",
|
||||
)
|
||||
@@ -97,19 +96,17 @@ class JupyterAppTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
self.assertEqual(
|
||||
next(
|
||||
message.content["data"]["text/plain"]
|
||||
message.content.data["text/plain"]
|
||||
for message in beta_view
|
||||
if isinstance(message, KernelMessage)
|
||||
and message.msg_type == "execute_result"
|
||||
if isinstance(message, ExecuteResultMessage)
|
||||
),
|
||||
"False",
|
||||
)
|
||||
self.assertEqual(
|
||||
next(
|
||||
message.content["data"]["text/plain"]
|
||||
message.content.data["text/plain"]
|
||||
for message in alpha_view
|
||||
if isinstance(message, KernelMessage)
|
||||
and message.msg_type == "execute_result"
|
||||
if isinstance(message, ExecuteResultMessage)
|
||||
),
|
||||
"'alpha'",
|
||||
)
|
||||
@@ -143,10 +140,9 @@ class JupyterAppTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
self.assertEqual(
|
||||
next(
|
||||
message.content["data"]["text/plain"]
|
||||
message.content.data["text/plain"]
|
||||
for message in messages
|
||||
if isinstance(message, KernelMessage)
|
||||
and message.msg_type == "execute_result"
|
||||
if isinstance(message, ExecuteResultMessage)
|
||||
),
|
||||
"'Ada'",
|
||||
)
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
import subprocess
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from ipython_shell.jupyter.messages import (
|
||||
ExecuteResultMessage,
|
||||
UnknownJupyterMessage,
|
||||
)
|
||||
from st_demo.app import visible_records
|
||||
from st_demo.ui import event_to_record, preferred_mime
|
||||
|
||||
|
||||
class UIDisplayTests(unittest.TestCase):
|
||||
def test_completed_call_records_remain_visible_after_active_call_clears(self):
|
||||
records = {
|
||||
"finished-call": [{"kind": "execute_result", "text": "42"}],
|
||||
}
|
||||
|
||||
self.assertEqual(visible_records(records), records["finished-call"])
|
||||
|
||||
def test_app_module_loads_when_streamlit_executes_file_path(self):
|
||||
app_path = Path(__file__).parents[1] / "src" / "st_demo" / "app.py"
|
||||
result = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-c",
|
||||
"import runpy, sys; runpy.run_path(sys.argv[1], run_name='st_demo_script')",
|
||||
str(app_path),
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
self.assertEqual(result.returncode, 0, result.stderr)
|
||||
|
||||
def test_preferred_mime_chooses_html_before_plain_text(self):
|
||||
mime, value = preferred_mime(
|
||||
{
|
||||
"text/plain": "fallback",
|
||||
"text/html": "<b>rich</b>",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual((mime, value), ("text/html", "<b>rich</b>"))
|
||||
|
||||
def test_execute_result_record_keeps_mime_and_repr(self):
|
||||
event = ExecuteResultMessage(
|
||||
message_id="message-1",
|
||||
parent_id="call-1",
|
||||
content={
|
||||
"execution_count": 4,
|
||||
"data": {
|
||||
"text/plain": "Line2D(_line0)",
|
||||
"image/png": "ZmFrZQ==",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
record = event_to_record(event)
|
||||
|
||||
self.assertEqual(record["kind"], "execute_result")
|
||||
self.assertEqual(record["text"], "Line2D(_line0)")
|
||||
self.assertEqual(record["data"]["image/png"], "ZmFrZQ==")
|
||||
|
||||
def test_unknown_message_record_keeps_raw_protocol_data(self):
|
||||
event = UnknownJupyterMessage(
|
||||
channel="iopub",
|
||||
msg_type="vendor_extension",
|
||||
message_id="message-2",
|
||||
content={"value": 42},
|
||||
raw={"msg_type": "vendor_extension", "content": {"value": 42}},
|
||||
)
|
||||
|
||||
record = event_to_record(event)
|
||||
|
||||
self.assertEqual(record["kind"], "vendor_extension")
|
||||
self.assertEqual(record["content"]["value"], 42)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+132
-9
@@ -1,16 +1,20 @@
|
||||
import asyncio
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from ipython_shell.jupyter.transport import (
|
||||
InputRequest,
|
||||
JupyterTransport,
|
||||
KernelMessage,
|
||||
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 = KernelMessage(
|
||||
message = JupyterMessage(
|
||||
msg_type="display_data",
|
||||
parent_id="call-1",
|
||||
content={"data": {"text/plain": "42"}},
|
||||
@@ -21,6 +25,46 @@ class TransportTests(unittest.TestCase):
|
||||
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):
|
||||
@@ -54,7 +98,7 @@ class AsyncTransportTests(unittest.IsolatedAsyncioTestCase):
|
||||
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")
|
||||
self.assertEqual(messages[-1].content.status, "ok")
|
||||
|
||||
async def test_transport_preserves_mime_and_errors(self):
|
||||
transport = JupyterTransport()
|
||||
@@ -80,12 +124,68 @@ class AsyncTransportTests(unittest.IsolatedAsyncioTestCase):
|
||||
for message in display_messages
|
||||
if message.msg_type == "display_data"
|
||||
)
|
||||
self.assertEqual(display.content["data"]["text/plain"], "hello")
|
||||
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")
|
||||
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()
|
||||
@@ -101,7 +201,30 @@ class AsyncTransportTests(unittest.IsolatedAsyncioTestCase):
|
||||
display = next(
|
||||
message for message in messages if message.msg_type == "display_data"
|
||||
)
|
||||
self.assertIn("image/png", display.content["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\n"
|
||||
"plt.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()
|
||||
|
||||
Reference in New Issue
Block a user