diff --git a/packages/django-cf/django_cf/__init__.py b/packages/django-cf/django_cf/__init__.py index 26f3d951..b93d0014 100644 --- a/packages/django-cf/django_cf/__init__.py +++ b/packages/django-cf/django_cf/__init__.py @@ -1,11 +1,7 @@ -import os - from workers import wsgi async def handle_wsgi(request, app, env=None): - os.environ.setdefault("DJANGO_ALLOW_ASYNC_UNSAFE", "1") - return await wsgi.fetch(app, request, env) diff --git a/packages/django-cf/django_cf/db/_async_unsafe.py b/packages/django-cf/django_cf/db/_async_unsafe.py new file mode 100644 index 00000000..8c94f1d2 --- /dev/null +++ b/packages/django-cf/django_cf/db/_async_unsafe.py @@ -0,0 +1,32 @@ +"""Helper for removing Django's ``async_unsafe`` guard from backend methods.""" + +from django.utils import asyncio as _django_asyncio + + +def _async_unsafe_probe(): + pass + + +# Capture async_unsafe function so we don't accidentally unwrap something else +_ASYNC_UNSAFE_WRAPPER_CODE = _django_asyncio.async_unsafe(_async_unsafe_probe).__code__ + + +def _unwrap_async_unsafe(method): + """ + Unwrap Django's async_unsafe decorator to get the original method. + + This is needed because Django does not allow calling database operations + from within an async context, but Python workers always run in an async context. + + All the database backends that django-cf provides are async-safe so we need + to unwrap the async_unsafe decorator to allow calling database operations + from within an async context. + """ + method_code = getattr(method, "__code__", None) + is_async_unsafe = method_code is _ASYNC_UNSAFE_WRAPPER_CODE + wrapped = getattr(method, "__wrapped__", None) + + if is_async_unsafe: + return wrapped + + return method diff --git a/packages/django-cf/django_cf/db/base_engine.py b/packages/django-cf/django_cf/db/base_engine.py index b7a7e29b..9155745f 100644 --- a/packages/django-cf/django_cf/db/base_engine.py +++ b/packages/django-cf/django_cf/db/base_engine.py @@ -12,6 +12,7 @@ OperationalError, ProgrammingError, ) +from django.db.backends.base.base import BaseDatabaseWrapper from django.db.backends.sqlite3.base import DatabaseWrapper as SQLiteDatabaseWrapper from django.db.backends.sqlite3.client import DatabaseClient as SQLiteDatabaseClient from django.db.backends.sqlite3.creation import ( @@ -43,6 +44,8 @@ ) from django.db.models.sql.compiler import SQLCompiler +from ._async_unsafe import _unwrap_async_unsafe + def replace_date_trunc_in_sql(sql): """Replace django_date_trunc and django_datetime_trunc function calls with SQLite equivalents.""" @@ -478,6 +481,16 @@ class CFDatabaseWrapper(SQLiteDatabaseWrapper): transaction_modes = frozenset([]) + connect = _unwrap_async_unsafe(BaseDatabaseWrapper.connect) + ensure_connection = _unwrap_async_unsafe(BaseDatabaseWrapper.ensure_connection) + cursor = _unwrap_async_unsafe(BaseDatabaseWrapper.cursor) + commit = _unwrap_async_unsafe(BaseDatabaseWrapper.commit) + rollback = _unwrap_async_unsafe(BaseDatabaseWrapper.rollback) + savepoint = _unwrap_async_unsafe(BaseDatabaseWrapper.savepoint) + savepoint_rollback = _unwrap_async_unsafe(BaseDatabaseWrapper.savepoint_rollback) + savepoint_commit = _unwrap_async_unsafe(BaseDatabaseWrapper.savepoint_commit) + clean_savepoints = _unwrap_async_unsafe(BaseDatabaseWrapper.clean_savepoints) + def get_compiler(self, default_using=None, using=None, **kwargs): if using is None: using = default_using diff --git a/packages/django-cf/tests/in_worker/worker/src/test_base_engine.py b/packages/django-cf/tests/in_worker/worker/src/test_base_engine.py index cb01bbbb..a9fc0566 100644 --- a/packages/django-cf/tests/in_worker/worker/src/test_base_engine.py +++ b/packages/django-cf/tests/in_worker/worker/src/test_base_engine.py @@ -6,6 +6,10 @@ from unittest.mock import MagicMock, patch import pytest +from django.db.backends.base.base import BaseDatabaseWrapper +from django.utils.asyncio import async_unsafe + +from django_cf.db.base_engine import CFDatabaseWrapper, _unwrap_async_unsafe class TestCFResult: @@ -547,7 +551,42 @@ def test_bulk_insert_enabled(self): assert features.can_return_columns_from_insert is True +class TestUnwrapAsyncUnsafe: + def test_removes_async_unsafe_guard(self): + def target(): + return "called" + + guarded = async_unsafe(target) + + assert _unwrap_async_unsafe(guarded) is target + + def test_returns_undecorated_method_unchanged(self): + def target(): + return "called" + + assert _unwrap_async_unsafe(target) is target + + class TestCFDatabaseWrapper: + def test_connection_lifecycle_is_not_guarded_by_async_unsafe(self): + method_names = ( + "connect", + "ensure_connection", + "cursor", + "commit", + "rollback", + "savepoint", + "savepoint_rollback", + "savepoint_commit", + "clean_savepoints", + ) + + for method_name in method_names: + guarded_method = getattr(BaseDatabaseWrapper, method_name) + assert getattr(CFDatabaseWrapper, method_name) is _unwrap_async_unsafe( + guarded_method + ) + def test_get_database_version(self): from django_cf.db.base_engine import CFDatabaseWrapper diff --git a/packages/django-cf/tests/in_worker/worker/src/worker.py b/packages/django-cf/tests/in_worker/worker/src/worker.py index 2b163358..6bb0f4e1 100644 --- a/packages/django-cf/tests/in_worker/worker/src/worker.py +++ b/packages/django-cf/tests/in_worker/worker/src/worker.py @@ -1,6 +1,5 @@ # pyright: reportMissingImports=false -import os from pathlib import Path from urllib.parse import urlparse @@ -14,7 +13,6 @@ from worker_durable_object import TestDurableObject # noqa: F401 BASE_DIR = Path(__file__).parent -os.environ.setdefault("DJANGO_ALLOW_ASYNC_UNSAFE", "true") if not django.conf.settings.configured: