summaryrefslogtreecommitdiff
path: root/tests
diff options
context:
space:
mode:
authorArfey <Arfey17.mg@gmail.com>2025-11-10 01:10:32 +0200
committerJacob Walls <jacobtylerwalls@gmail.com>2025-12-29 09:48:11 -0500
commitcc0f6c4f74cc278fdab79b269401127f2d869334 (patch)
tree55c4ebf38ffbc8247d6bf1688bf9bc2e9a747852 /tests
parent1c34b8716afde049f95ad1c72c2f8e148f826662 (diff)
Fixed #36714 -- Fixed context sharing among async signal handlers.
Diffstat (limited to 'tests')
-rw-r--r--tests/signals/tests.py236
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)