diff options
| author | Arfey <Arfey17.mg@gmail.com> | 2025-11-10 01:10:32 +0200 |
|---|---|---|
| committer | Jacob Walls <jacobtylerwalls@gmail.com> | 2025-12-29 09:48:11 -0500 |
| commit | cc0f6c4f74cc278fdab79b269401127f2d869334 (patch) | |
| tree | 55c4ebf38ffbc8247d6bf1688bf9bc2e9a747852 /tests | |
| parent | 1c34b8716afde049f95ad1c72c2f8e148f826662 (diff) | |
Fixed #36714 -- Fixed context sharing among async signal handlers.
Diffstat (limited to 'tests')
| -rw-r--r-- | tests/signals/tests.py | 236 |
1 files changed, 236 insertions, 0 deletions
diff --git a/tests/signals/tests.py b/tests/signals/tests.py index 7cb64f6e05..907612459a 100644 --- a/tests/signals/tests.py +++ b/tests/signals/tests.py @@ -1,3 +1,4 @@ +import contextvars from unittest import mock from asgiref.sync import markcoroutinefunction @@ -645,3 +646,238 @@ class AsyncReceiversTests(SimpleTestCase): result = await signal.asend_robust(self.__class__) self.assertEqual(result, [(async_handler, 1)]) + + +class TestReceiversContextVarsSharing(SimpleTestCase): + def setUp(self): + self.ctx_var = contextvars.ContextVar("test_var", default=0) + + class CtxSyncHandler: + def __init__(self, ctx_var): + self.ctx_var = ctx_var + self.values = [] + + def __call__(self, **kwargs): + val = self.ctx_var.get() + self.ctx_var.set(val + 1) + self.values.append(self.ctx_var.get()) + return self.ctx_var.get() + + class CtxAsyncHandler: + def __init__(self, ctx_var): + self.ctx_var = ctx_var + self.values = [] + markcoroutinefunction(self) + + async def __call__(self, **kwargs): + val = self.ctx_var.get() + self.ctx_var.set(val + 1) + self.values.append(self.ctx_var.get()) + return self.ctx_var.get() + + self.CtxSyncHandler = CtxSyncHandler + self.CtxAsyncHandler = CtxAsyncHandler + + async def test_asend_correct_contextvars_sharing_async_receivers(self): + handler1 = self.CtxAsyncHandler(self.ctx_var) + handler2 = self.CtxAsyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + await signal.asend(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + async def test_asend_correct_contextvars_sharing_sync_receivers(self): + handler1 = self.CtxSyncHandler(self.ctx_var) + handler2 = self.CtxSyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + await signal.asend(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + async def test_asend_correct_contextvars_sharing_mix_receivers(self): + handler1 = self.CtxSyncHandler(self.ctx_var) + handler2 = self.CtxAsyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + await signal.asend(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + async def test_asend_robust_correct_contextvars_sharing_async_receivers(self): + handler1 = self.CtxAsyncHandler(self.ctx_var) + handler2 = self.CtxAsyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + await signal.asend_robust(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + async def test_asend_robust_correct_contextvars_sharing_sync_receivers(self): + handler1 = self.CtxSyncHandler(self.ctx_var) + handler2 = self.CtxSyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + await signal.asend_robust(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + async def test_asend_robust_correct_contextvars_sharing_mix_receivers(self): + handler1 = self.CtxSyncHandler(self.ctx_var) + handler2 = self.CtxAsyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + await signal.asend_robust(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + def test_send_correct_contextvars_sharing_async_receivers(self): + handler1 = self.CtxAsyncHandler(self.ctx_var) + handler2 = self.CtxAsyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + signal.send(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + def test_send_correct_contextvars_sharing_sync_receivers(self): + handler1 = self.CtxSyncHandler(self.ctx_var) + handler2 = self.CtxSyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + signal.send(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + def test_send_correct_contextvars_sharing_mix_receivers(self): + handler1 = self.CtxSyncHandler(self.ctx_var) + handler2 = self.CtxAsyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + signal.send(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + def test_send_robust_correct_contextvars_sharing_async_receivers(self): + handler1 = self.CtxAsyncHandler(self.ctx_var) + handler2 = self.CtxAsyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + signal.send_robust(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + def test_send_robust_correct_contextvars_sharing_sync_receivers(self): + handler1 = self.CtxSyncHandler(self.ctx_var) + handler2 = self.CtxSyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + signal.send_robust(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) + + def test_send_robust_correct_contextvars_sharing_mix_receivers(self): + handler1 = self.CtxSyncHandler(self.ctx_var) + handler2 = self.CtxAsyncHandler(self.ctx_var) + signal = dispatch.Signal() + signal.connect(handler1) + signal.connect(handler2) + + # set custom value outer signal + self.ctx_var.set(1) + + signal.send_robust(self.__class__) + + self.assertEqual(len(handler1.values), 1) + self.assertEqual(len(handler2.values), 1) + self.assertEqual(sorted([*handler1.values, *handler2.values]), [2, 3]) + self.assertEqual(self.ctx_var.get(), 3) |
