Fix MQTT network loop with high socket file descriptors (#24573)
CI / AMD64 Build (push) Canceled after 0s
CI / AMD64 Smoke Test (push) Canceled after 0s
CI / ARM Build (push) Canceled after 0s
CI / Jetson Jetpack 6 (push) Canceled after 0s
CI / AMD64 Extra Build (push) Canceled after 0s
CI / ARM Extra Build (push) Canceled after 0s
CI / Synaptics Build (push) Canceled after 0s
CI / Assemble and push default build (push) Canceled after 0s

This commit is contained in:
Ronnie
2026-10-06 07:38:40 -05:00
committed by GitHub
parent 9ab52c2be6
commit 6b1e084fdc
3 changed files with 184 additions and 17 deletions
+35 -3
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
import logging import logging
import queue import queue
import selectors
import threading import threading
import time import time
from collections.abc import Callable from collections.abc import Callable
@@ -358,7 +359,7 @@ class MqttClient(Communicator):
deadline = time.monotonic() + MQTT_SHUTDOWN_FLUSH_TIMEOUT deadline = time.monotonic() + MQTT_SHUTDOWN_FLUSH_TIMEOUT
while not message_info.is_published() and time.monotonic() < deadline: while not message_info.is_published() and time.monotonic() < deadline:
if ( if (
self.client.loop(timeout=MQTT_PUBLISH_WAIT_INTERVAL) self._loop_client(timeout=MQTT_PUBLISH_WAIT_INTERVAL)
!= mqtt.MQTT_ERR_SUCCESS != mqtt.MQTT_ERR_SUCCESS
): ):
break break
@@ -368,6 +369,37 @@ class MqttClient(Communicator):
exc_info=True, 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: def _mqtt_loop_worker(self) -> None:
# The worker owns all socket I/O so reconnect, subscribe, and publish # The worker owns all socket I/O so reconnect, subscribe, and publish
# ordering stays serialized in one place. # ordering stays serialized in one place.
@@ -384,7 +416,7 @@ class MqttClient(Communicator):
assert self.client is not None assert self.client is not None
try: try:
result = self.client.loop(timeout=MQTT_LOOP_TIMEOUT) result = self._loop_client(timeout=MQTT_LOOP_TIMEOUT)
except (OSError, mqtt.WebsocketConnectionError) as err: except (OSError, mqtt.WebsocketConnectionError) as err:
logger.warning("MQTT loop error: %s", err) logger.warning("MQTT loop error: %s", err)
self._schedule_reconnect() self._schedule_reconnect()
@@ -617,7 +649,7 @@ class MqttClient(Communicator):
return return
try: 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: except (OSError, mqtt.WebsocketConnectionError) as err:
logger.warning("MQTT publish wait failed: %s", err) logger.warning("MQTT publish wait failed: %s", err)
self._schedule_reconnect() self._schedule_reconnect()
+19 -14
View File
@@ -448,7 +448,6 @@ class TestMqttClientLifecycle(unittest.TestCase):
def test_publish_direct_waits_for_flush_barrier(self) -> None: def test_publish_direct_waits_for_flush_barrier(self) -> None:
mock_client = MagicMock() mock_client = MagicMock()
mock_client.loop.return_value = mqtt.MQTT_ERR_SUCCESS
self.client.client = mock_client self.client.client = mock_client
message_info = MagicMock(rc=mqtt.MQTT_ERR_SUCCESS, mid=1) message_info = MagicMock(rc=mqtt.MQTT_ERR_SUCCESS, mid=1)
# inflight tracking checks once, then _wait_for_publish polls # inflight tracking checks once, then _wait_for_publish polls
@@ -456,11 +455,14 @@ class TestMqttClientLifecycle(unittest.TestCase):
mock_client.publish.return_value = message_info mock_client.publish.return_value = message_info
barrier = MagicMock() barrier = MagicMock()
self.client._publish_direct( with patch.object(
QueuedPublish("frigate/available", "stopped", True, barrier) 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() barrier.set.assert_called_once()
def test_shutdown_barrier_releases_when_publish_raises(self) -> None: def test_shutdown_barrier_releases_when_publish_raises(self) -> None:
@@ -643,9 +645,8 @@ class TestMqttClientLifecycle(unittest.TestCase):
loop_calls[0] += 1 loop_calls[0] += 1
return mqtt.MQTT_ERR_SUCCESS return mqtt.MQTT_ERR_SUCCESS
mock_client.loop.side_effect = loop_side_effect with patch.object(self.client, "_loop_client", side_effect=loop_side_effect):
self.client._wait_for_publish(message_info)
self.client._wait_for_publish(message_info)
self.assertEqual(loop_calls[0], 1) self.assertEqual(loop_calls[0], 1)
self.assertIsNone(self.client.client) 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: def test_mqtt_loop_worker_reconnects_on_recoverable_loop_error(self) -> None:
self.client.client = MagicMock() self.client.client = MagicMock()
self.client.client.loop.side_effect = OSError("socket closed")
def stop_after_reconnect() -> None: def stop_after_reconnect() -> None:
self.client._stop_event.set() self.client._stop_event.set()
with patch.object( with (
self.client, patch.object(
"_schedule_reconnect", self.client, "_loop_client", side_effect=OSError("socket closed")
side_effect=stop_after_reconnect, ),
) as mock_schedule_reconnect: patch.object(
self.client,
"_schedule_reconnect",
side_effect=stop_after_reconnect,
) as mock_schedule_reconnect,
):
self.client._mqtt_loop_worker() self.client._mqtt_loop_worker()
mock_schedule_reconnect.assert_called_once() mock_schedule_reconnect.assert_called_once()
+130
View File
@@ -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)