Coverage for utilities/testing/base.py: 0%
126 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 18:35 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 18:35 +0000
1import json
2from contextlib import contextmanager
4from django.contrib.contenttypes.fields import GenericForeignKey
5from django.contrib.contenttypes.models import ContentType
6from django.contrib.postgres.fields import ArrayField, RangeField
7from django.core.exceptions import FieldDoesNotExist
8from django.db import transaction
9from django.db.models import DateField, DateTimeField, JSONField, ManyToManyField, ManyToManyRel
10from django.forms.models import model_to_dict
11from django.test import Client
12from django.test import TestCase as _TestCase
13from netaddr import IPNetwork
14from taggit.managers import TaggableManager
16from core.choices import ObjectChangeActionChoices
17from core.models import ObjectType
18from users.models import ObjectPermission, User
19from utilities.data import ranges_to_string
20from utilities.object_types import object_type_identifier
21from utilities.permissions import resolve_permission_type
23from .utils import DUMMY_CF_DATA, extract_form_failures
25__all__ = (
26 'ModelTestCase',
27 'TestCase',
28)
31class TestCase(_TestCase):
32 user_permissions = ()
34 def setUp(self):
36 # Create the test user and assign permissions
37 self.user = User.objects.create_user(username='testuser')
38 self.add_permissions(*self.user_permissions)
40 # Initialize the test client
41 self.client = Client()
42 self.client.force_login(self.user)
44 @contextmanager
45 def cleanupSubTest(self, **params):
46 """
47 Context manager that wraps subTest with automatic cleanup.
48 All database changes within the context will be rolled back.
49 """
50 sid = transaction.savepoint_create()
52 try:
53 with self.subTest(**params):
54 yield
55 finally:
56 transaction.savepoint_rollback(sid)
58 #
59 # Permissions management
60 #
62 def add_permissions(self, *names):
63 """
64 Assign a set of permissions to the test user. Accepts permission names in the form <app>.<action>_<model>.
65 """
66 for name in names:
67 object_type, action = resolve_permission_type(name)
68 obj_perm = ObjectPermission(name=name, actions=[action])
69 obj_perm.save()
70 obj_perm.users.add(self.user)
71 obj_perm.object_types.add(object_type)
73 def remove_permissions(self, *names):
74 """
75 Remove a set of permissions from the test user. Accepts permission names in the form <app>.<action>_<model>.
76 """
77 for name in names:
78 object_type, action = resolve_permission_type(name)
79 ObjectPermission.objects.filter(
80 actions__contains=[action], object_types=object_type, users=self.user
81 ).delete()
83 #
84 # Custom assertions
85 #
87 def assertObjectChange(self, objectchange, *, action, message=None):
88 """
89 Assert that an ObjectChange record has the expected attributes. If message is provided, it will be
90 compared against objectchange.message.
91 """
92 # Verify the change action (create, update, delete)
93 self.assertEqual(objectchange.action, action)
95 # Verify the changelog message if provided
96 if message is not None:
97 self.assertEqual(objectchange.message, message)
99 # Verify pre/postchange data presence and integrity based on action type
100 if action == ObjectChangeActionChoices.ACTION_CREATE:
101 self.assertIsNone(objectchange.prechange_data, "Expected prechange_data to be None for a create")
102 self.assertIsNotNone(objectchange.postchange_data, "Expected postchange_data to be populated for a create")
103 elif action == ObjectChangeActionChoices.ACTION_UPDATE:
104 self.assertIsNotNone(objectchange.prechange_data, "Expected prechange_data to be populated for an update")
105 self.assertIsNotNone(objectchange.postchange_data, "Expected postchange_data to be populated for an update")
106 self.assertNotEqual(objectchange.prechange_data, objectchange.postchange_data)
107 elif action == ObjectChangeActionChoices.ACTION_DELETE:
108 self.assertIsNotNone(objectchange.prechange_data, "Expected prechange_data to be populated for a delete")
109 self.assertIsNone(objectchange.postchange_data, "Expected postchange_data to be None for a delete")
111 def assertHttpStatus(self, response, expected_status):
112 """
113 TestCase method. Provide more detail in the event of an unexpected HTTP response.
114 """
115 err_message = None
116 # Construct an error message only if we know the test is going to fail
117 if response.status_code != expected_status:
118 if hasattr(response, 'data'):
119 # REST API response; pass the response data through directly
120 err = response.data
121 else:
122 # Attempt to extract form validation errors from the response HTML
123 form_errors = extract_form_failures(response.content)
124 err = form_errors or response.content or 'No data'
125 err_message = f"Expected HTTP status {expected_status}; received {response.status_code}: {err}"
126 self.assertEqual(response.status_code, expected_status, err_message)
128 def assertNotCacheable(self, response):
129 """
130 TestCase method. Assert that a response instructs the browser not to persist its content
131 to the local cache. Views which render potentially sensitive content (e.g. the contents of
132 a synced data file) must not be written to the browser's cache, where they would remain
133 readable after the session has ended.
134 """
135 cache_control = response.headers.get('Cache-Control', '')
136 self.assertIn(
137 'no-store',
138 cache_control,
139 f"Expected a no-store cache directive; received Cache-Control: '{cache_control}'"
140 )
143class ModelTestCase(TestCase):
144 """
145 Parent class for TestCases which deal with models.
146 """
147 model = None
149 def _get_queryset(self):
150 """
151 Return a base queryset suitable for use in test methods.
152 """
153 return self.model.objects.all()
155 def prepare_instance(self, instance):
156 """
157 Test cases can override this method to perform any necessary manipulation of an instance prior to its evaluation
158 against test data. For example, it can be used to decrypt a Secret's plaintext attribute.
159 """
160 return instance
162 def model_to_dict(self, instance, fields, api=False):
163 """
164 Return a dictionary representation of an instance.
165 """
166 # Prepare the instance and call Django's model_to_dict() to extract all fields
167 model_dict = model_to_dict(self.prepare_instance(instance), fields=fields)
169 # Map any additional (non-field) instance attributes that were specified
170 for attr in fields:
171 if hasattr(instance, attr) and attr not in model_dict:
172 model_dict[attr] = getattr(instance, attr)
174 for key, value in list(model_dict.items()):
175 try:
176 field = instance._meta.get_field(key)
177 except FieldDoesNotExist:
178 # Attribute is not a model field
179 continue
181 # Handle ManyToManyFields
182 if value and (
183 type(field) in (ManyToManyField, ManyToManyRel) or isinstance(field, TaggableManager)
184 ):
185 # Resolve reverse M2M relationships
186 if isinstance(field, ManyToManyRel):
187 value = getattr(instance, field.related_name).all()
188 if field.related_model in (ContentType, ObjectType) and api:
189 model_dict[key] = sorted([object_type_identifier(ot) for ot in value])
190 else:
191 model_dict[key] = sorted([obj.pk for obj in value])
193 # Handle GenericForeignKeys
194 elif value and type(field) is GenericForeignKey:
195 model_dict[key] = value.pk
197 # Handle API output
198 elif api:
199 # Replace ContentType numeric IDs with <app_label>.<model>
200 if type(getattr(instance, key)) in (ContentType, ObjectType):
201 object_type = ObjectType.objects.get(pk=value)
202 model_dict[key] = object_type_identifier(object_type)
204 # Convert IPNetwork instances to strings
205 elif type(value) is IPNetwork:
206 model_dict[key] = str(value)
208 # Convert date values to ISO 8601 strings (as rendered by the REST API). DateTimeField
209 # subclasses DateField, so exclude it here to preserve existing datetime handling.
210 elif isinstance(field, DateField) and not isinstance(field, DateTimeField) and value is not None:
211 model_dict[key] = value.isoformat()
213 # Normalize arrays of numeric ranges (e.g. VLAN IDs or port ranges).
214 # DB uses canonical half-open [lo, hi) via NumericRange; API uses inclusive [lo, hi].
215 # Convert to inclusive pairs for stable API comparisons.
216 elif type(field) is ArrayField and issubclass(type(field.base_field), RangeField):
217 model_dict[key] = [[r.lower, r.upper - 1] for r in value]
219 else:
220 # Convert ArrayFields (including subclasses) to CSV strings
221 if isinstance(field, ArrayField):
222 if getattr(field.base_field, 'choices', None):
223 # Values for fields with pre-defined choices can be returned as lists
224 model_dict[key] = value
225 elif type(field.base_field) is ArrayField:
226 # Handle nested arrays (e.g. choice sets)
227 model_dict[key] = '\n'.join([f'{k},{v}' for k, v in value])
228 elif issubclass(type(field.base_field), RangeField):
229 # Handle arrays of numeric ranges (e.g. VLANGroup VLAN ID ranges)
230 model_dict[key] = ranges_to_string(value)
231 else:
232 model_dict[key] = ','.join([str(v) for v in value])
234 # JSON
235 if type(field) is JSONField and value is not None:
236 model_dict[key] = json.dumps(value)
238 return model_dict
240 #
241 # Custom assertions
242 #
244 def assertInstanceEqual(self, instance, data, exclude=None, api=False):
245 """
246 Compare a model instance to a dictionary, checking that its attribute values match those specified
247 in the dictionary.
249 :param instance: Python object instance
250 :param data: Dictionary of test data used to define the instance
251 :param exclude: List of fields to exclude from comparison (e.g. passwords, which get hashed)
252 :param api: Set to True is the data is a JSON representation of the instance
253 """
254 if exclude is None:
255 exclude = []
257 fields = [k for k in data.keys() if k not in exclude]
258 model_dict = self.model_to_dict(instance, fields=fields, api=api)
260 # Omit any dictionary keys which are not instance attributes or have been excluded
261 model_data = {
262 k: v for k, v in data.items() if hasattr(instance, k) and k not in exclude
263 }
265 self.assertDictEqual(model_dict, model_data)
267 # Validate any custom field data, if present
268 if getattr(instance, 'custom_field_data', None):
269 self.assertDictEqual(instance.custom_field_data, DUMMY_CF_DATA)