Coverage for utilities/testing/api.py: 0%
1068 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 18:35 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 18:35 +0000
1import copy
2import importlib
3import inspect
4import json
5import types
6import typing
7from collections.abc import Callable
8from dataclasses import dataclass
9from decimal import Decimal
11import strawberry
12import strawberry_django
13from django.apps import apps
14from django.conf import settings
15from django.contrib.contenttypes.models import ContentType
16from django.contrib.postgres.fields import ArrayField
17from django.db import models
18from django.test import override_settings
19from django.urls import reverse
20from graphql import GraphQLList, GraphQLNonNull, GraphQLObjectType
21from rest_framework import status
22from rest_framework.test import APIClient
23from strawberry.schema.schema_converter import GraphQLCoreConverter
24from strawberry.types.base import StrawberryList, StrawberryOptional
25from strawberry.types.lazy_type import LazyType
26from strawberry.types.union import StrawberryUnion
27from strawberry_django import (
28 BaseFilterLookup,
29 ComparisonFilterLookup,
30 DateFilterLookup,
31 DatetimeFilterLookup,
32 FilterLookup,
33 RangeLookup,
34 StrFilterLookup,
35 TimeFilterLookup,
36)
38from core.choices import ObjectChangeActionChoices
39from core.models import ObjectChange, ObjectType
40from ipam.graphql.types import IPAddressFamilyType
41from netbox.api.exceptions import GraphQLTypeNotFound
42from netbox.graphql.filter_lookups import (
43 ArrayLookup,
44 BigIntegerLookup,
45 FloatLookup,
46 IntegerLookup,
47 IntegerRangeArrayLookup,
48 JSONFilter,
49 TreeNodeFilter,
50)
51from netbox.models.features import ChangeLoggingMixin
52from users.constants import TOKEN_PREFIX
53from users.models import ObjectPermission, Token, User
54from utilities.api import get_graphql_type_for_model
56from .base import ModelTestCase, TestCase
57from .query_counts import assert_expected_query_count
58from .utils import disable_logging, disable_warnings, get_random_string
60__all__ = (
61 'APITestCase',
62 'APIViewTestCases',
63 'GraphQLFilterTest',
64 'GraphQLQueryTest',
65)
68@dataclass(frozen=True)
69class GraphQLFilterTest:
70 """
71 Declarative GraphQL filter test case for APIViewTestCases.GraphQLTestCase.
73 ``filters`` is the raw content to place inside the GraphQL ``filters`` input,
74 e.g. ``name: {i_contains: "site"}``.
76 ``expected`` may be a callable accepting the model queryset, an ORM filter
77 dict, a queryset, an iterable of model instances, or an iterable of object
78 IDs. When omitted, the test only asserts that the filter returns at least one
79 result; this preserves compatibility with the legacy ``graphql_filter``
80 attribute.
81 """
82 name: str
83 filters: str
84 expected: object = None
85 permissions: tuple[str, ...] = ()
88@dataclass(frozen=True)
89class GraphQLQueryTest:
90 """
91 Declarative GraphQL query test case for model-specific complex queries.
93 ``assert_result`` is called as ``assert_result(testcase, data)`` where
94 ``testcase`` is the running ``GraphQLTestCase`` instance (use it for
95 ``testcase.assertEqual`` etc.) and ``data`` is the decoded GraphQL
96 ``data`` object (the inner ``response.json()['data']``, not the full HTTP
97 response).
98 """
99 name: str
100 query: str
101 assert_result: Callable
102 permissions: tuple[str, ...] = ()
105#
106# REST/GraphQL API Tests
107#
109class APITestCase(ModelTestCase):
110 """
111 Base test case for API requests.
113 client_class: Test client class
114 view_namespace: Namespace for API views. If None, the model's app_label will be used.
115 """
116 client_class = APIClient
117 view_namespace = None
119 def setUp(self):
120 """
121 Create a user and token for API calls.
122 """
123 # Create the test user and assign permissions
124 self.user = User.objects.create_user(username='testuser')
125 self.add_permissions(*self.user_permissions)
126 self.token = Token.objects.create(user=self.user)
127 self.header = {'HTTP_AUTHORIZATION': f'Bearer {TOKEN_PREFIX}{self.token.key}.{self.token.token}'}
129 def _get_view_namespace(self):
130 return f'{self.view_namespace or self.model._meta.app_label}-api'
132 def _get_detail_url(self, instance):
133 viewname = f'{self._get_view_namespace()}:{instance._meta.model_name}-detail'
134 return reverse(viewname, kwargs={'pk': instance.pk})
136 def _get_list_url(self):
137 viewname = f'{self._get_view_namespace()}:{self.model._meta.model_name}-list'
138 return reverse(viewname)
141class APIViewTestCases:
143 class GetObjectViewTestCase(APITestCase):
145 @override_settings(EXEMPT_VIEW_PERMISSIONS=['*'], LOGIN_REQUIRED=False)
146 def test_get_object_anonymous(self):
147 """
148 GET a single object as an unauthenticated user.
149 """
150 url = self._get_detail_url(self._get_queryset().first())
151 if (self.model._meta.app_label, self.model._meta.model_name) in settings.EXEMPT_EXCLUDE_MODELS:
152 # Models listed in EXEMPT_EXCLUDE_MODELS should not be accessible to anonymous users
153 with disable_warnings('django.request'):
154 self.assertHttpStatus(self.client.get(url, **self.header), status.HTTP_403_FORBIDDEN)
155 else:
156 response = self.client.get(url, **self.header)
157 self.assertHttpStatus(response, status.HTTP_200_OK)
159 def test_get_object_without_permission(self):
160 """
161 GET a single object as an authenticated user without the required permission.
162 """
163 url = self._get_detail_url(self._get_queryset().first())
165 # Try GET without permission
166 with disable_warnings('django.request'):
167 self.assertHttpStatus(self.client.get(url, **self.header), status.HTTP_403_FORBIDDEN)
169 def test_get_object(self):
170 """
171 GET a single object as an authenticated user with permission to view the object.
172 """
173 self.assertGreaterEqual(self._get_queryset().count(), 2,
174 f"Test requires the creation of at least two {self.model} instances")
175 instance1, instance2 = self._get_queryset()[:2]
177 # Add object-level permission
178 obj_perm = ObjectPermission(
179 name='Test permission',
180 constraints={'pk': instance1.pk},
181 actions=['view']
182 )
183 obj_perm.save()
184 obj_perm.users.add(self.user)
185 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
187 # Try GET to permitted object
188 url = self._get_detail_url(instance1)
189 response = self.client.get(url, **self.header)
190 self.assertHttpStatus(response, status.HTTP_200_OK)
192 # Verify ETag header is present for objects with timestamps
193 if issubclass(self.model, ChangeLoggingMixin):
194 self.assertIn('ETag', response, "ETag header missing from detail response")
196 # Try GET to non-permitted object
197 url = self._get_detail_url(instance2)
198 self.assertHttpStatus(self.client.get(url, **self.header), status.HTTP_404_NOT_FOUND)
200 @override_settings(EXEMPT_VIEW_PERMISSIONS=['*'])
201 def test_options_object(self):
202 """
203 Make an OPTIONS request for a single object.
204 """
205 url = self._get_detail_url(self._get_queryset().first())
206 response = self.client.options(url, **self.header)
207 self.assertHttpStatus(response, status.HTTP_200_OK)
209 class ListObjectsViewTestCase(APITestCase):
210 brief_fields = []
212 @override_settings(EXEMPT_VIEW_PERMISSIONS=['*'], LOGIN_REQUIRED=False)
213 def test_list_objects_anonymous(self):
214 """
215 GET a list of objects as an unauthenticated user.
216 """
217 url = self._get_list_url()
218 if (self.model._meta.app_label, self.model._meta.model_name) in settings.EXEMPT_EXCLUDE_MODELS:
219 # Models listed in EXEMPT_EXCLUDE_MODELS should not be accessible to anonymous users
220 with disable_warnings('django.request'):
221 self.assertHttpStatus(self.client.get(url, **self.header), status.HTTP_403_FORBIDDEN)
222 else:
223 response = self.client.get(url, **self.header)
224 self.assertHttpStatus(response, status.HTTP_200_OK)
225 self.assertEqual(len(response.data['results']), self._get_queryset().count())
227 def test_list_objects_brief(self):
228 """
229 GET a list of objects using the "brief" parameter.
230 """
231 self.add_permissions(f'{self.model._meta.app_label}.view_{self.model._meta.model_name}')
232 url = f'{self._get_list_url()}?brief=1'
233 response = self.client.get(url, **self.header)
235 self.assertHttpStatus(response, status.HTTP_200_OK)
236 self.assertEqual(len(response.data['results']), self._get_queryset().count())
237 self.assertEqual(sorted(response.data['results'][0]), self.brief_fields)
239 def test_list_objects_without_permission(self):
240 """
241 GET a list of objects as an authenticated user without the required permission.
242 """
243 url = self._get_list_url()
245 # Try GET without permission
246 with disable_warnings('django.request'):
247 self.assertHttpStatus(self.client.get(url, **self.header), status.HTTP_403_FORBIDDEN)
249 def test_list_objects(self):
250 """
251 GET a list of objects as an authenticated user with permission to view the objects.
252 """
253 self.assertGreaterEqual(self._get_queryset().count(), 3,
254 f"Test requires the creation of at least three {self.model} instances")
255 instance1, instance2 = self._get_queryset()[:2]
257 # Add object-level permission
258 obj_perm = ObjectPermission(
259 name='Test permission',
260 constraints={'pk__in': [instance1.pk, instance2.pk]},
261 actions=['view']
262 )
263 obj_perm.save()
264 obj_perm.users.add(self.user)
265 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
267 # Try GET to permitted objects
268 with assert_expected_query_count(self, 'api_list_objects'):
269 response = self.client.get(self._get_list_url(), **self.header)
270 self.assertHttpStatus(response, status.HTTP_200_OK)
271 self.assertEqual(len(response.data['results']), 2)
273 @override_settings(EXEMPT_VIEW_PERMISSIONS=['*'])
274 def test_options_objects(self):
275 """
276 Make an OPTIONS request for a list endpoint.
277 """
278 response = self.client.options(self._get_list_url(), **self.header)
279 self.assertHttpStatus(response, status.HTTP_200_OK)
281 class CreateObjectViewTestCase(APITestCase):
282 create_data = []
283 validation_excluded_fields = []
285 def test_create_object_without_permission(self):
286 """
287 POST a single object without permission.
288 """
289 url = self._get_list_url()
291 # Try POST without permission
292 with disable_warnings('django.request'):
293 response = self.client.post(url, self.create_data[0], format='json', **self.header)
294 self.assertHttpStatus(response, status.HTTP_403_FORBIDDEN)
296 def test_create_object(self):
297 """
298 POST a single object with permission.
299 """
300 # Add object-level permission
301 obj_perm = ObjectPermission(
302 name='Test permission',
303 actions=['add']
304 )
305 obj_perm.save()
306 obj_perm.users.add(self.user)
307 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
309 data = copy.deepcopy(self.create_data[0])
311 # If supported, add a changelog message
312 if issubclass(self.model, ChangeLoggingMixin):
313 data['changelog_message'] = get_random_string(10)
315 initial_count = self._get_queryset().count()
316 response = self.client.post(self._get_list_url(), data, format='json', **self.header)
317 self.assertHttpStatus(response, status.HTTP_201_CREATED)
318 self.assertEqual(self._get_queryset().count(), initial_count + 1)
319 instance = self._get_queryset().get(pk=response.data['id'])
320 self.assertInstanceEqual(
321 instance,
322 self.create_data[0],
323 exclude=self.validation_excluded_fields,
324 api=True
325 )
327 # Verify ObjectChange creation
328 if issubclass(self.model, ChangeLoggingMixin):
329 objectchange = ObjectChange.objects.get(
330 changed_object_type=ContentType.objects.get_for_model(instance),
331 changed_object_id=instance.pk,
332 action=ObjectChangeActionChoices.ACTION_CREATE,
333 )
334 self.assertObjectChange(objectchange, action=ObjectChangeActionChoices.ACTION_CREATE,
335 message=data['changelog_message'])
337 def test_bulk_create_objects(self):
338 """
339 POST a set of objects in a single request.
340 """
341 # Add object-level permission
342 obj_perm = ObjectPermission(
343 name='Test permission',
344 actions=['add']
345 )
346 obj_perm.save()
347 obj_perm.users.add(self.user)
348 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
350 # If supported, add a changelog message
351 changelog_message = get_random_string(10)
352 if issubclass(self.model, ChangeLoggingMixin):
353 for obj_data in self.create_data:
354 obj_data['changelog_message'] = changelog_message
356 initial_count = self._get_queryset().count()
357 response = self.client.post(self._get_list_url(), self.create_data, format='json', **self.header)
358 self.assertHttpStatus(response, status.HTTP_201_CREATED)
359 self.assertEqual(len(response.data), len(self.create_data))
360 self.assertEqual(self._get_queryset().count(), initial_count + len(self.create_data))
361 for i, obj in enumerate(response.data):
362 for field in self.create_data[i]:
363 if field in ('changelog_message', 'add_tags', 'remove_tags'):
364 # Write-only field
365 continue
366 if field not in self.validation_excluded_fields:
367 self.assertIn(field, obj, f"Bulk create field '{field}' missing from object {i} in response")
368 for i, obj in enumerate(response.data):
369 self.assertInstanceEqual(
370 self._get_queryset().get(pk=obj['id']),
371 self.create_data[i],
372 exclude=self.validation_excluded_fields,
373 api=True
374 )
376 # Verify ObjectChange creation
377 if issubclass(self.model, ChangeLoggingMixin):
378 id_list = [
379 obj['id'] for obj in response.data
380 ]
381 objectchanges = ObjectChange.objects.filter(
382 changed_object_type=ContentType.objects.get_for_model(self.model),
383 changed_object_id__in=id_list,
384 action=ObjectChangeActionChoices.ACTION_CREATE,
385 )
386 self.assertEqual(len(objectchanges), len(self.create_data))
387 for oc in objectchanges:
388 self.assertObjectChange(oc, action=ObjectChangeActionChoices.ACTION_CREATE,
389 message=changelog_message)
391 def test_bulk_create_objects_invalid_item(self):
392 """
393 POST a set of objects in which one item is invalid. The failure must be correlated to
394 that item's position in the request, and the entire batch must be rolled back.
395 """
396 obj_perm = ObjectPermission(
397 name='Test permission',
398 actions=['add']
399 )
400 obj_perm.save()
401 obj_perm.users.add(self.user)
402 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
404 initial_count = self._get_queryset().count()
406 # A non-dictionary is used as the invalid item because it is guaranteed to fail for every
407 # model, whereas which *fields* are required varies from one model to the next.
408 response = self.client.post(
409 self._get_list_url(),
410 [self.create_data[0], 'this is not an object'],
411 format='json',
412 **self.header,
413 )
415 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
416 self.assertEqual(
417 self._get_queryset().count(), initial_count,
418 'No objects should be created when any sibling fails validation'
419 )
420 self.assertIn('detail', response.data)
421 self.assertEqual(len(response.data['errors']), 1)
422 self.assertEqual(response.data['errors'][0]['index'], 1)
423 self.assertIn('errors', response.data['errors'][0])
425 class UpdateObjectViewTestCase(APITestCase):
426 update_data = {}
427 bulk_update_data = None
428 bulk_update_invalid_data = None
429 validation_excluded_fields = []
431 def test_update_object_without_permission(self):
432 """
433 PATCH a single object without permission.
434 """
435 url = self._get_detail_url(self._get_queryset().first())
436 update_data = self.update_data or getattr(self, 'create_data')[0]
438 # Try PATCH without permission
439 with disable_warnings('django.request'):
440 response = self.client.patch(url, update_data, format='json', **self.header)
441 self.assertHttpStatus(response, status.HTTP_403_FORBIDDEN)
443 def test_update_object(self):
444 """
445 PATCH a single object identified by its numeric ID.
446 """
447 instance = self._get_queryset().first()
448 url = self._get_detail_url(instance)
449 update_data = self.update_data or getattr(self, 'create_data')[0]
451 # Add object-level permission
452 obj_perm = ObjectPermission(
453 name='Test permission',
454 actions=['change']
455 )
456 obj_perm.save()
457 obj_perm.users.add(self.user)
458 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
460 data = copy.deepcopy(update_data)
462 # If supported, add a changelog message
463 if issubclass(self.model, ChangeLoggingMixin):
464 data['changelog_message'] = get_random_string(10)
466 response = self.client.patch(url, data, format='json', **self.header)
467 self.assertHttpStatus(response, status.HTTP_200_OK)
468 instance.refresh_from_db()
469 self.assertInstanceEqual(
470 instance,
471 data,
472 exclude=self.validation_excluded_fields,
473 api=True
474 )
476 # Verify ObjectChange creation
477 if hasattr(self.model, 'to_objectchange'):
478 objectchange = ObjectChange.objects.get(
479 changed_object_type=ContentType.objects.get_for_model(instance),
480 changed_object_id=instance.pk
481 )
482 self.assertObjectChange(objectchange, action=ObjectChangeActionChoices.ACTION_UPDATE,
483 message=data['changelog_message'])
485 def test_update_object_with_etag(self):
486 """
487 PATCH an object using a valid If-Match ETag → expect 200.
488 PATCH again with the now-stale ETag → expect 412.
489 """
490 if not issubclass(self.model, ChangeLoggingMixin):
491 self.skipTest("Model does not support ETags")
493 self.add_permissions(
494 f'{self.model._meta.app_label}.view_{self.model._meta.model_name}',
495 f'{self.model._meta.app_label}.change_{self.model._meta.model_name}',
496 )
497 instance = self._get_queryset().first()
498 url = self._get_detail_url(instance)
499 update_data = self.update_data or getattr(self, 'create_data')[0]
501 # Fetch current ETag
502 get_response = self.client.get(url, **self.header)
503 self.assertHttpStatus(get_response, status.HTTP_200_OK)
504 etag = get_response.get('ETag')
505 self.assertIsNotNone(etag, "No ETag returned by GET")
507 # PATCH with correct ETag → 200
508 response = self.client.patch(
509 url, update_data, format='json',
510 **{**self.header, 'HTTP_IF_MATCH': etag}
511 )
512 self.assertHttpStatus(response, status.HTTP_200_OK)
513 new_etag = response.get('ETag')
514 self.assertIsNotNone(new_etag)
515 self.assertNotEqual(etag, new_etag) # ETag must change after update
517 # PATCH with the old (stale) ETag → 412
518 with disable_warnings('django.request'):
519 response = self.client.patch(
520 url, update_data, format='json',
521 **{**self.header, 'HTTP_IF_MATCH': etag}
522 )
523 self.assertHttpStatus(response, status.HTTP_412_PRECONDITION_FAILED)
525 def test_bulk_update_objects(self):
526 """
527 PATCH a set of objects in a single request.
528 """
529 if self.bulk_update_data is None:
530 self.skipTest("Bulk update data not set")
532 # Add object-level permission
533 obj_perm = ObjectPermission(
534 name='Test permission',
535 actions=['change']
536 )
537 obj_perm.save()
538 obj_perm.users.add(self.user)
539 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
541 id_list = list(self._get_queryset().values_list('id', flat=True)[:3])
542 self.assertEqual(len(id_list), 3, "Insufficient number of objects to test bulk update")
543 data = [
544 {'id': id, **self.bulk_update_data} for id in id_list
545 ]
547 # If supported, add a changelog message
548 changelog_message = get_random_string(10)
549 if issubclass(self.model, ChangeLoggingMixin):
550 for obj_data in data:
551 obj_data['changelog_message'] = changelog_message
553 response = self.client.patch(self._get_list_url(), data, format='json', **self.header)
554 self.assertHttpStatus(response, status.HTTP_200_OK)
555 for i, obj in enumerate(response.data):
556 for field in self.bulk_update_data:
557 if field in ('changelog_message', 'add_tags', 'remove_tags'):
558 # Write-only field
559 continue
560 self.assertIn(field, obj, f"Bulk update field '{field}' missing from object {i} in response")
561 for instance in self._get_queryset().filter(pk__in=id_list):
562 self.assertInstanceEqual(instance, self.bulk_update_data, api=True)
564 # Verify ObjectChange creation
565 if issubclass(self.model, ChangeLoggingMixin):
566 objectchanges = ObjectChange.objects.filter(
567 changed_object_type=ContentType.objects.get_for_model(self.model),
568 changed_object_id__in=id_list
569 )
570 self.assertEqual(len(objectchanges), len(data))
571 for oc in objectchanges:
572 self.assertObjectChange(oc, action=ObjectChangeActionChoices.ACTION_UPDATE,
573 message=changelog_message)
575 def test_bulk_update_objects_string_id(self):
576 """
577 PATCH a set of objects whose IDs are given as strings rather than as numbers. The ID
578 field coerces such a value, so the object is identified and its attributes must be
579 applied -- rather than the entry being treated as though it carried no data.
580 """
581 if self.bulk_update_data is None:
582 self.skipTest("Bulk update data not set")
584 obj_perm = ObjectPermission(name='Test permission', actions=['change'])
585 obj_perm.save()
586 obj_perm.users.add(self.user)
587 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
589 id_list = list(self._get_queryset().values_list('id', flat=True)[:2])
590 self.assertEqual(len(id_list), 2, "Insufficient number of objects to test bulk update")
592 # Quote only the second ID, so that a batch mixing the two forms is covered as well
593 data = [
594 {'id': id_list[0], **self.bulk_update_data},
595 {'id': str(id_list[1]), **self.bulk_update_data},
596 ]
598 response = self.client.patch(self._get_list_url(), data, format='json', **self.header)
600 # The attributes must have been applied to both objects. Note that the response body is
601 # deliberately not inspected: for a model whose viewset narrows its own queryset (e.g.
602 # SavedFilter, which is restricted to shared or owned objects), an update which moves an
603 # object outside that queryset succeeds but is not echoed back.
604 self.assertHttpStatus(response, status.HTTP_200_OK)
605 for instance in self._get_queryset().filter(pk__in=id_list):
606 self.assertInstanceEqual(instance, self.bulk_update_data, api=True)
608 def test_bulk_update_objects_validation_error(self):
609 """
610 PATCH a set of objects where one fails validation. Verify the structured per-object error
611 response and that no objects are modified (atomic rollback).
612 """
613 if self.bulk_update_data is None or self.bulk_update_invalid_data is None:
614 self.skipTest('Bulk update data not set')
616 obj_perm = ObjectPermission(name='Test permission', actions=['change'])
617 obj_perm.save()
618 obj_perm.users.add(self.user)
619 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
621 id_list = list(self._get_queryset().values_list('id', flat=True)[:2])
622 self.assertEqual(len(id_list), 2, 'Insufficient number of objects to test bulk update validation error')
624 # First object: valid data; second: invalid data that must fail validation
625 data = [
626 {'id': id_list[0], **self.bulk_update_data},
627 {'id': id_list[1], **self.bulk_update_invalid_data},
628 ]
630 # Snapshot field values before the request so we can verify atomicity afterward
631 instance0_before = self._get_queryset().get(pk=id_list[0])
633 response = self.client.patch(self._get_list_url(), data, format='json', **self.header)
635 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
636 self.assertIn('detail', response.data)
637 self.assertIn('errors', response.data)
638 self.assertEqual(len(response.data['errors']), 1)
639 self.assertEqual(response.data['errors'][0]['id'], id_list[1])
640 self.assertIn('errors', response.data['errors'][0])
642 # Verify atomicity: object 0 passed validation but must not have been modified
643 instance0_after = self._get_queryset().get(pk=id_list[0])
644 for field in self.bulk_update_data:
645 if field in ('changelog_message', 'add_tags', 'remove_tags'):
646 continue
647 self.assertEqual(
648 getattr(instance0_after, field, None),
649 getattr(instance0_before, field, None),
650 f'Field {field!r} of object {id_list[0]} was modified — atomic rollback may be broken',
651 )
653 def test_bulk_update_objects_nonexistent_id(self):
654 """
655 PATCH a set of objects where one of the IDs does not identify an existing object. Verify
656 the structured per-object error response and that no objects are modified.
657 """
658 if self.bulk_update_data is None:
659 self.skipTest('Bulk update data not set')
661 obj_perm = ObjectPermission(name='Test permission', actions=['change'])
662 obj_perm.save()
663 obj_perm.users.add(self.user)
664 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
666 id_list = list(self._get_queryset().values_list('id', flat=True)[:2])
667 self.assertEqual(len(id_list), 2, 'Insufficient number of objects to test bulk update')
668 missing_id = self._get_queryset().order_by('-id').first().id + 1
670 data = [{'id': id, **self.bulk_update_data} for id in (*id_list, missing_id)]
672 # Snapshot the objects which would otherwise have been updated
673 instances_before = list(self._get_queryset().filter(pk__in=id_list))
675 response = self.client.patch(self._get_list_url(), data, format='json', **self.header)
677 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
678 self.assertIn('detail', response.data)
679 self.assertIn('errors', response.data)
680 self.assertEqual(len(response.data['errors']), 1)
681 self.assertEqual(response.data['errors'][0]['id'], missing_id)
682 self.assertIn('id', response.data['errors'][0]['errors'])
684 # The objects named alongside the missing one must not have been updated
685 for instance_before in instances_before:
686 instance_after = self._get_queryset().get(pk=instance_before.pk)
687 for field in self.bulk_update_data:
688 if field in ('changelog_message', 'add_tags', 'remove_tags'):
689 continue
690 self.assertEqual(
691 getattr(instance_after, field, None),
692 getattr(instance_before, field, None),
693 f'Field {field!r} of object {instance_before.pk} was modified despite an unresolvable '
694 f'sibling ID',
695 )
697 def test_bulk_update_objects_malformed_entry(self):
698 """
699 PATCH a set of objects in which one entry does not identify an object. The failure must be
700 reported in the same structured form as a per-object failure, correlated by position.
701 """
702 if self.bulk_update_data is None:
703 self.skipTest('Bulk update data not set')
705 obj_perm = ObjectPermission(name='Test permission', actions=['change'])
706 obj_perm.save()
707 obj_perm.users.add(self.user)
708 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
710 instance = self._get_queryset().first()
712 # The second entry omits the object ID, so it cannot be matched to an object
713 data = [{'id': instance.pk, **self.bulk_update_data}, {}]
715 response = self.client.patch(self._get_list_url(), data, format='json', **self.header)
717 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
718 self.assertIn('detail', response.data)
719 self.assertEqual(len(response.data['errors']), 1)
721 # Correlated by position, as no object was identified for this entry
722 self.assertEqual(response.data['errors'][0]['index'], 1)
723 self.assertIn('id', response.data['errors'][0]['errors'])
725 # The valid entry must not have been applied
726 instance_after = self._get_queryset().get(pk=instance.pk)
727 for field in self.bulk_update_data:
728 if field in ('changelog_message', 'add_tags', 'remove_tags'):
729 continue
730 self.assertEqual(
731 getattr(instance_after, field, None),
732 getattr(instance, field, None),
733 f'Field {field!r} of object {instance.pk} was modified despite a malformed sibling entry',
734 )
736 def test_bulk_update_objects_duplicate_id(self):
737 """
738 PATCH a set of objects in which the same object is named twice. The request must be
739 rejected rather than applying only one of the entries given for that object.
740 """
741 if self.bulk_update_data is None:
742 self.skipTest('Bulk update data not set')
744 obj_perm = ObjectPermission(name='Test permission', actions=['change'])
745 obj_perm.save()
746 obj_perm.users.add(self.user)
747 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
749 id_list = list(self._get_queryset().values_list('id', flat=True)[:2])
750 self.assertEqual(len(id_list), 2, 'Insufficient number of objects to test bulk update')
752 # Repeat the first ID at the end of the request
753 data = [{'id': id, **self.bulk_update_data} for id in (*id_list, id_list[0])]
755 # Snapshot the objects which would otherwise have been updated
756 instances_before = list(self._get_queryset().filter(pk__in=id_list))
758 response = self.client.patch(self._get_list_url(), data, format='json', **self.header)
760 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
761 self.assertIn('detail', response.data)
762 self.assertIn('errors', response.data)
764 # The repeated ID must be reported once, not once per occurrence
765 self.assertEqual(len(response.data['errors']), 1)
766 self.assertEqual(response.data['errors'][0]['id'], id_list[0])
767 self.assertIn('id', response.data['errors'][0]['errors'])
769 # No object named in the request may have been updated, including the one named only once
770 for instance_before in instances_before:
771 instance_after = self._get_queryset().get(pk=instance_before.pk)
772 for field in self.bulk_update_data:
773 if field in ('changelog_message', 'add_tags', 'remove_tags'):
774 continue
775 self.assertEqual(
776 getattr(instance_after, field, None),
777 getattr(instance_before, field, None),
778 f'Field {field!r} of object {instance_before.pk} was modified despite a duplicated '
779 f'sibling ID',
780 )
782 class DeleteObjectViewTestCase(APITestCase):
784 def test_delete_object_without_permission(self):
785 """
786 DELETE a single object without permission.
787 """
788 url = self._get_detail_url(self._get_queryset().first())
790 # Try DELETE without permission
791 with disable_warnings('django.request'):
792 response = self.client.delete(url, **self.header)
793 self.assertHttpStatus(response, status.HTTP_403_FORBIDDEN)
795 def test_delete_object(self):
796 """
797 DELETE a single object identified by its numeric ID.
798 """
799 instance = self._get_queryset().first()
800 url = self._get_detail_url(instance)
802 # Add object-level permission
803 obj_perm = ObjectPermission(
804 name='Test permission',
805 actions=['delete']
806 )
807 obj_perm.save()
808 obj_perm.users.add(self.user)
809 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
811 data = {}
813 # If supported, add a changelog message
814 if issubclass(self.model, ChangeLoggingMixin):
815 data['changelog_message'] = get_random_string(10)
817 response = self.client.delete(url, data, **self.header)
818 self.assertHttpStatus(response, status.HTTP_204_NO_CONTENT)
819 self.assertFalse(self._get_queryset().filter(pk=instance.pk).exists())
821 # Verify ObjectChange creation
822 if hasattr(self.model, 'to_objectchange'):
823 objectchange = ObjectChange.objects.get(
824 changed_object_type=ContentType.objects.get_for_model(instance),
825 changed_object_id=instance.pk
826 )
827 self.assertObjectChange(objectchange, action=ObjectChangeActionChoices.ACTION_DELETE,
828 message=data['changelog_message'])
830 def test_bulk_delete_objects(self):
831 """
832 DELETE a set of objects in a single request.
833 """
834 # Add object-level permission
835 obj_perm = ObjectPermission(
836 name='Test permission',
837 actions=['delete']
838 )
839 obj_perm.save()
840 obj_perm.users.add(self.user)
841 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
843 # Target the three most recently created objects to avoid triggering recursive deletions
844 # (e.g. with MPTT objects)
845 id_list = list(self._get_queryset().order_by('-id').values_list('id', flat=True)[:3])
846 self.assertEqual(len(id_list), 3, "Insufficient number of objects to test bulk deletion")
847 data = [{"id": id} for id in id_list]
849 # If supported, add a changelog message
850 changelog_message = get_random_string(10)
851 if issubclass(self.model, ChangeLoggingMixin):
852 for obj_data in data:
853 obj_data['changelog_message'] = changelog_message
855 initial_count = self._get_queryset().count()
856 response = self.client.delete(self._get_list_url(), data, format='json', **self.header)
857 self.assertHttpStatus(response, status.HTTP_204_NO_CONTENT)
858 self.assertEqual(self._get_queryset().count(), initial_count - 3)
860 # Verify ObjectChange creation
861 if issubclass(self.model, ChangeLoggingMixin):
862 objectchanges = ObjectChange.objects.filter(
863 changed_object_type=ContentType.objects.get_for_model(self.model),
864 changed_object_id__in=id_list
865 )
866 self.assertEqual(len(objectchanges), len(data))
867 for oc in objectchanges:
868 self.assertObjectChange(oc, action=ObjectChangeActionChoices.ACTION_DELETE,
869 message=changelog_message)
871 def test_bulk_delete_objects_nonexistent_id(self):
872 """
873 DELETE a set of objects where one of the IDs does not identify an existing object. Verify
874 the structured per-object error response and that no objects are deleted.
875 """
876 obj_perm = ObjectPermission(
877 name='Test permission',
878 actions=['delete']
879 )
880 obj_perm.save()
881 obj_perm.users.add(self.user)
882 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
884 # Target the most recently created objects to avoid triggering recursive deletions
885 id_list = list(self._get_queryset().order_by('-id').values_list('id', flat=True)[:3])
886 self.assertEqual(len(id_list), 3, 'Insufficient number of objects to test bulk deletion')
887 missing_id = max(id_list) + 1
888 data = [{'id': id} for id in (*id_list, missing_id)]
890 initial_count = self._get_queryset().count()
891 response = self.client.delete(self._get_list_url(), data, format='json', **self.header)
893 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
894 self.assertIn('detail', response.data)
895 self.assertIn('errors', response.data)
896 self.assertEqual(len(response.data['errors']), 1)
897 self.assertEqual(response.data['errors'][0]['id'], missing_id)
898 self.assertIn('id', response.data['errors'][0]['errors'])
900 # The objects named alongside the missing one must not have been deleted
901 self.assertEqual(self._get_queryset().count(), initial_count)
903 def test_bulk_delete_objects_malformed_entry(self):
904 """
905 DELETE a set of objects in which one entry does not identify an object. The failure must be
906 reported in the same structured form as a per-object failure, correlated by position.
907 """
908 obj_perm = ObjectPermission(
909 name='Test permission',
910 actions=['delete']
911 )
912 obj_perm.save()
913 obj_perm.users.add(self.user)
914 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
916 # Target the most recently created object to avoid triggering recursive deletions
917 instance = self._get_queryset().order_by('-id').first()
919 # The second entry omits the object ID, so it cannot be matched to an object
920 data = [{'id': instance.pk}, {}]
922 initial_count = self._get_queryset().count()
923 response = self.client.delete(self._get_list_url(), data, format='json', **self.header)
925 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
926 self.assertIn('detail', response.data)
927 self.assertEqual(len(response.data['errors']), 1)
929 # Correlated by position, as no object was identified for this entry
930 self.assertEqual(response.data['errors'][0]['index'], 1)
931 self.assertIn('id', response.data['errors'][0]['errors'])
933 # Nothing may have been deleted
934 self.assertEqual(self._get_queryset().count(), initial_count)
936 def test_bulk_delete_objects_duplicate_id(self):
937 """
938 DELETE a set of objects in which the same object is named twice. The request must be
939 rejected rather than reporting success for a batch it only partly acted on.
940 """
941 obj_perm = ObjectPermission(
942 name='Test permission',
943 actions=['delete']
944 )
945 obj_perm.save()
946 obj_perm.users.add(self.user)
947 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
949 # Target the most recently created objects to avoid triggering recursive deletions
950 id_list = list(self._get_queryset().order_by('-id').values_list('id', flat=True)[:2])
951 self.assertEqual(len(id_list), 2, 'Insufficient number of objects to test bulk deletion')
953 # Repeat the first ID at the end of the request
954 data = [{'id': id} for id in (*id_list, id_list[0])]
956 initial_count = self._get_queryset().count()
957 response = self.client.delete(self._get_list_url(), data, format='json', **self.header)
959 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
960 self.assertIn('detail', response.data)
961 self.assertIn('errors', response.data)
963 # The repeated ID must be reported once, not once per occurrence
964 self.assertEqual(len(response.data['errors']), 1)
965 self.assertEqual(response.data['errors'][0]['id'], id_list[0])
966 self.assertIn('id', response.data['errors'][0]['errors'])
968 # No object named in the request may have been deleted
969 self.assertEqual(self._get_queryset().count(), initial_count)
971 def test_bulk_delete_objects_no_body(self):
972 """
973 DELETE a list endpoint with no body at all. Nothing may be deleted -- the request names no
974 objects, so it cannot mean "all of them" -- and the response must say so intelligibly.
975 """
976 obj_perm = ObjectPermission(
977 name='Test permission',
978 actions=['delete']
979 )
980 obj_perm.save()
981 obj_perm.users.add(self.user)
982 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
984 initial_count = self._get_queryset().count()
985 self.assertNotEqual(initial_count, 0, 'No objects exist against which to test bulk deletion')
987 response = self.client.delete(self._get_list_url(), **self.header)
989 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
990 self.assertIn('detail', response.data)
991 # There are no entries to report against, so no per-object errors are returned
992 self.assertNotIn('errors', response.data)
993 self.assertEqual(
994 self._get_queryset().count(), initial_count,
995 'A bulk delete naming no objects must not delete anything'
996 )
998 class GraphQLTestCase(APITestCase):
999 graphql_auto_filter_tests = True
1000 graphql_auto_filter_exclude = ()
1002 # Cap fields per lookup kind to keep test counts balanced across kinds
1003 # (string fields shouldn't crowd out numeric/date/array fields).
1004 graphql_auto_filter_fields_per_kind = 2
1006 # Fail when auto mode is on and no tests were generated.
1007 graphql_auto_filter_required = True
1009 # Gate the negative constrained-permission check in the get/list tests; the positive
1010 # query still runs. Set False for types not enforcing object permissions (e.g. no BaseObjectType).
1011 graphql_object_permission_assertions = True
1013 # Additional explicit-list filter cases as GraphQLFilterTest instances.
1014 graphql_filter_tests = ()
1016 # Additional full-query cases (e.g. nested filters) as GraphQLQueryTest instances.
1017 graphql_query_tests = ()
1019 # GraphQL type under test. Defaults to the type derived from `model` via the naming
1020 # convention; set explicitly when the convention does not apply (e.g. plugin types).
1021 type_class = None
1023 # Exclude this test case from GraphQL schema coverage.
1024 graphql_test_exempt = False
1026 @classmethod
1027 def get_graphql_type_class(cls):
1028 if getattr(cls, 'type_class', None) is not None:
1029 return cls.type_class
1030 model = getattr(cls, 'model', None)
1031 if model is None:
1032 return None
1033 return get_graphql_type_for_model(model)
1035 def _get_graphql_base_name(self):
1036 """
1037 Return graphql_base_name, if set. Otherwise, construct the base name for the query
1038 field from the model's verbose name.
1039 """
1040 base_name = self.model._meta.verbose_name.lower().replace(' ', '_')
1041 return getattr(self, 'graphql_base_name', base_name)
1043 def _build_query_with_filter(self, name, filter_string):
1044 """
1045 Called by either _build_query or _build_filtered_query - construct the actual
1046 query given a name and filter string
1047 """
1048 type_class = self.get_graphql_type_class()
1050 # Compile list of fields to include
1051 fields_string = ''
1053 file_fields = (
1054 strawberry_django.fields.types.DjangoFileType,
1055 strawberry_django.fields.types.DjangoImageType,
1056 )
1057 for field in type_class.__strawberry_definition__.fields:
1058 if (
1059 field.type in file_fields or (
1060 type(field.type) is StrawberryOptional and field.type.of_type in file_fields
1061 )
1062 ):
1063 # image / file fields nullable or not...
1064 fields_string += f'{field.name} {{ name }}\n'
1065 elif type(field.type) is StrawberryList and type(field.type.of_type) is LazyType:
1066 # List of related objects (queryset)
1067 fields_string += f'{field.name} {{ id }}\n'
1068 elif type(field.type) is StrawberryList and type(field.type.of_type) is StrawberryUnion:
1069 # this would require a fragment query
1070 continue
1071 elif type(field.type) is StrawberryUnion:
1072 # this would require a fragment query
1073 continue
1074 elif type(field.type) is StrawberryOptional and type(field.type.of_type) is StrawberryUnion:
1075 # this would require a fragment query
1076 continue
1077 elif type(field.type) is StrawberryOptional and type(field.type.of_type) is LazyType:
1078 fields_string += f'{field.name} {{ id }}\n'
1079 elif hasattr(field, 'is_relation') and field.is_relation:
1080 # Ignore private fields
1081 if field.name.startswith('_'):
1082 continue
1083 # Note: StrawberryField types do not have is_relation
1084 fields_string += f'{field.name} {{ id }}\n'
1085 elif inspect.isclass(field.type) and issubclass(field.type, IPAddressFamilyType):
1086 fields_string += f'{field.name} {{ value, label }}\n'
1087 else:
1088 fields_string += f'{field.name}\n'
1090 query = f"""
1091 {{
1092 {name}{filter_string} {{
1093 {fields_string}
1094 }}
1095 }}
1096 """
1098 return query
1100 @staticmethod
1101 def _graphql_literal(value):
1102 """
1103 Render a Python value as a GraphQL literal.
1104 """
1105 if value is None:
1106 return 'null'
1107 if isinstance(value, bool):
1108 return 'true' if value else 'false'
1109 if isinstance(value, (int, float)):
1110 return str(value)
1111 if isinstance(value, Decimal):
1112 return str(float(value))
1113 if isinstance(value, (list, tuple)):
1114 items = ', '.join(
1115 APIViewTestCases.GraphQLTestCase._graphql_literal(v) for v in value
1116 )
1117 return f'[{items}]'
1118 if isinstance(value, str):
1119 return json.dumps(value)
1121 return json.dumps(str(value))
1123 def _render_graphql_filter_value(self, params):
1124 """
1125 Render the legacy graphql_filter dict value to a GraphQL filter value.
1126 """
1127 if isinstance(params, str):
1128 return params
1130 if not isinstance(params, dict):
1131 return self._graphql_literal(params)
1133 lookup = params.get('lookup')
1134 value = params['value']
1136 if lookup:
1137 return f'{{{lookup}: {self._graphql_literal(value)}}}'
1139 return self._graphql_literal(value)
1141 def _build_graphql_filter_string(self, **filters):
1142 if not filters:
1143 return ''
1145 filter_expressions = [
1146 f'{field_name}: {self._render_graphql_filter_value(params)}'
1147 for field_name, params in filters.items()
1148 ]
1150 return f'(filters: {{{", ".join(filter_expressions)}}})'
1152 def _build_filtered_query(self, name, **filters):
1153 """
1154 Create a filtered query: i.e. device_list(filters: {name: {i_contains: "akron"}}){.
1155 """
1156 filter_string = self._build_graphql_filter_string(**filters)
1158 return self._build_query_with_filter(name, filter_string)
1160 def _build_graphql_id_list_query(self, name, filters):
1161 filter_string = f'(filters: {{{filters}}})' if filters else ''
1162 selection = 'id' if self._graphql_type_exposes_id() else '__typename'
1164 return f"""
1165 {{
1166 {name}{filter_string} {{
1167 {selection}
1168 }}
1169 }}
1170 """
1172 def _graphql_type_exposes_id(self):
1173 """
1174 Return True when the model's GraphQL type exposes ``id`` as a
1175 queryable selection. Some NetBox types (e.g. Notification,
1176 Subscription) omit ``id`` from the output type; for those, the
1177 assertion path falls back to length-only comparison.
1178 """
1179 type_class = self.get_graphql_type_class()
1180 strawberry_definition = getattr(type_class, '__strawberry_definition__', None)
1181 if strawberry_definition is None:
1182 return False
1183 return any(field.name == 'id' for field in strawberry_definition.fields)
1185 def _get_model_graphql_filter_class(self, model=None):
1186 """
1187 Return the model's GraphQL filter class, if one follows NetBox's
1188 conventional <app>.graphql.filters.<Model>Filter path. ``None`` if
1189 the filter module (or any of its parent packages) is absent or the
1190 class is not present in the module. Import errors originating
1191 inside an existing filter module are re-raised.
1192 """
1193 model = model or self.model
1194 module_path = f'{model._meta.app_label}.graphql.filters'
1195 class_name = f'{model.__name__}Filter'
1197 try:
1198 module = importlib.import_module(module_path)
1199 except ModuleNotFoundError as exc:
1200 # Treat both "<app>.graphql.filters" absent and any missing
1201 # parent (e.g. "<app>.graphql" or "<app>") as "no conventional
1202 # filter class". Real ImportErrors from inside an existing
1203 # filter module still propagate.
1204 if exc.name == module_path or module_path.startswith(f'{exc.name}.'):
1205 return None
1206 raise
1208 return getattr(module, class_name, None)
1210 def _get_graphql_filter_field_names(self):
1211 """
1212 Return the names exposed by the model's GraphQL filter input, sourced
1213 only from the conventional <app>.graphql.filters.<Model>Filter path.
1214 """
1215 filter_class = self._get_model_graphql_filter_class()
1216 if filter_class is None:
1217 return set()
1219 return self._collect_filter_class_annotation_names(filter_class)
1221 @staticmethod
1222 def _collect_filter_class_annotation_names(filter_class):
1223 field_names = set()
1224 for cls in reversed(getattr(filter_class, '__mro__', ())):
1225 field_names.update(
1226 field_name for field_name in getattr(cls, '__annotations__', {})
1227 if not field_name.startswith('_')
1228 )
1229 return field_names
1231 def _assert_graphql_filter_class_present(self, filter_fields, handwritten_tests=()):
1232 """
1233 Raise when the model has no discoverable filter class or the class
1234 declares no fields. Skipped when auto-filter generation is disabled,
1235 the per-model opt-out attribute is set, or hand-written (legacy or
1236 explicit) filter tests are declared for the model.
1237 """
1238 if handwritten_tests:
1239 return
1240 if not getattr(self, 'graphql_auto_filter_required', True):
1241 return
1242 if not getattr(self, 'graphql_auto_filter_tests', True):
1243 return
1245 label = self.model._meta.label
1246 path = f'{self.model._meta.app_label}.graphql.filters.{self.model.__name__}Filter'
1248 filter_class = self._get_model_graphql_filter_class()
1249 self.assertIsNotNone(
1250 filter_class,
1251 f'No GraphQL filter class found for {label} at {path}. '
1252 f'Set graphql_auto_filter_required = False on this test case if intentional.'
1253 )
1254 self.assertTrue(
1255 filter_fields,
1256 f'GraphQL filter class for {label} declares no fields. '
1257 f'Set graphql_auto_filter_required = False on this test case if intentional.'
1258 )
1260 def _get_nonempty_field_value(self, field):
1261 queryset = self._get_queryset()
1263 if getattr(field, 'null', False):
1264 queryset = queryset.exclude(**{f'{field.name}__isnull': True})
1266 if isinstance(field, (models.CharField, models.TextField)):
1267 queryset = queryset.exclude(**{field.name: ''})
1269 return queryset.values_list(field.name, flat=True).first()
1271 def _get_model_field_for_filter_field(self, field_name):
1272 """
1273 Find the Django model field matching a filter field name. Filter
1274 fields are declared with either the model field name (e.g. `name`)
1275 or the FK attname (e.g. `tenant_id`).
1276 """
1277 for field in self.model._meta.fields:
1278 if field.name == field_name or getattr(field, 'attname', None) == field_name:
1279 return field
1280 return None
1282 def _iter_filter_class_annotations(self, filter_class):
1283 """
1284 Yield (field_name, annotation) pairs for the filter class, walking
1285 its MRO so inherited fields surface. Subclass annotations override
1286 inherited ones (private `_`-prefixed names are skipped).
1287 """
1288 annotations = {}
1289 for cls in reversed(filter_class.__mro__):
1290 annotations.update({
1291 name: ann for name, ann in getattr(cls, '__annotations__', {}).items()
1292 if not name.startswith('_')
1293 })
1294 yield from annotations.items()
1296 @staticmethod
1297 def _unwrap_filter_annotation(annotation):
1298 """
1299 Strip ``X | None`` / ``Optional[X]`` and ``Annotated[X, ...]``
1300 layers. Resolve `strawberry.lazy('...')` metadata so lazily-annotated
1301 lookup types (e.g. ``Annotated['FloatLookup', strawberry.lazy('mod')] | None``)
1302 are returned as the actual class. When an ``Annotated`` layer carries
1303 multiple metadata entries, the first ``module``-bearing entry wins.
1304 Returns None when the inner type cannot be resolved.
1305 """
1306 if annotation is None:
1307 return None
1309 lazy_module = None
1310 lazy_package = None
1311 # Cap iterations at 8: typical NetBox annotations nest at most 3 layers
1312 # (Union > Annotated > ForwardRef). 8 is a generous safety net to
1313 # prevent infinite loops on pathological / future annotation shapes.
1314 for _ in range(8):
1315 origin = typing.get_origin(annotation)
1316 args = typing.get_args(annotation)
1318 if origin in (typing.Union, types.UnionType):
1319 non_none = [a for a in args if a is not type(None)]
1320 if len(non_none) != 1:
1321 return None
1322 annotation = non_none[0]
1323 continue
1325 if hasattr(annotation, '__metadata__'):
1326 for meta in annotation.__metadata__:
1327 module_name = getattr(meta, 'module', None)
1328 if module_name:
1329 lazy_module = module_name
1330 # strawberry.lazy('.relative') records the anchor package
1331 # needed to resolve the leading-dot module path.
1332 lazy_package = getattr(meta, 'package', None)
1333 break
1334 inner = args[0] if args else None
1335 if inner is None:
1336 return None
1337 annotation = inner
1338 continue
1340 break
1342 if isinstance(annotation, (str, typing.ForwardRef)):
1343 if lazy_module is None:
1344 return None
1345 name = annotation.__forward_arg__ if isinstance(annotation, typing.ForwardRef) else annotation
1346 # Resolve via import_module(module, package) rather than import_string()
1347 # so relative lazy modules (e.g. strawberry.lazy('.filters')) resolve
1348 # against their anchor package, as strawberry itself does.
1349 try:
1350 module = importlib.import_module(lazy_module, lazy_package)
1351 return getattr(module, name)
1352 except (ImportError, AttributeError):
1353 return None
1355 return annotation
1357 @classmethod
1358 def _classify_filter_annotation(cls, annotation):
1359 """
1360 Resolve a filter field annotation to a (kind, kind_arg) tuple keyed
1361 on the declared GraphQL lookup type. Returns (None, None) for
1362 annotations the dispatcher does not handle (those fields are
1363 silently skipped).
1364 """
1365 annotation = cls._unwrap_filter_annotation(annotation)
1366 if annotation is None or isinstance(annotation, str):
1367 return None, None
1369 if annotation is strawberry.ID:
1370 return 'id', None
1372 origin = typing.get_origin(annotation)
1373 target = origin if isinstance(origin, type) else annotation
1374 type_args = typing.get_args(annotation)
1376 if not isinstance(target, type):
1377 return None, None
1379 if target in (IntegerLookup, BigIntegerLookup, FloatLookup):
1380 return 'numeric', target
1382 # TreeNodeFilter schema requires {id, match_type}; skip auto-emit.
1383 if target is TreeNodeFilter:
1384 return None, None
1386 if issubclass(target, (DateFilterLookup, DatetimeFilterLookup, TimeFilterLookup)):
1387 return 'date_lookup', None
1389 if target is RangeLookup or issubclass(target, RangeLookup):
1390 return 'range_lookup', type_args[0] if type_args else None
1392 if issubclass(target, ArrayLookup):
1393 return 'array_lookup', None
1394 if target is IntegerRangeArrayLookup or issubclass(target, IntegerRangeArrayLookup):
1395 return 'range_array_lookup', None
1396 if target is JSONFilter:
1397 # JSONFilter requires explicit (path, typed lookup); no general auto shape.
1398 return None, None
1400 if issubclass(target, StrFilterLookup):
1401 return 'str_lookup', None
1402 if issubclass(target, ComparisonFilterLookup):
1403 return 'comparison_lookup', type_args[0] if type_args else None
1404 if issubclass(target, FilterLookup):
1405 return 'filter_lookup', type_args[0] if type_args else None
1406 # Enum-typed BaseFilterLookup needs an enum literal; skip auto-emit.
1407 if issubclass(target, BaseFilterLookup):
1408 return None, None
1410 return None, None
1412 def _emit_id_filter_tests(self, field_name, _kind_arg):
1413 if field_name == 'id':
1414 instance = self._get_queryset().first()
1415 if instance is None:
1416 return
1417 yield GraphQLFilterTest(
1418 name='id__exact',
1419 filters=f'id: {self._graphql_literal(str(instance.pk))}',
1420 expected=lambda qs, pk=instance.pk: qs.filter(pk=pk),
1421 )
1422 return
1424 model_field = self._get_model_field_for_filter_field(field_name)
1425 if model_field is None or not isinstance(model_field, models.ForeignKey):
1426 return
1427 queryset = self._get_queryset().exclude(**{f'{model_field.name}__isnull': True})
1428 value = queryset.values_list(model_field.attname, flat=True).first()
1429 if value is None:
1430 return
1431 yield GraphQLFilterTest(
1432 name=f'{field_name}__exact',
1433 filters=f'{field_name}: {self._graphql_literal(str(value))}',
1434 expected=lambda qs, attname=model_field.attname, v=value: qs.filter(**{attname: v}),
1435 )
1437 def _emit_str_lookup_filter_tests(self, field_name, _kind_arg):
1438 model_field = self._get_model_field_for_filter_field(field_name)
1439 if model_field is None:
1440 return
1441 value = self._get_nonempty_field_value(model_field)
1442 if value in (None, ''):
1443 return
1444 value = str(value)
1445 token = max(1, min(3, len(value)))
1446 lookups = (
1447 ('exact', 'exact', value),
1448 ('i_contains', 'icontains', value[:token]),
1449 ('i_starts_with', 'istartswith', value[:token]),
1450 ('i_ends_with', 'iendswith', value[-token:]),
1451 )
1452 for lookup, orm_lookup, filter_value in lookups:
1453 yield GraphQLFilterTest(
1454 name=f'{field_name}__{lookup}',
1455 filters=f'{field_name}: {{{lookup}: {self._graphql_literal(filter_value)}}}',
1456 expected=(
1457 lambda qs, fn=model_field.name, ol=orm_lookup, v=filter_value:
1458 qs.filter(**{f'{fn}__{ol}': v})
1459 ),
1460 )
1462 def _emit_filter_lookup_filter_tests(self, field_name, type_arg):
1463 model_field = self._get_model_field_for_filter_field(field_name)
1464 if model_field is None:
1465 return
1466 value = self._get_nonempty_field_value(model_field)
1467 if value is None:
1468 return
1469 if type_arg is bool or isinstance(value, bool):
1470 yield GraphQLFilterTest(
1471 name=f'{field_name}__exact',
1472 filters=f'{field_name}: {{exact: {self._graphql_literal(value)}}}',
1473 expected=lambda qs, fn=model_field.name, v=value: qs.filter(**{fn: v}),
1474 )
1475 return
1476 yield GraphQLFilterTest(
1477 name=f'{field_name}__exact',
1478 filters=f'{field_name}: {{exact: {self._graphql_literal(value)}}}',
1479 expected=lambda qs, fn=model_field.name, v=value: qs.filter(**{f'{fn}__exact': v}),
1480 )
1482 def _emit_comparison_lookup_filter_tests(self, field_name, _type_arg):
1483 model_field = self._get_model_field_for_filter_field(field_name)
1484 if model_field is None:
1485 return
1486 value = self._get_nonempty_field_value(model_field)
1487 if value is None:
1488 return
1489 yield GraphQLFilterTest(
1490 name=f'{field_name}__exact',
1491 filters=f'{field_name}: {{exact: {self._graphql_literal(value)}}}',
1492 expected=lambda qs, fn=model_field.name, v=value: qs.filter(**{f'{fn}__exact': v}),
1493 )
1495 def _emit_numeric_filter_tests(self, field_name, _type_arg):
1496 # NetBox numeric wrapper: {filter_lookup: {exact: N}}.
1497 model_field = self._get_model_field_for_filter_field(field_name)
1498 if model_field is None:
1499 return
1500 if isinstance(model_field, ArrayField):
1501 return
1502 value = self._get_nonempty_field_value(model_field)
1503 if value is None:
1504 return
1505 if isinstance(value, Decimal):
1506 value = float(value)
1507 yield GraphQLFilterTest(
1508 name=f'{field_name}__filter_lookup__exact',
1509 filters=(
1510 f'{field_name}: {{filter_lookup: '
1511 f'{{exact: {self._graphql_literal(value)}}}}}'
1512 ),
1513 expected=lambda qs, fn=model_field.name, v=value: qs.filter(**{f'{fn}__exact': v}),
1514 )
1516 def _emit_date_lookup_filter_tests(self, field_name, _kind_arg):
1517 model_field = self._get_model_field_for_filter_field(field_name)
1518 if model_field is None:
1519 return
1520 value = self._get_nonempty_field_value(model_field)
1521 if value is None:
1522 return
1523 iso_value = value.isoformat() if hasattr(value, 'isoformat') else str(value)
1524 yield GraphQLFilterTest(
1525 name=f'{field_name}__exact',
1526 filters=f'{field_name}: {{exact: "{iso_value}"}}',
1527 expected=lambda qs, fn=model_field.name, v=value: qs.filter(**{fn: v}),
1528 )
1530 def _emit_range_lookup_filter_tests(self, field_name, _kind_arg):
1531 model_field = self._get_model_field_for_filter_field(field_name)
1532 if model_field is None:
1533 return
1534 aggregates = self._get_queryset().aggregate(
1535 _min=models.Min(model_field.name), _max=models.Max(model_field.name),
1536 )
1537 start, end = aggregates['_min'], aggregates['_max']
1538 if start is None or end is None or start == end:
1539 return
1540 yield GraphQLFilterTest(
1541 name=f'{field_name}__range_lookup',
1542 filters=(
1543 f'{field_name}: {{range_lookup: '
1544 f'{{start: {self._graphql_literal(start)}, end: {self._graphql_literal(end)}}}}}'
1545 ),
1546 expected=(
1547 lambda qs, fn=model_field.name, lo=start, hi=end:
1548 qs.filter(**{f'{fn}__gte': lo, f'{fn}__lte': hi})
1549 ),
1550 )
1552 def _emit_array_lookup_filter_tests(self, field_name, _kind_arg):
1553 model_field = self._get_model_field_for_filter_field(field_name)
1554 if model_field is None:
1555 return
1556 if not isinstance(model_field, ArrayField):
1557 return
1558 queryset = self._get_queryset().exclude(**{field_name: []})
1559 sample = queryset.values_list(field_name, flat=True).first()
1560 if not sample:
1561 return
1562 element = sample[0]
1563 yield GraphQLFilterTest(
1564 name=f'{field_name}__contains',
1565 filters=(
1566 f'{field_name}: {{contains: [{self._graphql_literal(element)}]}}'
1567 ),
1568 expected=(
1569 lambda qs, fn=model_field.name, v=element: qs.filter(**{f'{fn}__contains': [v]})
1570 ),
1571 )
1573 def _emit_range_array_lookup_filter_tests(self, field_name, _kind_arg):
1574 model_field = self._get_model_field_for_filter_field(field_name)
1575 if model_field is None:
1576 return
1577 queryset = self._get_queryset().exclude(**{f'{field_name}__isnull': True})
1578 sample = queryset.values_list(field_name, flat=True).first()
1579 if not sample:
1580 return
1581 first_range = sample[0]
1582 lower = getattr(first_range, 'lower', None)
1583 if lower is None:
1584 return
1585 yield GraphQLFilterTest(
1586 name=f'{field_name}__contains',
1587 filters=f'{field_name}: {{contains: {self._graphql_literal(lower)}}}',
1588 expected=(
1589 lambda qs, fn=model_field.name, v=lower: qs.filter(**{f'{fn}__range_contains': v})
1590 ),
1591 )
1593 def _iter_auto_graphql_filter_tests(self):
1594 if not getattr(self, 'graphql_auto_filter_tests', True):
1595 return
1597 filter_class = self._get_model_graphql_filter_class()
1598 if filter_class is None:
1599 return
1601 exclude = set(getattr(self, 'graphql_auto_filter_exclude', ()))
1602 per_kind = self.graphql_auto_filter_fields_per_kind
1604 # Bucket eligible fields by lookup kind so per-kind budgeting balances coverage.
1605 by_kind: dict[str, list[tuple[str, object]]] = {}
1606 for field_name, annotation in self._iter_filter_class_annotations(filter_class):
1607 if field_name in exclude:
1608 continue
1609 kind, kind_arg = self._classify_filter_annotation(annotation)
1610 if kind is None:
1611 continue
1612 by_kind.setdefault(kind, []).append((field_name, kind_arg))
1614 # Emit per-kind; the cap counts SUCCESSFUL emissions, not candidate fields, so
1615 # early null/empty fields don't shadow later fields with usable fixture data.
1616 for kind, fields in by_kind.items():
1617 emitter = getattr(self, f'_emit_{kind}_filter_tests', None)
1618 if emitter is None:
1619 continue
1621 emitted_fields = 0
1622 for field_name, kind_arg in fields:
1623 tests = list(emitter(field_name, kind_arg))
1624 if not tests:
1625 continue
1626 yield from tests
1627 emitted_fields += 1
1628 if emitted_fields >= per_kind:
1629 break
1631 def _iter_legacy_graphql_filter_tests(self):
1632 if not hasattr(self, 'graphql_filter'):
1633 return
1635 filter_expressions = [
1636 f'{field_name}: {self._render_graphql_filter_value(params)}'
1637 for field_name, params in self.graphql_filter.items()
1638 ]
1640 yield GraphQLFilterTest(
1641 name='graphql_filter',
1642 filters=', '.join(filter_expressions),
1643 )
1645 def _coerce_graphql_filter_test(self, filter_test):
1646 if isinstance(filter_test, GraphQLFilterTest):
1647 return filter_test
1649 filter_test = dict(filter_test)
1650 if 'filter' in filter_test and 'filters' not in filter_test:
1651 filter_test['filters'] = filter_test.pop('filter')
1653 return GraphQLFilterTest(**filter_test)
1655 def _iter_explicit_graphql_filter_tests(self):
1656 for filter_test in getattr(self, 'graphql_filter_tests', ()):
1657 yield self._coerce_graphql_filter_test(filter_test)
1659 def _get_expected_id_set(self, filter_test):
1660 expected = filter_test.expected
1662 if callable(expected):
1663 expected = expected(self._get_queryset())
1665 if isinstance(expected, dict):
1666 expected = self._get_queryset().filter(**expected)
1668 if hasattr(expected, 'values_list'):
1669 values = expected.distinct().values_list('pk', flat=True)
1670 else:
1671 values = [getattr(value, 'pk', value) for value in expected]
1673 return {str(value) for value in values}
1675 def _assert_graphql_filter_test(self, url, field_name, filter_test):
1676 query = self._build_graphql_id_list_query(field_name, filter_test.filters)
1678 for permission in filter_test.permissions:
1679 self.add_permissions(permission)
1681 response = self.client.post(url, data={'query': query}, format="json", **self.header)
1682 self.assertHttpStatus(response, status.HTTP_200_OK)
1684 data = json.loads(response.content)
1685 self.assertNotIn('errors', data)
1687 results = data['data'][field_name]
1689 if filter_test.expected is None:
1690 self.assertGreater(len(results), 0)
1691 return
1693 expected_ids = self._get_expected_id_set(filter_test)
1695 self.assertGreater(
1696 len(expected_ids), 0,
1697 msg=(
1698 f'{self.model._meta.label}: filter "{filter_test.name}" produced an empty '
1699 f'expected set; the test would tautologically pass. Adjust fixtures or the '
1700 f'filter so the expected ORM queryset is non-empty.'
1701 ),
1702 )
1704 if self._graphql_type_exposes_id():
1705 result_ids = [str(result['id']) for result in results]
1706 self.assertEqual(
1707 set(result_ids), expected_ids,
1708 msg=f'{self.model._meta.label}: filter "{filter_test.name}" ID set mismatch',
1709 )
1711 self.assertEqual(
1712 len(results), len(expected_ids),
1713 msg=(
1714 f'{self.model._meta.label}: filter "{filter_test.name}" result count mismatch '
1715 f'(GraphQL type does not expose id; comparing by length).'
1716 ),
1717 )
1719 def _coerce_graphql_query_test(self, query_test):
1720 if isinstance(query_test, GraphQLQueryTest):
1721 return query_test
1723 query_test = dict(query_test)
1724 if 'assertion' in query_test and 'assert_result' not in query_test:
1725 query_test['assert_result'] = query_test.pop('assertion')
1727 return GraphQLQueryTest(**query_test)
1729 def _build_query(self, name, **filters):
1730 """
1731 Create a normal query - unfiltered or with a string query: i.e. site(name: "aaa"){.
1732 """
1733 if filters:
1734 filter_string = ', '.join(f'{k}:{v}' for k, v in filters.items())
1735 filter_string = f'({filter_string})'
1736 else:
1737 filter_string = ''
1739 return self._build_query_with_filter(name, filter_string)
1741 @override_settings(LOGIN_REQUIRED=True)
1742 def test_graphql_get_object(self):
1743 url = reverse('graphql')
1744 field_name = self._get_graphql_base_name()
1745 object_id = self._get_queryset().first().pk
1746 query = self._build_query(field_name, id=object_id)
1748 # Non-authenticated requests should fail
1749 header = {
1750 'HTTP_ACCEPT': 'application/json',
1751 }
1752 with disable_warnings('django.request'):
1753 response = self.client.post(url, data={'query': query}, format="json", **header)
1754 self.assertHttpStatus(response, status.HTTP_403_FORBIDDEN)
1756 # Add constrained permission
1757 obj_perm = ObjectPermission(
1758 name='Test permission',
1759 actions=['view'],
1760 constraints={'id': 0} # Impossible constraint
1761 )
1762 obj_perm.save()
1763 obj_perm.users.add(self.user)
1764 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
1766 if self.graphql_object_permission_assertions:
1767 # Request should succeed but return empty result
1768 with disable_logging():
1769 response = self.client.post(url, data={'query': query}, format="json", **self.header)
1770 self.assertHttpStatus(response, status.HTTP_200_OK)
1771 data = json.loads(response.content)
1772 self.assertIn('errors', data)
1773 self.assertIsNone(data['data'])
1775 # Remove permission constraint
1776 obj_perm.constraints = None
1777 obj_perm.save()
1779 # Request should return requested object
1780 response = self.client.post(url, data={'query': query}, format="json", **self.header)
1781 self.assertHttpStatus(response, status.HTTP_200_OK)
1782 data = json.loads(response.content)
1783 self.assertNotIn('errors', data)
1784 self.assertIsNotNone(data['data'])
1786 @override_settings(LOGIN_REQUIRED=True)
1787 def test_graphql_list_objects(self):
1788 url = reverse('graphql')
1789 field_name = f'{self._get_graphql_base_name()}_list'
1790 query = self._build_query(field_name)
1792 # Non-authenticated requests should fail
1793 header = {
1794 'HTTP_ACCEPT': 'application/json',
1795 }
1796 with disable_warnings('django.request'):
1797 response = self.client.post(url, data={'query': query}, format="json", **header)
1798 self.assertHttpStatus(response, status.HTTP_403_FORBIDDEN)
1800 # Add constrained permission
1801 obj_perm = ObjectPermission(
1802 name='Test permission',
1803 actions=['view'],
1804 constraints={'id': 0} # Impossible constraint
1805 )
1806 obj_perm.save()
1807 obj_perm.users.add(self.user)
1808 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
1810 if self.graphql_object_permission_assertions:
1811 # Request should succeed but return empty results list
1812 response = self.client.post(url, data={'query': query}, format="json", **self.header)
1813 self.assertHttpStatus(response, status.HTTP_200_OK)
1814 data = json.loads(response.content)
1815 self.assertNotIn('errors', data)
1816 self.assertEqual(len(data['data'][field_name]), 0)
1818 # Remove permission constraint
1819 obj_perm.constraints = None
1820 obj_perm.save()
1822 # Request should return all objects
1823 response = self.client.post(url, data={'query': query}, format="json", **self.header)
1824 self.assertHttpStatus(response, status.HTTP_200_OK)
1825 data = json.loads(response.content)
1826 self.assertNotIn('errors', data)
1827 self.assertEqual(len(data['data'][field_name]), self.model.objects.count())
1829 def _assert_graphql_filter_tests_exist(self, auto_tests, legacy_tests, explicit_tests):
1830 """
1831 Fail loudly when auto mode is required and no GraphQL filter tests
1832 (auto, legacy, or explicit) exist for the current model.
1833 """
1834 if (
1835 getattr(self, 'graphql_auto_filter_tests', True)
1836 and getattr(self, 'graphql_auto_filter_required', True)
1837 and not auto_tests
1838 and not legacy_tests
1839 and not explicit_tests
1840 ):
1841 self.fail(
1842 f'No GraphQL filter tests were generated for {self.model._meta.label}. '
1843 f'Set graphql_auto_filter_required = False or add explicit graphql_filter_tests '
1844 f'if intentional.'
1845 )
1847 @override_settings(LOGIN_REQUIRED=True)
1848 def test_graphql_filter_objects(self):
1849 legacy_tests = list(self._iter_legacy_graphql_filter_tests())
1850 explicit_tests = list(self._iter_explicit_graphql_filter_tests())
1852 filter_fields = self._get_graphql_filter_field_names()
1853 self._assert_graphql_filter_class_present(
1854 filter_fields, handwritten_tests=[*legacy_tests, *explicit_tests]
1855 )
1857 auto_tests = list(self._iter_auto_graphql_filter_tests())
1859 self._assert_graphql_filter_tests_exist(auto_tests, legacy_tests, explicit_tests)
1861 filter_tests = [*auto_tests, *legacy_tests, *explicit_tests]
1862 if not filter_tests:
1863 return
1865 url = reverse('graphql')
1866 field_name = f'{self._get_graphql_base_name()}_list'
1868 # Add object-level permission
1869 obj_perm = ObjectPermission(
1870 name='Test permission',
1871 actions=['view']
1872 )
1873 obj_perm.save()
1874 obj_perm.users.add(self.user)
1875 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
1877 for filter_test in filter_tests:
1878 with self.subTest(filter=filter_test.name):
1879 self._assert_graphql_filter_test(url, field_name, filter_test)
1881 @override_settings(LOGIN_REQUIRED=True)
1882 def test_graphql_extra_queries(self):
1883 query_tests = [
1884 self._coerce_graphql_query_test(query_test)
1885 for query_test in getattr(self, 'graphql_query_tests', ())
1886 ]
1888 if not query_tests:
1889 return
1891 url = reverse('graphql')
1893 # Add object-level permission for this model. Additional permissions
1894 # required by the query can be declared on the GraphQLQueryTest.
1895 obj_perm = ObjectPermission(
1896 name='Test permission',
1897 actions=['view']
1898 )
1899 obj_perm.save()
1900 obj_perm.users.add(self.user)
1901 obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
1903 for query_test in query_tests:
1904 with self.subTest(query=query_test.name):
1905 for permission in query_test.permissions:
1906 self.add_permissions(permission)
1908 response = self.client.post(url, data={'query': query_test.query}, format="json", **self.header)
1909 self.assertHttpStatus(response, status.HTTP_200_OK)
1911 data = json.loads(response.content)
1912 self.assertNotIn('errors', data)
1913 query_test.assert_result(self, data['data'])
1915 class APIViewTestCase(
1916 GetObjectViewTestCase,
1917 ListObjectsViewTestCase,
1918 CreateObjectViewTestCase,
1919 UpdateObjectViewTestCase,
1920 DeleteObjectViewTestCase,
1921 GraphQLTestCase
1922 ):
1923 pass
1925 class GraphQLSchemaCoverageTestCase(TestCase):
1926 """
1927 Assert every model-backed GraphQL type exposed as a root query field is covered by a
1928 concrete GraphQLTestCase subclass. Subclass this in a test module to run the audit.
1930 Scope is intentionally limited to types reachable as root query fields (e.g. ``site``,
1931 ``site_list``); these are exactly the types the detail/list GraphQLTestCase methods can
1932 exercise. Types reachable only as nested object fields are out of scope.
1933 """
1934 # Per-app test submodules to import so their GraphQLTestCase subclasses are defined.
1935 graphql_test_modules = ('test_api', 'test_graphql')
1937 # GraphQL type classes intentionally excluded from coverage.
1938 graphql_exempt_type_classes = ()
1940 def get_graphql_schema(self):
1941 # Imported lazily so importing this testing utility does not eagerly build the schema.
1942 from netbox.graphql.schema import schema
1943 return schema._schema
1945 def iter_test_module_names(self):
1946 # Import test modules only for apps exposing model-backed root query types;
1947 # coverage classes are expected to live with the app whose type they cover.
1948 app_labels = {model._meta.app_label for model in self.get_schema_type_classes().values()}
1949 for app_label in sorted(app_labels):
1950 app_config = apps.get_app_config(app_label)
1951 for module_name in self.graphql_test_modules:
1952 yield f'{app_config.name}.tests.{module_name}'
1954 def import_graphql_test_modules(self):
1955 for module_name in self.iter_test_module_names():
1956 self.import_graphql_test_module(module_name)
1958 def import_graphql_test_module(self, module_name):
1959 try:
1960 importlib.import_module(module_name)
1961 except ModuleNotFoundError as exc:
1962 # A missing test module, or a missing parent package (e.g. `<app>.tests`),
1963 # is fine. An import error raised from inside an existing test module
1964 # should still fail loudly.
1965 if exc.name == module_name or module_name.startswith(f'{exc.name}.'):
1966 return
1967 raise
1969 def unwrap_graphql_type(self, graphql_type):
1970 while isinstance(graphql_type, (GraphQLNonNull, GraphQLList)):
1971 graphql_type = graphql_type.of_type
1972 return graphql_type
1974 def get_schema_field_type_class(self, field):
1975 graphql_type = self.unwrap_graphql_type(field.type)
1976 if not isinstance(graphql_type, GraphQLObjectType):
1977 return None
1978 extensions = getattr(graphql_type, 'extensions', None) or {}
1979 definition = extensions.get(GraphQLCoreConverter.DEFINITION_BACKREF)
1980 return getattr(definition, 'origin', None)
1982 def get_graphql_type_model(self, type_class):
1983 django_definition = getattr(type_class, '__strawberry_django_definition__', None)
1984 return getattr(django_definition, 'model', None)
1986 def get_schema_type_classes(self):
1987 """Return {type_class: model} for every model-backed root query type (cached per instance)."""
1988 cached = getattr(self, '_schema_type_classes', None)
1989 if cached is not None:
1990 return cached
1991 type_classes = {}
1992 for field in self.get_graphql_schema().query_type.fields.values():
1993 type_class = self.get_schema_field_type_class(field)
1994 if type_class is None:
1995 continue
1996 model = self.get_graphql_type_model(type_class)
1997 if model is None:
1998 continue
1999 type_classes[type_class] = model
2000 self._schema_type_classes = type_classes
2001 return type_classes
2003 def iter_graphql_testcase_classes(self, base_class=None):
2004 base_class = base_class or APIViewTestCases.GraphQLTestCase
2005 for subclass in base_class.__subclasses__():
2006 yield subclass
2007 yield from self.iter_graphql_testcase_classes(subclass)
2009 def get_testcase_type_class(self, testcase):
2010 if getattr(testcase, 'graphql_test_exempt', False):
2011 return None
2012 try:
2013 return testcase.get_graphql_type_class()
2014 except GraphQLTypeNotFound as exc:
2015 model = getattr(testcase, 'model', None)
2016 model_label = model._meta.label if model is not None else 'unknown model'
2017 self.fail(
2018 f'{testcase.__module__}.{testcase.__name__} sets model = {model_label} '
2019 f'but no GraphQL type could be resolved. Set type_class if the type lives '
2020 f'outside the conventional <app>.graphql.types.<Model>Type path, or set '
2021 f'graphql_test_exempt = True if this test case should not count toward '
2022 f'schema coverage. Original error: {exc}'
2023 )
2025 def get_testcase_type_classes(self):
2026 self.import_graphql_test_modules()
2027 type_classes = set()
2028 for testcase in self.iter_graphql_testcase_classes():
2029 type_class = self.get_testcase_type_class(testcase)
2030 if type_class is not None:
2031 type_classes.add(type_class)
2032 return type_classes
2034 def format_type_class(self, type_class):
2035 model = self.get_graphql_type_model(type_class)
2036 label = f' ({model._meta.label})' if model is not None else ''
2037 return f'{type_class.__module__}.{type_class.__name__}{label}'
2039 def test_schema_types_have_graphql_test_coverage(self):
2040 """Every model-backed root query type is covered by a GraphQLTestCase."""
2041 expected = set(self.get_schema_type_classes())
2042 self.assertGreater(
2043 len(expected), 0,
2044 'No model-backed root query GraphQL types were discovered; schema '
2045 'introspection may have broken.'
2046 )
2047 actual = self.get_testcase_type_classes()
2048 exempt = set(self.graphql_exempt_type_classes)
2049 missing = sorted(self.format_type_class(tc) for tc in expected - actual - exempt)
2050 self.assertEqual(missing, [])