Coverage for src/backend/InvenTree/InvenTree/unit_test.py: 16%

389 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 17:47 +0000

1"""Helper functions for unit testing / CI.""" 

2 

3import csv 

4import io 

5import json 

6import os 

7import re 

8import time 

9from collections.abc import Callable 

10from contextlib import contextmanager 

11from pathlib import Path 

12from typing import Optional 

13from unittest import mock 

14 

15from django.contrib.auth import get_user_model 

16from django.contrib.auth.models import Group, Permission, User 

17from django.db import connections, models 

18from django.http.response import StreamingHttpResponse 

19from django.test import TestCase, tag 

20from django.test.utils import CaptureQueriesContext, override_settings 

21from django.urls import reverse 

22 

23from djmoney.contrib.exchange.models import ExchangeBackend, Rate 

24from rest_framework.test import APITestCase 

25 

26from plugin import registry 

27from plugin.models import PluginConfig 

28 

29 

30@contextmanager 

31def count_queries( 

32 msg: Optional[str] = None, 

33 log_to_file: bool = False, 

34 using: str = 'default', 

35 threshold: int = 10, 

36): # pragma: no cover 

37 """Helper function to count the number of queries executed. 

38 

39 Arguments: 

40 msg: Optional message to print after counting queries 

41 log_to_file: If True, log the queries to a file (default = False) 

42 using: The database connection to use (default = 'default') 

43 threshold: Minimum number of queries to log (default = 10) 

44 """ 

45 t1 = time.time() 

46 

47 with CaptureQueriesContext(connections[using]) as context: 

48 yield 

49 

50 dt = time.time() - t1 

51 

52 n = len(context.captured_queries) 

53 

54 if log_to_file: 

55 with open('queries.txt', 'w', encoding='utf-8') as f: 

56 for q in context.captured_queries: 

57 f.write(str(q['sql']) + '\n\n') 

58 

59 output = f'Executed {n} queries in {dt:.4f}s' 

60 

61 if threshold and n >= threshold: 

62 if msg: 

63 print(f'{msg}: {output}') 

64 else: 

65 print(output) 

66 

67 

68def addUserPermission(user: User, app_name: str, model_name: str, perm: str) -> None: 

69 """Add a specific permission for the provided user. 

70 

71 Arguments: 

72 user: The user to add the permission to 

73 app_name: The name of the app (e.g. 'part') 

74 model_name: The name of the model (e.g. 'location') 

75 perm: The permission to add (e.g. 'add', 'change', 'delete', 'view') 

76 """ 

77 # Get the permission object 

78 permission = Permission.objects.get( 

79 content_type__model=model_name, codename=f'{perm}_{model_name}' 

80 ) 

81 

82 # Add the permission to the user 

83 user.user_permissions.add(permission) 

84 user.save() 

85 

86 

87def getMigrationFileNames(app): 

88 """Return a list of all migration filenames for provided app.""" 

89 local_dir = Path(__file__).parent 

90 files = local_dir.joinpath('..', app, 'migrations').iterdir() 

91 

92 # Regex pattern for migration files 

93 regex = re.compile(r'^[\d]+_.*\.py$') 

94 

95 migration_files = [] 

96 

97 for f in files: 

98 if regex.match(f.name): 

99 migration_files.append(f.name) 

100 

101 return migration_files 

102 

103 

104def getOldestMigrationFile(app, exclude_extension=True, ignore_initial=True): 

105 """Return the filename associated with the oldest migration.""" 

106 oldest_num = -1 

107 oldest_file = None 

108 

109 for f in getMigrationFileNames(app): 

110 if ignore_initial and f.startswith('0001_initial'): 

111 continue 

112 

113 num = int(f.split('_')[0]) 

114 

115 if oldest_file is None or num < oldest_num: 

116 oldest_num = num 

117 oldest_file = f 

118 

119 if exclude_extension and oldest_file: 

120 oldest_file = oldest_file.replace('.py', '') 

121 

122 return oldest_file 

123 

124 

125def getNewestMigrationFile(app, exclude_extension=True): 

126 """Return the filename associated with the newest migration.""" 

127 newest_file = None 

128 newest_num = -1 

129 

130 for f in getMigrationFileNames(app): 

131 num = int(f.split('_')[0]) 

132 

133 if newest_file is None or num > newest_num: 

134 newest_num = num 

135 newest_file = f 

136 

137 if not newest_file: # pragma: no cover 

138 return newest_file 

139 

140 if exclude_extension: 

141 newest_file = newest_file.replace('.py', '') 

142 

143 return newest_file 

144 

145 

146def findOffloadedTask( 

147 task_name: str, 

148 clear_after: bool = False, 

149 reverse: bool = False, 

150 matching_args=None, 

151 matching_kwargs=None, 

152): 

153 """Find an offloaded tasks in the background worker queue. 

154 

155 Arguments: 

156 task_name: The name of the task to search for 

157 clear_after: Clear the task queue after searching 

158 reverse: Search in reverse order (most recent first) 

159 matching_args: List of argument names to match against 

160 matching_kwargs: List of keyword argument names to match against 

161 """ 

162 from django_q.models import OrmQ 

163 

164 tasks = OrmQ.objects.all() 

165 

166 if reverse: 

167 tasks = tasks.order_by('-pk') 

168 

169 task = None 

170 

171 for t in tasks: 

172 if t.func() == task_name: 

173 found = True 

174 

175 if matching_args: 

176 for arg in matching_args: 

177 if arg not in t.args(): 

178 found = False 

179 break 

180 

181 if matching_kwargs: 

182 for kwarg in matching_kwargs: 

183 if kwarg not in t.kwargs(): 

184 found = False 

185 break 

186 

187 if found: 

188 task = t 

189 break 

190 

191 if clear_after: 

192 OrmQ.objects.all().delete() 

193 

194 return task 

195 

196 

197def findOffloadedEvent( 

198 event_name: str, 

199 clear_after: bool = False, 

200 reverse: bool = False, 

201 matching_kwargs=None, 

202): 

203 """Find an offloaded event in the background worker queue.""" 

204 return findOffloadedTask( 

205 'plugin.base.event.events.register_event', 

206 matching_args=[str(event_name)], 

207 matching_kwargs=matching_kwargs, 

208 clear_after=clear_after, 

209 reverse=reverse, 

210 ) 

211 

212 

213class UserMixin: 

214 """Mixin to setup a user and login for tests. 

215 

216 Use parameters to set username, password, email, roles and permissions. 

217 """ 

218 

219 # User information 

220 username = 'testuser' 

221 password = 'mypassword' 

222 email = 'test@testing.com' 

223 

224 superuser = False 

225 is_staff = True 

226 auto_login = True 

227 

228 # Set list of roles automatically associated with the user 

229 roles = [] 

230 

231 @classmethod 

232 def setUpTestData(cls): 

233 """Run setup for all tests in a given class.""" 

234 super().setUpTestData() 

235 

236 # Create a user to log in with 

237 cls.user = get_user_model().objects.create_user( 

238 username=cls.username, password=cls.password, email=cls.email 

239 ) 

240 

241 # Create a group for the user 

242 cls.group = Group.objects.create(name='my_test_group') 

243 cls.user.groups.add(cls.group) 

244 

245 if cls.superuser: 

246 cls.user.is_superuser = True 

247 

248 if cls.is_staff: 

249 cls.user.is_staff = True 

250 

251 cls.user.save() 

252 

253 # Assign all roles if set 

254 if cls.roles == 'all': 

255 cls.assignRole(group=cls.group, assign_all=True) 

256 

257 # else filter the roles 

258 else: 

259 for role in cls.roles: 

260 cls.assignRole(role=role, group=cls.group) 

261 

262 def setUp(self): 

263 """Run setup for individual test methods.""" 

264 if self.auto_login: 

265 self.login() 

266 

267 def login(self): 

268 """Login with the current user credentials.""" 

269 self.client.login(username=self.username, password=self.password) 

270 

271 def logout(self): 

272 """Lougout current user.""" 

273 self.client.logout() 

274 

275 @classmethod 

276 def clearRoles(cls): 

277 """Remove all user roles from the registered user.""" 

278 for ruleset in cls.group.rule_sets.all(): 

279 ruleset.can_view = False 

280 ruleset.can_change = False 

281 ruleset.can_delete = False 

282 ruleset.can_add = False 

283 

284 ruleset.save() 

285 

286 @classmethod 

287 def assignRole(cls, role=None, assign_all: bool = False, group=None): 

288 """Set the user roles for the registered user. 

289 

290 Arguments: 

291 role: Role of the format 'rule.permission' e.g. 'part.add' 

292 assign_all: Set to True to assign *all* roles 

293 group: The group to assign roles to (or leave None to use the group assigned to this class) 

294 """ 

295 if group is None: 

296 group = cls.group 

297 

298 if type(assign_all) is not bool: 

299 # Raise exception if common mistake is made! 

300 raise TypeError( 

301 'assignRole: assign_all must be a boolean value' 

302 ) # pragma: no cover 

303 

304 if not role and not assign_all: 

305 raise ValueError( 

306 'assignRole: either role must be provided, or assign_all must be set' 

307 ) # pragma: no cover 

308 

309 if not assign_all and role: 

310 rule, perm = role.split('.') 

311 

312 for ruleset in group.rule_sets.all(): 

313 if assign_all or ruleset.name == rule: 

314 if assign_all or perm == 'view': 

315 ruleset.can_view = True 

316 elif assign_all or perm == 'change': 

317 ruleset.can_change = True 

318 elif assign_all or perm == 'delete': 

319 ruleset.can_delete = True 

320 elif assign_all or perm == 'add': 

321 ruleset.can_add = True 

322 

323 ruleset.save() 

324 if not assign_all: 

325 break 

326 

327 

328class PluginMixin: 

329 """Mixin to ensure that all plugins are loaded for tests.""" 

330 

331 def setUp(self): 

332 """Setup for plugin tests.""" 

333 super().setUp() 

334 

335 # Load plugin configs 

336 self.plugin_confs = PluginConfig.objects.all() 

337 # Reload if not present 

338 if not self.plugin_confs: 

339 registry.reload_plugins() 

340 self.plugin_confs = PluginConfig.objects.all() 

341 

342 

343class ExchangeRateMixin: 

344 """Mixin class for generating exchange rate data.""" 

345 

346 def generate_exchange_rates(self): 

347 """Helper function which generates some exchange rates to work with.""" 

348 rates = {'AUD': 1.5, 'CAD': 1.7, 'GBP': 0.9, 'USD': 1.0} 

349 

350 # Create a dummy backend 

351 ExchangeBackend.objects.create(name='InvenTreeExchange', base_currency='USD') 

352 

353 backend = ExchangeBackend.objects.get(name='InvenTreeExchange') 

354 

355 items = [] 

356 

357 for currency, rate in rates.items(): 

358 items.append(Rate(currency=currency, value=rate, backend=backend)) 

359 

360 Rate.objects.bulk_create(items) 

361 

362 

363class TestQueryMixin: 

364 """Mixin class for testing query counts.""" 

365 

366 # Default query count threshold value 

367 # TODO: This value should be reduced 

368 MAX_QUERY_COUNT = 250 

369 

370 WARNING_QUERY_THRESHOLD = 100 

371 

372 # Default query time threshold value 

373 # TODO: This value should be reduced 

374 # Note: There is a lot of variability in the query time in unit testing... 

375 MAX_QUERY_TIME = 7.5 

376 

377 @contextmanager 

378 def assertNumQueriesLessThan( 

379 self, value, using='default', verbose=False, url=None, log_to_file=False 

380 ): 

381 """Context manager to check that the number of queries is less than a certain value. 

382 

383 Example: 

384 with self.assertNumQueriesLessThan(10): 

385 # Do some stuff 

386 Ref: https://stackoverflow.com/questions/1254170/django-is-there-a-way-to-count-sql-queries-from-an-unit-test/59089020#59089020 

387 """ 

388 with CaptureQueriesContext(connections[using]) as context: 

389 yield # your test will be run here 

390 

391 n = len(context.captured_queries) 

392 

393 if url and n >= value: 

394 print( 

395 f'Query count exceeded at {url}: Expected < {value} queries, got {n}' 

396 ) # pragma: no cover 

397 

398 # Useful for debugging, disabled by default 

399 if log_to_file: 

400 with open('queries.txt', 'w', encoding='utf-8') as f: 

401 for q in context.captured_queries: 

402 f.write(str(q['sql']) + '\n') 

403 

404 if verbose and n >= value: 

405 msg = f'\r\n{json.dumps(context.captured_queries, indent=4)}' # pragma: no cover 

406 else: 

407 msg = None 

408 

409 if url and n > self.WARNING_QUERY_THRESHOLD: 

410 print(f'Warning: {n} queries executed at {url}') 

411 

412 self.assertLess(n, value, msg=msg) 

413 

414 

415class PluginRegistryMixin: 

416 """Mixin to ensure that the plugin registry is ready for tests.""" 

417 

418 @classmethod 

419 def setUpTestData(cls): 

420 """Ensure that the plugin registry is ready for tests.""" 

421 from time import sleep 

422 

423 from common.models import InvenTreeSetting 

424 from plugin.registry import registry 

425 

426 while not registry.is_ready: 

427 print('Waiting for plugin registry to be ready...') 

428 sleep(0.1) 

429 

430 assert registry.is_ready, 'Plugin registry is not ready' 

431 

432 InvenTreeSetting.build_default_values() 

433 super().setUpTestData() 

434 

435 def ensurePluginsLoaded(self, force: bool = False): 

436 """Helper function to ensure that plugins are loaded.""" 

437 from plugin.models import PluginConfig 

438 

439 if force or PluginConfig.objects.count() == 0: 

440 # Reload the plugin registry at this point to ensure all PluginConfig objects are created 

441 # This is because the django test system may have re-initialized the database (to an empty state) 

442 registry.reload_plugins(full_reload=True, force_reload=True, collect=True) 

443 

444 assert PluginConfig.objects.count() > 0, 'No plugins are installed' 

445 

446 

447class InvenTreeTestCase(ExchangeRateMixin, PluginRegistryMixin, UserMixin, TestCase): 

448 """Testcase with user setup build in.""" 

449 

450 

451class InvenTreeAPITestCase( 

452 ExchangeRateMixin, PluginRegistryMixin, TestQueryMixin, UserMixin, APITestCase 

453): 

454 """Base class for running InvenTree API tests.""" 

455 

456 def check_response(self, url, response, expected_code=None, msg=None): 

457 """Debug output for an unexpected response.""" 

458 # Check that the response returned the expected status code 

459 

460 if expected_code is not None: 

461 if expected_code != response.status_code: # pragma: no cover 

462 print( 

463 f"Unexpected response at '{url}': status_code = {response.status_code} (expected {expected_code})" 

464 ) 

465 

466 if hasattr(response, 'data'): 

467 print('data:', response.data) 

468 if hasattr(response, 'body'): 

469 print('body:', response.body) 

470 if hasattr(response, 'content'): 

471 print('content:', response.content) 

472 

473 self.assertEqual(response.status_code, expected_code, msg) 

474 

475 def getActions(self, url): 

476 """Return a dict of the 'actions' available at a given endpoint. 

477 

478 Makes use of the HTTP 'OPTIONS' method to request this. 

479 """ 

480 response = self.client.options(url) 

481 self.assertEqual(response.status_code, 200) 

482 

483 actions = response.data.get('actions', {}) 

484 return actions 

485 

486 def query(self, url, method, data=None, **kwargs): 

487 """Perform a generic API query.""" 

488 if data is None: 

489 data = {} 

490 

491 kwargs['format'] = kwargs.get('format', 'json') 

492 

493 expected_code = kwargs.pop('expected_code', None) 

494 msg = kwargs.pop('msg', None) 

495 max_queries = kwargs.pop('max_query_count', self.MAX_QUERY_COUNT) 

496 max_query_time = kwargs.pop('max_query_time', self.MAX_QUERY_TIME) 

497 

498 t1 = time.time() 

499 

500 with self.assertNumQueriesLessThan(max_queries, url=url): 

501 response = method(url, data, **kwargs) 

502 

503 t2 = time.time() 

504 dt = t2 - t1 

505 

506 self.check_response(url, response, expected_code=expected_code, msg=msg) 

507 

508 if dt > max_query_time: 

509 print( 

510 f'Query time exceeded at {url}: Expected {max_query_time}s, got {dt}s' 

511 ) 

512 

513 self.assertLessEqual(dt, max_query_time) 

514 

515 return response 

516 

517 def get(self, url, data=None, expected_code=200, **kwargs): 

518 """Issue a GET request.""" 

519 kwargs['data'] = data 

520 

521 return self.query(url, self.client.get, expected_code=expected_code, **kwargs) 

522 

523 def post(self, url, data=None, expected_code=201, **kwargs): 

524 """Issue a POST request.""" 

525 # Default query limit is higher for POST requests, due to extra event processing 

526 kwargs['max_query_count'] = kwargs.get( 

527 'max_query_count', self.MAX_QUERY_COUNT + 100 

528 ) 

529 

530 kwargs['data'] = data 

531 

532 return self.query(url, self.client.post, expected_code=expected_code, **kwargs) 

533 

534 def delete(self, url, data=None, expected_code=204, **kwargs): 

535 """Issue a DELETE request.""" 

536 kwargs['data'] = data 

537 

538 return self.query( 

539 url, self.client.delete, expected_code=expected_code, **kwargs 

540 ) 

541 

542 def patch(self, url, data=None, expected_code=200, **kwargs): 

543 """Issue a PATCH request.""" 

544 kwargs['data'] = data or {} 

545 

546 return self.query(url, self.client.patch, expected_code=expected_code, **kwargs) 

547 

548 def put(self, url, data=None, expected_code=200, **kwargs): 

549 """Issue a PUT request.""" 

550 kwargs['data'] = data or {} 

551 

552 return self.query(url, self.client.put, expected_code=expected_code, **kwargs) 

553 

554 def options(self, url, expected_code=None, **kwargs): 

555 """Issue an OPTIONS request.""" 

556 kwargs['data'] = kwargs.get('data') 

557 

558 return self.query( 

559 url, self.client.options, expected_code=expected_code, **kwargs 

560 ) 

561 

562 def download_file( 

563 self, 

564 url, 

565 data=None, 

566 expected_code=None, 

567 expected_fn=None, 

568 decode=True, 

569 **kwargs, 

570 ): 

571 """Download a file from the server, and return an in-memory file.""" 

572 response = self.client.get(url, data=data, format='json') 

573 

574 self.check_response(url, response, expected_code=expected_code) 

575 

576 # Check that the response is of the correct type 

577 if not isinstance(response, StreamingHttpResponse): 

578 raise ValueError( 

579 'Response is not a StreamingHttpResponse object as expected' 

580 ) 

581 

582 # Extract filename 

583 disposition = response.headers['Content-Disposition'] 

584 

585 result = re.search( 

586 r'(attachment|inline); filename=[\'"]([\w\d\-.]+)[\'"]', disposition 

587 ) 

588 if not result: 

589 raise ValueError( 

590 'No filename match found in disposition' 

591 ) # pragma: no cover 

592 

593 fn = result.groups()[1] 

594 

595 if expected_fn is not None: 

596 self.assertRegex(fn, expected_fn) 

597 

598 if decode: 

599 # Decode data and return as StringIO file object 

600 file = io.StringIO() 

601 file.name = file 

602 file.write(response.getvalue().decode('UTF-8')) 

603 else: 

604 # Return a a BytesIO file object 

605 file = io.BytesIO() 

606 file.name = fn 

607 file.write(response.getvalue()) 

608 

609 file.seek(0) 

610 

611 return file 

612 

613 def export_data( 

614 self, 

615 url, 

616 params=None, 

617 export_format='csv', 

618 export_plugin='inventree-exporter', 

619 **kwargs, 

620 ): 

621 """Perform a data export operation against the provided URL. 

622 

623 Uses the 'data_exporter' functionality to override the POST response. 

624 

625 Arguments: 

626 url: URL to perform the export operation against 

627 params: Dictionary of parameters to pass to the export operation 

628 export_format: Export format (default = 'csv') 

629 export_plugin: Export plugin (default = 'inventree-exporter') 

630 

631 Returns: 

632 A file object containing the exported dataset 

633 """ 

634 # Ensure that the plugin registry is up-to-date 

635 registry.reload_plugins(full_reload=True, force_reload=True, collect=True) 

636 

637 download = kwargs.pop('download', True) 

638 expected_code = kwargs.pop('expected_code', 200) 

639 

640 if not params: 

641 params = {} 

642 

643 params = { 

644 **params, 

645 'export': True, 

646 'export_format': export_format, 

647 'export_plugin': export_plugin, 

648 } 

649 

650 # Add in any other export specific kwargs 

651 for key, value in kwargs.items(): 

652 if key.startswith('export_'): 

653 params[key] = value 

654 

655 # Append URL params 

656 url += '?' + '&'.join([f'{key}={value}' for key, value in params.items()]) 

657 

658 response = self.get( 

659 url, data=None, format='json', expected_code=expected_code, **kwargs 

660 ) 

661 self.check_response(url, response, expected_code=expected_code) 

662 

663 # Check that the response is of the correct type 

664 data = response.data 

665 

666 if expected_code != 200: 

667 # Response failed 

668 return response.data 

669 

670 self.assertEqual(data['plugin'], export_plugin) 

671 self.assertTrue(data['complete']) 

672 filename = data.get('output') 

673 self.assertIsNotNone(filename) 

674 

675 if download: 

676 return self.download_file(filename, **kwargs) 

677 

678 else: 

679 return response.data 

680 

681 def process_csv( 

682 self, 

683 file_object, 

684 delimiter=',', 

685 required_cols=None, 

686 excluded_cols=None, 

687 required_rows=None, 

688 ): 

689 """Helper function to process and validate a downloaded csv file.""" 

690 # Check that the correct object type has been passed 

691 self.assertIsInstance(file_object, io.StringIO) 

692 

693 file_object.seek(0) 

694 

695 reader = csv.reader(file_object, delimiter=delimiter) 

696 

697 headers = [] 

698 rows = [] 

699 

700 for idx, row in enumerate(reader): 

701 if idx == 0: 

702 headers = row 

703 else: 

704 rows.append(row) 

705 

706 if required_cols is not None: 

707 for col in required_cols: 

708 self.assertIn(col, headers) 

709 

710 if excluded_cols is not None: 

711 for col in excluded_cols: 

712 self.assertNotIn(col, headers) 

713 

714 if required_rows is not None: 

715 self.assertEqual(len(rows), required_rows) 

716 

717 # Return the file data as a list of dict items, based on the headers 

718 data = [] 

719 

720 for row in rows: 

721 entry = {} 

722 

723 for idx, col in enumerate(headers): 

724 entry[col] = row[idx] 

725 

726 data.append(entry) 

727 

728 return data 

729 

730 def assertDictContainsSubset(self, a, b): 

731 """Assert that dictionary 'a' is a subset of dictionary 'b'.""" 

732 self.assertEqual(b, b | a) 

733 

734 def run_ordering_test( 

735 self, url: str, ordering_field: str, params: Optional[dict] = None 

736 ): 

737 """Run a test to check that the results are ordered correctly. 

738 

739 Arguments: 

740 url: The URL to test 

741 ordering_field: The field to order by (e.g. 'name') 

742 params: Additional parameters to include in the request (e.g. filters) 

743 

744 Process: 

745 - Run a GET request against the provided URL with the appropriate ordering parameter 

746 - Run a separate GET request with the opposite ordering parameter (e.g. '-name') 

747 - Check that the results are ordered differently in each case 

748 """ 

749 query_params = {**(params or {})} 

750 

751 pk_values = set() 

752 

753 for ordering in [None, ordering_field, f'-{ordering_field}']: 

754 response = self.get( 

755 url, 

756 data={**query_params, 'ordering': ordering} 

757 if ordering 

758 else query_params, 

759 expected_code=200, 

760 ) 

761 

762 self.assertGreater( 

763 len(response.data), 

764 1, 

765 f'No data returned from {url} with ordering={ordering}', 

766 ) 

767 

768 pk_values.add(response.data[0]['pk']) 

769 

770 self.assertGreater( 

771 len(pk_values), 

772 1, 

773 f"Ordering by '{ordering_field}' does not change the order of results at {url}", 

774 ) 

775 

776 def run_output_test( 

777 self, 

778 url: str, 

779 test_cases: list[tuple[str, str] | str], 

780 additional_params: Optional[dict] = None, 

781 assert_subset: bool = False, 

782 assert_fnc: Optional[Callable] = None, 

783 ): 

784 """Run a series of tests against the provided URL. 

785 

786 Arguments: 

787 url: The URL to test 

788 test_cases: A list of tuples of the form (parameter_name, response_field_name) 

789 additional_params: Additional request parameters to include in the request 

790 assert_subset: If True, make the assertion against the first item in the response rather than the entire response 

791 assert_fnc: If provided, call this function with the response data and make the assertion against the return value 

792 """ 

793 

794 def get_response(response): 

795 if assert_subset: 

796 return response.data[0] 

797 if assert_fnc: 

798 return assert_fnc(response) 

799 return response.data 

800 

801 for case in test_cases: 

802 if isinstance(case, str): 

803 param = case 

804 field = case 

805 else: 

806 param, field = case 

807 # Test with parameter set to 'true' 

808 response = self.get( 

809 url, 

810 {param: 'true', **(additional_params or {})}, 

811 expected_code=200, 

812 msg=f'Testing {param}=true returns anything but 200', 

813 ) 

814 self.assertIn( 

815 field, 

816 get_response(response), 

817 f"Field '{field}' should be present when {param}=true", 

818 ) 

819 

820 # Test with parameter set to 'false' 

821 response = self.get( 

822 url, 

823 {param: 'false', **(additional_params or {})}, 

824 expected_code=200, 

825 msg=f'Testing {param}=false returns anything but 200', 

826 ) 

827 self.assertNotIn( 

828 field, 

829 get_response(response), 

830 f"Field '{field}' should NOT be present when {param}=false", 

831 ) 

832 

833 

834@override_settings( 

835 SITE_URL='http://testserver', CSRF_TRUSTED_ORIGINS=['http://testserver'] 

836) 

837class AdminTestCase(InvenTreeAPITestCase): 

838 """Tests for the admin interface integration.""" 

839 

840 superuser = True 

841 

842 def helper(self, model: type[models.Model], model_kwargs=None): 

843 """Test the admin URL.""" 

844 if model_kwargs is None: 

845 model_kwargs = {} 

846 

847 # Add object 

848 obj = model.objects.create(**model_kwargs) 

849 app_app, app_mdl = model._meta.app_label, model._meta.model_name 

850 

851 # 'Test listing 

852 response = self.get( 

853 reverse(f'admin:{app_app}_{app_mdl}_changelist'), max_query_count=300 

854 ) 

855 self.assertEqual(response.status_code, 200) 

856 

857 # Test change view 

858 response = self.get( 

859 reverse(f'admin:{app_app}_{app_mdl}_change', kwargs={'object_id': obj.pk}), 

860 max_query_count=300, 

861 ) 

862 self.assertEqual(response.status_code, 200) 

863 self.assertContains(response, 'Django site admin') 

864 

865 return obj 

866 

867 

868def in_env_context(envs): 

869 """Patch the env to include the given dict.""" 

870 return mock.patch.dict(os.environ, envs) 

871 

872 

873@tag('performance_test') 

874class InvenTreeAPIPerformanceTestCase(InvenTreeAPITestCase): 

875 """Base class for InvenTree API performance tests.""" 

876 

877 MAX_QUERY_COUNT = 50 

878 MAX_QUERY_TIME = 60