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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 17:47 +0000
1"""Helper functions for unit testing / CI."""
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
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
23from djmoney.contrib.exchange.models import ExchangeBackend, Rate
24from rest_framework.test import APITestCase
26from plugin import registry
27from plugin.models import PluginConfig
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.
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()
47 with CaptureQueriesContext(connections[using]) as context:
48 yield
50 dt = time.time() - t1
52 n = len(context.captured_queries)
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')
59 output = f'Executed {n} queries in {dt:.4f}s'
61 if threshold and n >= threshold:
62 if msg:
63 print(f'{msg}: {output}')
64 else:
65 print(output)
68def addUserPermission(user: User, app_name: str, model_name: str, perm: str) -> None:
69 """Add a specific permission for the provided user.
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 )
82 # Add the permission to the user
83 user.user_permissions.add(permission)
84 user.save()
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()
92 # Regex pattern for migration files
93 regex = re.compile(r'^[\d]+_.*\.py$')
95 migration_files = []
97 for f in files:
98 if regex.match(f.name):
99 migration_files.append(f.name)
101 return migration_files
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
109 for f in getMigrationFileNames(app):
110 if ignore_initial and f.startswith('0001_initial'):
111 continue
113 num = int(f.split('_')[0])
115 if oldest_file is None or num < oldest_num:
116 oldest_num = num
117 oldest_file = f
119 if exclude_extension and oldest_file:
120 oldest_file = oldest_file.replace('.py', '')
122 return oldest_file
125def getNewestMigrationFile(app, exclude_extension=True):
126 """Return the filename associated with the newest migration."""
127 newest_file = None
128 newest_num = -1
130 for f in getMigrationFileNames(app):
131 num = int(f.split('_')[0])
133 if newest_file is None or num > newest_num:
134 newest_num = num
135 newest_file = f
137 if not newest_file: # pragma: no cover
138 return newest_file
140 if exclude_extension:
141 newest_file = newest_file.replace('.py', '')
143 return newest_file
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.
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
164 tasks = OrmQ.objects.all()
166 if reverse:
167 tasks = tasks.order_by('-pk')
169 task = None
171 for t in tasks:
172 if t.func() == task_name:
173 found = True
175 if matching_args:
176 for arg in matching_args:
177 if arg not in t.args():
178 found = False
179 break
181 if matching_kwargs:
182 for kwarg in matching_kwargs:
183 if kwarg not in t.kwargs():
184 found = False
185 break
187 if found:
188 task = t
189 break
191 if clear_after:
192 OrmQ.objects.all().delete()
194 return task
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 )
213class UserMixin:
214 """Mixin to setup a user and login for tests.
216 Use parameters to set username, password, email, roles and permissions.
217 """
219 # User information
220 username = 'testuser'
221 password = 'mypassword'
222 email = 'test@testing.com'
224 superuser = False
225 is_staff = True
226 auto_login = True
228 # Set list of roles automatically associated with the user
229 roles = []
231 @classmethod
232 def setUpTestData(cls):
233 """Run setup for all tests in a given class."""
234 super().setUpTestData()
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 )
241 # Create a group for the user
242 cls.group = Group.objects.create(name='my_test_group')
243 cls.user.groups.add(cls.group)
245 if cls.superuser:
246 cls.user.is_superuser = True
248 if cls.is_staff:
249 cls.user.is_staff = True
251 cls.user.save()
253 # Assign all roles if set
254 if cls.roles == 'all':
255 cls.assignRole(group=cls.group, assign_all=True)
257 # else filter the roles
258 else:
259 for role in cls.roles:
260 cls.assignRole(role=role, group=cls.group)
262 def setUp(self):
263 """Run setup for individual test methods."""
264 if self.auto_login:
265 self.login()
267 def login(self):
268 """Login with the current user credentials."""
269 self.client.login(username=self.username, password=self.password)
271 def logout(self):
272 """Lougout current user."""
273 self.client.logout()
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
284 ruleset.save()
286 @classmethod
287 def assignRole(cls, role=None, assign_all: bool = False, group=None):
288 """Set the user roles for the registered user.
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
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
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
309 if not assign_all and role:
310 rule, perm = role.split('.')
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
323 ruleset.save()
324 if not assign_all:
325 break
328class PluginMixin:
329 """Mixin to ensure that all plugins are loaded for tests."""
331 def setUp(self):
332 """Setup for plugin tests."""
333 super().setUp()
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()
343class ExchangeRateMixin:
344 """Mixin class for generating exchange rate data."""
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}
350 # Create a dummy backend
351 ExchangeBackend.objects.create(name='InvenTreeExchange', base_currency='USD')
353 backend = ExchangeBackend.objects.get(name='InvenTreeExchange')
355 items = []
357 for currency, rate in rates.items():
358 items.append(Rate(currency=currency, value=rate, backend=backend))
360 Rate.objects.bulk_create(items)
363class TestQueryMixin:
364 """Mixin class for testing query counts."""
366 # Default query count threshold value
367 # TODO: This value should be reduced
368 MAX_QUERY_COUNT = 250
370 WARNING_QUERY_THRESHOLD = 100
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
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.
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
391 n = len(context.captured_queries)
393 if url and n >= value:
394 print(
395 f'Query count exceeded at {url}: Expected < {value} queries, got {n}'
396 ) # pragma: no cover
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')
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
409 if url and n > self.WARNING_QUERY_THRESHOLD:
410 print(f'Warning: {n} queries executed at {url}')
412 self.assertLess(n, value, msg=msg)
415class PluginRegistryMixin:
416 """Mixin to ensure that the plugin registry is ready for tests."""
418 @classmethod
419 def setUpTestData(cls):
420 """Ensure that the plugin registry is ready for tests."""
421 from time import sleep
423 from common.models import InvenTreeSetting
424 from plugin.registry import registry
426 while not registry.is_ready:
427 print('Waiting for plugin registry to be ready...')
428 sleep(0.1)
430 assert registry.is_ready, 'Plugin registry is not ready'
432 InvenTreeSetting.build_default_values()
433 super().setUpTestData()
435 def ensurePluginsLoaded(self, force: bool = False):
436 """Helper function to ensure that plugins are loaded."""
437 from plugin.models import PluginConfig
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)
444 assert PluginConfig.objects.count() > 0, 'No plugins are installed'
447class InvenTreeTestCase(ExchangeRateMixin, PluginRegistryMixin, UserMixin, TestCase):
448 """Testcase with user setup build in."""
451class InvenTreeAPITestCase(
452 ExchangeRateMixin, PluginRegistryMixin, TestQueryMixin, UserMixin, APITestCase
453):
454 """Base class for running InvenTree API tests."""
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
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 )
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)
473 self.assertEqual(response.status_code, expected_code, msg)
475 def getActions(self, url):
476 """Return a dict of the 'actions' available at a given endpoint.
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)
483 actions = response.data.get('actions', {})
484 return actions
486 def query(self, url, method, data=None, **kwargs):
487 """Perform a generic API query."""
488 if data is None:
489 data = {}
491 kwargs['format'] = kwargs.get('format', 'json')
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)
498 t1 = time.time()
500 with self.assertNumQueriesLessThan(max_queries, url=url):
501 response = method(url, data, **kwargs)
503 t2 = time.time()
504 dt = t2 - t1
506 self.check_response(url, response, expected_code=expected_code, msg=msg)
508 if dt > max_query_time:
509 print(
510 f'Query time exceeded at {url}: Expected {max_query_time}s, got {dt}s'
511 )
513 self.assertLessEqual(dt, max_query_time)
515 return response
517 def get(self, url, data=None, expected_code=200, **kwargs):
518 """Issue a GET request."""
519 kwargs['data'] = data
521 return self.query(url, self.client.get, expected_code=expected_code, **kwargs)
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 )
530 kwargs['data'] = data
532 return self.query(url, self.client.post, expected_code=expected_code, **kwargs)
534 def delete(self, url, data=None, expected_code=204, **kwargs):
535 """Issue a DELETE request."""
536 kwargs['data'] = data
538 return self.query(
539 url, self.client.delete, expected_code=expected_code, **kwargs
540 )
542 def patch(self, url, data=None, expected_code=200, **kwargs):
543 """Issue a PATCH request."""
544 kwargs['data'] = data or {}
546 return self.query(url, self.client.patch, expected_code=expected_code, **kwargs)
548 def put(self, url, data=None, expected_code=200, **kwargs):
549 """Issue a PUT request."""
550 kwargs['data'] = data or {}
552 return self.query(url, self.client.put, expected_code=expected_code, **kwargs)
554 def options(self, url, expected_code=None, **kwargs):
555 """Issue an OPTIONS request."""
556 kwargs['data'] = kwargs.get('data')
558 return self.query(
559 url, self.client.options, expected_code=expected_code, **kwargs
560 )
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')
574 self.check_response(url, response, expected_code=expected_code)
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 )
582 # Extract filename
583 disposition = response.headers['Content-Disposition']
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
593 fn = result.groups()[1]
595 if expected_fn is not None:
596 self.assertRegex(fn, expected_fn)
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())
609 file.seek(0)
611 return file
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.
623 Uses the 'data_exporter' functionality to override the POST response.
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')
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)
637 download = kwargs.pop('download', True)
638 expected_code = kwargs.pop('expected_code', 200)
640 if not params:
641 params = {}
643 params = {
644 **params,
645 'export': True,
646 'export_format': export_format,
647 'export_plugin': export_plugin,
648 }
650 # Add in any other export specific kwargs
651 for key, value in kwargs.items():
652 if key.startswith('export_'):
653 params[key] = value
655 # Append URL params
656 url += '?' + '&'.join([f'{key}={value}' for key, value in params.items()])
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)
663 # Check that the response is of the correct type
664 data = response.data
666 if expected_code != 200:
667 # Response failed
668 return response.data
670 self.assertEqual(data['plugin'], export_plugin)
671 self.assertTrue(data['complete'])
672 filename = data.get('output')
673 self.assertIsNotNone(filename)
675 if download:
676 return self.download_file(filename, **kwargs)
678 else:
679 return response.data
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)
693 file_object.seek(0)
695 reader = csv.reader(file_object, delimiter=delimiter)
697 headers = []
698 rows = []
700 for idx, row in enumerate(reader):
701 if idx == 0:
702 headers = row
703 else:
704 rows.append(row)
706 if required_cols is not None:
707 for col in required_cols:
708 self.assertIn(col, headers)
710 if excluded_cols is not None:
711 for col in excluded_cols:
712 self.assertNotIn(col, headers)
714 if required_rows is not None:
715 self.assertEqual(len(rows), required_rows)
717 # Return the file data as a list of dict items, based on the headers
718 data = []
720 for row in rows:
721 entry = {}
723 for idx, col in enumerate(headers):
724 entry[col] = row[idx]
726 data.append(entry)
728 return data
730 def assertDictContainsSubset(self, a, b):
731 """Assert that dictionary 'a' is a subset of dictionary 'b'."""
732 self.assertEqual(b, b | a)
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.
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)
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 {})}
751 pk_values = set()
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 )
762 self.assertGreater(
763 len(response.data),
764 1,
765 f'No data returned from {url} with ordering={ordering}',
766 )
768 pk_values.add(response.data[0]['pk'])
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 )
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.
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 """
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
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 )
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 )
834@override_settings(
835 SITE_URL='http://testserver', CSRF_TRUSTED_ORIGINS=['http://testserver']
836)
837class AdminTestCase(InvenTreeAPITestCase):
838 """Tests for the admin interface integration."""
840 superuser = True
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 = {}
847 # Add object
848 obj = model.objects.create(**model_kwargs)
849 app_app, app_mdl = model._meta.app_label, model._meta.model_name
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)
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')
865 return obj
868def in_env_context(envs):
869 """Patch the env to include the given dict."""
870 return mock.patch.dict(os.environ, envs)
873@tag('performance_test')
874class InvenTreeAPIPerformanceTestCase(InvenTreeAPITestCase):
875 """Base class for InvenTree API performance tests."""
877 MAX_QUERY_COUNT = 50
878 MAX_QUERY_TIME = 60