From 9e06baf2e4dd17834b5deaf3a3e6ae3bff2fcf26 Mon Sep 17 00:00:00 2001 From: Thomas Grainger Date: Sat, 25 Apr 2026 15:20:21 +0100 Subject: [PATCH] Fixed #37066 -- Wrapped async iteration sites in (maybe_)aclosing. MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Replaced bare `async for` over a not-locally-defined async iterable with `async with (maybe_)aclosing(aiter(...)) as it: async for item in it:` at every consumption site, so the iterator's aclose() runs deterministically when the consumer exits — instead of being deferred to asyncio's asyncgen finalizer hook. Added django.utils.asyncio.maybe_aclosing as a small wrapper that uses contextlib.aclosing when the iterator has aclose() and contextlib.nullcontext otherwise (for caller-supplied iterators that may only implement __aiter__ and __anext__). Sites: - contrib/auth/backends.py:_aget_permissions (aclosing) - core/paginator.py:AsyncPage.__aiter__ (maybe_aclosing) - core/paginator.py:AsyncPage.aget_object_list (maybe_aclosing) - db/models/query.py:QuerySet.aiterator (aclosing, hoisted across both prefetch and no-prefetch branches) - http/response.py:StreamingHttpResponse.streaming_content awrapper (maybe_aclosing) - http/response.py:StreamingHttpResponse.__iter__ to_list (maybe_aclosing) - http/response.py:StreamingHttpResponse.__aiter__ (aclosing) - shortcuts.py:aget_list_or_404 (aclosing) - test/client.py:aclosing_iterator_wrapper (maybe_aclosing) - utils/text.py:acompress_sequence (maybe_aclosing) --- django/contrib/auth/backends.py | 13 ++++++----- django/core/paginator.py | 9 +++++--- django/db/models/query.py | 37 +++++++++++++++--------------- django/http/response.py | 17 +++++++++----- django/shortcuts.py | 5 ++++- django/test/client.py | 6 +++-- django/utils/asyncio.py | 19 ++++++++++++++++ django/utils/text.py | 14 +++++++----- docs/ref/utils.txt | 26 +++++++++++++++++++++ docs/releases/6.2.txt | 3 +++ tests/async/tests.py | 40 ++++++++++++++++++++++++++++++++- tests/httpwrappers/tests.py | 17 ++++++++++++++ tests/pagination/tests.py | 19 ++++++++++++++++ tests/test_client/tests.py | 19 ++++++++++++++++ tests/utils_tests/test_text.py | 19 +++++++++++++++- 15 files changed, 220 insertions(+), 43 deletions(-) diff --git a/django/contrib/auth/backends.py b/django/contrib/auth/backends.py index 514c67060204..68603fa0455b 100644 --- a/django/contrib/auth/backends.py +++ b/django/contrib/auth/backends.py @@ -1,3 +1,5 @@ +from contextlib import aclosing + from asgiref.sync import sync_to_async from django.contrib.auth import ( @@ -139,11 +141,12 @@ async def _aget_permissions(self, user_obj, obj, from_name): else: perms = getattr(self, "_get_%s_permissions" % from_name)(user_obj) perms = perms.values_list("content_type__app_label", "codename").order_by() - setattr( - user_obj, - perm_cache_name, - {"%s.%s" % (ct, name) async for ct, name in perms}, - ) + async with aclosing(aiter(perms)) as it: + setattr( + user_obj, + perm_cache_name, + {"%s.%s" % (ct, name) async for ct, name in it}, + ) return getattr(user_obj, perm_cache_name) def get_user_permissions(self, user_obj, obj=None): diff --git a/django/core/paginator.py b/django/core/paginator.py index bb5f1987c9f1..4302eeb7ccc1 100644 --- a/django/core/paginator.py +++ b/django/core/paginator.py @@ -5,6 +5,7 @@ from asgiref.sync import sync_to_async +from django.utils.asyncio import maybe_aclosing from django.utils.deprecation import RemovedInDjango2028Warning from django.utils.functional import cached_property from django.utils.inspect import method_has_no_args @@ -370,8 +371,9 @@ def __repr__(self): async def __aiter__(self): if hasattr(self.object_list, "__aiter__"): - async for obj in self.object_list: - yield obj + async with maybe_aclosing(aiter(self.object_list)) as it: + async for obj in it: + yield obj else: for obj in self.object_list: yield obj @@ -414,7 +416,8 @@ async def aget_object_list(self): """ if not isinstance(self.object_list, list): if hasattr(self.object_list, "__aiter__"): - self.object_list = [obj async for obj in self.object_list] + async with maybe_aclosing(aiter(self.object_list)) as it: + self.object_list = [obj async for obj in it] else: self.object_list = await sync_to_async(list)(self.object_list) return self.object_list diff --git a/django/db/models/query.py b/django/db/models/query.py index 9891001caea6..3db85eed26df 100644 --- a/django/db/models/query.py +++ b/django/db/models/query.py @@ -5,7 +5,7 @@ import copy import operator import warnings -from contextlib import nullcontext +from contextlib import aclosing, nullcontext from functools import partial, reduce from itertools import chain, islice from weakref import ref as weak_ref @@ -656,28 +656,29 @@ async def aiterator(self, chunk_size=None): chunked_fetch=use_chunked_fetch, chunk_size=chunk_size or 2000, ) - if self._prefetch_related_lookups: - results = [] - - async for item in iterable: - results.append(item) - if len(results) >= chunk_size: + async with aclosing(aiter(iterable)) as iterator: + if self._prefetch_related_lookups: + results = [] + + async for item in iterator: + results.append(item) + if len(results) >= chunk_size: + await aprefetch_related_objects( + results, *self._prefetch_related_lookups + ) + for result in results: + yield result + results.clear() + + if results: await aprefetch_related_objects( results, *self._prefetch_related_lookups ) for result in results: yield result - results.clear() - - if results: - await aprefetch_related_objects( - results, *self._prefetch_related_lookups - ) - for result in results: - yield result - else: - async for item in iterable: - yield item + else: + async for item in iterator: + yield item def aggregate(self, *args, **kwargs): """ diff --git a/django/http/response.py b/django/http/response.py index b1d70b4047b4..022cd3741388 100644 --- a/django/http/response.py +++ b/django/http/response.py @@ -7,6 +7,7 @@ import sys import time import warnings +from contextlib import aclosing from email.header import Header from http.client import responses from urllib.parse import urlsplit @@ -19,6 +20,7 @@ from django.core.serializers.json import DjangoJSONEncoder from django.http.cookie import SimpleCookie from django.utils import timezone +from django.utils.asyncio import maybe_aclosing from django.utils.datastructures import CaseInsensitiveMapping from django.utils.deprecation import RemovedInDjango2029Warning from django.utils.encoding import iri_to_uri @@ -492,8 +494,9 @@ def streaming_content(self): _iterator = self._iterator async def awrapper(): - async for part in _iterator: - yield self.make_bytes(part) + async with maybe_aclosing(_iterator) as it: + async for part in it: + yield self.make_bytes(part) return awrapper() else: @@ -528,16 +531,18 @@ def __iter__(self): # async iterator. Consume in async_to_sync and map back. async def to_list(_iterator): as_list = [] - async for chunk in _iterator: - as_list.append(chunk) + async with maybe_aclosing(_iterator) as it: + async for chunk in it: + as_list.append(chunk) return as_list return map(self.make_bytes, iter(async_to_sync(to_list)(self._iterator))) async def __aiter__(self): try: - async for part in self.streaming_content: - yield part + async with aclosing(aiter(self.streaming_content)) as content: + async for part in content: + yield part except TypeError: warnings.warn( "StreamingHttpResponse must consume synchronous iterators in order to " diff --git a/django/shortcuts.py b/django/shortcuts.py index 84288a6e6f89..0f6f9d9fa65a 100644 --- a/django/shortcuts.py +++ b/django/shortcuts.py @@ -4,6 +4,8 @@ for convenience's sake. """ +from contextlib import aclosing + from django.http import ( Http404, HttpResponse, @@ -162,7 +164,8 @@ async def aget_list_or_404(klass, *args, **kwargs): "First argument to aget_list_or_404() must be a Model, Manager, or " f"QuerySet, not '{klass__name}'." ) - obj_list = [obj async for obj in queryset.filter(*args, **kwargs)] + async with aclosing(aiter(queryset.filter(*args, **kwargs))) as it: + obj_list = [obj async for obj in it] if not obj_list: raise Http404( _("No %s matches the given query.") % queryset.model._meta.object_name diff --git a/django/test/client.py b/django/test/client.py index d13bc28fc6ab..c258ec3635ed 100644 --- a/django/test/client.py +++ b/django/test/client.py @@ -23,6 +23,7 @@ from django.test import signals from django.test.utils import ContextList from django.urls import resolve +from django.utils.asyncio import maybe_aclosing from django.utils.encoding import force_bytes from django.utils.functional import SimpleLazyObject from django.utils.http import urlencode @@ -128,8 +129,9 @@ def closing_iterator_wrapper(iterable, close): async def aclosing_iterator_wrapper(iterable, close): try: - async for chunk in iterable: - yield chunk + async with maybe_aclosing(aiter(iterable)) as it: + async for chunk in it: + yield chunk finally: request_finished.disconnect(close_old_connections) close() # will fire request_finished diff --git a/django/utils/asyncio.py b/django/utils/asyncio.py index 1e79f90c2c1b..676d0e15e514 100644 --- a/django/utils/asyncio.py +++ b/django/utils/asyncio.py @@ -1,3 +1,4 @@ +import contextlib import os from asyncio import get_running_loop from functools import wraps @@ -5,6 +6,24 @@ from django.core.exceptions import SynchronousOnlyOperation +def maybe_aclosing(iterator): + """ + Return a context manager that calls ``aclose()`` on *iterator* on exit + if it has one, otherwise a no-op context manager that yields *iterator* + unchanged. + + Use to consume a caller-supplied async iterable that may or may not be + a real async generator. Wrapping with :func:`contextlib.aclosing` + unconditionally would raise ``AttributeError`` at ``__aexit__`` for any + iterator without ``aclose()``. + """ + return ( + contextlib.aclosing(iterator) + if hasattr(iterator, "aclose") + else contextlib.nullcontext(iterator) + ) + + def async_unsafe(message): """ Decorator to mark functions as async-unsafe. Someone trying to access diff --git a/django/utils/text.py b/django/utils/text.py index d1306f9c6fed..6297e4b6c145 100644 --- a/django/utils/text.py +++ b/django/utils/text.py @@ -11,6 +11,7 @@ from io import BytesIO from django.core.exceptions import SuspiciousFileOperation +from django.utils.asyncio import maybe_aclosing from django.utils.functional import ( SimpleLazyObject, cached_property, @@ -397,12 +398,13 @@ async def acompress_sequence(sequence, *, max_random_bytes=None): ) as zfile: # Output headers... yield buf.read() - async for item in sequence: - zfile.write(item) - zfile.flush() - data = buf.read() - if data: - yield data + async with maybe_aclosing(aiter(sequence)) as it: + async for item in it: + zfile.write(item) + zfile.flush() + data = buf.read() + if data: + yield data yield buf.read() diff --git a/docs/ref/utils.txt b/docs/ref/utils.txt index 08a264c1e041..ae5d5249d4d4 100644 --- a/docs/ref/utils.txt +++ b/docs/ref/utils.txt @@ -11,6 +11,32 @@ following parts can be considered stable and thus backwards compatible as per the :ref:`internal release deprecation policy `. +``django.utils.asyncio`` +======================== + +.. module:: django.utils.asyncio + :synopsis: Helper functions for use with asyncio. + +.. function:: maybe_aclosing(iterator) + +.. versionadded:: 6.2 + + Return a context manager that calls ``aclose()`` on *iterator* on exit if + it has one, otherwise a no-op context manager that yields *iterator* + unchanged. Use it to consume a caller-supplied async iterable that may not + be a real async generator:: + + from django.utils.asyncio import maybe_aclosing + + + async def consume(iterable): + async with maybe_aclosing(aiter(iterable)) as it: + async for item in it: + ... + + Wrapping with :func:`contextlib.aclosing` unconditionally would raise + ``AttributeError`` at ``__aexit__`` for any iterator without ``aclose()``. + ``django.utils.cache`` ====================== diff --git a/docs/releases/6.2.txt b/docs/releases/6.2.txt index b135e8b5188c..31eb7a33f061 100644 --- a/docs/releases/6.2.txt +++ b/docs/releases/6.2.txt @@ -258,6 +258,9 @@ Utilities top-level modules did not work, and submodules only worked if already imported. +* The new :func:`django.utils.asyncio.maybe_aclosing` returns a context manager + that calls ``aclose()`` on a caller-provided iterator only if it defines one. + * ... Validators diff --git a/tests/async/tests.py b/tests/async/tests.py index b7b9dfbe7b16..813dd2719f18 100644 --- a/tests/async/tests.py +++ b/tests/async/tests.py @@ -9,7 +9,7 @@ from django.core.exceptions import ImproperlyConfigured, SynchronousOnlyOperation from django.http import HttpResponse, HttpResponseNotAllowed from django.test import RequestFactory, SimpleTestCase -from django.utils.asyncio import async_unsafe +from django.utils.asyncio import async_unsafe, maybe_aclosing from django.views.generic.base import View from .models import SimpleModel @@ -65,6 +65,44 @@ async def test_async_unsafe_suppressed(self): self.fail("SynchronousOnlyOperation should not be raised.") +class MaybeAClosingTest(SimpleTestCase): + async def test_calls_aclose_on_iterator_with_aclose(self): + finally_ran = False + + async def gen(): + nonlocal finally_ran + try: + yield 1 + yield 2 + finally: + finally_ran = True + + iterator = gen() + async with maybe_aclosing(iterator) as it: + self.assertIs(it, iterator) + self.assertEqual(await anext(it), 1) + self.assertTrue(finally_ran) + + async def test_no_aclose_on_iterator_without_aclose(self): + class PlainAsyncIterator: + def __init__(self): + self.values = iter([1, 2]) + + def __aiter__(self): + return self + + async def __anext__(self): + try: + return next(self.values) + except StopIteration: + raise StopAsyncIteration + + iterator = PlainAsyncIterator() + async with maybe_aclosing(iterator) as it: + self.assertIs(it, iterator) + self.assertEqual(await anext(it), 1) + + class SyncView(View): def get(self, request, *args, **kwargs): return HttpResponse("Hello (sync) world!") diff --git a/tests/httpwrappers/tests.py b/tests/httpwrappers/tests.py index 5e57d372a03b..4495191508b0 100644 --- a/tests/httpwrappers/tests.py +++ b/tests/httpwrappers/tests.py @@ -5,6 +5,7 @@ import pickle import unittest import uuid +from contextlib import aclosing from django.core.exceptions import DisallowedRedirect from django.core.serializers.json import DjangoJSONEncoder @@ -857,6 +858,22 @@ def test_text_attribute_error(self): with self.assertRaisesMessage(AttributeError, msg): r.text + async def test_streaming_response_closes_user_generator_promptly(self): + finally_ran = False + + async def body(): + nonlocal finally_ran + try: + yield b"chunk" + finally: + finally_ran = True + + response = StreamingHttpResponse(body()) + async with aclosing(aiter(response)) as content: + async for _ in content: + break + self.assertTrue(finally_ran) + class FileCloseTests(SimpleTestCase): def setUp(self): diff --git a/tests/pagination/tests.py b/tests/pagination/tests.py index 4a35495ef771..4417dcfa0fd4 100644 --- a/tests/pagination/tests.py +++ b/tests/pagination/tests.py @@ -3,9 +3,11 @@ import pathlib import unittest.mock import warnings +from contextlib import aclosing from datetime import datetime from django.core.paginator import ( + AsyncPage, AsyncPaginator, BasePaginator, EmptyPage, @@ -990,3 +992,20 @@ async def test_aget_object_list(self): # It returns the same list that was converted on the first call. second_called_objs = await p.aget_object_list() self.assertEqual(id(first_called_objs), id(second_called_objs)) + + async def test_async_page_aiteration_closes_object_list_promptly(self): + finally_ran = False + + async def object_list(): + nonlocal finally_ran + try: + yield 1 + yield 2 + finally: + finally_ran = True + + page = AsyncPage(object_list(), number=1, paginator=None) + async with aclosing(aiter(page)) as it: + async for _ in it: + break + self.assertTrue(finally_ran) diff --git a/tests/test_client/tests.py b/tests/test_client/tests.py index cc66a157b0a5..0437bc4bed13 100644 --- a/tests/test_client/tests.py +++ b/tests/test_client/tests.py @@ -23,6 +23,7 @@ import copy import itertools import tempfile +from contextlib import aclosing from unittest import mock from django.contrib.auth.models import Permission, User @@ -38,6 +39,7 @@ modify_settings, override_settings, ) +from django.test.client import aclosing_iterator_wrapper from django.urls import reverse_lazy from django.utils.datastructures import MultiValueDict from django.utils.decorators import async_only_middleware @@ -1264,6 +1266,23 @@ async def test_response_resolver_match(self): self.assertTrue(hasattr(response, "resolver_match")) self.assertEqual(response.resolver_match.url_name, "async_get_view") + async def test_aclosing_iterator_wrapper_closes_inner_iterable_promptly(self): + finally_ran = False + + async def iterable(): + nonlocal finally_ran + try: + yield b"chunk" + finally: + finally_ran = True + + async with aclosing( + aclosing_iterator_wrapper(iterable(), close=lambda: None) + ) as wrapped: + async for _ in wrapped: + break + self.assertTrue(finally_ran) + @modify_settings( MIDDLEWARE={"prepend": "test_client.tests.async_middleware_urlconf"}, ) diff --git a/tests/utils_tests/test_text.py b/tests/utils_tests/test_text.py index 101943957c0c..7fd767c00ab3 100644 --- a/tests/utils_tests/test_text.py +++ b/tests/utils_tests/test_text.py @@ -1,12 +1,13 @@ import gzip import json import sys +from contextlib import aclosing from django.core.exceptions import SuspiciousFileOperation from django.test import SimpleTestCase from django.utils import text from django.utils.functional import lazystr -from django.utils.text import format_lazy +from django.utils.text import acompress_sequence, format_lazy from django.utils.translation import gettext_lazy, override IS_WIDE_BUILD = len("\U0001f4a9") == 1 @@ -418,6 +419,22 @@ def test_compress_sequence(self): self.assertLess(len(out), len(original)) self.assertGreater(len(compressed_chunks), 2) + async def test_acompress_sequence_closes_inner_sequence_promptly(self): + finally_ran = False + + async def sequence(): + nonlocal finally_ran + try: + yield b"chunk1" * 1024 + yield b"chunk2" * 1024 + finally: + finally_ran = True + + async with aclosing(acompress_sequence(sequence())) as compressed: + await anext(compressed) # gzip header + await anext(compressed) # first sequence-derived chunk + self.assertTrue(finally_ran) + def test_format_lazy(self): self.assertEqual("django/test", format_lazy("{}/{}", "django", lazystr("test"))) self.assertEqual("django/test", format_lazy("{0}/{1}", *("django", "test")))