diff --git a/Lib/asyncio/tasks.py b/Lib/asyncio/tasks.py index f432cf0afa895a..cb01f71382f4a4 100644 --- a/Lib/asyncio/tasks.py +++ b/Lib/asyncio/tasks.py @@ -248,9 +248,12 @@ def __eager_start(self): prev_task = _py_swap_current_task(self._loop, self) try: _py_register_eager_task(self) + # gh-157299: the eager step must record the awaited_by edge + futures.future_add_to_awaited_by(self, prev_task) try: self._context.run(self.__step_run_and_handle_result, None) finally: + futures.future_discard_from_awaited_by(self, prev_task) _py_unregister_eager_task(self) finally: try: diff --git a/Lib/test/test_asyncio/test_eager_task_factory.py b/Lib/test/test_asyncio/test_eager_task_factory.py index 594a48baa32e52..6831e322081768 100644 --- a/Lib/test/test_asyncio/test_eager_task_factory.py +++ b/Lib/test/test_asyncio/test_eager_task_factory.py @@ -69,6 +69,22 @@ async def run(): self.assertEqual(self.run_coro(run()), 'my message') + def test_awaited_by_during_eager_step(self): + # gh-157299 + + async def run(): + parent = asyncio.current_task() + + async def coro(): + self.assertEqual(asyncio.current_task()._asyncio_awaited_by, + frozenset({parent})) + + t = self.loop.create_task(coro()) + self.assertFalse(t._asyncio_awaited_by) + await t + + self.run_coro(run()) + def test_eager_completion(self): async def coro(): @@ -290,11 +306,29 @@ def setUp(self): self._current_task = asyncio.current_task asyncio.current_task = asyncio.tasks.current_task = asyncio.tasks._py_current_task asyncio.all_tasks = asyncio.tasks.all_tasks = asyncio.tasks._py_all_tasks + + futures = asyncio.futures + + self._future_add_to_awaited_by = asyncio.future_add_to_awaited_by + futures.future_add_to_awaited_by = futures._py_future_add_to_awaited_by + asyncio.future_add_to_awaited_by = futures.future_add_to_awaited_by + + self._future_discard_from_awaited_by = asyncio.future_discard_from_awaited_by + futures.future_discard_from_awaited_by = futures._py_future_discard_from_awaited_by + asyncio.future_discard_from_awaited_by = futures.future_discard_from_awaited_by return super().setUp() def tearDown(self): asyncio.current_task = asyncio.tasks.current_task = self._current_task asyncio.all_tasks = asyncio.tasks.all_tasks = self._all_tasks + + futures = asyncio.futures + + futures.future_discard_from_awaited_by = self._future_discard_from_awaited_by + asyncio.future_discard_from_awaited_by = self._future_discard_from_awaited_by + + futures.future_add_to_awaited_by = self._future_add_to_awaited_by + asyncio.future_add_to_awaited_by = self._future_add_to_awaited_by return super().tearDown() @@ -309,11 +343,29 @@ def setUp(self): self._all_tasks = asyncio.all_tasks asyncio.current_task = asyncio.tasks.current_task = asyncio.tasks._c_current_task asyncio.all_tasks = asyncio.tasks.all_tasks = asyncio.tasks._c_all_tasks + + futures = asyncio.futures + + self._future_add_to_awaited_by = asyncio.future_add_to_awaited_by + futures.future_add_to_awaited_by = futures._c_future_add_to_awaited_by + asyncio.future_add_to_awaited_by = futures.future_add_to_awaited_by + + self._future_discard_from_awaited_by = asyncio.future_discard_from_awaited_by + futures.future_discard_from_awaited_by = futures._c_future_discard_from_awaited_by + asyncio.future_discard_from_awaited_by = futures.future_discard_from_awaited_by return super().setUp() def tearDown(self): asyncio.current_task = asyncio.tasks.current_task = self._current_task asyncio.all_tasks = asyncio.tasks.all_tasks = self._all_tasks + + futures = asyncio.futures + + futures.future_discard_from_awaited_by = self._future_discard_from_awaited_by + asyncio.future_discard_from_awaited_by = self._future_discard_from_awaited_by + + futures.future_add_to_awaited_by = self._future_add_to_awaited_by + asyncio.future_add_to_awaited_by = self._future_add_to_awaited_by return super().tearDown() def test_issue105987(self): diff --git a/Lib/test/test_asyncio/test_graph.py b/Lib/test/test_asyncio/test_graph.py index 36841672e1f0f6..544971b4d0e490 100644 --- a/Lib/test/test_asyncio/test_graph.py +++ b/Lib/test/test_asyncio/test_graph.py @@ -613,6 +613,24 @@ async def main(): await main() self.assertEqual(stack[:3], ['gen', 'middle', 'main']) + def set_eager_task_factory(self): + loop = asyncio.get_running_loop() + loop.set_task_factory(asyncio.create_eager_task_factory(asyncio.Task)) + self.addCleanup(loop.set_task_factory, None) + + async def test_stack_eager_task(self): + # gh-157299 + self.set_eager_task_factory() + + async def child(): + nonlocal stack + stack = capture_test_stack() + + stack = None + await asyncio.gather(child()) + + self.assertEqual(stack[0][2], [['T', ['a test_stack_eager_task'], []]]) + @unittest.skipIf( not hasattr(asyncio.futures, "_c_future_add_to_awaited_by"), diff --git a/Misc/NEWS.d/next/Library/2026-09-11-12-43-04.gh-issue-157299.HQ3I4F.rst b/Misc/NEWS.d/next/Library/2026-09-11-12-43-04.gh-issue-157299.HQ3I4F.rst new file mode 100644 index 00000000000000..8674c45eef0c59 --- /dev/null +++ b/Misc/NEWS.d/next/Library/2026-09-11-12-43-04.gh-issue-157299.HQ3I4F.rst @@ -0,0 +1,2 @@ +Fix eagerly started :class:`asyncio.Task` objects missing their awaited-by +edge during the first step. diff --git a/Modules/_asynciomodule.c b/Modules/_asynciomodule.c index a380f8ac72b32f..2aaf39832917d4 100644 --- a/Modules/_asynciomodule.c +++ b/Modules/_asynciomodule.c @@ -3460,10 +3460,32 @@ task_eager_start(_PyThreadStateImpl *ts, asyncio_state *state, TaskObj *task) int retval = 0; + // gh-157299: the eager step must record the awaited_by edge + int eager_edge = (prevtask != Py_None); + if (eager_edge) { + int res; + Py_BEGIN_CRITICAL_SECTION(task); + res = future_awaited_by_add(state, (FutureObj *)task, prevtask); + Py_END_CRITICAL_SECTION(); + if (res) { + eager_edge = 0; + retval = -1; + } + } + PyObject *stepres; Py_BEGIN_CRITICAL_SECTION(task); stepres = task_step_impl(state, task, NULL); Py_END_CRITICAL_SECTION(); + if (eager_edge) { + int res; + Py_BEGIN_CRITICAL_SECTION(task); + res = future_awaited_by_discard(state, (FutureObj *)task, prevtask); + Py_END_CRITICAL_SECTION(); + if (res) { + retval = -1; + } + } if (stepres == NULL) { PyObject *exc = PyErr_GetRaisedException(); _PyErr_ChainExceptions1(exc);