mirror of
https://github.com/blakeblackshear/frigate.git
synced 2026-10-10 00:32:48 +03:00
Lock ZMQ sockets shared across threads (#24607)
CI / ARM Extra Build (push) Blocked by required conditions
CI / AMD64 Build (push) Waiting to run
CI / ARM Build (push) Waiting to run
CI / Jetson Jetpack 6 (push) Waiting to run
CI / AMD64 Extra Build (push) Blocked by required conditions
CI / Synaptics Build (push) Blocked by required conditions
CI / Assemble and push default build (push) Blocked by required conditions
CI / ARM Extra Build (push) Blocked by required conditions
CI / AMD64 Build (push) Waiting to run
CI / ARM Build (push) Waiting to run
CI / Jetson Jetpack 6 (push) Waiting to run
CI / AMD64 Extra Build (push) Blocked by required conditions
CI / Synaptics Build (push) Blocked by required conditions
CI / Assemble and push default build (push) Blocked by required conditions
* lock zmq sockets shared across threads Several zmq sockets were used from more than one thread at a time, which zmq does not allow. In the embeddings process, GenAI review and object description threads shared one requestor, so overlapping calls silently dropped the second write and could eventually leave a thread waiting forever on a reply that had already arrived, after which nothing was saved until restart. The API's embeddings requestor had the same problem, concurrent config publishes could mix the topic of one update with the payload of another, and the webpush config subscribers were drained from every publishing thread, which could abort the process. Each socket now has a lock. The embeddings requestor lock does not wait, so an overlapping call still returns empty right away instead of stalling the API behind a slow reply. * add test
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
"""Facilitates communication between processes."""
|
||||
|
||||
import multiprocessing as mp
|
||||
import threading
|
||||
from _pickle import UnpicklingError
|
||||
from multiprocessing.synchronize import Event as MpEvent
|
||||
from typing import Any
|
||||
@@ -18,11 +19,13 @@ class ConfigPublisher:
|
||||
self.socket = self.context.socket(zmq.PUB)
|
||||
self.socket.bind(SOCKET_PUB_SUB)
|
||||
self.stop_event: MpEvent = mp.Event()
|
||||
self.lock = threading.Lock()
|
||||
|
||||
def publish(self, topic: str, payload: Any) -> None:
|
||||
"""There is no communication back to the processes."""
|
||||
self.socket.send_string(topic, flags=zmq.SNDMORE)
|
||||
self.socket.send_pyobj(payload)
|
||||
with self.lock:
|
||||
self.socket.send_string(topic, flags=zmq.SNDMORE)
|
||||
self.socket.send_pyobj(payload)
|
||||
|
||||
def stop(self) -> None:
|
||||
self.stop_event.set()
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""Facilitates communication between processes."""
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
@@ -78,14 +79,21 @@ class EmbeddingsRequestor:
|
||||
self.context = zmq.Context()
|
||||
self.socket = self.context.socket(zmq.REQ)
|
||||
self.socket.connect(SOCKET_REP_REQ)
|
||||
self.lock = threading.Lock()
|
||||
|
||||
def send_data(self, topic: str, data: Any) -> Any:
|
||||
"""Sends data and then waits for reply."""
|
||||
# an overlapping call fails fast so a slow reply can't stall the API
|
||||
if not self.lock.acquire(blocking=False):
|
||||
return ""
|
||||
|
||||
try:
|
||||
self.socket.send_json((topic, data))
|
||||
return self.socket.recv_json()
|
||||
except zmq.ZMQError:
|
||||
return ""
|
||||
finally:
|
||||
self.lock.release()
|
||||
|
||||
def stop(self) -> None:
|
||||
self.socket.close()
|
||||
|
||||
@@ -73,14 +73,16 @@ class InterProcessRequestor:
|
||||
self.context = zmq.Context()
|
||||
self.socket = self.context.socket(zmq.REQ)
|
||||
self.socket.connect(SOCKET_REP_REQ)
|
||||
self.lock = threading.Lock()
|
||||
|
||||
def send_data(self, topic: str, data: Any) -> Any:
|
||||
"""Sends data and then waits for reply."""
|
||||
try:
|
||||
self.socket.send_json((topic, data))
|
||||
return self.socket.recv_json()
|
||||
except zmq.ZMQError:
|
||||
return ""
|
||||
with self.lock:
|
||||
try:
|
||||
self.socket.send_json((topic, data))
|
||||
return self.socket.recv_json()
|
||||
except zmq.ZMQError:
|
||||
return ""
|
||||
|
||||
def stop(self) -> None:
|
||||
self.socket.close(linger=0)
|
||||
|
||||
+23
-21
@@ -70,6 +70,7 @@ class WebPushClient(Communicator):
|
||||
for c in self.config.cameras.values()
|
||||
}
|
||||
self.suspension_broadcaster: Callable[[str, Any, bool], None] | None = None
|
||||
self.config_lock = threading.Lock()
|
||||
self.last_camera_notification_time: dict[str, float] = {
|
||||
c.name: 0 # type: ignore[misc]
|
||||
for c in self.config.cameras.values()
|
||||
@@ -205,29 +206,30 @@ class WebPushClient(Communicator):
|
||||
|
||||
def publish(self, topic: str, payload: Any, retain: bool = False) -> None:
|
||||
"""Wrapper for publishing when client is in valid state."""
|
||||
# check for updated global config (notifications, auth)
|
||||
while True:
|
||||
config_topic, config_payload = (
|
||||
self.global_config_subscriber.check_for_update()
|
||||
)
|
||||
if config_topic is None:
|
||||
break
|
||||
if config_topic == "config/notifications" and config_payload:
|
||||
self.config.notifications = config_payload
|
||||
elif config_topic == "config/auth":
|
||||
if isinstance(config_payload, AuthConfig):
|
||||
self.config.auth = config_payload
|
||||
with self.config_lock:
|
||||
# check for updated global config (notifications, auth)
|
||||
while True:
|
||||
config_topic, config_payload = (
|
||||
self.global_config_subscriber.check_for_update()
|
||||
)
|
||||
if config_topic is None:
|
||||
break
|
||||
if config_topic == "config/notifications" and config_payload:
|
||||
self.config.notifications = config_payload
|
||||
elif config_topic == "config/auth":
|
||||
if isinstance(config_payload, AuthConfig):
|
||||
self.config.auth = config_payload
|
||||
self._refresh_user_cameras()
|
||||
|
||||
updates = self.config_subscriber.check_for_updates()
|
||||
|
||||
if "add" in updates:
|
||||
for camera in updates["add"]:
|
||||
self.suspended_cameras[camera] = 0
|
||||
self.last_camera_notification_time[camera] = 0
|
||||
|
||||
self._refresh_user_cameras()
|
||||
|
||||
updates = self.config_subscriber.check_for_updates()
|
||||
|
||||
if "add" in updates:
|
||||
for camera in updates["add"]:
|
||||
self.suspended_cameras[camera] = 0
|
||||
self.last_camera_notification_time[camera] = 0
|
||||
|
||||
self._refresh_user_cameras()
|
||||
|
||||
if topic == "reviews":
|
||||
decoded = json.loads(payload)
|
||||
camera = decoded["before"]["camera"]
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
"""Tests for sharing one zmq socket wrapper across threads."""
|
||||
|
||||
import os
|
||||
import random
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
import zmq
|
||||
|
||||
from frigate.comms import config_updater, embeddings_updater, inter_process
|
||||
from frigate.comms.config_updater import ConfigPublisher, ConfigSubscriber
|
||||
from frigate.comms.embeddings_updater import EmbeddingsRequestor
|
||||
from frigate.comms.inter_process import InterProcessRequestor
|
||||
from frigate.comms.webpush import WebPushClient
|
||||
|
||||
THREADS = 8
|
||||
CALLS = 100
|
||||
|
||||
# nothing reads until the end, so the total stays under the default zmq HWM
|
||||
PUBLISH_CALLS = 100
|
||||
|
||||
|
||||
class TestSharedRequestor(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.tmp = tempfile.TemporaryDirectory()
|
||||
self.address = f"ipc://{os.path.join(self.tmp.name, 'comms')}"
|
||||
self.stop = threading.Event()
|
||||
self.context = zmq.Context()
|
||||
self.responder = self.context.socket(zmq.REP)
|
||||
self.responder.bind(self.address)
|
||||
self.responder_thread = threading.Thread(target=self._respond)
|
||||
self.responder_thread.start()
|
||||
|
||||
def tearDown(self) -> None:
|
||||
self.stop.set()
|
||||
self.responder_thread.join()
|
||||
self.responder.close(linger=0)
|
||||
self.context.destroy(linger=0)
|
||||
self.tmp.cleanup()
|
||||
|
||||
def _respond(self) -> None:
|
||||
while not self.stop.is_set():
|
||||
ready, _, _ = zmq.select([self.responder], [], [], 0.1)
|
||||
|
||||
if ready:
|
||||
self.responder.recv_json()
|
||||
time.sleep(random.uniform(0, 0.002))
|
||||
self.responder.send_json(["ok"])
|
||||
|
||||
def _call_from_threads(self, requestor) -> list:
|
||||
results: list = []
|
||||
|
||||
def call() -> None:
|
||||
for _ in range(CALLS):
|
||||
# jitter lands some calls on the instant another reply arrives
|
||||
time.sleep(random.uniform(0, 0.002))
|
||||
results.append(requestor.send_data("topic", {"key": "value"}))
|
||||
|
||||
threads = [threading.Thread(target=call, daemon=True) for _ in range(THREADS)]
|
||||
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
|
||||
for thread in threads:
|
||||
thread.join(30)
|
||||
|
||||
self.assertFalse(any(thread.is_alive() for thread in threads))
|
||||
return results
|
||||
|
||||
def test_inter_process_requestor_waits_for_overlapping_calls(self) -> None:
|
||||
with patch.object(inter_process, "SOCKET_REP_REQ", self.address):
|
||||
requestor = InterProcessRequestor()
|
||||
|
||||
results = self._call_from_threads(requestor)
|
||||
self.assertEqual(results, [["ok"]] * THREADS * CALLS)
|
||||
requestor.stop()
|
||||
|
||||
def test_embeddings_requestor_survives_overlapping_calls(self) -> None:
|
||||
with patch.object(embeddings_updater, "SOCKET_REP_REQ", self.address):
|
||||
requestor = EmbeddingsRequestor()
|
||||
|
||||
# an overlapping call may fail fast, but the socket has to stay usable
|
||||
results = self._call_from_threads(requestor)
|
||||
self.assertEqual(len(results), THREADS * CALLS)
|
||||
self.assertEqual(requestor.send_data("topic", {}), ["ok"])
|
||||
requestor.stop()
|
||||
|
||||
|
||||
class TestSharedConfigPublisher(unittest.TestCase):
|
||||
def test_frames_stay_paired_across_threads(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
address = f"ipc://{os.path.join(tmp, 'config')}"
|
||||
|
||||
with patch.object(config_updater, "SOCKET_PUB_SUB", address):
|
||||
publisher = ConfigPublisher()
|
||||
subscriber = ConfigSubscriber("config/")
|
||||
|
||||
# a subscription that is still connecting drops messages
|
||||
subscriber.socket.setsockopt(zmq.RCVTIMEO, 5000)
|
||||
|
||||
while True:
|
||||
publisher.publish("config/ready", None)
|
||||
|
||||
if zmq.select([subscriber.socket], [], [], 0.05)[0]:
|
||||
break
|
||||
|
||||
def publish(index: int) -> None:
|
||||
topic = f"config/cameras/camera_{index}"
|
||||
|
||||
for _ in range(PUBLISH_CALLS):
|
||||
publisher.publish(topic, topic)
|
||||
|
||||
threads = [
|
||||
threading.Thread(target=publish, args=(i,)) for i in range(THREADS)
|
||||
]
|
||||
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
|
||||
publisher.publish("config/done", None)
|
||||
received = 0
|
||||
|
||||
while True:
|
||||
topic = subscriber.socket.recv_string()
|
||||
payload = subscriber.socket.recv_pyobj()
|
||||
|
||||
if topic == "config/done":
|
||||
break
|
||||
|
||||
if topic != "config/ready":
|
||||
self.assertEqual(topic, payload)
|
||||
received += 1
|
||||
|
||||
self.assertEqual(received, THREADS * PUBLISH_CALLS)
|
||||
subscriber.stop()
|
||||
publisher.stop()
|
||||
|
||||
|
||||
class FakeConfigSubscriber:
|
||||
"""Records whether two threads read at the same time."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.reading = threading.Lock()
|
||||
self.overlapped = False
|
||||
|
||||
def _read(self) -> None:
|
||||
if not self.reading.acquire(blocking=False):
|
||||
self.overlapped = True
|
||||
return
|
||||
|
||||
time.sleep(0.001)
|
||||
self.reading.release()
|
||||
|
||||
def check_for_update(self) -> tuple[None, None]:
|
||||
self._read()
|
||||
return (None, None)
|
||||
|
||||
def check_for_updates(self) -> dict:
|
||||
self._read()
|
||||
return {}
|
||||
|
||||
|
||||
class TestSharedWebPushClient(unittest.TestCase):
|
||||
def test_config_subscribers_read_by_one_thread_at_a_time(self) -> None:
|
||||
client = WebPushClient.__new__(WebPushClient)
|
||||
client.config_lock = threading.Lock()
|
||||
client.global_config_subscriber = FakeConfigSubscriber()
|
||||
client.config_subscriber = FakeConfigSubscriber()
|
||||
|
||||
def publish() -> None:
|
||||
for _ in range(20):
|
||||
client.publish("topic", "payload")
|
||||
|
||||
threads = [threading.Thread(target=publish) for _ in range(THREADS)]
|
||||
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
|
||||
self.assertFalse(client.global_config_subscriber.overlapped)
|
||||
self.assertFalse(client.config_subscriber.overlapped)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user