fix: keep Jupyter channel readers alive
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
from collections.abc import AsyncIterator
|
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from jupyter_client import AsyncKernelManager
|
from jupyter_client import AsyncKernelManager
|
||||||
|
|
||||||
@@ -33,14 +34,65 @@ class JupyterTransport:
|
|||||||
self._active_call: str | None = None
|
self._active_call: str | None = None
|
||||||
self._call_lock = asyncio.Lock()
|
self._call_lock = asyncio.Lock()
|
||||||
self._waiting_for_input: str | None = None
|
self._waiting_for_input: str | None = None
|
||||||
|
self._message_queue: asyncio.Queue[
|
||||||
|
tuple[str, dict[str, Any] | BaseException]
|
||||||
|
] | None = None
|
||||||
|
self._reader_tasks: set[asyncio.Task[None]] = set()
|
||||||
|
|
||||||
async def start(self) -> None:
|
async def start(self) -> None:
|
||||||
"""Start the kernel process and wait until it accepts requests."""
|
"""Start the kernel and its persistent channel readers."""
|
||||||
|
if self.client is not None:
|
||||||
|
return
|
||||||
|
|
||||||
await self.manager.start_kernel()
|
await self.manager.start_kernel()
|
||||||
self.client = self.manager.client()
|
self.client = self.manager.client()
|
||||||
self.client.start_channels()
|
self.client.start_channels()
|
||||||
await self.client.wait_for_ready()
|
await self.client.wait_for_ready()
|
||||||
|
|
||||||
|
self._message_queue = asyncio.Queue()
|
||||||
|
self._reader_tasks = {
|
||||||
|
asyncio.create_task(
|
||||||
|
self._read_channel("iopub", self.client.get_iopub_msg),
|
||||||
|
name="jupyter-iopub-reader",
|
||||||
|
),
|
||||||
|
asyncio.create_task(
|
||||||
|
self._read_channel("shell", self.client.get_shell_msg),
|
||||||
|
name="jupyter-shell-reader",
|
||||||
|
),
|
||||||
|
asyncio.create_task(
|
||||||
|
self._read_channel("stdin", self.client.get_stdin_msg),
|
||||||
|
name="jupyter-stdin-reader",
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
async def _read_channel(
|
||||||
|
self,
|
||||||
|
channel: str,
|
||||||
|
get_message: Callable[[], Awaitable[dict[str, Any]]],
|
||||||
|
) -> None:
|
||||||
|
"""Read one ZMQ channel continuously into the transport queue."""
|
||||||
|
if self._message_queue is None:
|
||||||
|
raise RuntimeError("JupyterTransport has not been started")
|
||||||
|
|
||||||
|
try:
|
||||||
|
while True:
|
||||||
|
message = await get_message()
|
||||||
|
await self._message_queue.put((channel, message))
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except BaseException as error:
|
||||||
|
# Surface reader failures instead of leaving messages_for() stuck.
|
||||||
|
await self._message_queue.put((channel, error))
|
||||||
|
|
||||||
|
async def _stop_readers(self) -> None:
|
||||||
|
"""Cancel channel readers once, during transport shutdown."""
|
||||||
|
readers = self._reader_tasks
|
||||||
|
self._reader_tasks = set()
|
||||||
|
for reader in readers:
|
||||||
|
reader.cancel()
|
||||||
|
await asyncio.gather(*readers, return_exceptions=True)
|
||||||
|
self._message_queue = None
|
||||||
|
|
||||||
async def execute(self, code: str) -> str:
|
async def execute(self, code: str) -> str:
|
||||||
"""Submit one cell and return its Jupyter message ID."""
|
"""Submit one cell and return its Jupyter message ID."""
|
||||||
if self.client is None:
|
if self.client is None:
|
||||||
@@ -60,59 +112,53 @@ class JupyterTransport:
|
|||||||
raise RuntimeError("JupyterTransport has not been started")
|
raise RuntimeError("JupyterTransport has not been started")
|
||||||
if self._active_call != call_id:
|
if self._active_call != call_id:
|
||||||
raise RuntimeError(f"Unknown or inactive call: {call_id}")
|
raise RuntimeError(f"Unknown or inactive call: {call_id}")
|
||||||
|
if self._message_queue is None:
|
||||||
|
raise RuntimeError("JupyterTransport message readers are not running")
|
||||||
|
|
||||||
reply: KernelMessage | None = None
|
reply: KernelMessage | None = None
|
||||||
idle = False
|
idle = False
|
||||||
try:
|
try:
|
||||||
while True:
|
while True:
|
||||||
iopub_task = asyncio.create_task(self.client.get_iopub_msg())
|
_, queued_message = await self._message_queue.get()
|
||||||
shell_task = asyncio.create_task(self.client.get_shell_msg())
|
if isinstance(queued_message, BaseException):
|
||||||
stdin_task = asyncio.create_task(self.client.get_stdin_msg())
|
raise queued_message
|
||||||
done, pending = await asyncio.wait(
|
|
||||||
{iopub_task, shell_task, stdin_task},
|
|
||||||
return_when=asyncio.FIRST_COMPLETED,
|
|
||||||
)
|
|
||||||
for task in pending:
|
|
||||||
task.cancel()
|
|
||||||
await asyncio.gather(*pending, return_exceptions=True)
|
|
||||||
|
|
||||||
for task in done:
|
message = queued_message
|
||||||
message = task.result()
|
parent_id = message.get("parent_header", {}).get("msg_id")
|
||||||
parent_id = message.get("parent_header", {}).get("msg_id")
|
if parent_id != call_id:
|
||||||
if parent_id != call_id:
|
continue
|
||||||
continue
|
|
||||||
|
|
||||||
if message["msg_type"] == "input_request":
|
if message["msg_type"] == "input_request":
|
||||||
self._waiting_for_input = call_id
|
self._waiting_for_input = call_id
|
||||||
yield InputRequest(
|
yield InputRequest(
|
||||||
call_id=call_id,
|
call_id=call_id,
|
||||||
prompt=str(message.get("content", {}).get("prompt", "")),
|
prompt=str(message.get("content", {}).get("prompt", "")),
|
||||||
password=bool(
|
password=bool(
|
||||||
message.get("content", {}).get("password", False)
|
message.get("content", {}).get("password", False)
|
||||||
),
|
),
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
decoded = KernelMessage(
|
|
||||||
msg_type=message["msg_type"],
|
|
||||||
parent_id=parent_id,
|
|
||||||
content=dict(message.get("content", {})),
|
|
||||||
buffers=[
|
|
||||||
bytes(buffer) for buffer in message.get("buffers", [])
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
if decoded.msg_type == "execute_reply":
|
continue
|
||||||
# The shell reply and IOPub messages use different
|
|
||||||
# channels. Buffer the reply so it remains terminal
|
decoded = KernelMessage(
|
||||||
# even when both channels become ready together.
|
msg_type=message["msg_type"],
|
||||||
reply = decoded
|
parent_id=parent_id,
|
||||||
else:
|
content=dict(message.get("content", {})),
|
||||||
yield decoded
|
buffers=[
|
||||||
if (
|
bytes(buffer) for buffer in message.get("buffers", [])
|
||||||
decoded.msg_type == "status"
|
],
|
||||||
and decoded.content.get("execution_state") == "idle"
|
)
|
||||||
):
|
if decoded.msg_type == "execute_reply":
|
||||||
idle = True
|
# The shell reply and IOPub messages use different
|
||||||
|
# channels. Buffer the reply so it remains terminal
|
||||||
|
# even when both channels arrive close together.
|
||||||
|
reply = decoded
|
||||||
|
else:
|
||||||
|
yield decoded
|
||||||
|
if (
|
||||||
|
decoded.msg_type == "status"
|
||||||
|
and decoded.content.get("execution_state") == "idle"
|
||||||
|
):
|
||||||
|
idle = True
|
||||||
|
|
||||||
if reply is not None and idle:
|
if reply is not None and idle:
|
||||||
yield reply
|
yield reply
|
||||||
@@ -124,7 +170,7 @@ class JupyterTransport:
|
|||||||
self._active_call = None
|
self._active_call = None
|
||||||
|
|
||||||
async def reply_to_input(self, value: str) -> None:
|
async def reply_to_input(self, value: str) -> None:
|
||||||
"""Reply to the active kernel input request without using process stdin."""
|
"""Reply to the active kernel input request without process stdin."""
|
||||||
if self.client is None:
|
if self.client is None:
|
||||||
raise RuntimeError("JupyterTransport has not been started")
|
raise RuntimeError("JupyterTransport has not been started")
|
||||||
if self._waiting_for_input is None:
|
if self._waiting_for_input is None:
|
||||||
@@ -134,7 +180,8 @@ class JupyterTransport:
|
|||||||
self._waiting_for_input = None
|
self._waiting_for_input = None
|
||||||
|
|
||||||
async def shutdown(self) -> None:
|
async def shutdown(self) -> None:
|
||||||
"""Stop channels and terminate the kernel process."""
|
"""Stop readers, channels, and the kernel process."""
|
||||||
|
await self._stop_readers()
|
||||||
if self.client is not None:
|
if self.client is not None:
|
||||||
self.client.stop_channels()
|
self.client.stop_channels()
|
||||||
self.client = None
|
self.client = None
|
||||||
|
|||||||
@@ -23,6 +23,26 @@ class TransportTests(unittest.TestCase):
|
|||||||
|
|
||||||
|
|
||||||
class AsyncTransportTests(unittest.IsolatedAsyncioTestCase):
|
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):
|
async def test_transport_executes_and_returns_execute_reply(self):
|
||||||
transport = JupyterTransport()
|
transport = JupyterTransport()
|
||||||
await transport.start()
|
await transport.start()
|
||||||
|
|||||||
Reference in New Issue
Block a user