diff --git a/doc/conf.py b/doc/conf.py index e6216605..e54d3d0e 100644 --- a/doc/conf.py +++ b/doc/conf.py @@ -4,6 +4,7 @@ # sphinx-quickstart on Wed Nov 9 16:42:53 2016. import os import sys + sys.path.insert(0, os.path.abspath("../lewis")) @@ -22,11 +23,11 @@ ] templates_path = ["_templates"] # General information about the project. -project = u"lewis" -language = 'en' +project = "lewis" +language = "en" exclude_patterns = ["_build", "Thumbs.db", ".DS_Store"] # -- Options for HTML output --------------------------------------------- -suppress_warnings =["docutils"] +suppress_warnings = ["docutils"] html_theme = "sphinx_rtd_theme" html_logo = "resources/logo/lewis-logo.png" html_context = { diff --git a/doc/developer_guide/framework_details.md b/doc/developer_guide/framework_details.md index 01dbe786..b8293dfd 100644 --- a/doc/developer_guide/framework_details.md +++ b/doc/developer_guide/framework_details.md @@ -72,3 +72,24 @@ statemachine: - Implicit: Implement handlers in the device class, with standard names like `on_entry_init` for a state called "init", and call `bindHandlersByName()` + +## Adapter Concurrency + +Adapters performing network I/O for communicating with client applications +make use of python's [asyncio](https://docs.python.org/3/library/asyncio.html) +library. + +- Lewis is a multi-threaded application, each adapter is moved on its own + dedicated thread, which is isolated from the main simulation thread. +- The main thread uses the following two synchronization tools: + - device lock: ensures that the device is only accessed from one + thread at a time + - is_running event: sends stop request to the adapter thread +- Adapters have to implement the following three + [async coroutines](https://docs.python.org/3/library/asyncio-task.html), + which will be scheduled as tasks by their respective async event loops: + - start_server: starts the server, handles client connections + - stop_server: gracefully closes client connections and stops the server + - handle: synchronizes with the simulation steps + +![The adapter concurrency diagram.](../resources/diagrams/AdapterConcurrency.drawio.png) diff --git a/doc/resources/diagrams/AdapterConcurrency.drawio.png b/doc/resources/diagrams/AdapterConcurrency.drawio.png new file mode 100644 index 00000000..bc50ffd3 Binary files /dev/null and b/doc/resources/diagrams/AdapterConcurrency.drawio.png differ diff --git a/lewis/adapters/epics.py b/lewis/adapters/epics.py index 950e766a..49d45820 100644 --- a/lewis/adapters/epics.py +++ b/lewis/adapters/epics.py @@ -441,14 +441,14 @@ def write(self, pv, value) -> bool: return True except LimitViolationException as e: self.log.warning( - "Rejected writing value %s to PV %s due to limit " "violation. %s", + "Rejected writing value %s to PV %s due to limit violation. %s", value, pv, e, ) except AccessViolationException: self.log.warning( - "Rejected writing value %s to PV %s due to access " "violation, PV is read-only.", + "Rejected writing value %s to PV %s due to access violation, PV is read-only.", value, pv, ) @@ -573,7 +573,7 @@ def documentation(self): return "\n\n".join([inspect.getdoc(self.interface) or "", "PVs\n==="] + pvs) - def start_server(self) -> None: + async def start_server(self) -> None: """ Creates a pcaspy-server. @@ -597,7 +597,7 @@ def start_server(self) -> None: ", ".join((self._options.prefix + pv for pv in self.interface.bound_pvs.keys())), ) - def stop_server(self) -> None: + async def stop_server(self) -> None: self._driver = None self._server = None @@ -605,7 +605,7 @@ def stop_server(self) -> None: def is_running(self): return self._server is not None - def handle(self, cycle_delay=0.1) -> None: + async def handle(self, cycle_delay=0.1) -> None: """ Call this method to spend about ``cycle_delay`` seconds processing requests in the pcaspy server. Under load, for example when running ``caget`` at a diff --git a/lewis/adapters/modbus.py b/lewis/adapters/modbus.py index 055bc400..c2c2d996 100644 --- a/lewis/adapters/modbus.py +++ b/lewis/adapters/modbus.py @@ -32,8 +32,7 @@ at lewis/examples/modbus_device. """ -import asyncore -import socket +import asyncio import struct from copy import deepcopy from math import ceil @@ -237,7 +236,7 @@ def create_exception(self, code): frame = deepcopy(self) frame.length = 3 frame.fcode += 0x80 - frame.data = bytearray(chr(code)) + frame.data = bytearray([code]) return frame def create_response(self, data=None): @@ -260,23 +259,22 @@ class ModbusProtocol: This class implements the Modbus TCP Protocol. The user of this class should provide a ModbusDataStore instance that will be used to - fulfill read and write requests, and a callable `sender` which accepts one bytearray - parameter. The `sender` will be called whenever a response frame is generated, with a - bytearray containing the response frame as the parameter. + fulfill read and write requests. The `writer` will be called whenever a response frame + is generated, with a bytearray containing the response frame as the parameter. Processing occurs when the user calls ModbusProtocol.process(), passing in the raw frame data to process as a bytearray. The data may include multiple frames and partial frame fragments. Any data that could not be processed (due to incomplete frames) is buffered for the next call to process. - :param sender: callable that accepts one bytearray parameter, called to send responses. + :param writer: asyncio.StreamWriter, called to send responses. :param datastore: ModbusDataStore instance to reference when processing requests """ - def __init__(self, sender, datastore) -> None: + def __init__(self, writer: asyncio.StreamWriter, datastore: ModbusDataStore) -> None: self._buffer = bytearray() self._datastore = datastore - self._send = lambda req: sender(req.to_bytearray()) + self._writer = writer # Lookup table to handle requests as per Modbus Application Protocol v1.1b3, Section 6. self._fcode_handler_map = { @@ -290,7 +288,11 @@ def __init__(self, sender, datastore) -> None: 0x10: self._handle_write_multiple_registers, } - def process(self, data, device_lock) -> None: + async def _send(self, response) -> None: + self._writer.write(response.to_bytearray()) + await self._writer.drain() + + async def process(self, data, device_lock) -> None: """ Process as much of given data as possible. @@ -302,22 +304,22 @@ def process(self, data, device_lock) -> None: """ self._buffer.extend(bytearray(data)) + responses = [] with device_lock: for request in self._buffered_requests(): - self.log.debug( - "Request: %s", - str(["{:#04x}".format(c) for c in request.to_bytearray()]), - ) - handler = self._get_handler(request.fcode) - response = handler(request) - - self.log.debug( - "Response: %s", - str(["{:#04x}".format(c) for c in response.to_bytearray()]), - ) + responses.append((request, handler(request))) - self._send(response) + for request, response in responses: + self.log.debug( + "Request: %s", + str(["{:#04x}".format(c) for c in request.to_bytearray()]), + ) + self.log.debug( + "Response: %s", + str(["{:#04x}".format(c) for c in response.to_bytearray()]), + ) + await self._send(response) def _buffered_requests(self): """Generator to yield all complete modbus requests in the internal buffer""" @@ -529,59 +531,104 @@ def _handle_write_multiple_registers(self, request): @has_log -class ModbusHandler(asyncore.dispatcher_with_send): - def __init__(self, sock, interface, server) -> None: - asyncore.dispatcher_with_send.__init__(self, sock=sock) +class ModbusHandler: + def __init__( + self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, interface, server + ) -> None: self._datastore = ModbusDataStore(interface.di, interface.co, interface.ir, interface.hr) - self._modbus = ModbusProtocol(self.send, self._datastore) + self._modbus = ModbusProtocol(writer, self._datastore) self._server = server + self._reader = reader + self._writer = writer + self._closing = False self._set_logging_context(interface) - self.log.info("Client connected from %s:%s", *sock.getpeername()) - def handle_read(self) -> None: - data = self.recv(8192) - self._modbus.process(data, self._server.device_lock) - - def handle_close(self) -> None: - self.log.info("Closing connection to client %s:%s", *self.socket.getpeername()) + async def handle_client(self) -> None: + try: + while True: + data = await self._reader.read(8192) + if data: + await self._modbus.process(data, self._server.device_lock) + else: + break + except OSError as e: + self.log.error("Connection error: %s", e) + finally: + await self.handle_close() + + async def handle_close(self) -> None: + if self._closing: + return + self._closing = True + sock = self._writer.get_extra_info("socket") + if sock is not None: + try: + self.log.info("Closing connection to client %s:%s", *sock.getpeername()) + except OSError: + self.log.info("Closing connection to client (peer address unavailable)") + if not self._writer.is_closing(): + self._writer.close() + try: + await self._writer.wait_closed() + except OSError: + self.log.debug("Connection reset by peer while waiting for close") self._server.remove_handler(self) - self.close() @has_log -class ModbusServer(asyncore.dispatcher): +class ModbusServer: def __init__(self, host, port, interface, device_lock) -> None: - asyncore.dispatcher.__init__(self) + self.host = host + self.port = port self.device_lock = device_lock self.interface = interface - self.create_socket(socket.AF_INET, socket.SOCK_STREAM) - self.set_reuse_addr() - self.bind((host, port)) - self.listen(5) + self._server = None self._set_logging_context(interface) - self.log.info("Listening on %s:%s", host, port) self._accepted_connections = [] - def handle_accept(self) -> None: - pair = self.accept() - if pair is not None: - sock, _ = pair - handler = ModbusHandler(sock, self.interface, self) - self._accepted_connections.append(handler) + async def start(self): + self._server = await asyncio.start_server( + self._handle_accept, + host=self.host, + port=self.port, + backlog=5, + reuse_address=True, + start_serving=True, + ) + self.log.info("Listening on %s:%s", self.host, self.port) + + async def _handle_accept( + self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter + ) -> None: + sock = writer.get_extra_info("socket") + if sock is not None: + try: + self.log.info("Client connected from %s:%s", *sock.getpeername()) + except OSError: + self.log.info("Client connected (peer address unavailable)") + handler = ModbusHandler(reader, writer, self.interface, self) + self._accepted_connections.append(handler) + await handler.handle_client() def remove_handler(self, handler) -> None: - self._accepted_connections.remove(handler) + try: + self._accepted_connections.remove(handler) + except ValueError: + pass # Removed from another path - def handle_close(self) -> None: - self.log.info("Shutting down server, closing all remaining client connections.") + async def close(self) -> None: + if self._server is not None: + self.log.info("Shutting down server, closing all remaining client connections.") + self._server.close() - for handler in self._accepted_connections: - handler.close() - self._accepted_connections = [] - self.close() + for handler in list(self._accepted_connections): + await handler.handle_close() + + self._accepted_connections = [] + await self._server.wait_closed() class ModbusAdapter(Adapter): @@ -591,7 +638,7 @@ def __init__(self, options=None) -> None: super(ModbusAdapter, self).__init__(options) self._server = None - def start_server(self) -> None: + async def start_server(self) -> None: self._server = ModbusServer( self._options.bind_address, self._options.port, @@ -599,17 +646,19 @@ def start_server(self) -> None: self.device_lock, ) - def stop_server(self) -> None: + await self._server.start() + + async def stop_server(self) -> None: if self._server is not None: - self._server.close() + await self._server.close() self._server = None @property def is_running(self): return self._server is not None - def handle(self, cycle_delay=0.1) -> None: - asyncore.loop(cycle_delay, count=1) + async def handle(self, cycle_delay=0.1) -> None: + await asyncio.sleep(cycle_delay) class ModbusInterface(InterfaceBase): diff --git a/lewis/adapters/stream.py b/lewis/adapters/stream.py index 35f65dfd..1ed4c066 100644 --- a/lewis/adapters/stream.py +++ b/lewis/adapters/stream.py @@ -17,11 +17,9 @@ # along with this program. If not, see . # ********************************************************************* -import asynchat -import asyncore +import asyncio import inspect import re -import socket from typing import NoReturn from scanf import scanf_compile @@ -33,10 +31,11 @@ @has_log -class StreamHandler(asynchat.async_chat): - def __init__(self, sock, target, stream_server) -> None: - asynchat.async_chat.__init__(self, sock=sock) - self.set_terminator(target.in_terminator.encode()) +class StreamHandler: + def __init__( + self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, target, stream_server + ) -> None: + self._in_terminator = target.in_terminator.encode() self._readtimeout = target.readtimeout self._readtimer = 0 self._target = target @@ -44,25 +43,54 @@ def __init__(self, sock, target, stream_server) -> None: self._stream_server = stream_server self._target.handler = self + self._reader = reader + self._writer = writer + self._pending_read: asyncio.Task | None = None + self._closing = False self._set_logging_context(target) - self.log.info("Client connected from %s:%s", *sock.getpeername()) - def process(self, msec) -> None: + async def process(self, msec) -> None: + # Start a read operation if none is in flight + if self._pending_read is None: + self._pending_read = asyncio.ensure_future(self._reader.read(4096)) + + # Process data if the read completed since the last tick + if self._pending_read.done(): + try: + chunk = self._pending_read.result() + except Exception as e: + self._pending_read = None + self.log.error("Error reading from client: %s", e) + await self.handle_close() + return + self._pending_read = None + + if not chunk: # EOF - client disconnected + await self.handle_close() + return + + self.collect_incoming_data(chunk) + + if self._in_terminator: + while not self._closing and b"".join(self._buffer).find(self._in_terminator) != -1: + await self.found_terminator() + + # Timeout processing if not self._buffer: return if self._readtimer >= self._readtimeout and self._readtimeout != 0: - if not self.get_terminator(): + if not self._in_terminator: # If no terminator is set, this timeout is the terminator - self.found_terminator() + await self.found_terminator() else: self._readtimer = 0 request = self._get_request() with self._stream_server.device_lock: error = RuntimeError("ReadTimeout while waiting for command terminator.") reply = self._handle_error(request, error) - self._send_reply(reply) + await self._send_reply(reply) if self._buffer: self._readtimer += msec @@ -71,13 +99,26 @@ def collect_incoming_data(self, data) -> None: self._buffer.append(data) self._readtimer = 0 - def _get_request(self): - request = b"".join(self._buffer) - self._buffer = [] + def _get_request(self) -> bytes: + data = b"".join(self._buffer) + if self._in_terminator: + term_pos = data.find(self._in_terminator) + if term_pos != -1: + request = data[:term_pos] + remainder = data[term_pos + len(self._in_terminator) :] + self._buffer = [remainder] if remainder else [] + else: + request = data + self._buffer = [] + else: + request = data + self._buffer = [] self.log.debug("Got request %s", request) return request - def _push(self, reply) -> None: + async def _push(self, reply) -> None: + if self._closing: + return try: if isinstance(reply, str): reply = reply.encode() @@ -86,20 +127,24 @@ def _push(self, reply) -> None: if isinstance(self._target.out_terminator, str) else self._target.out_terminator ) - self.push(reply + out_terminator) + self._writer.write(reply + out_terminator) + await self._writer.drain() except TypeError as e: self.log.error("Problem creating reply, type error {}!".format(e)) + except OSError as e: + self.log.error("Connection error while sending reply: %s", e) + await self.handle_close() - def _send_reply(self, reply) -> None: + async def _send_reply(self, reply) -> None: if reply is not None: self.log.debug("Sending reply %s", reply) - self._push(reply) + await self._push(reply) def _handle_error(self, request, error): self.log.debug("Error while processing request", exc_info=error) return self._target.handle_error(request, error) - def found_terminator(self) -> None: + async def found_terminator(self) -> None: self._readtimer = 0 request = self._get_request() @@ -124,61 +169,103 @@ def found_terminator(self) -> None: except Exception as error: reply = self._handle_error(request, error) - self._send_reply(reply) + await self._send_reply(reply) def unsolicited_reply(self, reply) -> None: + if self._closing: + return self.log.debug("Sending unsolicited reply %s", reply) - self._push(reply) + if self._stream_server._loop is None: + raise RuntimeError("Cannot send unsolicited reply: server not started.") + asyncio.run_coroutine_threadsafe(self._push(reply), self._stream_server._loop).result( + timeout=5.0 + ) - def handle_close(self) -> None: - self.log.info("Closing connection to client %s:%s", *self.socket.getpeername()) + async def handle_close(self) -> None: + if self._closing: + return + self._closing = True + if self._target.handler is self: + del self._target.handler + if self._pending_read is not None and not self._pending_read.done(): + self._pending_read.cancel() + try: + await self._pending_read + except asyncio.CancelledError: + pass # Suppress RuntimeWarning for not awaiting a cancelled task + self._pending_read = None + sock = self._writer.get_extra_info("socket") + if sock is not None: + try: + self.log.info("Closing connection to client %s:%s", *sock.getpeername()) + except OSError: + self.log.info("Closing connection to client (peer address unavailable)") + if not self._writer.is_closing(): + self._writer.close() + try: + await self._writer.wait_closed() + except OSError: + self.log.debug("Connection reset by peer while waiting for close") self._stream_server.remove_handler(self) - asynchat.async_chat.handle_close(self) @has_log -class StreamServer(asyncore.dispatcher): +class StreamServer: def __init__(self, host, port, target, device_lock) -> None: - asyncore.dispatcher.__init__(self) + self.host = host + self.port = port self.target = target self.device_lock = device_lock - self.create_socket(socket.AF_INET, socket.SOCK_STREAM) - self.set_reuse_addr() - self.bind((host, port)) - self.listen(5) + self._loop: asyncio.AbstractEventLoop | None = None + self._server = None + self._accepted_connections: list[StreamHandler] = [] self._set_logging_context(target) - self.log.info("Listening on %s:%s", host, port) - - self._accepted_connections = [] - def handle_accept(self) -> None: - pair = self.accept() - if pair is not None: - sock, addr = pair - handler = StreamHandler(sock, self.target, self) + async def start(self): + self._loop = asyncio.get_running_loop() + self._server = await asyncio.start_server( + self._handle_accept, + host=self.host, + port=self.port, + backlog=5, + reuse_address=True, + start_serving=True, + ) + self.log.info("Listening on %s:%s", self.host, self.port) - self._accepted_connections.append(handler) + def _handle_accept(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + sock = writer.get_extra_info("socket") + if sock is not None: + try: + self.log.info("Client connected from %s:%s", *sock.getpeername()) + except OSError: + self.log.info("Client connected (peer address unavailable)") + handler = StreamHandler(reader, writer, self.target, self) + self._accepted_connections.append(handler) def remove_handler(self, handler) -> None: - self._accepted_connections.remove(handler) + try: + self._accepted_connections.remove(handler) + except ValueError: + pass # Removed from another path - def close(self) -> None: - # As this is an old style class, the base class method must - # be called directly. This is important to still perform all - # the teardown-work that asyncore.dispatcher does. - self.log.info("Shutting down server, closing all remaining client connections.") - asyncore.dispatcher.close(self) + async def close(self) -> None: + if self._server is not None: + self.log.info("Shutting down server, closing all remaining client connections.") + self._server.close() - # But in addition, close all open sockets and clear the connection list. - for handler in self._accepted_connections: - handler.close() + # Close all open sockets and clear the connection list. + for handler in list(self._accepted_connections): + await handler.handle_close() - self._accepted_connections = [] + self._accepted_connections = [] + await self._server.wait_closed() + self._loop = None - def process(self, msec) -> None: - for handler in self._accepted_connections: - handler.process(msec) + async def process(self, msec) -> None: + for handler in list(self._accepted_connections): + await handler.process(msec) class PatternMatcher: @@ -713,7 +800,7 @@ def documentation(self): + commands ) - def start_server(self) -> None: + async def start_server(self) -> None: """ Starts the TCP stream server, binding to the configured host and port. Host and port are configured via the command line arguments. @@ -734,23 +821,25 @@ def start_server(self) -> None: self.device_lock, ) - def stop_server(self) -> None: + await self._server.start() + + async def stop_server(self) -> None: if self._server is not None: - self._server.close() + await self._server.close() self._server = None @property def is_running(self): return self._server is not None - def handle(self, cycle_delay=0.1) -> None: + async def handle(self, cycle_delay=0.1) -> None: """ Spend approximately ``cycle_delay`` seconds to process requests to the server. :param cycle_delay: S """ - asyncore.loop(cycle_delay, count=1) - self._server.process(int(cycle_delay * 1000)) + await self._server.process(int(cycle_delay * 1000)) + await asyncio.sleep(cycle_delay) class StreamInterface(InterfaceBase): @@ -842,7 +931,7 @@ def _bind_device(self) -> None: pattern = bound_cmd.matcher.pattern if pattern in patterns: raise RuntimeError( - "The regular expression {} is " "associated with multiple commands.".format( + "The regular expression {} is associated with multiple commands.".format( pattern ) ) diff --git a/lewis/core/adapters.py b/lewis/core/adapters.py index 2c7565c5..ce6533d0 100644 --- a/lewis/core/adapters.py +++ b/lewis/core/adapters.py @@ -23,6 +23,7 @@ be used to store multiple adapters and manage them together. """ +import asyncio import inspect import logging import threading @@ -45,8 +46,7 @@ class NoLock: def __enter__(self) -> None: raise RuntimeError( - "The attempted action requires a proper threading.Lock-object, " - "but none was available." + "The attempted action requires a proper threading.Lock-object, but none was available." ) def __exit__( @@ -150,7 +150,7 @@ def documentation(self) -> str: """ return inspect.getdoc(self) or "" - def start_server(self) -> None: + async def start_server(self) -> None: """ This method must be re-implemented to start the infrastructure required for the protocol in question. These startup operations are not supposed to be carried out on @@ -169,7 +169,7 @@ def start_server(self) -> None: "required for network communication." ) - def stop_server(self) -> None: + async def stop_server(self) -> None: """ This method must be re-implemented to stop and tear down anything that has been setup in :meth:`start_server`. This method should close all connections to clients that have @@ -196,7 +196,7 @@ def is_running(self) -> bool: "a server is currently running and listening for requests." ) - def handle(self, cycle_delay: float = 0.1) -> None: + async def handle(self, cycle_delay: float = 0.1) -> None: """ This function is called on each cycle of a simulation. It should process requests that are made via the protocol that exposes the device. The time spent processing should be @@ -305,7 +305,9 @@ def _start_server(self, adapter: Adapter) -> None: if adapter.protocol not in self._threads: self.log.info("Connecting device interface for protocol '%s'", adapter.protocol) - adapter_thread = threading.Thread(target=self._adapter_loop, args=(adapter, 0.01)) + adapter_thread = threading.Thread( + target=lambda: asyncio.run(self._adapter_loop(adapter, 0.01)) + ) adapter_thread.daemon = True self._threads[adapter.protocol] = adapter_thread @@ -318,17 +320,22 @@ def _start_server(self, adapter: Adapter) -> None: if not self._running[adapter.protocol].is_set(): raise LewisException("Adapter for '%s' failed to start!" % adapter.protocol) - def _adapter_loop(self, adapter: Adapter, dt: float) -> None: + async def _adapter_loop(self, adapter: Adapter, dt: float) -> None: adapter.device_lock = self._lock # This ensures that the adapter is using the correct lock - adapter.start_server() + await adapter.start_server() self._running[adapter.protocol].set() self.log.debug("Starting adapter loop for protocol %s.", adapter.protocol) - while self._running[adapter.protocol].is_set(): - adapter.handle(dt) - - adapter.stop_server() + try: + while self._running[adapter.protocol].is_set(): + await adapter.handle(dt) + except Exception: + self.log.exception("Adapter loop for protocol '%s' crashed.", adapter.protocol) + self._running[adapter.protocol].clear() + raise + finally: + await adapter.stop_server() def disconnect(self, *args: str) -> None: """ diff --git a/lewis/core/devices.py b/lewis/core/devices.py index 9dc86865..4f5bc03f 100644 --- a/lewis/core/devices.py +++ b/lewis/core/devices.py @@ -349,8 +349,7 @@ def create_device(self, setup=None): if setup_name not in self.setups: raise LewisException( - "Failed to find setup '{}' for device '{}'. " - "Available setups are:\n {}".format( + "Failed to find setup '{}' for device '{}'. Available setups are:\n {}".format( setup, self.name, "\n ".join(self.setups.keys()) ) ) diff --git a/lewis/core/statemachine.py b/lewis/core/statemachine.py index 5aadcac2..9f3d7de3 100644 --- a/lewis/core/statemachine.py +++ b/lewis/core/statemachine.py @@ -191,7 +191,7 @@ def __init__(self, cfg, context=None) -> None: # Specifying an initial state is not optional if "initial" not in cfg: raise StateMachineException( - "StateMachine configuration must include " "'initial' to specify starting state." + "StateMachine configuration must include 'initial' to specify starting state." ) self._initial = cfg["initial"] self._set_handlers(self._initial) diff --git a/lewis/scripts/control.py b/lewis/scripts/control.py index ce7db619..3cca58f8 100644 --- a/lewis/scripts/control.py +++ b/lewis/scripts/control.py @@ -117,7 +117,7 @@ def call_method(remote, object_name, method, arguments): positional_args.add_argument( "arguments", nargs="*", - help="Arguments to method call. For setting a property, " "supply the property value. ", + help="Arguments to method call. For setting a property, supply the property value. ", ) optional_args = parser.add_argument_group("Optional arguments") diff --git a/lewis/scripts/run.py b/lewis/scripts/run.py index 4e31dfbc..a75d8ed6 100644 --- a/lewis/scripts/run.py +++ b/lewis/scripts/run.py @@ -160,8 +160,7 @@ "-I", "--ignore-versions", action="store_true", - help="Ignore version mismatches between device and framework. A warning will still " - "be logged.", + help="Ignore version mismatches between device and framework. A warning will still be logged.", ) other_args.add_argument( "-v", "--version", action="store_true", help="Prints the version and exits." diff --git a/pyproject.toml b/pyproject.toml index d144cb8b..33cfb9a0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,7 +36,6 @@ dependencies=[ "semantic_version", "PyYAML", "scanf", - "pyasynchat;python_version >= '3.12'", ] [project.optional-dependencies] diff --git a/ruff.toml b/ruff.toml new file mode 100644 index 00000000..9f18fea3 --- /dev/null +++ b/ruff.toml @@ -0,0 +1,33 @@ +# Exclude a variety of commonly ignored directories. +exclude = [ + ".bzr", + ".direnv", + ".eggs", + ".git", + ".git-rewrite", + ".hg", + ".ipynb_checkpoints", + ".mypy_cache", + ".nox", + ".pants.d", + ".pyenv", + ".pytest_cache", + ".pytype", + ".ruff_cache", + ".svn", + ".tox", + ".venv", + ".vscode", + "__pypackages__", + "_build", + "buck-out", + "build", + "dist", + "node_modules", + "site-packages", + "venv", +] + +# Set the maximum line length to 100. +line-length = 100 +indent-width = 4 diff --git a/tests/test_StateMachine.py b/tests/test_StateMachine.py index 1829bfa8..54281e52 100644 --- a/tests/test_StateMachine.py +++ b/tests/test_StateMachine.py @@ -45,7 +45,7 @@ def test_first_cycle_transitions_to_initial(self): self.assertEqual( sm.state, "foobar", - "StateMachine failed to transition into " "initial state on first cycle", + "StateMachine failed to transition into initial state on first cycle", ) def test_can_transition_with_lambda(self): diff --git a/tests/test_core_adapters.py b/tests/test_core_adapters.py index 8be50a34..96ee78a2 100644 --- a/tests/test_core_adapters.py +++ b/tests/test_core_adapters.py @@ -27,10 +27,10 @@ def __init__(self, protocol, running=False, options=None): def protocol(self): return self._protocol - def start_server(self): + async def start_server(self): self._running = True - def stop_server(self): + async def stop_server(self): self._running = False @property @@ -47,19 +47,28 @@ def failing_function(): self.assertRaises(RuntimeError, failing_function) -class TestAdapter(unittest.TestCase): +class TestAdapter(unittest.IsolatedAsyncioTestCase): def test_documentation(self): adapter = DummyAdapter("foo") self.assertEqual(inspect.cleandoc(adapter.__doc__), adapter.documentation) - def test_not_implemented_errors(self): + async def test_not_implemented_errors(self): adapter = Adapter() - self.assertRaises(NotImplementedError, adapter.start_server) - self.assertRaises(NotImplementedError, adapter.stop_server) + with self.assertRaises(NotImplementedError): + await adapter.start_server() + with self.assertRaises(NotImplementedError): + await adapter.stop_server() self.assertRaises(NotImplementedError, getattr, adapter, "is_running") - assertRaisesNothing(self, adapter.handle, 0) + + try: + await adapter.handle(0) + except Exception as exc: + self.fail( + "Assertion error. An exception was caught where none " + "was expected in %s. Message: %s" % (adapter.handle.__name__, str(exc)) + ) def test_interface_property(self): adapter = Adapter() diff --git a/tests/test_stream_adapter.py b/tests/test_stream_adapter.py index 6b1704d7..e17aa76e 100644 --- a/tests/test_stream_adapter.py +++ b/tests/test_stream_adapter.py @@ -1,33 +1,169 @@ -import unittest -from unittest.mock import MagicMock, patch +import asyncio +from unittest import IsolatedAsyncioTestCase +from unittest.mock import MagicMock, AsyncMock from parameterized import parameterized from lewis.adapters.stream import StreamHandler -@patch("asynchat.async_chat") -class TestStreamHandler(unittest.TestCase): +class TestStreamHandler(IsolatedAsyncioTestCase): def setUp(self): - """Create a mock for the async_chat class""" self.target = MagicMock() self.stream_server = MagicMock() - self.socket = MagicMock() - self.handler = StreamHandler(self.socket, self.target, self.stream_server) + self.stream_reader = AsyncMock() + self.stream_writer = MagicMock() + self.stream_writer.drain = AsyncMock() + self.stream_writer.wait_closed = AsyncMock() + self.stream_writer.is_closing.return_value = False + self.handler = StreamHandler( + reader=self.stream_reader, + writer=self.stream_writer, + target=self.target, + stream_server=self.stream_server, + ) + self.handler._readtimeout = 0 + self.handler._in_terminator = b"\r\n" + self.target.out_terminator = "\r\n" + + def _create_mock_command(self, can_process=lambda x: True, response="OK"): + cmd_mock = MagicMock() + cmd_mock.can_process.side_effect = can_process + cmd_mock.process_request.return_value = response + return cmd_mock @parameterized.expand( [ (b"\n", "test", b"test\n"), (b"\n", b"test", b"test\n"), ("\n", "test", b"test\n"), - ("\n", "test", b"test\n"), + ("\r\n", "test", b"test\r\n"), ] ) - @patch("asynchat.async_chat.push") - def test_terminator_and_replies_of_different_types_can_be_concatenated( - self, terminator, message, expected, async_push, _ + async def test_terminator_and_replies_of_different_types_can_be_concatenated( + self, terminator, message, expected ): self.target.out_terminator = terminator - self.handler.unsolicited_reply(message) + # unsolicited_reply is sync and uses run_coroutine_threadsafe; it must be called + # from a worker thread so that .result() does not block the running event loop. + self.stream_server._loop = asyncio.get_running_loop() + await asyncio.get_running_loop().run_in_executor( + None, self.handler.unsolicited_reply, message + ) + + self.stream_writer.write.assert_called_once_with(expected) + + async def test_process_starts_pending_read_on_first_call(self): + await self.handler.process(10) + + self.assertIsNotNone(self.handler._pending_read) + + async def test_process_eof_triggers_handle_close(self): + self.handler._reader.read.return_value = b"" + + await self.handler.process(10) + await asyncio.sleep(0) + await self.handler.process(10) + + self.stream_server.remove_handler.assert_called_with(self.handler) + + async def test_process_dispatches_single_command_with_terminator(self): + cmd_mock = self._create_mock_command() + self.target.bound_commands = [cmd_mock] + self.handler._reader.read.return_value = b"CMD\r\n" + + await self.handler.process(10) + await asyncio.sleep(0) + await self.handler.process(10) + + cmd_mock.can_process.assert_called_with(b"CMD") + cmd_mock.process_request.assert_called_with(b"CMD") + + self.stream_writer.write.assert_called_once_with(b"OK\r\n") + + async def test_process_dispatches_two_commands_in_one_chunk(self): + cmd1_mock = self._create_mock_command(can_process=lambda x: x == b"CMD1", response="OK1") + cmd2_mock = self._create_mock_command(can_process=lambda x: x == b"CMD2", response="OK2") + self.target.bound_commands = [cmd1_mock, cmd2_mock] + self.handler._reader.read.return_value = b"CMD1\r\nCMD2\r\n" + + await self.handler.process(10) + await asyncio.sleep(0) + await self.handler.process(10) + + cmd1_mock.can_process.assert_called() + cmd2_mock.can_process.assert_called() + cmd1_mock.process_request.assert_called_once_with(b"CMD1") + cmd2_mock.process_request.assert_called_once_with(b"CMD2") + + self.stream_writer.write.assert_any_call(b"OK1\r\n") + self.stream_writer.write.assert_any_call(b"OK2\r\n") + self.assertEqual(self.stream_writer.write.call_count, 2) + + async def test_process_timeout_with_incomplete_command_sends_error(self): + self.handler._readtimeout = 10 + self.handler._reader.read.return_value = b"INCOMPLETE" + + # First call: starts the read task + await self.handler.process(10) + await asyncio.sleep(0) + # Second call: collects data; _readtimer resets to 0, then increments to 10 + await self.handler.process(10) + # Third call: _readtimer (10) >= _readtimeout (10) -> timeout fires, error reply sent + await self.handler.process(10) + + self.target.handle_error.assert_called_once() + self.stream_writer.write.assert_called_once() + + async def test_handle_close_is_idempotent(self): + await self.handler.handle_close() + await self.handler.handle_close() + + self.stream_server.remove_handler.assert_called_once_with(self.handler) + + async def test_unsolicited_reply_is_silent_noop_after_close(self): + await self.handler.handle_close() + + self.stream_server._loop = asyncio.get_running_loop() + # unsolicited_reply must be called from a worker thread (it calls .result() internally) + await asyncio.get_running_loop().run_in_executor( + None, self.handler.unsolicited_reply, "hello" + ) + + self.stream_writer.write.assert_not_called() + self.stream_server.remove_handler.assert_called_once_with(self.handler) + + async def test_push_does_not_write_after_close(self): + await self.handler.handle_close() + self.stream_writer.write.reset_mock() + + await self.handler._push("hello") + + self.stream_writer.write.assert_not_called() + + async def test_push_oserror_triggers_handle_close(self): + cmd_mock = self._create_mock_command() + self.target.bound_commands = [cmd_mock] + self.handler._reader.read.return_value = b"CMD\r\n" + self.stream_writer.drain.side_effect = OSError("connection broken") + + await self.handler.process(10) + await asyncio.sleep(0) + await self.handler.process(10) + + self.stream_server.remove_handler.assert_called_once_with(self.handler) + + async def test_process_stops_dispatching_on_broken_connection(self): + cmd1_mock = self._create_mock_command(can_process=lambda x: x == b"CMD1", response="OK1") + cmd2_mock = self._create_mock_command(can_process=lambda x: x == b"CMD2", response="OK2") + self.target.bound_commands = [cmd1_mock, cmd2_mock] + self.handler._reader.read.return_value = b"CMD1\r\nCMD2\r\n" + self.stream_writer.drain.side_effect = OSError("connection broken") + + await self.handler.process(10) + await asyncio.sleep(0) + await self.handler.process(10) - self.assertEqual(expected, async_push.call_args[0][0]) + cmd1_mock.process_request.assert_called_once_with(b"CMD1") + cmd2_mock.process_request.assert_not_called() + self.stream_server.remove_handler.assert_called_once_with(self.handler)