refactor(graphql): Use factories for schema extension initialization
Change `get_schema_extensions()` to return extension factories instead of instances. This defers extension initialization and prevents stale references to settings captured at import time. Lambdas capture settings values when extensions are constructed, and tests now instantiate extensions from factories to verify configuration. Fixes #22451
This commit is contained in:
parent
c889e58bee
commit
eaed2a7f8e
|
|
@ -1,3 +1,5 @@
|
|||
from collections.abc import Callable
|
||||
|
||||
import strawberry
|
||||
from django.conf import settings
|
||||
from strawberry.extensions import MaxAliasesLimiter, QueryDepthLimiter, SchemaExtension
|
||||
|
|
@ -18,6 +20,8 @@ from wireless.graphql.schema import WirelessQuery
|
|||
|
||||
from .scalars import BigInt, BigIntScalar
|
||||
|
||||
SchemaExtensionFactory = type[SchemaExtension] | Callable[[], SchemaExtension]
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class Query(
|
||||
|
|
@ -36,14 +40,16 @@ class Query(
|
|||
pass
|
||||
|
||||
|
||||
def get_schema_extensions() -> list[SchemaExtension]:
|
||||
extensions: list[SchemaExtension] = [
|
||||
DjangoOptimizerExtension(prefetch_custom_queryset=True),
|
||||
MaxAliasesLimiter(max_alias_count=settings.GRAPHQL_MAX_ALIASES),
|
||||
]
|
||||
def get_schema_extensions() -> list[SchemaExtensionFactory]:
|
||||
max_aliases = settings.GRAPHQL_MAX_ALIASES
|
||||
max_depth = settings.GRAPHQL_MAX_QUERY_DEPTH
|
||||
|
||||
extensions: list[SchemaExtensionFactory] = [
|
||||
lambda: DjangoOptimizerExtension(prefetch_custom_queryset=True),
|
||||
lambda: MaxAliasesLimiter(max_alias_count=max_aliases),
|
||||
]
|
||||
if max_depth and max_depth > 0:
|
||||
extensions.append(QueryDepthLimiter(max_depth=max_depth))
|
||||
extensions.append(lambda: QueryDepthLimiter(max_depth=max_depth))
|
||||
return extensions
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -19,6 +19,9 @@ from utilities.testing import APITestCase, TestCase, disable_warnings
|
|||
|
||||
class GraphQLTestCase(TestCase):
|
||||
|
||||
def _schema_extension_instances(self):
|
||||
return [factory() for factory in get_schema_extensions()]
|
||||
|
||||
@override_settings(GRAPHQL_ENABLED=False)
|
||||
def test_graphql_enabled(self):
|
||||
"""
|
||||
|
|
@ -32,21 +35,21 @@ class GraphQLTestCase(TestCase):
|
|||
"""
|
||||
QueryDepthLimiter should not be installed when GRAPHQL_MAX_QUERY_DEPTH is unset.
|
||||
"""
|
||||
self.assertFalse(any(isinstance(ext, QueryDepthLimiter) for ext in get_schema_extensions()))
|
||||
self.assertFalse(any(isinstance(ext, QueryDepthLimiter) for ext in self._schema_extension_instances()))
|
||||
|
||||
@override_settings(GRAPHQL_MAX_QUERY_DEPTH=0)
|
||||
def test_graphql_max_query_depth_disabled_when_zero(self):
|
||||
"""
|
||||
QueryDepthLimiter should not be installed when GRAPHQL_MAX_QUERY_DEPTH is zero.
|
||||
"""
|
||||
self.assertFalse(any(isinstance(ext, QueryDepthLimiter) for ext in get_schema_extensions()))
|
||||
self.assertFalse(any(isinstance(ext, QueryDepthLimiter) for ext in self._schema_extension_instances()))
|
||||
|
||||
@override_settings(GRAPHQL_MAX_QUERY_DEPTH=-1)
|
||||
def test_graphql_max_query_depth_disabled_when_negative(self):
|
||||
"""
|
||||
QueryDepthLimiter should not be installed when GRAPHQL_MAX_QUERY_DEPTH is negative.
|
||||
"""
|
||||
self.assertFalse(any(isinstance(ext, QueryDepthLimiter) for ext in get_schema_extensions()))
|
||||
self.assertFalse(any(isinstance(ext, QueryDepthLimiter) for ext in self._schema_extension_instances()))
|
||||
|
||||
@override_settings(GRAPHQL_MAX_QUERY_DEPTH=3)
|
||||
def test_graphql_max_query_depth_enforced(self):
|
||||
|
|
@ -54,9 +57,9 @@ class GraphQLTestCase(TestCase):
|
|||
Queries exceeding GRAPHQL_MAX_QUERY_DEPTH should be rejected.
|
||||
"""
|
||||
extensions = get_schema_extensions()
|
||||
self.assertTrue(any(isinstance(ext, QueryDepthLimiter) for ext in extensions))
|
||||
self.assertTrue(any(isinstance(ext, QueryDepthLimiter) for ext in self._schema_extension_instances()))
|
||||
|
||||
# Build a temporary schema with the configured extensions and execute a deep query
|
||||
# Build a temporary schema with the configured extension factories and execute a deep query
|
||||
test_schema = strawberry.Schema(
|
||||
query=Query,
|
||||
config=StrawberryConfig(auto_camel_case=False, scalar_map={BigInt: BigIntScalar}),
|
||||
|
|
|
|||
Loading…
Reference in New Issue