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

1import json 

2from contextlib import contextmanager 

3 

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 

15 

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 

22 

23from .utils import DUMMY_CF_DATA, extract_form_failures 

24 

25__all__ = ( 

26 'ModelTestCase', 

27 'TestCase', 

28) 

29 

30 

31class TestCase(_TestCase): 

32 user_permissions = () 

33 

34 def setUp(self): 

35 

36 # Create the test user and assign permissions 

37 self.user = User.objects.create_user(username='testuser') 

38 self.add_permissions(*self.user_permissions) 

39 

40 # Initialize the test client 

41 self.client = Client() 

42 self.client.force_login(self.user) 

43 

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

51 

52 try: 

53 with self.subTest(**params): 

54 yield 

55 finally: 

56 transaction.savepoint_rollback(sid) 

57 

58 # 

59 # Permissions management 

60 # 

61 

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) 

72 

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

82 

83 # 

84 # Custom assertions 

85 # 

86 

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) 

94 

95 # Verify the changelog message if provided 

96 if message is not None: 

97 self.assertEqual(objectchange.message, message) 

98 

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

110 

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) 

127 

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 ) 

141 

142 

143class ModelTestCase(TestCase): 

144 """ 

145 Parent class for TestCases which deal with models. 

146 """ 

147 model = None 

148 

149 def _get_queryset(self): 

150 """ 

151 Return a base queryset suitable for use in test methods. 

152 """ 

153 return self.model.objects.all() 

154 

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 

161 

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) 

168 

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) 

173 

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 

180 

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

192 

193 # Handle GenericForeignKeys 

194 elif value and type(field) is GenericForeignKey: 

195 model_dict[key] = value.pk 

196 

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) 

203 

204 # Convert IPNetwork instances to strings 

205 elif type(value) is IPNetwork: 

206 model_dict[key] = str(value) 

207 

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

212 

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] 

218 

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

233 

234 # JSON 

235 if type(field) is JSONField and value is not None: 

236 model_dict[key] = json.dumps(value) 

237 

238 return model_dict 

239 

240 # 

241 # Custom assertions 

242 # 

243 

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. 

248 

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 = [] 

256 

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) 

259 

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 } 

264 

265 self.assertDictEqual(model_dict, model_data) 

266 

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)