Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 8 additions & 5 deletions django/contrib/auth/backends.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from contextlib import aclosing

from asgiref.sync import sync_to_async

from django.contrib.auth import (
Expand Down Expand Up @@ -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):
Expand Down
9 changes: 6 additions & 3 deletions django/core/paginator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
37 changes: 19 additions & 18 deletions django/db/models/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
"""
Expand Down
17 changes: 11 additions & 6 deletions django/http/response.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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 "
Expand Down
5 changes: 4 additions & 1 deletion django/shortcuts.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@
for convenience's sake.
"""

from contextlib import aclosing

from django.http import (
Http404,
HttpResponse,
Expand Down Expand Up @@ -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
Expand Down
6 changes: 4 additions & 2 deletions django/test/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
19 changes: 19 additions & 0 deletions django/utils/asyncio.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,29 @@
import contextlib
import os
from asyncio import get_running_loop
from functools import wraps

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
Expand Down
14 changes: 8 additions & 6 deletions django/utils/text.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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()


Expand Down
26 changes: 26 additions & 0 deletions docs/ref/utils.txt
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,32 @@ following parts can be considered stable and thus backwards compatible as per
the :ref:`internal release deprecation policy
<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``
======================

Expand Down
3 changes: 3 additions & 0 deletions docs/releases/6.2.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
40 changes: 39 additions & 1 deletion tests/async/tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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!")
Expand Down
17 changes: 17 additions & 0 deletions tests/httpwrappers/tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down
Loading
Loading