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

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 

10 

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) 

37 

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 

55 

56from .base import ModelTestCase, TestCase 

57from .query_counts import assert_expected_query_count 

58from .utils import disable_logging, disable_warnings, get_random_string 

59 

60__all__ = ( 

61 'APITestCase', 

62 'APIViewTestCases', 

63 'GraphQLFilterTest', 

64 'GraphQLQueryTest', 

65) 

66 

67 

68@dataclass(frozen=True) 

69class GraphQLFilterTest: 

70 """ 

71 Declarative GraphQL filter test case for APIViewTestCases.GraphQLTestCase. 

72 

73 ``filters`` is the raw content to place inside the GraphQL ``filters`` input, 

74 e.g. ``name: {i_contains: "site"}``. 

75 

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, ...] = () 

86 

87 

88@dataclass(frozen=True) 

89class GraphQLQueryTest: 

90 """ 

91 Declarative GraphQL query test case for model-specific complex queries. 

92 

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, ...] = () 

103 

104 

105# 

106# REST/GraphQL API Tests 

107# 

108 

109class APITestCase(ModelTestCase): 

110 """ 

111 Base test case for API requests. 

112 

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 

118 

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}'} 

128 

129 def _get_view_namespace(self): 

130 return f'{self.view_namespace or self.model._meta.app_label}-api' 

131 

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}) 

135 

136 def _get_list_url(self): 

137 viewname = f'{self._get_view_namespace()}:{self.model._meta.model_name}-list' 

138 return reverse(viewname) 

139 

140 

141class APIViewTestCases: 

142 

143 class GetObjectViewTestCase(APITestCase): 

144 

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) 

158 

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()) 

164 

165 # Try GET without permission 

166 with disable_warnings('django.request'): 

167 self.assertHttpStatus(self.client.get(url, **self.header), status.HTTP_403_FORBIDDEN) 

168 

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] 

176 

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)) 

186 

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) 

191 

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") 

195 

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) 

199 

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) 

208 

209 class ListObjectsViewTestCase(APITestCase): 

210 brief_fields = [] 

211 

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()) 

226 

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) 

234 

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) 

238 

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() 

244 

245 # Try GET without permission 

246 with disable_warnings('django.request'): 

247 self.assertHttpStatus(self.client.get(url, **self.header), status.HTTP_403_FORBIDDEN) 

248 

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] 

256 

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)) 

266 

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) 

272 

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) 

280 

281 class CreateObjectViewTestCase(APITestCase): 

282 create_data = [] 

283 validation_excluded_fields = [] 

284 

285 def test_create_object_without_permission(self): 

286 """ 

287 POST a single object without permission. 

288 """ 

289 url = self._get_list_url() 

290 

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) 

295 

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)) 

308 

309 data = copy.deepcopy(self.create_data[0]) 

310 

311 # If supported, add a changelog message 

312 if issubclass(self.model, ChangeLoggingMixin): 

313 data['changelog_message'] = get_random_string(10) 

314 

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 ) 

326 

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']) 

336 

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)) 

349 

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 

355 

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 ) 

375 

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) 

390 

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)) 

403 

404 initial_count = self._get_queryset().count() 

405 

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 ) 

414 

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]) 

424 

425 class UpdateObjectViewTestCase(APITestCase): 

426 update_data = {} 

427 bulk_update_data = None 

428 bulk_update_invalid_data = None 

429 validation_excluded_fields = [] 

430 

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] 

437 

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) 

442 

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] 

450 

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)) 

459 

460 data = copy.deepcopy(update_data) 

461 

462 # If supported, add a changelog message 

463 if issubclass(self.model, ChangeLoggingMixin): 

464 data['changelog_message'] = get_random_string(10) 

465 

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 ) 

475 

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']) 

484 

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") 

492 

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] 

500 

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") 

506 

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 

516 

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) 

524 

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") 

531 

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)) 

540 

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 ] 

546 

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 

552 

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) 

563 

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) 

574 

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") 

583 

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)) 

588 

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") 

591 

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 ] 

597 

598 response = self.client.patch(self._get_list_url(), data, format='json', **self.header) 

599 

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) 

607 

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') 

615 

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)) 

620 

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') 

623 

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 ] 

629 

630 # Snapshot field values before the request so we can verify atomicity afterward 

631 instance0_before = self._get_queryset().get(pk=id_list[0]) 

632 

633 response = self.client.patch(self._get_list_url(), data, format='json', **self.header) 

634 

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]) 

641 

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 ) 

652 

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') 

660 

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)) 

665 

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 

669 

670 data = [{'id': id, **self.bulk_update_data} for id in (*id_list, missing_id)] 

671 

672 # Snapshot the objects which would otherwise have been updated 

673 instances_before = list(self._get_queryset().filter(pk__in=id_list)) 

674 

675 response = self.client.patch(self._get_list_url(), data, format='json', **self.header) 

676 

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']) 

683 

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 ) 

696 

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') 

704 

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)) 

709 

710 instance = self._get_queryset().first() 

711 

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}, {}] 

714 

715 response = self.client.patch(self._get_list_url(), data, format='json', **self.header) 

716 

717 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST) 

718 self.assertIn('detail', response.data) 

719 self.assertEqual(len(response.data['errors']), 1) 

720 

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']) 

724 

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 ) 

735 

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') 

743 

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)) 

748 

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') 

751 

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])] 

754 

755 # Snapshot the objects which would otherwise have been updated 

756 instances_before = list(self._get_queryset().filter(pk__in=id_list)) 

757 

758 response = self.client.patch(self._get_list_url(), data, format='json', **self.header) 

759 

760 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST) 

761 self.assertIn('detail', response.data) 

762 self.assertIn('errors', response.data) 

763 

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']) 

768 

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 ) 

781 

782 class DeleteObjectViewTestCase(APITestCase): 

783 

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()) 

789 

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) 

794 

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) 

801 

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)) 

810 

811 data = {} 

812 

813 # If supported, add a changelog message 

814 if issubclass(self.model, ChangeLoggingMixin): 

815 data['changelog_message'] = get_random_string(10) 

816 

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()) 

820 

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']) 

829 

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)) 

842 

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] 

848 

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 

854 

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) 

859 

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) 

870 

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)) 

883 

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)] 

889 

890 initial_count = self._get_queryset().count() 

891 response = self.client.delete(self._get_list_url(), data, format='json', **self.header) 

892 

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']) 

899 

900 # The objects named alongside the missing one must not have been deleted 

901 self.assertEqual(self._get_queryset().count(), initial_count) 

902 

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)) 

915 

916 # Target the most recently created object to avoid triggering recursive deletions 

917 instance = self._get_queryset().order_by('-id').first() 

918 

919 # The second entry omits the object ID, so it cannot be matched to an object 

920 data = [{'id': instance.pk}, {}] 

921 

922 initial_count = self._get_queryset().count() 

923 response = self.client.delete(self._get_list_url(), data, format='json', **self.header) 

924 

925 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST) 

926 self.assertIn('detail', response.data) 

927 self.assertEqual(len(response.data['errors']), 1) 

928 

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']) 

932 

933 # Nothing may have been deleted 

934 self.assertEqual(self._get_queryset().count(), initial_count) 

935 

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)) 

948 

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') 

952 

953 # Repeat the first ID at the end of the request 

954 data = [{'id': id} for id in (*id_list, id_list[0])] 

955 

956 initial_count = self._get_queryset().count() 

957 response = self.client.delete(self._get_list_url(), data, format='json', **self.header) 

958 

959 self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST) 

960 self.assertIn('detail', response.data) 

961 self.assertIn('errors', response.data) 

962 

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']) 

967 

968 # No object named in the request may have been deleted 

969 self.assertEqual(self._get_queryset().count(), initial_count) 

970 

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)) 

983 

984 initial_count = self._get_queryset().count() 

985 self.assertNotEqual(initial_count, 0, 'No objects exist against which to test bulk deletion') 

986 

987 response = self.client.delete(self._get_list_url(), **self.header) 

988 

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 ) 

997 

998 class GraphQLTestCase(APITestCase): 

999 graphql_auto_filter_tests = True 

1000 graphql_auto_filter_exclude = () 

1001 

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 

1005 

1006 # Fail when auto mode is on and no tests were generated. 

1007 graphql_auto_filter_required = True 

1008 

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 

1012 

1013 # Additional explicit-list filter cases as GraphQLFilterTest instances. 

1014 graphql_filter_tests = () 

1015 

1016 # Additional full-query cases (e.g. nested filters) as GraphQLQueryTest instances. 

1017 graphql_query_tests = () 

1018 

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 

1022 

1023 # Exclude this test case from GraphQL schema coverage. 

1024 graphql_test_exempt = False 

1025 

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) 

1034 

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) 

1042 

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() 

1049 

1050 # Compile list of fields to include 

1051 fields_string = '' 

1052 

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' 

1089 

1090 query = f""" 

1091 {{ 

1092 {name}{filter_string} {{ 

1093 {fields_string} 

1094 }} 

1095 }} 

1096 """ 

1097 

1098 return query 

1099 

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) 

1120 

1121 return json.dumps(str(value)) 

1122 

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 

1129 

1130 if not isinstance(params, dict): 

1131 return self._graphql_literal(params) 

1132 

1133 lookup = params.get('lookup') 

1134 value = params['value'] 

1135 

1136 if lookup: 

1137 return f'{{{lookup}: {self._graphql_literal(value)}}}' 

1138 

1139 return self._graphql_literal(value) 

1140 

1141 def _build_graphql_filter_string(self, **filters): 

1142 if not filters: 

1143 return '' 

1144 

1145 filter_expressions = [ 

1146 f'{field_name}: {self._render_graphql_filter_value(params)}' 

1147 for field_name, params in filters.items() 

1148 ] 

1149 

1150 return f'(filters: {{{", ".join(filter_expressions)}}})' 

1151 

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) 

1157 

1158 return self._build_query_with_filter(name, filter_string) 

1159 

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' 

1163 

1164 return f""" 

1165 {{ 

1166 {name}{filter_string} {{ 

1167 {selection} 

1168 }} 

1169 }} 

1170 """ 

1171 

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) 

1184 

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' 

1196 

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 

1207 

1208 return getattr(module, class_name, None) 

1209 

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() 

1218 

1219 return self._collect_filter_class_annotation_names(filter_class) 

1220 

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 

1230 

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 

1244 

1245 label = self.model._meta.label 

1246 path = f'{self.model._meta.app_label}.graphql.filters.{self.model.__name__}Filter' 

1247 

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 ) 

1259 

1260 def _get_nonempty_field_value(self, field): 

1261 queryset = self._get_queryset() 

1262 

1263 if getattr(field, 'null', False): 

1264 queryset = queryset.exclude(**{f'{field.name}__isnull': True}) 

1265 

1266 if isinstance(field, (models.CharField, models.TextField)): 

1267 queryset = queryset.exclude(**{field.name: ''}) 

1268 

1269 return queryset.values_list(field.name, flat=True).first() 

1270 

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 

1281 

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() 

1295 

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 

1308 

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) 

1317 

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 

1324 

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 

1339 

1340 break 

1341 

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 

1354 

1355 return annotation 

1356 

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 

1368 

1369 if annotation is strawberry.ID: 

1370 return 'id', None 

1371 

1372 origin = typing.get_origin(annotation) 

1373 target = origin if isinstance(origin, type) else annotation 

1374 type_args = typing.get_args(annotation) 

1375 

1376 if not isinstance(target, type): 

1377 return None, None 

1378 

1379 if target in (IntegerLookup, BigIntegerLookup, FloatLookup): 

1380 return 'numeric', target 

1381 

1382 # TreeNodeFilter schema requires {id, match_type}; skip auto-emit. 

1383 if target is TreeNodeFilter: 

1384 return None, None 

1385 

1386 if issubclass(target, (DateFilterLookup, DatetimeFilterLookup, TimeFilterLookup)): 

1387 return 'date_lookup', None 

1388 

1389 if target is RangeLookup or issubclass(target, RangeLookup): 

1390 return 'range_lookup', type_args[0] if type_args else None 

1391 

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 

1399 

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 

1409 

1410 return None, None 

1411 

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 

1423 

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 ) 

1436 

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 ) 

1461 

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 ) 

1481 

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 ) 

1494 

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 ) 

1515 

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 ) 

1529 

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 ) 

1551 

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 ) 

1572 

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 ) 

1592 

1593 def _iter_auto_graphql_filter_tests(self): 

1594 if not getattr(self, 'graphql_auto_filter_tests', True): 

1595 return 

1596 

1597 filter_class = self._get_model_graphql_filter_class() 

1598 if filter_class is None: 

1599 return 

1600 

1601 exclude = set(getattr(self, 'graphql_auto_filter_exclude', ())) 

1602 per_kind = self.graphql_auto_filter_fields_per_kind 

1603 

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)) 

1613 

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 

1620 

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 

1630 

1631 def _iter_legacy_graphql_filter_tests(self): 

1632 if not hasattr(self, 'graphql_filter'): 

1633 return 

1634 

1635 filter_expressions = [ 

1636 f'{field_name}: {self._render_graphql_filter_value(params)}' 

1637 for field_name, params in self.graphql_filter.items() 

1638 ] 

1639 

1640 yield GraphQLFilterTest( 

1641 name='graphql_filter', 

1642 filters=', '.join(filter_expressions), 

1643 ) 

1644 

1645 def _coerce_graphql_filter_test(self, filter_test): 

1646 if isinstance(filter_test, GraphQLFilterTest): 

1647 return filter_test 

1648 

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') 

1652 

1653 return GraphQLFilterTest(**filter_test) 

1654 

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) 

1658 

1659 def _get_expected_id_set(self, filter_test): 

1660 expected = filter_test.expected 

1661 

1662 if callable(expected): 

1663 expected = expected(self._get_queryset()) 

1664 

1665 if isinstance(expected, dict): 

1666 expected = self._get_queryset().filter(**expected) 

1667 

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] 

1672 

1673 return {str(value) for value in values} 

1674 

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) 

1677 

1678 for permission in filter_test.permissions: 

1679 self.add_permissions(permission) 

1680 

1681 response = self.client.post(url, data={'query': query}, format="json", **self.header) 

1682 self.assertHttpStatus(response, status.HTTP_200_OK) 

1683 

1684 data = json.loads(response.content) 

1685 self.assertNotIn('errors', data) 

1686 

1687 results = data['data'][field_name] 

1688 

1689 if filter_test.expected is None: 

1690 self.assertGreater(len(results), 0) 

1691 return 

1692 

1693 expected_ids = self._get_expected_id_set(filter_test) 

1694 

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 ) 

1703 

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 ) 

1710 

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 ) 

1718 

1719 def _coerce_graphql_query_test(self, query_test): 

1720 if isinstance(query_test, GraphQLQueryTest): 

1721 return query_test 

1722 

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') 

1726 

1727 return GraphQLQueryTest(**query_test) 

1728 

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 = '' 

1738 

1739 return self._build_query_with_filter(name, filter_string) 

1740 

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) 

1747 

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) 

1755 

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)) 

1765 

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']) 

1774 

1775 # Remove permission constraint 

1776 obj_perm.constraints = None 

1777 obj_perm.save() 

1778 

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']) 

1785 

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) 

1791 

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) 

1799 

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)) 

1809 

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) 

1817 

1818 # Remove permission constraint 

1819 obj_perm.constraints = None 

1820 obj_perm.save() 

1821 

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()) 

1828 

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 ) 

1846 

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()) 

1851 

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 ) 

1856 

1857 auto_tests = list(self._iter_auto_graphql_filter_tests()) 

1858 

1859 self._assert_graphql_filter_tests_exist(auto_tests, legacy_tests, explicit_tests) 

1860 

1861 filter_tests = [*auto_tests, *legacy_tests, *explicit_tests] 

1862 if not filter_tests: 

1863 return 

1864 

1865 url = reverse('graphql') 

1866 field_name = f'{self._get_graphql_base_name()}_list' 

1867 

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)) 

1876 

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) 

1880 

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 ] 

1887 

1888 if not query_tests: 

1889 return 

1890 

1891 url = reverse('graphql') 

1892 

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)) 

1902 

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) 

1907 

1908 response = self.client.post(url, data={'query': query_test.query}, format="json", **self.header) 

1909 self.assertHttpStatus(response, status.HTTP_200_OK) 

1910 

1911 data = json.loads(response.content) 

1912 self.assertNotIn('errors', data) 

1913 query_test.assert_result(self, data['data']) 

1914 

1915 class APIViewTestCase( 

1916 GetObjectViewTestCase, 

1917 ListObjectsViewTestCase, 

1918 CreateObjectViewTestCase, 

1919 UpdateObjectViewTestCase, 

1920 DeleteObjectViewTestCase, 

1921 GraphQLTestCase 

1922 ): 

1923 pass 

1924 

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. 

1929 

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') 

1936 

1937 # GraphQL type classes intentionally excluded from coverage. 

1938 graphql_exempt_type_classes = () 

1939 

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 

1944 

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}' 

1953 

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) 

1957 

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 

1968 

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 

1973 

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) 

1981 

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) 

1985 

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 

2002 

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) 

2008 

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 ) 

2024 

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 

2033 

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}' 

2038 

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, [])