@@ -83,23 +83,31 @@ async def test_task() -> None: ...
8383
8484
8585@pytest .mark .anyio
86- @pytest .mark .parametrize ('is_worker_process' , [True , False ])
86+ @pytest .mark .parametrize (
87+ ('is_worker_process' , 'startup' , 'shutdown' ),
88+ [
89+ (True , TaskiqEvents .WORKER_STARTUP , TaskiqEvents .WORKER_SHUTDOWN ),
90+ (False , TaskiqEvents .CLIENT_STARTUP , TaskiqEvents .CLIENT_SHUTDOWN ),
91+ ],
92+ )
8793async def test_async_context_manager_enter (
8894 * ,
8995 is_worker_process : bool ,
96+ startup : TaskiqEvents ,
97+ shutdown : TaskiqEvents ,
9098) -> None :
9199 """Test that `__aenter__` and `__aexit__` calls work."""
92100 broker = _TestBroker ()
93101 broker .is_worker_process = is_worker_process
94102 startup_called = False
95103 shutdown_called = False
96104
97- @broker .on_event (TaskiqEvents . CLIENT_STARTUP )
105+ @broker .on_event (startup )
98106 async def track_startup (state : TaskiqState ) -> None :
99107 nonlocal startup_called
100108 startup_called = True
101109
102- @broker .on_event (TaskiqEvents . CLIENT_SHUTDOWN )
110+ @broker .on_event (shutdown )
103111 async def track_shutdown (state : TaskiqState ) -> None :
104112 nonlocal shutdown_called
105113 shutdown_called = True
@@ -113,23 +121,31 @@ async def track_shutdown(state: TaskiqState) -> None:
113121
114122
115123@pytest .mark .anyio
116- @pytest .mark .parametrize ('is_worker_process' , [True , False ])
124+ @pytest .mark .parametrize (
125+ ('is_worker_process' , 'startup' , 'shutdown' ),
126+ [
127+ (True , TaskiqEvents .WORKER_STARTUP , TaskiqEvents .WORKER_SHUTDOWN ),
128+ (False , TaskiqEvents .CLIENT_STARTUP , TaskiqEvents .CLIENT_SHUTDOWN ),
129+ ],
130+ )
117131async def test_async_context_manager_exit_on_exception (
118132 * ,
119133 is_worker_process : bool ,
134+ startup : TaskiqEvents ,
135+ shutdown : TaskiqEvents ,
120136) -> None :
121137 """Test that __aexit__ calls shutdown even if exception is raised."""
122138 broker = _TestBroker ()
123139 broker .is_worker_process = is_worker_process
124140 startup_called = False
125141 shutdown_called = False
126142
127- @broker .on_event (TaskiqEvents . CLIENT_STARTUP )
143+ @broker .on_event (startup )
128144 async def track_startup (state : TaskiqState ) -> None :
129145 nonlocal startup_called
130146 startup_called = True
131147
132- @broker .on_event (TaskiqEvents . CLIENT_SHUTDOWN )
148+ @broker .on_event (shutdown )
133149 async def track_shutdown (state : TaskiqState ) -> None :
134150 nonlocal shutdown_called
135151 shutdown_called = True
0 commit comments