wake the mqtt worker when a publish is queued

publish() only queued the message, and the worker was blocked waiting on the broker socket for up to a second, so on a quiet connection messages went out up to 1s late. A socketpair registered in the worker's selector now interrupts the wait when a publish is queued or stop() is called.
This commit is contained in:
Josh Hawkins
2026-10-06 08:30:16 -05:00
parent 6b1e084fdc
commit df4e4860b5
3 changed files with 64 additions and 12 deletions
+29 -6
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import logging import logging
import queue import queue
import selectors import selectors
import socket
import threading import threading
import time import time
from collections.abc import Callable from collections.abc import Callable
@@ -57,6 +58,11 @@ class MqttClient(Communicator):
self._next_connect_time = 0.0 self._next_connect_time = 0.0
self._last_on_connect_dispatch = 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: def subscribe(self, receiver: Callable) -> None:
"""Wrapper for allowing dispatcher to subscribe.""" """Wrapper for allowing dispatcher to subscribe."""
self._dispatcher = receiver self._dispatcher = receiver
@@ -86,6 +92,7 @@ class MqttClient(Communicator):
return return
self._publish_queue.put(QueuedPublish(full_topic, payload, retain)) self._publish_queue.put(QueuedPublish(full_topic, payload, retain))
self._wake_worker()
def stop(self) -> None: def stop(self) -> None:
if self._worker is None: if self._worker is None:
@@ -101,9 +108,11 @@ class MqttClient(Communicator):
publish_done, publish_done,
) )
) )
self._wake_worker()
publish_done.wait(MQTT_SHUTDOWN_FLUSH_TIMEOUT) publish_done.wait(MQTT_SHUTDOWN_FLUSH_TIMEOUT)
self._stop_event.set() self._stop_event.set()
self._wake_worker()
if self.client is not None: if self.client is not None:
try: try:
@@ -369,11 +378,17 @@ class MqttClient(Communicator):
exc_info=True, 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: def _loop_client(self, timeout: float) -> int:
"""Drive Paho without select()'s limit on socket file descriptors.""" """Drive Paho without select()'s limit on socket file descriptors."""
assert self.client is not None
client = self.client client = self.client
if client is None:
return mqtt.MQTT_ERR_NO_CONN
sock = client.socket() sock = client.socket()
if sock is None: if sock is None:
return mqtt.MQTT_ERR_NO_CONN return mqtt.MQTT_ERR_NO_CONN
@@ -385,11 +400,19 @@ class MqttClient(Communicator):
pending = hasattr(sock, "pending") and sock.pending() > 0 pending = hasattr(sock, "pending") and sock.pending() > 0
with selectors.DefaultSelector() as selector: with selectors.DefaultSelector() as selector:
selector.register(sock, events) 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 if self._wake_recv in ready:
for _, mask in ready: try:
ready_events |= mask self._wake_recv.recv(4096)
except BlockingIOError:
pass
ready_events = ready.get(sock, 0)
if pending or ready_events & selectors.EVENT_READ: if pending or ready_events & selectors.EVENT_READ:
result = client.loop_read() result = client.loop_read()
if result != mqtt.MQTT_ERR_SUCCESS or client.socket() is None: if result != mqtt.MQTT_ERR_SUCCESS or client.socket() is None:
+16 -3
View File
@@ -126,12 +126,18 @@ class TestMqttClientLifecycle(unittest.TestCase):
os.makedirs(MODEL_CACHE_DIR) os.makedirs(MODEL_CACHE_DIR)
self.config = build_config() self.config = build_config()
self.client = MqttClient(self.config) self.client = self._build_client()
self.receiver = RuntimeSnapshotReceiver() self.receiver = RuntimeSnapshotReceiver()
self.client.attach_dispatcher(build_dispatcher(self.config, [])) 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) 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: with patch.object(client, "_start_worker") as mock_start_worker:
client.subscribe(self.receiver._receive) client.subscribe(self.receiver._receive)
@@ -142,7 +148,7 @@ class TestMqttClientLifecycle(unittest.TestCase):
mock_start_worker.assert_not_called() mock_start_worker.assert_not_called()
def test_attach_dispatcher_supplies_command_surface(self) -> None: 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")) 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.assertEqual(self.client._subscription_mid, 42)
self.client.client.subscribe.assert_called_once_with("frigate/#", qos=0) 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( def test_handle_connect_event_reconnects_on_recoverable_subscribe_error(
self, self,
) -> None: ) -> None:
+19 -3
View File
@@ -2,6 +2,7 @@ import fcntl
import resource import resource
import selectors import selectors
import socket import socket
import time
import unittest import unittest
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
@@ -23,6 +24,11 @@ class TestMqttNetworkLoop(unittest.TestCase):
self.sock, self.peer = socket.socketpair() self.sock, self.peer = socket.socketpair()
self.addCleanup(self.sock.close) self.addCleanup(self.sock.close)
self.addCleanup(self.peer.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 self.client.socket.return_value = self.sock
def test_high_fd_handles_connack_suback_publish_and_puback(self) -> None: 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.assertEqual(self.transport._loop_client(0), mqtt.MQTT_ERR_CONN_LOST)
self.client.loop_misc.assert_not_called() 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.client.socket.return_value = None
self.assertEqual(self.transport._loop_client(0), mqtt.MQTT_ERR_NO_CONN) 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)