Coverage for rest_api/test/mixins.py: 37%
225 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-10 03:22 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-10 03:22 +0000
1import json
2import os
3from urllib.parse import urlencode
5from django.contrib.auth import models as django_models
6from django.test import TestCase
7from django.urls import reverse
8from rest_framework import status
9from rest_framework.test import APIRequestFactory
10from rest_framework.test import APITestCase
11from rest_framework.test import APITransactionTestCase
12from rest_framework.test import force_authenticate
15class FixtureModelTestCase(TestCase):
16 fixtures = [
17 "initial_auth",
18 "test_data_10",
19 ]
20 model = None
21 obj = None
23 def get_object(self, *args):
24 if self.obj is None and self.model is not None:
25 self.obj = self.model.objects.first()
26 return self.obj
29class FixtureApiMixin(object):
30 def get_request_user(self, *args):
31 if self.request_user_name is None:
32 self.request_user_name = "root" # see test_data.json
33 return django_models.User.objects.get(username=self.request_user_name)
35 def get_user(self, username, password=None, *args):
36 user = self.usermap.get(username)
37 if user is None:
38 try:
39 user = django_models.User.objects.get(username=username)
40 except django_models.User.DoesNotExist:
41 print(
42 f"Users: {', '.join(map(lambda u: u.username, django_models.User.objects.all()))}"
43 )
44 self.fail(f"no user found with username {username}")
46 self.usermap[username] = user
47 if password is not None:
48 user.set_password(password)
49 user.save()
50 return user
53class FixtureAPITestCase(APITestCase, FixtureApiMixin):
54 fixtures = ["initial_auth", "test_data_10"]
56 factory = APIRequestFactory()
57 request_user_name = None
58 usermap = dict()
61class FixtureAPITransactionTestCase(APITransactionTestCase, FixtureApiMixin):
62 fixtures = ["initial_auth", "test_data_10"]
64 factory = APIRequestFactory()
65 request_user_name = None
66 usermap = dict()
69class CustomTestMixin:
70 model = None
71 viewset = None
72 arguments = {}
73 reverse_urls = []
75 def test_reverse_url_lookup(self, *args):
76 for name, args, url in self.reverse_urls:
77 self.assertEqual(reverse(name, args=args), url)
79 def get_instance_or_none(self, pk, *args):
80 try:
81 return self.model.objects.get(pk=pk)
82 except self.model.DoesNotExist:
83 return None
86class ListViewSetTestMixin(CustomTestMixin):
87 expected_list_count = 0
89 def get_list_response(self, parameter, *args):
90 view = self.viewset.as_view({"get": "list"})
92 request = self.factory.get(parameter)
93 force_authenticate(request, user=self.get_request_user())
94 return view(request, **self.arguments)
96 def test_list(self, *args):
97 response = self.get_list_response("")
98 self.assertEqual(response.status_code, status.HTTP_200_OK)
99 self.assertEqual(response.data["count"], self.expected_list_count)
101 def _test_list_filtering(self, filter, *args):
102 response = self.get_list_response(f"?{filter}")
103 self.assertEqual(response.status_code, status.HTTP_200_OK, response.data)
104 self.assertIn("count", response.data)
105 count = (int)(response.data["count"])
106 self.assertLess(count, self.expected_list_count)
109class RetrieveViewSetTestMixin(CustomTestMixin):
110 pk_name = "id"
112 def get_instance_and_lookup_field(self, pk=None, *args):
113 if pk is not None:
114 instance = self.model.objects.get(pk=pk)
115 else:
116 instance = self.model.objects.first()
117 lookup_field = {
118 self.viewset.lookup_field: instance.serializable_value(
119 self.viewset.lookup_field
120 )
121 }
122 return instance, lookup_field
124 def test_retrieve(self, *args):
125 self._test_retrieve()
127 def _test_retrieve(self, pk=None, *args, **kwargs):
128 view = self.viewset.as_view({"get": "retrieve"})
130 instance, lookup_field = self.get_instance_and_lookup_field(pk=pk)
131 request = self.factory.get("")
132 force_authenticate(request, user=self.get_request_user())
133 response = view(request, **lookup_field, **kwargs)
134 self.assertEqual(response.status_code, status.HTTP_200_OK)
135 self.assertIn(self.pk_name, response.data, response.data)
136 self.assertEqual(response.data[self.pk_name], instance.pk, response.data)
139class UpdateViewSetTestMixin(RetrieveViewSetTestMixin):
140 def _test_update(self, name, filename, pk=None, url="", *args):
141 with open(filename, mode="r") as f:
142 params = json.loads(f.read())
143 if not params.get("filters") or not params.get("data"):
144 self.fail(f"File {filename} needs filters and data to test update.")
145 filters = params.get("filters")
146 data = params.get("data")
147 if pk is None:
148 if "pk" in filters:
149 pk = filters["pk"]
150 elif "id" in filters:
151 pk = filters["id"]
152 view = self.viewset.as_view({"patch": "partial_update"})
153 request = self.factory.patch(
154 f"{url}?{urlencode(filters)}", data=data, format="json"
155 )
157 force_authenticate(request, user=self.get_request_user())
158 response = view(request, pk=pk)
159 response.render()
160 responseJSON = json.loads(response.content)
161 self.assertEqual(response.status_code, status.HTTP_200_OK, response.data)
162 for key in filters:
163 self.assertIn(key, responseJSON)
164 self.assertEqual(responseJSON[key], filters[key])
165 for key in data:
166 if key == "version" or key == "change_reason":
167 continue
168 self.assertIn(key, responseJSON)
169 self.assertEqual(responseJSON[key], data[key])
170 self.assertEqual(response.status_code, status.HTTP_200_OK, response.data)
172 def _test_update_error_examples(self, name, filename, pk=None, *args):
173 with open(filename, mode="r") as f:
174 test_content = json.loads(f.read())
175 if not test_content["params"]["filters"] or not test_content["params"]["data"]:
176 self.fail(f"File {filename} needs filters and data to test update.")
177 filters = test_content["params"]["filters"]
178 data = test_content["params"]["data"]
179 if pk is None:
180 if "pk" in filters:
181 pk = filters["pk"]
182 elif "id" in filters:
183 pk = filters["id"]
184 view = self.viewset.as_view({"patch": "partial_update"})
185 request = self.factory.patch(f"?{urlencode(filters)}", data=data, format="json")
187 force_authenticate(request, user=self.get_request_user())
188 response = view(request, pk=pk)
189 response.render()
190 responseJSON = json.loads(response.content)
191 error_example = test_content["expected_response"]
192 self.assertEqual(
193 response.status_code,
194 error_example["expected_status_code"],
195 msg=response.data,
196 )
197 if "expected_response_data" in error_example:
198 self.assertDictEqual(responseJSON, error_example["expected_response_data"])
199 if "expected_response_schema" in error_example:
200 try:
201 json_validate(error_example["expected_response_schema"], responseJSON)
202 except JsonSchemaException as e:
203 self.fail(e)
206class DestroyViewSetTestMixin(CustomTestMixin):
207 def test_destroy(self, *args):
208 self._test_destroy("", 1)
210 def _test_destroy(self, name, pk, username=None, *args):
211 view = self.viewset.as_view({"delete": "destroy"})
212 instance, lookup_field = self.get_instance_and_lookup_field(pk=pk)
214 request = self.factory.delete("", format="json")
215 if username is None:
216 force_authenticate(request, user=self.get_request_user())
217 else:
218 force_authenticate(request, user=self.get_user(username))
219 response = view(request, **lookup_field)
220 self.assertEqual(
221 response.status_code, status.HTTP_204_NO_CONTENT, response.data
222 )
223 instance = self.get_instance_or_none(pk)
224 self.assertIsNone(instance)
227class CreateViewSetTestMixin(CustomTestMixin):
228 def _test_create(self, _, filename, url="", *args):
229 view = self.viewset.as_view({"post": "create"})
230 with open(filename, mode="r") as f:
231 data = json.loads(f.read())
232 request = self.factory.post(url, data=data, format="json")
233 force_authenticate(request, user=self.get_request_user())
235 objects_before = self.model.objects.count()
236 response = view(request)
237 self.assertEqual(response.status_code, status.HTTP_201_CREATED, response.data)
238 self.assertEqual(objects_before + 1, self.model.objects.count())
240 def _test_error_examples(self, _, filename, *args, action="create", pk=None):
241 view = self.viewset.as_view({"post": action})
242 with open(filename, mode="r") as f:
243 error_example = json.loads(f.read())
244 request = self.factory.post("", data=error_example["data"], format="json")
245 force_authenticate(request, user=self.get_request_user())
246 response = view(request, pk=pk)
247 self.assertEqual(
248 response.status_code,
249 error_example["expected_status_code"],
250 msg=response.data,
251 )
252 if "expected_response_data" in error_example:
253 self.assertDictEqual(response.data, error_example["expected_response_data"])
254 if "expected_response_schema" in error_example:
255 try:
256 json_validate(error_example["expected_response_schema"], response.data)
257 except JsonSchemaException as e:
258 self.fail(e)
261class ListSchemaTestMixin(CustomTestMixin):
262 def _test_list_schema(self, schema_file_prefix, version, *args):
263 schema_file = os.path.join(
264 "transfer-protocol", "schemata", f"{schema_file_prefix}_v{version}.json"
265 )
266 view = self.viewset.as_view({"get": "list"})
267 request = self.factory.get(
268 "", HTTP_ACCEPT=f"application/json; version={version}"
269 )
270 force_authenticate(request, user=self.get_request_user())
272 response = view(request, **self.arguments)
273 self.assertEqual(response.status_code, status.HTTP_200_OK)
274 with open(schema_file, mode="r") as f:
275 json_schema = json.loads(f.read())
277 try:
278 json_validate(json_schema, response.data)
279 except JsonSchemaException as e:
280 self.fail(f"{e}: {response.data}")
283class OrderingTestMixin(CustomTestMixin):
284 def _test_list_ordering(self, orderingfield, *args):
285 view = self.viewset.as_view({"get": "list"})
286 request_asc = self.factory.get(f"?ordering={orderingfield}")
287 request_desc = self.factory.get(f"?ordering=-{orderingfield}")
288 force_authenticate(request_asc, user=self.get_request_user())
289 force_authenticate(request_desc, user=self.get_request_user())
290 response_asc = view(request_asc, **self.arguments)
291 response_desc = view(request_desc, **self.arguments)
292 self.assertEqual(
293 response_asc.status_code, status.HTTP_200_OK, response_asc.data
294 )
295 self.assertEqual(
296 response_desc.status_code, status.HTTP_200_OK, response_desc.data
297 )
298 self.assertIn("count", response_asc.data, response_asc.data)
299 if response_asc.data["count"] < 2:
300 self.skipTest("Not enough objects in test_data")
301 else:
302 self.assertIn("results", response_asc.data, response_asc.data)
303 self.assertIn("results", response_desc.data, response_desc.data)
304 first_result = response_asc.data["results"][0]
305 last_result = response_desc.data["results"][0]
306 self.assertNotEqual(first_result, last_result)