From 6b1e084fdce05a6dcf6a97f9c6ca7ea77b8d7a63 Mon Sep 17 00:00:00 2001 From: Ronnie Date: Tue, 6 Oct 2026 05:38:40 -0700 Subject: [PATCH] Fix MQTT network loop with high socket file descriptors (#24573) --- frigate/comms/mqtt.py | 38 +++++++- frigate/test/test_mqtt_lifecycle.py | 33 ++++--- frigate/test/test_mqtt_network_loop.py | 130 +++++++++++++++++++++++++ 3 files changed, 184 insertions(+), 17 deletions(-) create mode 100644 frigate/test/test_mqtt_network_loop.py diff --git a/frigate/comms/mqtt.py b/frigate/comms/mqtt.py index f2ce4f7c6d..28ef010893 100644 --- a/frigate/comms/mqtt.py +++ b/frigate/comms/mqtt.py @@ -2,6 +2,7 @@ from __future__ import annotations import logging import queue +import selectors import threading import time from collections.abc import Callable @@ -358,7 +359,7 @@ class MqttClient(Communicator): deadline = time.monotonic() + MQTT_SHUTDOWN_FLUSH_TIMEOUT while not message_info.is_published() and time.monotonic() < deadline: if ( - self.client.loop(timeout=MQTT_PUBLISH_WAIT_INTERVAL) + self._loop_client(timeout=MQTT_PUBLISH_WAIT_INTERVAL) != mqtt.MQTT_ERR_SUCCESS ): break @@ -368,6 +369,37 @@ class MqttClient(Communicator): exc_info=True, ) + def _loop_client(self, timeout: float) -> int: + """Drive Paho without select()'s limit on socket file descriptors.""" + 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 + + events = selectors.EVENT_READ + if client.want_write(): + events |= selectors.EVENT_WRITE + # TLS can have decrypted bytes buffered even when the socket is not ready. + 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) + + ready_events = 0 + for _, mask in ready: + ready_events |= mask + if pending or ready_events & selectors.EVENT_READ: + result = client.loop_read() + if result != mqtt.MQTT_ERR_SUCCESS or client.socket() is None: + return result + if ready_events & selectors.EVENT_WRITE: + result = client.loop_write() + if result != mqtt.MQTT_ERR_SUCCESS or client.socket() is None: + return result + return client.loop_misc() + def _mqtt_loop_worker(self) -> None: # The worker owns all socket I/O so reconnect, subscribe, and publish # ordering stays serialized in one place. @@ -384,7 +416,7 @@ class MqttClient(Communicator): assert self.client is not None try: - result = self.client.loop(timeout=MQTT_LOOP_TIMEOUT) + result = self._loop_client(timeout=MQTT_LOOP_TIMEOUT) except (OSError, mqtt.WebsocketConnectionError) as err: logger.warning("MQTT loop error: %s", err) self._schedule_reconnect() @@ -617,7 +649,7 @@ class MqttClient(Communicator): return try: - result = self.client.loop(timeout=MQTT_PUBLISH_WAIT_INTERVAL) + result = self._loop_client(timeout=MQTT_PUBLISH_WAIT_INTERVAL) except (OSError, mqtt.WebsocketConnectionError) as err: logger.warning("MQTT publish wait failed: %s", err) self._schedule_reconnect() diff --git a/frigate/test/test_mqtt_lifecycle.py b/frigate/test/test_mqtt_lifecycle.py index 7352cd1578..d9c9831096 100644 --- a/frigate/test/test_mqtt_lifecycle.py +++ b/frigate/test/test_mqtt_lifecycle.py @@ -448,7 +448,6 @@ class TestMqttClientLifecycle(unittest.TestCase): def test_publish_direct_waits_for_flush_barrier(self) -> None: mock_client = MagicMock() - mock_client.loop.return_value = mqtt.MQTT_ERR_SUCCESS self.client.client = mock_client message_info = MagicMock(rc=mqtt.MQTT_ERR_SUCCESS, mid=1) # inflight tracking checks once, then _wait_for_publish polls @@ -456,11 +455,14 @@ class TestMqttClientLifecycle(unittest.TestCase): mock_client.publish.return_value = message_info barrier = MagicMock() - self.client._publish_direct( - QueuedPublish("frigate/available", "stopped", True, barrier) - ) + with patch.object( + self.client, "_loop_client", return_value=mqtt.MQTT_ERR_SUCCESS + ) as mock_loop: + self.client._publish_direct( + QueuedPublish("frigate/available", "stopped", True, barrier) + ) - mock_client.loop.assert_called_once() + mock_loop.assert_called_once() barrier.set.assert_called_once() def test_shutdown_barrier_releases_when_publish_raises(self) -> None: @@ -643,9 +645,8 @@ class TestMqttClientLifecycle(unittest.TestCase): loop_calls[0] += 1 return mqtt.MQTT_ERR_SUCCESS - mock_client.loop.side_effect = loop_side_effect - - self.client._wait_for_publish(message_info) + with patch.object(self.client, "_loop_client", side_effect=loop_side_effect): + self.client._wait_for_publish(message_info) self.assertEqual(loop_calls[0], 1) self.assertIsNone(self.client.client) @@ -669,16 +670,20 @@ class TestMqttClientLifecycle(unittest.TestCase): def test_mqtt_loop_worker_reconnects_on_recoverable_loop_error(self) -> None: self.client.client = MagicMock() - self.client.client.loop.side_effect = OSError("socket closed") def stop_after_reconnect() -> None: self.client._stop_event.set() - with patch.object( - self.client, - "_schedule_reconnect", - side_effect=stop_after_reconnect, - ) as mock_schedule_reconnect: + with ( + patch.object( + self.client, "_loop_client", side_effect=OSError("socket closed") + ), + patch.object( + self.client, + "_schedule_reconnect", + side_effect=stop_after_reconnect, + ) as mock_schedule_reconnect, + ): self.client._mqtt_loop_worker() mock_schedule_reconnect.assert_called_once() diff --git a/frigate/test/test_mqtt_network_loop.py b/frigate/test/test_mqtt_network_loop.py new file mode 100644 index 0000000000..ea28e9cd49 --- /dev/null +++ b/frigate/test/test_mqtt_network_loop.py @@ -0,0 +1,130 @@ +import fcntl +import resource +import selectors +import socket +import unittest +from unittest.mock import MagicMock, patch + +import paho.mqtt.client as mqtt +from paho.mqtt.enums import CallbackAPIVersion + +from frigate.comms.mqtt import MqttClient + + +class TestMqttNetworkLoop(unittest.TestCase): + def setUp(self) -> None: + self.transport = object.__new__(MqttClient) + self.client = MagicMock() + self.client.want_write.return_value = False + self.client.loop_read.return_value = mqtt.MQTT_ERR_SUCCESS + self.client.loop_write.return_value = mqtt.MQTT_ERR_SUCCESS + self.client.loop_misc.return_value = mqtt.MQTT_ERR_SUCCESS + self.transport.client = self.client + self.sock, self.peer = socket.socketpair() + self.addCleanup(self.sock.close) + self.addCleanup(self.peer.close) + self.client.socket.return_value = self.sock + + def test_high_fd_handles_connack_suback_publish_and_puback(self) -> None: + """Process real MQTT packets on a socket beyond select()'s FD limit.""" + if selectors.DefaultSelector is selectors.SelectSelector: + self.skipTest("This platform has no selector supporting high socket FDs") + original_limit = resource.getrlimit(resource.RLIMIT_NOFILE) + soft, hard = original_limit + if soft <= 1024: + if hard != resource.RLIM_INFINITY and hard <= 1024: + self.skipTest("The hard file descriptor limit is too low") + new_soft = 2048 if hard == resource.RLIM_INFINITY else min(2048, hard) + resource.setrlimit(resource.RLIMIT_NOFILE, (new_soft, hard)) + self.addCleanup(resource.setrlimit, resource.RLIMIT_NOFILE, original_limit) + + fd = fcntl.fcntl(self.sock.fileno(), fcntl.F_DUPFD, 1024) + with socket.socket(fileno=fd) as high_sock: + high_sock.setblocking(False) + self.peer.settimeout(1) + client = mqtt.Client(CallbackAPIVersion.VERSION2, client_id="high-fd-test") + client._sock = high_sock + self.transport.client = client + connected, subscribed, received = [], [], [] + client.on_connect = lambda *args: connected.append(args[3]) + client.on_subscribe = lambda *args: subscribed.append(args[2]) + client.on_message = lambda client, userdata, message: received.append( + message.payload + ) + + self.assertGreaterEqual(high_sock.fileno(), 1024) + self.peer.sendall(b"\x20\x02\x00\x00") + self.assertEqual(self.transport._loop_client(0.1), mqtt.MQTT_ERR_SUCCESS) + self.assertEqual(len(connected), 1) + self.assertTrue(client.is_connected()) + + result, mid = client.subscribe("diagnostic", qos=1) + self.assertEqual(result, mqtt.MQTT_ERR_SUCCESS) + self.peer.recv(1024) + self.peer.sendall(b"\x90\x03" + mid.to_bytes(2, "big") + b"\x01") + self.assertEqual(self.transport._loop_client(0.1), mqtt.MQTT_ERR_SUCCESS) + self.assertEqual(subscribed, [mid]) + + info = client.publish("diagnostic", b"outgoing", qos=1) + self.peer.recv(1024) + self.peer.sendall(b"\x40\x02" + info.mid.to_bytes(2, "big")) + self.assertEqual(self.transport._loop_client(0.1), mqtt.MQTT_ERR_SUCCESS) + self.assertTrue(info.is_published()) + + payload = b"\x00\x0adiagnosticincoming" + self.peer.sendall(b"\x30" + bytes([len(payload)]) + payload) + self.assertEqual(self.transport._loop_client(0.1), mqtt.MQTT_ERR_SUCCESS) + self.assertEqual(received, [b"incoming"]) + + def test_idle_socket_still_runs_keepalive(self) -> None: + self.assertEqual(self.transport._loop_client(0), mqtt.MQTT_ERR_SUCCESS) + self.client.loop_read.assert_not_called() + self.client.loop_write.assert_not_called() + self.client.loop_misc.assert_called_once() + + def test_writable_socket_flushes_pending_packets(self) -> None: + self.client.want_write.return_value = True + self.assertEqual(self.transport._loop_client(0), mqtt.MQTT_ERR_SUCCESS) + self.client.loop_write.assert_called_once() + self.client.loop_read.assert_not_called() + + def test_tls_buffer_is_read_without_waiting_for_socket_readiness(self) -> None: + tls_sock = MagicMock() + tls_sock.fileno.return_value = self.sock.fileno() + tls_sock.pending.return_value = 1 + self.client.socket.return_value = tls_sock + with patch("frigate.comms.mqtt.selectors.DefaultSelector") as selector: + selector.return_value.__enter__.return_value.select.return_value = [] + self.assertEqual(self.transport._loop_client(1), mqtt.MQTT_ERR_SUCCESS) + selector.return_value.__enter__.return_value.select.assert_called_once_with( + 0.0 + ) + self.client.loop_read.assert_called_once() + + def test_read_failure_does_not_write_or_run_keepalive(self) -> None: + self.peer.sendall(b"ready") + self.client.want_write.return_value = True + self.client.loop_read.return_value = mqtt.MQTT_ERR_CONN_LOST + self.assertEqual(self.transport._loop_client(0), mqtt.MQTT_ERR_CONN_LOST) + self.client.loop_write.assert_not_called() + self.client.loop_misc.assert_not_called() + + def test_socket_closed_during_read_does_not_write(self) -> None: + self.peer.sendall(b"ready") + self.client.want_write.return_value = True + self.client.socket.side_effect = [self.sock, None] + self.transport._loop_client(0) + self.client.loop_write.assert_not_called() + self.client.loop_misc.assert_not_called() + + def test_write_failure_does_not_run_keepalive(self) -> None: + self.client.want_write.return_value = True + self.client.loop_write.return_value = mqtt.MQTT_ERR_CONN_LOST + 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: + 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)