mirror of
https://github.com/blakeblackshear/frigate.git
synced 2026-10-07 07:12:50 +03:00
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
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:
+35
-3
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
self.client, "_loop_client", return_value=mqtt.MQTT_ERR_SUCCESS
|
||||||
|
) as mock_loop:
|
||||||
self.client._publish_direct(
|
self.client._publish_direct(
|
||||||
QueuedPublish("frigate/available", "stopped", True, barrier)
|
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,8 +645,7 @@ 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)
|
||||||
@@ -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 (
|
||||||
|
patch.object(
|
||||||
|
self.client, "_loop_client", side_effect=OSError("socket closed")
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
self.client,
|
self.client,
|
||||||
"_schedule_reconnect",
|
"_schedule_reconnect",
|
||||||
side_effect=stop_after_reconnect,
|
side_effect=stop_after_reconnect,
|
||||||
) as mock_schedule_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()
|
||||||
|
|||||||
@@ -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