diff --git a/frigate/comms/mqtt.py b/frigate/comms/mqtt.py index 28ef010893..16bd19d08a 100644 --- a/frigate/comms/mqtt.py +++ b/frigate/comms/mqtt.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging import queue import selectors +import socket import threading import time from collections.abc import Callable @@ -57,6 +58,11 @@ class MqttClient(Communicator): self._next_connect_time = 0.0 self._last_on_connect_dispatch = 0.0 + # lets other threads interrupt the worker's socket wait + self._wake_recv, self._wake_send = socket.socketpair() + self._wake_recv.setblocking(False) + self._wake_send.setblocking(False) + def subscribe(self, receiver: Callable) -> None: """Wrapper for allowing dispatcher to subscribe.""" self._dispatcher = receiver @@ -86,6 +92,7 @@ class MqttClient(Communicator): return self._publish_queue.put(QueuedPublish(full_topic, payload, retain)) + self._wake_worker() def stop(self) -> None: if self._worker is None: @@ -101,9 +108,11 @@ class MqttClient(Communicator): publish_done, ) ) + self._wake_worker() publish_done.wait(MQTT_SHUTDOWN_FLUSH_TIMEOUT) self._stop_event.set() + self._wake_worker() if self.client is not None: try: @@ -369,11 +378,17 @@ class MqttClient(Communicator): exc_info=True, ) + def _wake_worker(self) -> None: + try: + self._wake_send.send(b"\0") + except BlockingIOError: + # the buffer is full, so a wake is already pending + pass + def _loop_client(self, timeout: float) -> int: """Drive Paho without select()'s limit on socket file descriptors.""" + assert self.client is not None client = self.client - if client is None: - return mqtt.MQTT_ERR_NO_CONN sock = client.socket() if sock is None: return mqtt.MQTT_ERR_NO_CONN @@ -385,11 +400,19 @@ class MqttClient(Communicator): pending = hasattr(sock, "pending") and sock.pending() > 0 with selectors.DefaultSelector() as selector: selector.register(sock, events) - ready = selector.select(0.0 if pending else timeout) + selector.register(self._wake_recv, selectors.EVENT_READ) + ready = { + key.fileobj: mask + for key, mask in selector.select(0.0 if pending else timeout) + } - ready_events = 0 - for _, mask in ready: - ready_events |= mask + if self._wake_recv in ready: + try: + self._wake_recv.recv(4096) + except BlockingIOError: + pass + + ready_events = ready.get(sock, 0) if pending or ready_events & selectors.EVENT_READ: result = client.loop_read() if result != mqtt.MQTT_ERR_SUCCESS or client.socket() is None: diff --git a/frigate/test/test_mqtt_lifecycle.py b/frigate/test/test_mqtt_lifecycle.py index d9c9831096..7cd55c28a2 100644 --- a/frigate/test/test_mqtt_lifecycle.py +++ b/frigate/test/test_mqtt_lifecycle.py @@ -126,12 +126,18 @@ class TestMqttClientLifecycle(unittest.TestCase): os.makedirs(MODEL_CACHE_DIR) self.config = build_config() - self.client = MqttClient(self.config) + self.client = self._build_client() self.receiver = RuntimeSnapshotReceiver() self.client.attach_dispatcher(build_dispatcher(self.config, [])) - def test_subscribe_stores_receiver_without_starting_worker(self) -> None: + def _build_client(self) -> MqttClient: client = MqttClient(self.config) + self.addCleanup(client._wake_recv.close) + self.addCleanup(client._wake_send.close) + return client + + def test_subscribe_stores_receiver_without_starting_worker(self) -> None: + client = self._build_client() with patch.object(client, "_start_worker") as mock_start_worker: client.subscribe(self.receiver._receive) @@ -142,7 +148,7 @@ class TestMqttClientLifecycle(unittest.TestCase): mock_start_worker.assert_not_called() def test_attach_dispatcher_supplies_command_surface(self) -> None: - client = MqttClient(self.config) + client = self._build_client() self.assertFalse(client._is_supported_command_topic("front/detect/set")) @@ -295,6 +301,13 @@ class TestMqttClientLifecycle(unittest.TestCase): self.assertEqual(self.client._subscription_mid, 42) self.client.client.subscribe.assert_called_once_with("frigate/#", qos=0) + def test_publish_wakes_worker(self) -> None: + self.client.connected = True + + self.client.publish("events", "payload") + + self.assertEqual(self.client._wake_recv.recv(16), b"\0") + def test_handle_connect_event_reconnects_on_recoverable_subscribe_error( self, ) -> None: diff --git a/frigate/test/test_mqtt_network_loop.py b/frigate/test/test_mqtt_network_loop.py index ea28e9cd49..175ec8e969 100644 --- a/frigate/test/test_mqtt_network_loop.py +++ b/frigate/test/test_mqtt_network_loop.py @@ -2,6 +2,7 @@ import fcntl import resource import selectors import socket +import time import unittest from unittest.mock import MagicMock, patch @@ -23,6 +24,11 @@ class TestMqttNetworkLoop(unittest.TestCase): self.sock, self.peer = socket.socketpair() self.addCleanup(self.sock.close) self.addCleanup(self.peer.close) + self.transport._wake_recv, self.transport._wake_send = socket.socketpair() + self.transport._wake_recv.setblocking(False) + self.transport._wake_send.setblocking(False) + self.addCleanup(self.transport._wake_recv.close) + self.addCleanup(self.transport._wake_send.close) self.client.socket.return_value = self.sock def test_high_fd_handles_connack_suback_publish_and_puback(self) -> None: @@ -123,8 +129,18 @@ class TestMqttNetworkLoop(unittest.TestCase): self.assertEqual(self.transport._loop_client(0), mqtt.MQTT_ERR_CONN_LOST) self.client.loop_misc.assert_not_called() - def test_missing_client_or_socket_reports_no_connection(self) -> None: + def test_missing_socket_reports_no_connection(self) -> None: self.client.socket.return_value = None self.assertEqual(self.transport._loop_client(0), mqtt.MQTT_ERR_NO_CONN) - self.transport.client = None - self.assertEqual(self.transport._loop_client(0), mqtt.MQTT_ERR_NO_CONN) + + def test_wake_interrupts_wait_and_is_consumed(self) -> None: + self.transport._wake_worker() + + start = time.monotonic() + self.assertEqual(self.transport._loop_client(5), mqtt.MQTT_ERR_SUCCESS) + self.assertLess(time.monotonic() - start, 1) + self.client.loop_read.assert_not_called() + + start = time.monotonic() + self.transport._loop_client(0.2) + self.assertGreater(time.monotonic() - start, 0.15)