mirror of
https://github.com/blakeblackshear/frigate.git
synced 2026-10-10 08:42:49 +03:00
Compare commits
1
Commits
dot-username
...
master
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cc86519603 |
@@ -1,6 +1,7 @@
|
|||||||
"""Facilitates communication between processes."""
|
"""Facilitates communication between processes."""
|
||||||
|
|
||||||
import multiprocessing as mp
|
import multiprocessing as mp
|
||||||
|
import threading
|
||||||
from _pickle import UnpicklingError
|
from _pickle import UnpicklingError
|
||||||
from multiprocessing.synchronize import Event as MpEvent
|
from multiprocessing.synchronize import Event as MpEvent
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -18,11 +19,13 @@ class ConfigPublisher:
|
|||||||
self.socket = self.context.socket(zmq.PUB)
|
self.socket = self.context.socket(zmq.PUB)
|
||||||
self.socket.bind(SOCKET_PUB_SUB)
|
self.socket.bind(SOCKET_PUB_SUB)
|
||||||
self.stop_event: MpEvent = mp.Event()
|
self.stop_event: MpEvent = mp.Event()
|
||||||
|
self.lock = threading.Lock()
|
||||||
|
|
||||||
def publish(self, topic: str, payload: Any) -> None:
|
def publish(self, topic: str, payload: Any) -> None:
|
||||||
"""There is no communication back to the processes."""
|
"""There is no communication back to the processes."""
|
||||||
self.socket.send_string(topic, flags=zmq.SNDMORE)
|
with self.lock:
|
||||||
self.socket.send_pyobj(payload)
|
self.socket.send_string(topic, flags=zmq.SNDMORE)
|
||||||
|
self.socket.send_pyobj(payload)
|
||||||
|
|
||||||
def stop(self) -> None:
|
def stop(self) -> None:
|
||||||
self.stop_event.set()
|
self.stop_event.set()
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Facilitates communication between processes."""
|
"""Facilitates communication between processes."""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import threading
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -78,14 +79,21 @@ class EmbeddingsRequestor:
|
|||||||
self.context = zmq.Context()
|
self.context = zmq.Context()
|
||||||
self.socket = self.context.socket(zmq.REQ)
|
self.socket = self.context.socket(zmq.REQ)
|
||||||
self.socket.connect(SOCKET_REP_REQ)
|
self.socket.connect(SOCKET_REP_REQ)
|
||||||
|
self.lock = threading.Lock()
|
||||||
|
|
||||||
def send_data(self, topic: str, data: Any) -> Any:
|
def send_data(self, topic: str, data: Any) -> Any:
|
||||||
"""Sends data and then waits for reply."""
|
"""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:
|
try:
|
||||||
self.socket.send_json((topic, data))
|
self.socket.send_json((topic, data))
|
||||||
return self.socket.recv_json()
|
return self.socket.recv_json()
|
||||||
except zmq.ZMQError:
|
except zmq.ZMQError:
|
||||||
return ""
|
return ""
|
||||||
|
finally:
|
||||||
|
self.lock.release()
|
||||||
|
|
||||||
def stop(self) -> None:
|
def stop(self) -> None:
|
||||||
self.socket.close()
|
self.socket.close()
|
||||||
|
|||||||
@@ -73,14 +73,16 @@ class InterProcessRequestor:
|
|||||||
self.context = zmq.Context()
|
self.context = zmq.Context()
|
||||||
self.socket = self.context.socket(zmq.REQ)
|
self.socket = self.context.socket(zmq.REQ)
|
||||||
self.socket.connect(SOCKET_REP_REQ)
|
self.socket.connect(SOCKET_REP_REQ)
|
||||||
|
self.lock = threading.Lock()
|
||||||
|
|
||||||
def send_data(self, topic: str, data: Any) -> Any:
|
def send_data(self, topic: str, data: Any) -> Any:
|
||||||
"""Sends data and then waits for reply."""
|
"""Sends data and then waits for reply."""
|
||||||
try:
|
with self.lock:
|
||||||
self.socket.send_json((topic, data))
|
try:
|
||||||
return self.socket.recv_json()
|
self.socket.send_json((topic, data))
|
||||||
except zmq.ZMQError:
|
return self.socket.recv_json()
|
||||||
return ""
|
except zmq.ZMQError:
|
||||||
|
return ""
|
||||||
|
|
||||||
def stop(self) -> None:
|
def stop(self) -> None:
|
||||||
self.socket.close(linger=0)
|
self.socket.close(linger=0)
|
||||||
|
|||||||
+23
-21
@@ -70,6 +70,7 @@ class WebPushClient(Communicator):
|
|||||||
for c in self.config.cameras.values()
|
for c in self.config.cameras.values()
|
||||||
}
|
}
|
||||||
self.suspension_broadcaster: Callable[[str, Any, bool], None] | None = None
|
self.suspension_broadcaster: Callable[[str, Any, bool], None] | None = None
|
||||||
|
self.config_lock = threading.Lock()
|
||||||
self.last_camera_notification_time: dict[str, float] = {
|
self.last_camera_notification_time: dict[str, float] = {
|
||||||
c.name: 0 # type: ignore[misc]
|
c.name: 0 # type: ignore[misc]
|
||||||
for c in self.config.cameras.values()
|
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:
|
def publish(self, topic: str, payload: Any, retain: bool = False) -> None:
|
||||||
"""Wrapper for publishing when client is in valid state."""
|
"""Wrapper for publishing when client is in valid state."""
|
||||||
# check for updated global config (notifications, auth)
|
with self.config_lock:
|
||||||
while True:
|
# check for updated global config (notifications, auth)
|
||||||
config_topic, config_payload = (
|
while True:
|
||||||
self.global_config_subscriber.check_for_update()
|
config_topic, config_payload = (
|
||||||
)
|
self.global_config_subscriber.check_for_update()
|
||||||
if config_topic is None:
|
)
|
||||||
break
|
if config_topic is None:
|
||||||
if config_topic == "config/notifications" and config_payload:
|
break
|
||||||
self.config.notifications = config_payload
|
if config_topic == "config/notifications" and config_payload:
|
||||||
elif config_topic == "config/auth":
|
self.config.notifications = config_payload
|
||||||
if isinstance(config_payload, AuthConfig):
|
elif config_topic == "config/auth":
|
||||||
self.config.auth = config_payload
|
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()
|
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":
|
if topic == "reviews":
|
||||||
decoded = json.loads(payload)
|
decoded = json.loads(payload)
|
||||||
camera = decoded["before"]["camera"]
|
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