Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
193 changes: 193 additions & 0 deletions pylabrobot/io/executor_setup_tests.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,193 @@
import asyncio
import threading
import unittest
from types import SimpleNamespace
from typing import Any, Callable, Coroutine, Optional
from unittest import mock

from pylabrobot.io import ftdi as ftdi_module
from pylabrobot.io import hid as hid_module
from pylabrobot.io import serial as serial_module
from pylabrobot.io import usb as usb_module


class _ClosableDevice:
def __init__(self) -> None:
self.setup_thread: Optional[int] = None
self.close_thread: Optional[int] = None

def close(self) -> None:
self.close_thread = threading.get_ident()


class _DisposableDevice:
def __init__(self) -> None:
self.setup_thread: Optional[int] = None
self.dispose_thread: Optional[int] = None

def set_configuration(self) -> None:
self.setup_thread = threading.get_ident()
raise RuntimeError("configuration failed")


class ExecutorSetupTests(unittest.IsolatedAsyncioTestCase):
async def _wait_for_task(self, task: "asyncio.Task[None]") -> None:
while not task.done():
await asyncio.wait({task}, timeout=0.01)
await task

async def _cancel_while_worker_is_running(
self,
setup: Callable[[], Coroutine[Any, Any, None]],
started: threading.Event,
release: threading.Event,
) -> None:
task: "asyncio.Task[None]" = asyncio.create_task(setup())
while not started.is_set():
await asyncio.sleep(0)

try:
task.cancel()
await asyncio.sleep(0)
self.assertFalse(task.done(), "setup returned before its worker finished")
finally:
release.set()

with self.assertRaises(asyncio.CancelledError):
await self._wait_for_task(task)

async def test_cancelled_ftdi_setup_closes_device_before_executor_shutdown(self) -> None:
io = ftdi_module.FTDI.__new__(ftdi_module.FTDI)
io.human_readable_device_name = "mock FTDI"
io._device_id = "mock"
io._dev = None
io._executor = None
device = _ClosableDevice()
started = threading.Event()
release = threading.Event()

def setup_sync() -> None:
device.setup_thread = threading.get_ident()
started.set()
release.wait()
io._dev = device # type: ignore[assignment]

io._setup_sync = setup_sync # type: ignore[method-assign]
ftdi_error = type("MockFtdiError", (Exception,), {})
with mock.patch.object(ftdi_module, "FtdiError", ftdi_error, create=True):
await self._cancel_while_worker_is_running(io.setup, started, release)

self.assertEqual(device.close_thread, device.setup_thread)
self.assertIsNone(io._dev)
self.assertIsNone(io._executor)

async def test_cancelled_hid_setup_closes_device_before_executor_shutdown(self) -> None:
io = hid_module.HID.__new__(hid_module.HID)
io.human_readable_device_name = "mock HID"
io._unique_id = "mock"
io.device = None
io._executor = None
device = _ClosableDevice()
started = threading.Event()
release = threading.Event()

def setup_sync() -> None:
device.setup_thread = threading.get_ident()
started.set()
release.wait()
io.device = device # type: ignore[assignment]

io._setup_sync = setup_sync # type: ignore[method-assign]
with mock.patch.object(hid_module, "USE_HID", True):
await self._cancel_while_worker_is_running(io.setup, started, release)

self.assertEqual(device.close_thread, device.setup_thread)
self.assertIsNone(io.device)
self.assertIsNone(io._executor)

async def test_cancelled_serial_setup_closes_port_before_executor_shutdown(self) -> None:
io = serial_module.Serial.__new__(serial_module.Serial)
io.human_readable_device_name = "mock serial"
io._ser = None
io._executor = None
device = _ClosableDevice()
started = threading.Event()
release = threading.Event()

def setup_sync() -> str:
device.setup_thread = threading.get_ident()
started.set()
release.wait()
io._ser = device # type: ignore[assignment]
return "/dev/mock"

io._setup_sync = setup_sync # type: ignore[method-assign]
with mock.patch.object(serial_module, "HAS_SERIAL", True):
await self._cancel_while_worker_is_running(io.setup, started, release)

self.assertEqual(device.close_thread, device.setup_thread)
self.assertIsNone(io._ser)
self.assertIsNone(io._executor)

async def test_cancelled_usb_setup_disposes_device_before_executor_shutdown(self) -> None:
io = usb_module.USB(
id_vendor=1,
id_product=2,
human_readable_device_name="mock USB",
packet_read_timeout=1,
read_timeout=2,
)
device = _DisposableDevice()
started = threading.Event()
release = threading.Event()

def setup_sync(empty_buffer: bool) -> None:
device.setup_thread = threading.get_ident()
started.set()
release.wait()
io.dev = device # type: ignore[assignment]

def dispose_resources(dev: object) -> None:
self.assertIs(dev, device)
device.dispose_thread = threading.get_ident()

io._setup_sync = setup_sync # type: ignore[method-assign]
fake_usb = SimpleNamespace(util=SimpleNamespace(dispose_resources=dispose_resources))
with (
mock.patch.object(usb_module, "USE_USB", True),
mock.patch.object(usb_module, "usb", fake_usb, create=True),
):
await self._cancel_while_worker_is_running(io.setup, started, release)

self.assertEqual(device.dispose_thread, device.setup_thread)
self.assertIsNone(io.dev)
self.assertIsNone(io._read_executor)
self.assertIsNone(io._write_executor)

async def test_usb_setup_error_disposes_acquired_device_on_worker(self) -> None:
io = usb_module.USB(
id_vendor=1,
id_product=2,
human_readable_device_name="mock USB",
packet_read_timeout=1,
read_timeout=2,
)
device = _DisposableDevice()

def dispose_resources(dev: object) -> None:
self.assertIs(dev, device)
device.dispose_thread = threading.get_ident()

io.get_available_devices = mock.Mock(return_value=[device]) # type: ignore[method-assign]
fake_usb = SimpleNamespace(util=SimpleNamespace(dispose_resources=dispose_resources))
with (
mock.patch.object(usb_module, "USE_USB", True),
mock.patch.object(usb_module, "usb", fake_usb, create=True),
self.assertRaisesRegex(RuntimeError, "configuration failed"),
):
await self._wait_for_task(asyncio.create_task(io.setup()))

self.assertEqual(device.dispose_thread, device.setup_thread)
self.assertIsNone(io.dev)
self.assertIsNone(io._read_executor)
self.assertIsNone(io._write_executor)
97 changes: 68 additions & 29 deletions pylabrobot/io/ftdi.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,32 +178,68 @@ def _resolve_device_serial(self) -> str:
device_serial_number = cast(str, usb.util.get_string(device, device.iSerialNumber))
return device_serial_number

async def setup(self):
"""Initialize the FTDI device connection with device resolution."""
def _setup_sync(self) -> None:
"""Resolve and open the device. Runs on the executor that owns all device calls."""
if self._dev is not None and not self._dev.closed:
self._dev.close()
self._dev = None

# Resolve which device to connect to
self._device_id = self._resolve_device_serial()

# Create and open device
dev = Device(
lazy_open=True,
device_id=self.device_id,
pid=self._pid,
vid=self._vid,
interface_select=self._interface_select,
)
try:
# Resolve which device to connect to
self._device_id = self._resolve_device_serial()

# Create and open device
self._dev = Device(
lazy_open=True,
device_id=self.device_id,
pid=self._pid,
vid=self._vid,
interface_select=self._interface_select,
)
self._dev.open()
logger.info(f"Successfully opened FTDI device: {self.device_id}")
except FtdiError as e:
raise RuntimeError(
f"Failed to open FTDI device for '{self.human_readable_device_name}': {e}. "
"Is the device connected? Is it in use by another process? "
"Try restarting the kernel."
) from e
dev.open()
except BaseException:
try:
dev.close()
except Exception:
logger.warning("Failed to close FTDI device after setup failure", exc_info=True)
raise
self._dev = dev

self._executor = ThreadPoolExecutor(max_workers=1)
async def setup(self):
"""Initialize the FTDI device connection with device resolution."""
if self._executor is None:
self._executor = ThreadPoolExecutor(max_workers=1)
loop = asyncio.get_running_loop()
setup_future = loop.run_in_executor(self._executor, self._setup_sync)
try:
await asyncio.shield(setup_future)
except BaseException as exc:
if isinstance(exc, asyncio.CancelledError):
try:
await setup_future
except BaseException:
pass
if self._dev is not None:
try:
await loop.run_in_executor(self._executor, self._dev.close)
except Exception:
logger.warning("Failed to close FTDI device after setup failure", exc_info=True)
self._dev = None
self._shutdown_executor()
if isinstance(exc, FtdiError):
raise RuntimeError(
f"Failed to open FTDI device for '{self.human_readable_device_name}': {exc}. "
"Is the device connected? Is it in use by another process? "
"Try restarting the kernel."
) from exc
raise
logger.info(f"Successfully opened FTDI device: {self.device_id}")

def _shutdown_executor(self) -> None:
if self._executor is not None:
# the worker is idle here, so this does not block the event loop
self._executor.shutdown(wait=False, cancel_futures=True)
self._executor = None

@property
def device_id(self) -> str:
Expand Down Expand Up @@ -297,20 +333,22 @@ async def request_serial(self) -> str:
return self.device_id

async def stop(self):
loop = asyncio.get_running_loop()
if self._dev is not None:
self.dev.close()
if self._executor is not None:
self._executor.shutdown(wait=True)
self._executor = None
await loop.run_in_executor(self._executor, self.dev.close)
self._dev = None
self._shutdown_executor()

async def write(self, data: bytes) -> int:
"""Write data to the device. Returns the number of bytes written."""
logger.log(LOG_LEVEL_IO, "[%s] write %s", self._device_id, data)
capturer.record(FTDICommand(device_id=self.device_id, action="write", data=data.hex()))
return cast(int, self.dev.write(data))
loop = asyncio.get_running_loop()
return cast(int, await loop.run_in_executor(self._executor, self.dev.write, data))

async def read(self, num_bytes: int = 1) -> bytes:
data = self.dev.read(num_bytes)
loop = asyncio.get_running_loop()
data = await loop.run_in_executor(self._executor, self.dev.read, num_bytes)
if len(data) != 0:
logger.log(LOG_LEVEL_IO, "[%s] read %s", self._device_id, data)
capturer.record(
Expand All @@ -323,7 +361,8 @@ async def read(self, num_bytes: int = 1) -> bytes:
return cast(bytes, data)

async def readline(self) -> bytes: # type: ignore # very dumb it's reading from pyserial
data = self.dev.readline()
loop = asyncio.get_running_loop()
data = await loop.run_in_executor(self._executor, self.dev.readline)
if len(data) != 0:
logger.log(LOG_LEVEL_IO, "[%s] readline %s", self._device_id, data)
capturer.record(FTDICommand(device_id=self.device_id, action="readline", data=data.hex()))
Expand Down
Loading
Loading