mirror of
https://github.com/blakeblackshear/frigate.git
synced 2026-10-06 23:02:49 +03:00
Fix MQTT network loop with high socket file descriptors (#24573)
CI / AMD64 Extra Build (push) Blocked by required conditions
CI / ARM Extra Build (push) Blocked by required conditions
CI / Jetson Jetpack 6 (push) Waiting to run
CI / Assemble and push default build (push) Blocked by required conditions
CI / Synaptics Build (push) Blocked by required conditions
CI / AMD64 Build (push) Waiting to run
CI / AMD64 Smoke Test (push) Blocked by required conditions
CI / ARM Build (push) Waiting to run
CI / AMD64 Extra Build (push) Blocked by required conditions
CI / ARM Extra Build (push) Blocked by required conditions
CI / Jetson Jetpack 6 (push) Waiting to run
CI / Assemble and push default build (push) Blocked by required conditions
CI / Synaptics Build (push) Blocked by required conditions
CI / AMD64 Build (push) Waiting to run
CI / AMD64 Smoke Test (push) Blocked by required conditions
CI / ARM Build (push) Waiting to run
This commit is contained in:
+35
-3
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user