Skip to content

Commit e294784

Browse files
authored
Merge branch 'main' into agents/dps-group-registration-timeout
2 parents a527246 + a5d5590 commit e294784

3 files changed

Lines changed: 73 additions & 18 deletions

File tree

‎azure-iot-device/azure/iot/device/iothub/aio/loop_management.py‎

Lines changed: 28 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,8 @@
33
# Licensed under the MIT License. See License.txt in the project root for
44
# license information.
55
# --------------------------------------------------------------------------
6-
""" This module contains functions of managing event loops for the IoTHub client
7-
"""
6+
"""This module contains functions of managing event loops for the IoTHub client"""
7+
88
import asyncio
99
import threading
1010
import logging
@@ -16,21 +16,25 @@
1616
"CLIENT_INTERNAL_LOOP": None,
1717
"CLIENT_HANDLER_RUNNER_LOOP": None,
1818
}
19+
# Janus queues bind to the first loop they use, so concurrent callers must receive the same loop.
20+
_loop_creation_lock = threading.Lock()
1921

2022

2123
def _cleanup():
2224
"""Clear all running loops and end respective threads.
2325
ONLY FOR TESTING USAGE
2426
By using this function, you can wipe all global loops.
27+
Do not call while clients or inboxes are still in use.
2528
DO NOT USE THIS IN PRODUCTION CODE
2629
"""
27-
for loop_name, loop in loops.items():
28-
if loop is not None:
29-
logger.debug("Stopping event loop - {}".format(loop_name))
30-
loop.call_soon_threadsafe(loop.stop)
31-
# NOTE: Stopping the loop will also end the thread, because the only thing keeping
32-
# the thread alive was the loop running
33-
loops[loop_name] = None
30+
with _loop_creation_lock:
31+
for loop_name, loop in loops.items():
32+
if loop is not None:
33+
logger.debug("Stopping event loop - {}".format(loop_name))
34+
loop.call_soon_threadsafe(loop.stop)
35+
# NOTE: Stopping the loop will also end the thread, because the only thing keeping
36+
# the thread alive was the loop running
37+
loops[loop_name] = None
3438

3539

3640
def _make_new_loop(loop_name):
@@ -45,23 +49,29 @@ def _make_new_loop(loop_name):
4549
loops[loop_name] = new_loop
4650

4751

52+
def _get_or_create_loop(loop_name):
53+
loop = loops[loop_name]
54+
if loop is None:
55+
with _loop_creation_lock:
56+
# Another caller may have created the loop while this caller waited for the lock.
57+
loop = loops[loop_name]
58+
if loop is None:
59+
_make_new_loop(loop_name)
60+
loop = loops[loop_name]
61+
return loop
62+
63+
4864
def get_client_internal_loop():
4965
"""Return the loop for internal client operations"""
50-
if loops["CLIENT_INTERNAL_LOOP"] is None:
51-
_make_new_loop("CLIENT_INTERNAL_LOOP")
52-
return loops["CLIENT_INTERNAL_LOOP"]
66+
return _get_or_create_loop("CLIENT_INTERNAL_LOOP")
5367

5468

5569
def get_client_handler_runner_loop():
5670
"""Return the loop for handler runners"""
57-
if loops["CLIENT_HANDLER_RUNNER_LOOP"] is None:
58-
_make_new_loop("CLIENT_HANDLER_RUNNER_LOOP")
59-
return loops["CLIENT_HANDLER_RUNNER_LOOP"]
71+
return _get_or_create_loop("CLIENT_HANDLER_RUNNER_LOOP")
6072

6173

6274
def get_client_handler_loop():
6375
"""Return the loop for invoking user-provided handlers on the client"""
6476
# TODO: Try and store the user loop somehow
65-
if loops["CLIENT_HANDLER_LOOP"] is None:
66-
_make_new_loop("CLIENT_HANDLER_LOOP")
67-
return loops["CLIENT_HANDLER_LOOP"]
77+
return _get_or_create_loop("CLIENT_HANDLER_LOOP")

‎tests/unit/iothub/aio/test_async_inbox.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -105,6 +105,14 @@ async def test_removes_item_from_inbox_if_already_there(self, mocker, inbox):
105105
assert retrieved_item is item
106106
assert inbox.empty()
107107

108+
@pytest.mark.it("Runs Janus async operations on the shared internal loop")
109+
async def test_uses_shared_internal_loop(self, mocker, inbox):
110+
inbox.put(mocker.MagicMock())
111+
112+
await asyncio.wait_for(inbox.get(), timeout=PROMPT_TIMEOUT)
113+
114+
assert inbox._queue._loop is loop_management.get_client_internal_loop()
115+
108116
@pytest.mark.it(
109117
"Blocks on an empty inbox until an item is available to remove and return, if using blocking mode"
110118
)

‎tests/unit/iothub/aio/test_loop_management.py‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,11 @@
66

77
import pytest
88
import asyncio
9+
import concurrent.futures
910
import logging
11+
import threading
1012
from azure.iot.device.iothub.aio import loop_management
13+
from tests.unit.helpers import BATCH_COMPLETION_TIMEOUT
1114

1215
logging.basicConfig(level=logging.DEBUG)
1316

@@ -49,6 +52,40 @@ def test_same_loop(self, fn_under_test):
4952
loop2 = fn_under_test()
5053
assert loop1 is loop2
5154

55+
@pytest.mark.it("Creates only one event loop when first called concurrently")
56+
def test_threadsafe_first_call(self, mocker, fn_under_test):
57+
class CoordinatedLoopMap(dict):
58+
def __init__(self, loops):
59+
super().__init__(loops)
60+
self._read_barrier = threading.Barrier(2)
61+
self._read_lock = threading.Lock()
62+
self._reads_to_coordinate = 2
63+
64+
def __getitem__(self, loop_name):
65+
loop = super().__getitem__(loop_name)
66+
with self._read_lock:
67+
coordinate_read = self._reads_to_coordinate > 0
68+
if coordinate_read:
69+
self._reads_to_coordinate -= 1
70+
if coordinate_read:
71+
self._read_barrier.wait(timeout=BATCH_COMPLETION_TIMEOUT)
72+
return loop
73+
74+
def make_loop(loop_name):
75+
loop_management.loops[loop_name] = mocker.MagicMock()
76+
77+
mocker.patch.object(loop_management, "loops", CoordinatedLoopMap(loop_management.loops))
78+
make_loop_mock = mocker.patch.object(
79+
loop_management, "_make_new_loop", side_effect=make_loop
80+
)
81+
82+
with concurrent.futures.ThreadPoolExecutor(max_workers=2) as executor:
83+
futures = [executor.submit(fn_under_test) for _ in range(2)]
84+
returned_loops = [future.result(timeout=BATCH_COMPLETION_TIMEOUT) for future in futures]
85+
86+
assert make_loop_mock.call_count == 1
87+
assert returned_loops[0] is returned_loops[1]
88+
5289

5390
@pytest.mark.describe(".get_client_internal_loop()")
5491
class TestGetClientInternalLoop(SharedCustomLoopTests):

0 commit comments

Comments
 (0)