From a1406cddf2ed4ba457370421b43e3b50e4762a79 Mon Sep 17 00:00:00 2001 From: Aarif Date: Thu, 18 Feb 2021 18:09:45 +0500 Subject: [PATCH] replaced unittest assertions pytest assertions (#26528) --- .../terrain/stubs/tests/test_edxnotes.py | 94 +++++++++---------- .../terrain/stubs/tests/test_http.py | 24 ++--- .../terrain/stubs/tests/test_lti_stub.py | 18 ++-- .../terrain/stubs/tests/test_video.py | 6 +- .../terrain/stubs/tests/test_xqueue_stub.py | 16 ++-- .../terrain/stubs/tests/test_youtube_stub.py | 36 +++---- .../third_party_auth/tests/test_utils.py | 16 +--- .../track/backends/tests/test_mongodb.py | 6 +- .../management/tests/test_tracked_command.py | 2 +- common/djangoapps/track/tests/__init__.py | 4 +- .../djangoapps/track/tests/test_contexts.py | 17 +--- .../djangoapps/track/tests/test_middleware.py | 48 ++++------ common/djangoapps/track/tests/test_segment.py | 34 +++---- common/djangoapps/track/tests/test_shim.py | 24 ++--- common/djangoapps/track/tests/test_tracker.py | 16 ++-- common/djangoapps/track/tests/test_util.py | 12 +-- .../track/views/tests/test_segmentio.py | 18 ++-- common/djangoapps/util/testing.py | 6 +- common/djangoapps/util/tests/test_course.py | 6 +- .../djangoapps/util/tests/test_date_utils.py | 18 ++-- common/djangoapps/util/tests/test_db.py | 14 +-- .../util/tests/test_disable_rate_limit.py | 3 +- .../util/tests/test_django_utils.py | 2 +- common/djangoapps/util/tests/test_file.py | 93 ++++++++---------- .../util/tests/test_json_request.py | 86 ++++++++--------- .../util/tests/test_keyword_sub_utils.py | 22 ++--- common/djangoapps/util/tests/test_memcache.py | 26 +++-- .../util/tests/test_milestones_helpers.py | 30 +++--- .../tests/test_password_policy_validators.py | 13 +-- .../djangoapps/util/tests/test_sandboxing.py | 16 ++-- .../util/tests/test_string_utils.py | 24 ++--- .../xblock_django/tests/test_api.py | 18 ++-- .../xblock_django/tests/test_user_service.py | 37 ++++---- 33 files changed, 367 insertions(+), 438 deletions(-) diff --git a/common/djangoapps/terrain/stubs/tests/test_edxnotes.py b/common/djangoapps/terrain/stubs/tests/test_edxnotes.py index adef731833..aa27ddac83 100644 --- a/common/djangoapps/terrain/stubs/tests/test_edxnotes.py +++ b/common/djangoapps/terrain/stubs/tests/test_edxnotes.py @@ -77,38 +77,38 @@ class StubEdxNotesServiceTest(unittest.TestCase): ], } response = requests.post(self._get_url("api/v1/annotations"), data=json.dumps(dummy_note)) - self.assertTrue(response.ok) + assert response.ok response_content = response.json() - self.assertIn("id", response_content) - self.assertIn("created", response_content) - self.assertIn("updated", response_content) - self.assertIn("annotator_schema_version", response_content) + assert 'id' in response_content + assert 'created' in response_content + assert 'updated' in response_content + assert 'annotator_schema_version' in response_content self.assertDictContainsSubset(dummy_note, response_content) def test_note_read(self): notes = self._get_notes() for note in notes: response = requests.get(self._get_url("api/v1/annotations/" + note["id"])) - self.assertTrue(response.ok) + assert response.ok self.assertDictEqual(note, response.json()) response = requests.get(self._get_url("api/v1/annotations/does_not_exist")) - self.assertEqual(response.status_code, 404) + assert response.status_code == 404 def test_note_update(self): notes = self._get_notes() for note in notes: response = requests.get(self._get_url("api/v1/annotations/" + note["id"])) - self.assertTrue(response.ok) + assert response.ok self.assertDictEqual(note, response.json()) response = requests.get(self._get_url("api/v1/annotations/does_not_exist")) - self.assertEqual(response.status_code, 404) + assert response.status_code == 404 def test_search(self): # Without user response = requests.get(self._get_url("api/v1/search")) - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 # get response with default page and page size response = requests.get(self._get_url("api/v1/search"), params={ @@ -116,7 +116,7 @@ class StubEdxNotesServiceTest(unittest.TestCase): "course_id": "dummy-course-id", }) - self.assertTrue(response.ok) + assert response.ok self._verify_pagination_info( response=response.json(), total_notes=5, @@ -135,7 +135,7 @@ class StubEdxNotesServiceTest(unittest.TestCase): "text": "world war 2" }) - self.assertTrue(response.ok) + assert response.ok self._verify_pagination_info( response=response.json(), total_notes=0, @@ -160,42 +160,42 @@ class StubEdxNotesServiceTest(unittest.TestCase): 'user': 'dummy-user-id', 'course_id': 'dummy-course-id' }) - self.assertTrue(response.ok) + assert response.ok response = response.json() parsed = six.moves.urllib.parse.urlparse(url) query_params = six.moves.urllib.parse.parse_qs(parsed.query) query_params['usage_id'].reverse() - self.assertEqual(len(response), len(query_params['usage_id'])) + assert len(response) == len(query_params['usage_id']) for index, usage_id in enumerate(query_params['usage_id']): - self.assertEqual(response[index]['usage_id'], usage_id) + assert response[index]['usage_id'] == usage_id def test_delete(self): notes = self._get_notes() response = requests.delete(self._get_url("api/v1/annotations/does_not_exist")) - self.assertEqual(response.status_code, 404) + assert response.status_code == 404 for note in notes: response = requests.delete(self._get_url("api/v1/annotations/" + note["id"])) - self.assertEqual(response.status_code, 204) + assert response.status_code == 204 remaining_notes = self.server.get_all_notes() - self.assertNotIn(note["id"], [note["id"] for note in remaining_notes]) + assert note['id'] not in [note['id'] for note in remaining_notes] - self.assertEqual(len(remaining_notes), 0) + assert len(remaining_notes) == 0 def test_update(self): note = self._get_notes()[0] response = requests.put(self._get_url("api/v1/annotations/" + note["id"]), data=json.dumps({ "text": "new test text" })) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 updated_note = self._get_notes()[0] - self.assertEqual("new test text", updated_note["text"]) - self.assertEqual(note["id"], updated_note["id"]) + assert 'new test text' == updated_note['text'] + assert note['id'] == updated_note['id'] six.assertCountEqual(self, note, updated_note) response = requests.get(self._get_url("api/v1/annotations/does_not_exist")) - self.assertEqual(response.status_code, 404) + assert response.status_code == 404 # pylint: disable=too-many-arguments def _verify_pagination_info( @@ -235,13 +235,13 @@ class StubEdxNotesServiceTest(unittest.TestCase): page = query_params["page"][0] return page if page is None else int(page) - self.assertEqual(response["total"], total_notes) - self.assertEqual(response["num_pages"], num_pages) - self.assertEqual(len(response["rows"]), notes_per_page) - self.assertEqual(response["current_page"], current_page) - self.assertEqual(get_page_value(response["previous"]), previous_page) - self.assertEqual(get_page_value(response["next"]), next_page) - self.assertEqual(response["start"], start) + assert response['total'] == total_notes + assert response['num_pages'] == num_pages + assert len(response['rows']) == notes_per_page + assert response['current_page'] == current_page + assert get_page_value(response['previous']) == previous_page + assert get_page_value(response['next']) == next_page + assert response['start'] == start def test_notes_collection(self): """ @@ -250,12 +250,12 @@ class StubEdxNotesServiceTest(unittest.TestCase): # Without user response = requests.get(self._get_url("api/v1/annotations")) - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 # Without any pagination parameters response = requests.get(self._get_url("api/v1/annotations"), params={"user": "dummy-user-id"}) - self.assertTrue(response.ok) + assert response.ok self._verify_pagination_info( response=response.json(), total_notes=5, @@ -274,7 +274,7 @@ class StubEdxNotesServiceTest(unittest.TestCase): "page_size": 3 }) - self.assertTrue(response.ok) + assert response.ok self._verify_pagination_info( response=response.json(), total_notes=5, @@ -296,7 +296,7 @@ class StubEdxNotesServiceTest(unittest.TestCase): "page_size": 10 }) - self.assertTrue(response.ok) + assert response.ok self._verify_pagination_info( response=response.json(), total_notes=5, @@ -318,7 +318,7 @@ class StubEdxNotesServiceTest(unittest.TestCase): # Get default page response = requests.get(self._get_url("api/v1/annotations"), params={"user": "dummy-user-id"}) - self.assertTrue(response.ok) + assert response.ok self._verify_pagination_info( response=response.json(), total_notes=0, @@ -332,36 +332,36 @@ class StubEdxNotesServiceTest(unittest.TestCase): def test_cleanup(self): response = requests.put(self._get_url("cleanup")) - self.assertTrue(response.ok) - self.assertEqual(len(self.server.get_all_notes()), 0) + assert response.ok + assert len(self.server.get_all_notes()) == 0 def test_create_notes(self): dummy_notes = self._get_dummy_notes(count=2) response = requests.post(self._get_url("create_notes"), data=json.dumps(dummy_notes)) - self.assertTrue(response.ok) - self.assertEqual(len(self._get_notes()), 7) + assert response.ok + assert len(self._get_notes()) == 7 response = requests.post(self._get_url("create_notes")) - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 def test_headers(self): note = self._get_notes()[0] response = requests.get(self._get_url("api/v1/annotations/" + note["id"])) - self.assertTrue(response.ok) - self.assertEqual(response.headers.get("access-control-allow-origin"), "*") + assert response.ok + assert response.headers.get('access-control-allow-origin') == '*' response = requests.options(self._get_url("api/v1/annotations/")) - self.assertTrue(response.ok) - self.assertEqual(response.headers.get("access-control-allow-origin"), "*") - self.assertEqual(response.headers.get("access-control-allow-methods"), "GET, POST, PUT, DELETE, OPTIONS") - self.assertIn("X-CSRFToken", response.headers.get("access-control-allow-headers")) + assert response.ok + assert response.headers.get('access-control-allow-origin') == '*' + assert response.headers.get('access-control-allow-methods') == 'GET, POST, PUT, DELETE, OPTIONS' + assert 'X-CSRFToken' in response.headers.get('access-control-allow-headers') def _get_notes(self): """ Return a list of notes from the stub EdxNotes service. """ notes = self.server.get_all_notes() - self.assertGreater(len(notes), 0, "Notes are empty.") + assert len(notes) > 0, 'Notes are empty.' return notes def _get_url(self, path): diff --git a/common/djangoapps/terrain/stubs/tests/test_http.py b/common/djangoapps/terrain/stubs/tests/test_http.py index b88f9c7ec9..906cb5db19 100644 --- a/common/djangoapps/terrain/stubs/tests/test_http.py +++ b/common/djangoapps/terrain/stubs/tests/test_http.py @@ -44,31 +44,31 @@ class StubHttpServiceTest(unittest.TestCase): # lint-amnesty, pylint: disable=m # JSON-encode each parameter post_params = {key: json.dumps(val)} response = requests.put(self.url, data=post_params) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 # Check that the expected values were set in the configuration for key, val in six.iteritems(params): - self.assertEqual(self.server.config.get(key), val) + assert self.server.config.get(key) == val def test_bad_json(self): response = requests.put(self.url, data="{,}") - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 def test_no_post_data(self): response = requests.put(self.url, data={}) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 def test_unicode_non_json(self): # Send unicode without json-encoding it response = requests.put(self.url, data={'test_unicode': u'\u2603 the snowman'}) - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 def test_unknown_path(self): response = requests.put( "http://127.0.0.1:{0}/invalid_url".format(self.server.port), data="{}" ) - self.assertEqual(response.status_code, 404) + assert response.status_code == 404 class RequireRequestHandler(StubHttpRequestHandler): # lint-amnesty, pylint: disable=missing-class-docstring @@ -100,26 +100,26 @@ class RequireParamTest(unittest.TestCase): # Expect success when we provide the required param response = requests.get(self.url, params={"test_param": 2}) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 # Expect failure when we do not proivde the param response = requests.get(self.url) - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 # Expect failure when we provide an empty param response = requests.get(self.url + "?test_param=") - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 def test_require_post_param(self): # Expect success when we provide the required param response = requests.post(self.url, data={"test_param": 2}) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 # Expect failure when we do not proivde the param response = requests.post(self.url) - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 # Expect failure when we provide an empty param response = requests.post(self.url, data={"test_param": None}) - self.assertEqual(response.status_code, 400) + assert response.status_code == 400 diff --git a/common/djangoapps/terrain/stubs/tests/test_lti_stub.py b/common/djangoapps/terrain/stubs/tests/test_lti_stub.py index ba1775f1b0..14bcb8bc94 100644 --- a/common/djangoapps/terrain/stubs/tests/test_lti_stub.py +++ b/common/djangoapps/terrain/stubs/tests/test_lti_stub.py @@ -49,7 +49,7 @@ class StubLtiServiceTest(unittest.TestCase): """ self.launch_uri = self.uri + 'wrong_lti_endpoint' response = requests.post(self.launch_uri, data=self.payload) - self.assertIn(b'Invalid request URL', response.content) + assert b'Invalid request URL' in response.content def test_wrong_signature(self): """ @@ -57,7 +57,7 @@ class StubLtiServiceTest(unittest.TestCase): path and responses with incorrect signature. """ response = requests.post(self.launch_uri, data=self.payload) - self.assertIn(b'Wrong LTI signature', response.content) + assert b'Wrong LTI signature' in response.content @patch('common.djangoapps.terrain.stubs.lti.signature.verify_hmac_sha1', return_value=True) def test_success_response_launch_lti(self, check_oauth): # lint-amnesty, pylint: disable=unused-argument @@ -65,34 +65,34 @@ class StubLtiServiceTest(unittest.TestCase): Success lti launch. """ response = requests.post(self.launch_uri, data=self.payload) - self.assertIn(b'This is LTI tool. Success.', response.content) + assert b'This is LTI tool. Success.' in response.content @patch('common.djangoapps.terrain.stubs.lti.signature.verify_hmac_sha1', return_value=True) def test_send_graded_result(self, verify_hmac): # pylint: disable=unused-argument response = requests.post(self.launch_uri, data=self.payload) - self.assertIn(b'This is LTI tool. Success.', response.content) + assert b'This is LTI tool. Success.' in response.content grade_uri = self.uri + 'grade' with patch('common.djangoapps.terrain.stubs.lti.requests.post') as mocked_post: mocked_post.return_value = Mock(content='Test response', status_code=200) response = six.moves.urllib.request.urlopen(grade_uri, data=b'') - self.assertIn(b'Test response', response.read()) + assert b'Test response' in response.read() @patch('common.djangoapps.terrain.stubs.lti.signature.verify_hmac_sha1', return_value=True) def test_lti20_outcomes_put(self, verify_hmac): # pylint: disable=unused-argument response = requests.post(self.launch_uri, data=self.payload) - self.assertIn(b'This is LTI tool. Success.', response.content) + assert b'This is LTI tool. Success.' in response.content grade_uri = self.uri + 'lti2_outcome' with patch('common.djangoapps.terrain.stubs.lti.requests.put') as mocked_put: mocked_put.return_value = Mock(status_code=200) response = six.moves.urllib.request.urlopen(grade_uri, data=b'') - self.assertIn(b'LTI consumer (edX) responded with HTTP 200', response.read()) + assert b'LTI consumer (edX) responded with HTTP 200' in response.read() @patch('common.djangoapps.terrain.stubs.lti.signature.verify_hmac_sha1', return_value=True) def test_lti20_outcomes_put_like_delete(self, verify_hmac): # pylint: disable=unused-argument response = requests.post(self.launch_uri, data=self.payload) - self.assertIn(b'This is LTI tool. Success.', response.content) + assert b'This is LTI tool. Success.' in response.content grade_uri = self.uri + 'lti2_delete' with patch('common.djangoapps.terrain.stubs.lti.requests.put') as mocked_put: mocked_put.return_value = Mock(status_code=200) response = six.moves.urllib.request.urlopen(grade_uri, data=b'') - self.assertIn(b'LTI consumer (edX) responded with HTTP 200', response.read()) + assert b'LTI consumer (edX) responded with HTTP 200' in response.read() diff --git a/common/djangoapps/terrain/stubs/tests/test_video.py b/common/djangoapps/terrain/stubs/tests/test_video.py index edfe33c697..685aa7f0b0 100644 --- a/common/djangoapps/terrain/stubs/tests/test_video.py +++ b/common/djangoapps/terrain/stubs/tests/test_video.py @@ -43,6 +43,6 @@ class StubVideoServiceTest(unittest.TestCase): Verify that correct hls manifest is received. """ response = requests.get("http://127.0.0.1:{port}/hls/history.m3u8".format(port=self.server.port)) - self.assertTrue(response.ok) - self.assertEqual(response.text, HLS_MANIFEST_TEXT.lstrip()) - self.assertEqual(response.headers['Access-Control-Allow-Origin'], '*') + assert response.ok + assert response.text == HLS_MANIFEST_TEXT.lstrip() + assert response.headers['Access-Control-Allow-Origin'] == '*' diff --git a/common/djangoapps/terrain/stubs/tests/test_xqueue_stub.py b/common/djangoapps/terrain/stubs/tests/test_xqueue_stub.py index 2f36a9d6dd..28e2cbe306 100644 --- a/common/djangoapps/terrain/stubs/tests/test_xqueue_stub.py +++ b/common/djangoapps/terrain/stubs/tests/test_xqueue_stub.py @@ -116,8 +116,8 @@ class StubXQueueServiceTest(unittest.TestCase): # lint-amnesty, pylint: disable # Expect that we do NOT receive a response # and that an error message is logged - self.assertFalse(self.post.called) - self.assertTrue(logger.error.called) + assert not self.post.called + assert logger.error.called def _post_submission(self, callback_url, lms_key, queue_name, xqueue_body): # lint-amnesty, pylint: disable=unused-argument """ @@ -144,7 +144,7 @@ class StubXQueueServiceTest(unittest.TestCase): # lint-amnesty, pylint: disable resp = requests.post(self.url, data=grade_request) # Expect that the response is success - self.assertEqual(resp.status_code, 200) + assert resp.status_code == 200 # Return back the header, so we can authenticate the response we receive return grade_request['xqueue_header'] @@ -167,9 +167,7 @@ class StubXQueueServiceTest(unittest.TestCase): # lint-amnesty, pylint: disable 'xqueue_body': expected_body, } # Check that the POST request was made with the correct params - self.assertEqual(self.post.call_args[1]['data']['xqueue_body'], expected_callback_dict['xqueue_body']) - self.assertEqual( - ast.literal_eval(self.post.call_args[1]['data']['xqueue_header']), - ast.literal_eval(expected_callback_dict['xqueue_header']) - ) - self.assertEqual(self.post.call_args[0][0], callback_url) + assert self.post.call_args[1]['data']['xqueue_body'] == expected_callback_dict['xqueue_body'] + assert ast.literal_eval(self.post.call_args[1]['data']['xqueue_header']) ==\ + ast.literal_eval(expected_callback_dict['xqueue_header']) + assert self.post.call_args[0][0] == callback_url diff --git a/common/djangoapps/terrain/stubs/tests/test_youtube_stub.py b/common/djangoapps/terrain/stubs/tests/test_youtube_stub.py index a1fe1d8329..146525ff2b 100644 --- a/common/djangoapps/terrain/stubs/tests/test_youtube_stub.py +++ b/common/djangoapps/terrain/stubs/tests/test_youtube_stub.py @@ -21,7 +21,7 @@ class StubYouTubeServiceTest(unittest.TestCase): # lint-amnesty, pylint: disabl def test_unused_url(self): response = requests.get(self.url + 'unused_url') - self.assertEqual(b"Unused url", response.content) + assert b'Unused url' == response.content @unittest.skip('Failing intermittently due to inconsistent responses from YT. See TE-871') def test_video_url(self): @@ -30,41 +30,31 @@ class StubYouTubeServiceTest(unittest.TestCase): # lint-amnesty, pylint: disabl ) # YouTube metadata for video `OEoXaMPEzfM` states that duration is 116. - self.assertEqual( - b'callback_func({"data": {"duration": 116, "message": "I\'m youtube.", "id": "OEoXaMPEzfM"}})', - response.content - ) + assert b'callback_func({"data": {"duration": 116, "message": "I\'m youtube.", "id": "OEoXaMPEzfM"}})' ==\ + response.content def test_transcript_url_equal(self): response = requests.get( self.url + 'test_transcripts_youtube/t__eq_exist' ) - self.assertEqual( - "".join([ - '', - '', - 'Equal transcripts' - ]).encode('utf-8'), response.content - ) + assert ''.join(['', + '', + 'Equal transcripts']).encode('utf-8') == response.content def test_transcript_url_not_equal(self): response = requests.get( self.url + 'test_transcripts_youtube/t_neq_exist', ) - self.assertEqual( - "".join([ - '', - '', - 'Transcripts sample, different that on server', - '' - ]).encode('utf-8'), response.content - ) + assert ''.join(['', + '', + 'Transcripts sample, different that on server', + '']).encode('utf-8') == response.content def test_transcript_not_found(self): response = requests.get(self.url + 'test_transcripts_youtube/some_id') - self.assertEqual(404, response.status_code) + assert 404 == response.status_code def test_reset_configuration(self): @@ -75,7 +65,7 @@ class StubYouTubeServiceTest(unittest.TestCase): # lint-amnesty, pylint: disabl # reset server configuration response = requests.delete(reset_config_url) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 # ensure that server config dict is empty after successful reset - self.assertEqual(self.server.config, {}) + assert self.server.config == {} diff --git a/common/djangoapps/third_party_auth/tests/test_utils.py b/common/djangoapps/third_party_auth/tests/test_utils.py index 1e4b90a5a3..8af2909415 100644 --- a/common/djangoapps/third_party_auth/tests/test_utils.py +++ b/common/djangoapps/third_party_auth/tests/test_utils.py @@ -55,12 +55,8 @@ class TestUtils(TestCase): """ # Create users from factory UserFactory(username='test_user', email='test_user@example.com') - self.assertTrue( - get_user_from_email({'email': 'test_user@example.com'}), - ) - self.assertFalse( - get_user_from_email({'email': 'invalid@example.com'}), - ) + assert get_user_from_email({'email': 'test_user@example.com'}) + assert not get_user_from_email({'email': 'invalid@example.com'}) def test_is_enterprise_customer_user(self): """ @@ -79,9 +75,5 @@ class TestUtils(TestCase): user_id=user.id, ) - self.assertTrue( - is_enterprise_customer_user('the-provider', user), - ) - self.assertFalse( - is_enterprise_customer_user('the-provider', other_user), - ) + assert is_enterprise_customer_user('the-provider', user) + assert not is_enterprise_customer_user('the-provider', other_user) diff --git a/common/djangoapps/track/backends/tests/test_mongodb.py b/common/djangoapps/track/backends/tests/test_mongodb.py index 0610c7fdd3..2a715519d6 100644 --- a/common/djangoapps/track/backends/tests/test_mongodb.py +++ b/common/djangoapps/track/backends/tests/test_mongodb.py @@ -25,7 +25,7 @@ class TestMongoBackend(TestCase): # lint-amnesty, pylint: disable=missing-class calls = self.backend.collection.insert.mock_calls - self.assertEqual(len(calls), 2) + assert len(calls) == 2 # Unpack the arguments and check if the events were used # as the first argument to collection.insert @@ -34,5 +34,5 @@ class TestMongoBackend(TestCase): # lint-amnesty, pylint: disable=missing-class _, args, _ = call return args[0] - self.assertEqual(events[0], first_argument(calls[0])) - self.assertEqual(events[1], first_argument(calls[1])) + assert events[0] == first_argument(calls[0]) + assert events[1] == first_argument(calls[1]) diff --git a/common/djangoapps/track/management/tests/test_tracked_command.py b/common/djangoapps/track/management/tests/test_tracked_command.py index 2ce671c15d..95d94c9ec0 100644 --- a/common/djangoapps/track/management/tests/test_tracked_command.py +++ b/common/djangoapps/track/management/tests/test_tracked_command.py @@ -24,4 +24,4 @@ class CommandsTestBase(TestCase): args = ['whee'] kwargs = {'key1': 'default', 'key2': True} json_out = self._run_dummy_command(*args, **kwargs) - self.assertEqual(json_out['command'].strip(), 'tracked_dummy_command') + assert json_out['command'].strip() == 'tracked_dummy_command' diff --git a/common/djangoapps/track/tests/__init__.py b/common/djangoapps/track/tests/__init__.py index bf20b09052..11efce9569 100644 --- a/common/djangoapps/track/tests/__init__.py +++ b/common/djangoapps/track/tests/__init__.py @@ -72,8 +72,8 @@ class EventTrackingTestCase(TestCase): def assert_no_events_emitted(self): """Ensure no events were emitted at this point in the test.""" - self.assertEqual(len(self.backend.events), 0) + assert len(self.backend.events) == 0 def assert_events_emitted(self): """Ensure at least one event has been emitted at this point in the test.""" - self.assertGreaterEqual(len(self.backend.events), 1) + assert len(self.backend.events) >= 1 diff --git a/common/djangoapps/track/tests/test_contexts.py b/common/djangoapps/track/tests/test_contexts.py index 0f6d79dd50..ffcd8a6100 100644 --- a/common/djangoapps/track/tests/test_contexts.py +++ b/common/djangoapps/track/tests/test_contexts.py @@ -27,25 +27,14 @@ class TestContexts(TestCase): # lint-amnesty, pylint: disable=missing-class-doc self.assert_parses_course_id_from_url(url, course_id) def assert_parses_course_id_from_url(self, format_string, course_id): - self.assertEqual( - contexts.course_context_from_url(format_string.format(course_id=course_id)), - { - 'course_id': course_id, - 'org_id': self.ORG_ID - } - ) + assert contexts.course_context_from_url(format_string.format(course_id=course_id)) ==\ + {'course_id': course_id, 'org_id': self.ORG_ID} def test_no_course_id_in_url(self): self.assert_empty_context_for_url('http://foo.bar.com/dashboard') def assert_empty_context_for_url(self, url): - self.assertEqual( - contexts.course_context_from_url(url), - { - 'course_id': '', - 'org_id': '' - } - ) + assert contexts.course_context_from_url(url) == {'course_id': '', 'org_id': ''} @ddt.data('', '/', '/?', '?format=json') def test_malformed_course_id(self, postfix): diff --git a/common/djangoapps/track/tests/test_middleware.py b/common/djangoapps/track/tests/test_middleware.py index 749b325ff6..7ed30faa93 100644 --- a/common/djangoapps/track/tests/test_middleware.py +++ b/common/djangoapps/track/tests/test_middleware.py @@ -31,7 +31,7 @@ class TrackMiddlewareTestCase(TestCase): def test_normal_request(self): request = self.request_factory.get('/somewhere') self.track_middleware.process_request(request) - self.assertTrue(self.mock_server_track.called) + assert self.mock_server_track.called @ddt.unpack @ddt.data( @@ -48,48 +48,37 @@ class TrackMiddlewareTestCase(TestCase): request.META[meta_key] = 'test latin1 \xd3 \xe9 \xf1' context = self.get_context_for_request(request) - self.assertEqual(context[context_key], u'test latin1 Ó é ñ') + assert context[context_key] == u'test latin1 Ó é ñ' def test_default_filters_do_not_render_view(self): for url in ['/event', '/event/1', '/login', '/heartbeat']: request = self.request_factory.get(url) self.track_middleware.process_request(request) - self.assertFalse(self.mock_server_track.called) + assert not self.mock_server_track.called self.mock_server_track.reset_mock() @override_settings(TRACKING_IGNORE_URL_PATTERNS=[]) def test_reading_filtered_urls_from_settings(self): request = self.request_factory.get('/event') self.track_middleware.process_request(request) - self.assertTrue(self.mock_server_track.called) + assert self.mock_server_track.called @override_settings(TRACKING_IGNORE_URL_PATTERNS=[r'^/some/excluded.*']) def test_anchoring_of_patterns_at_beginning(self): request = self.request_factory.get('/excluded') self.track_middleware.process_request(request) - self.assertTrue(self.mock_server_track.called) + assert self.mock_server_track.called self.mock_server_track.reset_mock() request = self.request_factory.get('/some/excluded/url') self.track_middleware.process_request(request) - self.assertFalse(self.mock_server_track.called) + assert not self.mock_server_track.called def test_default_request_context(self): context = self.get_context_for_path('/courses/') - self.assertEqual(context, { - 'accept_language': '', - 'referer': '', - 'user_id': '', - 'session': '', - 'username': '', - 'ip': '127.0.0.1', - 'host': 'testserver', - 'agent': '', - 'path': '/courses/', - 'org_id': '', - 'course_id': '', - 'client_id': None, - }) + assert context == {'accept_language': '', 'referer': '', 'user_id': '', 'session': '', 'username': '', + 'ip': '127.0.0.1', 'host': 'testserver', 'agent': '', 'path': '/courses/', 'org_id': '', + 'course_id': '', 'client_id': None} def test_no_forward_for_header_ip_context(self): request = self.request_factory.get('/courses/') @@ -98,7 +87,7 @@ class TrackMiddlewareTestCase(TestCase): request.META['REMOTE_ADDR'] = remote_addr context = self.get_context_for_request(request) - self.assertEqual(context['ip'], remote_addr) + assert context['ip'] == remote_addr def test_single_forward_for_header_ip_context(self): request = self.request_factory.get('/courses/') @@ -109,7 +98,7 @@ class TrackMiddlewareTestCase(TestCase): request.META['HTTP_X_FORWARDED_FOR'] = forwarded_ip context = self.get_context_for_request(request) - self.assertEqual(context['ip'], forwarded_ip) + assert context['ip'] == forwarded_ip def test_multiple_forward_for_header_ip_context(self): request = self.request_factory.get('/courses/') @@ -120,7 +109,7 @@ class TrackMiddlewareTestCase(TestCase): request.META['HTTP_X_FORWARDED_FOR'] = forwarded_ip context = self.get_context_for_request(request) - self.assertEqual(context['ip'], '11.22.33.44') + assert context['ip'] == '11.22.33.44' def get_context_for_path(self, path): """Extract the generated event tracking context for a given request for the given path.""" @@ -135,10 +124,7 @@ class TrackMiddlewareTestCase(TestCase): finally: self.track_middleware.process_response(request, None) - self.assertEqual( - tracker.get_tracker().resolve_context(), - {} - ) + assert tracker.get_tracker().resolve_context() == {} return captured_context @@ -153,7 +139,7 @@ class TrackMiddlewareTestCase(TestCase): def assert_dict_subset(self, superset, subset): """Assert that the superset dict contains all of the key-value pairs found in the subset dict.""" for key, expected_value in six.iteritems(subset): - self.assertEqual(superset[key], expected_value) + assert superset[key] == expected_value def test_request_with_user(self): user_id = 1 @@ -174,7 +160,7 @@ class TrackMiddlewareTestCase(TestCase): request.session.save() session_key = request.session.session_key expected_session_key = self.track_middleware.substitute_session_key(session_key) - self.assertEqual(len(session_key), len(expected_session_key)) + assert len(session_key) == len(expected_session_key) context = self.get_context_for_request(request) self.assert_dict_subset(context, { 'session': expected_session_key, @@ -186,12 +172,12 @@ class TrackMiddlewareTestCase(TestCase): # Output value pinned to alert on unintended changes to generator expected_session_key = 'b4103566fc80d20da1970cbb4380bccd' substitute_session_key = self.track_middleware.substitute_session_key(session_key) - self.assertEqual(substitute_session_key, expected_session_key) + assert substitute_session_key == expected_session_key # Confirm that we get *different* outputs for different inputs expected_session_key_2 = "6f0c784c1087c6bc4624b7eac982fedf" substitute_session_key_2 = self.track_middleware.substitute_session_key(session_key + "different") - self.assertEqual(expected_session_key_2, substitute_session_key_2) + assert expected_session_key_2 == substitute_session_key_2 def test_request_headers(self): ip_address = '10.0.0.0' diff --git a/common/djangoapps/track/tests/test_segment.py b/common/djangoapps/track/tests/test_segment.py index 432ea9026e..ac357a7338 100644 --- a/common/djangoapps/track/tests/test_segment.py +++ b/common/djangoapps/track/tests/test_segment.py @@ -27,25 +27,25 @@ class SegmentTrackTestCase(TestCase): def test_missing_key(self): segment.track(sentinel.user_id, sentinel.name, self.properties) - self.assertFalse(self.mock_segment_track.called) + assert not self.mock_segment_track.called @override_settings(LMS_SEGMENT_KEY=None) def test_null_key(self): segment.track(sentinel.user_id, sentinel.name, self.properties) - self.assertFalse(self.mock_segment_track.called) + assert not self.mock_segment_track.called @override_settings(LMS_SEGMENT_KEY="testkey") def test_missing_name(self): segment.track(sentinel.user_id, None, self.properties) - self.assertFalse(self.mock_segment_track.called) + assert not self.mock_segment_track.called @override_settings(LMS_SEGMENT_KEY="testkey") def test_track_without_tracking_context(self): segment.track(sentinel.user_id, sentinel.name, self.properties) - self.assertTrue(self.mock_segment_track.called) + assert self.mock_segment_track.called args, kwargs = self.mock_segment_track.call_args # lint-amnesty, pylint: disable=unused-variable expected_segment_context = {} - self.assertEqual((sentinel.user_id, sentinel.name, self.properties, expected_segment_context), args) + assert (sentinel.user_id, sentinel.name, self.properties, expected_segment_context) == args @ddt.unpack @ddt.data( @@ -62,19 +62,19 @@ class SegmentTrackTestCase(TestCase): with self.tracker.context('test', tracking_context): segment.track(sentinel.user_id, sentinel.name, self.properties) args, kwargs = self.mock_segment_track.call_args # lint-amnesty, pylint: disable=unused-variable - self.assertEqual((sentinel.user_id, sentinel.name, self.properties, expected_segment_context), args) + assert (sentinel.user_id, sentinel.name, self.properties, expected_segment_context) == args # Test with provided context and no tracking context. segment.track(sentinel.user_id, sentinel.name, self.properties, provided_context) args, kwargs = self.mock_segment_track.call_args - self.assertEqual((sentinel.user_id, sentinel.name, self.properties, provided_context), args) + assert (sentinel.user_id, sentinel.name, self.properties, provided_context) == args # Test with provided context and also tracking context. with self.tracker.context('test', tracking_context): segment.track(sentinel.user_id, sentinel.name, self.properties, provided_context) - self.assertTrue(self.mock_segment_track.called) + assert self.mock_segment_track.called args, kwargs = self.mock_segment_track.call_args - self.assertEqual((sentinel.user_id, sentinel.name, self.properties, provided_context), args) + assert (sentinel.user_id, sentinel.name, self.properties, provided_context) == args @override_settings(LMS_SEGMENT_KEY="testkey") def test_track_with_standard_context(self): @@ -97,7 +97,7 @@ class SegmentTrackTestCase(TestCase): with self.tracker.context('test', tracking_context): segment.track(sentinel.user_id, sentinel.name, self.properties) - self.assertTrue(self.mock_segment_track.called) + assert self.mock_segment_track.called args, kwargs = self.mock_segment_track.call_args # lint-amnesty, pylint: disable=unused-variable expected_segment_context = { @@ -112,7 +112,7 @@ class SegmentTrackTestCase(TestCase): 'url': 'https://hostname/this/is/a/path' # Synthesized URL value. } } - self.assertEqual((sentinel.user_id, sentinel.name, self.properties, expected_segment_context), args) + assert (sentinel.user_id, sentinel.name, self.properties, expected_segment_context) == args class SegmentIdentifyTestCase(TestCase): @@ -127,24 +127,24 @@ class SegmentIdentifyTestCase(TestCase): def test_missing_key(self): segment.identify(sentinel.user_id, self.properties) - self.assertFalse(self.mock_segment_identify.called) + assert not self.mock_segment_identify.called @override_settings(LMS_SEGMENT_KEY=None) def test_null_key(self): segment.identify(sentinel.user_id, self.properties) - self.assertFalse(self.mock_segment_identify.called) + assert not self.mock_segment_identify.called @override_settings(LMS_SEGMENT_KEY="testkey") def test_normal_call(self): segment.identify(sentinel.user_id, self.properties) - self.assertTrue(self.mock_segment_identify.called) + assert self.mock_segment_identify.called args, kwargs = self.mock_segment_identify.call_args # lint-amnesty, pylint: disable=unused-variable - self.assertEqual((sentinel.user_id, self.properties, {}), args) + assert (sentinel.user_id, self.properties, {}) == args @override_settings(LMS_SEGMENT_KEY="testkey") def test_call_with_context(self): provided_context = {sentinel.context_key: sentinel.context_value} segment.identify(sentinel.user_id, self.properties, provided_context) - self.assertTrue(self.mock_segment_identify.called) + assert self.mock_segment_identify.called args, kwargs = self.mock_segment_identify.call_args # lint-amnesty, pylint: disable=unused-variable - self.assertEqual((sentinel.user_id, self.properties, provided_context), args) + assert (sentinel.user_id, self.properties, provided_context) == args diff --git a/common/djangoapps/track/tests/test_shim.py b/common/djangoapps/track/tests/test_shim.py index 8c163f3a83..3b137d4f5b 100644 --- a/common/djangoapps/track/tests/test_shim.py +++ b/common/djangoapps/track/tests/test_shim.py @@ -2,7 +2,7 @@ from collections import namedtuple - +import pytest import ddt from django.test.utils import override_settings from mock import sentinel @@ -247,7 +247,7 @@ class EventTransformerRegistryTestCase(EventTrackingTestCase): def test_event_registry_dispatch(self, event_name, expected_transformer): event = {'name': event_name} transformer = self.registry.create_transformer(event) - self.assertIsInstance(transformer, expected_transformer) + assert isinstance(transformer, expected_transformer) @ddt.data( 'edx.ui.lms.sequence.next_selected.what', @@ -256,7 +256,7 @@ class EventTransformerRegistryTestCase(EventTrackingTestCase): ) def test_dispatch_to_nonexistent_events(self, event_name): event = {'name': event_name} - with self.assertRaises(KeyError): + with pytest.raises(KeyError): self.registry.create_transformer(event) @@ -293,13 +293,13 @@ class PrefixedEventProcessorTestCase(EventTrackingTestCase): else: offset = -1 if sequence_ddt.legacy_event_type: - self.assertEqual(result[u'event_type'], sequence_ddt.legacy_event_type) - self.assertEqual(result[u'event'][u'old'], sequence_ddt.current_tab) - self.assertEqual(result[u'event'][u'new'], sequence_ddt.current_tab + offset) + assert result[u'event_type'] == sequence_ddt.legacy_event_type + assert result[u'event'][u'old'] == sequence_ddt.current_tab + assert result[u'event'][u'new'] == (sequence_ddt.current_tab + offset) else: - self.assertNotIn(u'event_type', result) - self.assertNotIn(u'old', result[u'event']) - self.assertNotIn(u'new', result[u'event']) + assert u'event_type' not in result + assert u'old' not in result[u'event'] + assert u'new' not in result[u'event'] def test_sequence_tab_navigation(self): event_name = u'edx.ui.lms.sequence.tab_selected' @@ -316,6 +316,6 @@ class PrefixedEventProcessorTestCase(EventTrackingTestCase): process_event_shim = PrefixedEventProcessor() result = process_event_shim(event) - self.assertEqual(result[u'event_type'], u'seq_goto') - self.assertEqual(result[u'event'][u'old'], 2) - self.assertEqual(result[u'event'][u'new'], 5) + assert result[u'event_type'] == u'seq_goto' + assert result[u'event'][u'old'] == 2 + assert result[u'event'][u'new'] == 5 diff --git a/common/djangoapps/track/tests/test_tracker.py b/common/djangoapps/track/tests/test_tracker.py index e7cdd772ec..eb67cf2508 100644 --- a/common/djangoapps/track/tests/test_tracker.py +++ b/common/djangoapps/track/tests/test_tracker.py @@ -39,8 +39,8 @@ class TestTrackerInstantiation(TestCase): options = {'flag': True} backend = self.get_backend(name, options) - self.assertIsInstance(backend, DummyBackend) - self.assertTrue(backend.flag) + assert isinstance(backend, DummyBackend) + assert backend.flag def test_instatiate_backends_with_invalid_values(self): def get_invalid_backend(name, parameters): @@ -69,11 +69,11 @@ class TestTrackerDjangoInstantiation(TestCase): backends = self._reload_backends() - self.assertEqual(len(backends), 1) + assert len(backends) == 1 tracker.send({}) - self.assertEqual(list(backends.values())[0].count, 1) + assert list(backends.values())[0].count == 1 @override_settings(TRACKING_BACKENDS=MULTI_SETTINGS.copy()) def test_django_multi_settings(self): @@ -81,14 +81,14 @@ class TestTrackerDjangoInstantiation(TestCase): backends = list(self._reload_backends().values()) - self.assertEqual(len(backends), 2) + assert len(backends) == 2 event_count = 10 for _ in range(event_count): tracker.send({}) - self.assertEqual(backends[0].count, event_count) - self.assertEqual(backends[1].count, event_count) + assert backends[0].count == event_count + assert backends[1].count == event_count @override_settings(TRACKING_BACKENDS=MULTI_SETTINGS.copy()) def test_django_remove_settings(self): @@ -98,7 +98,7 @@ class TestTrackerDjangoInstantiation(TestCase): backends = self._reload_backends() - self.assertEqual(len(backends), 1) + assert len(backends) == 1 def _reload_backends(self): # lint-amnesty, pylint: disable=missing-function-docstring # pylint: disable=protected-access diff --git a/common/djangoapps/track/tests/test_util.py b/common/djangoapps/track/tests/test_util.py index 57a4afb356..351349b341 100644 --- a/common/djangoapps/track/tests/test_util.py +++ b/common/djangoapps/track/tests/test_util.py @@ -29,10 +29,10 @@ class TestDateTimeJSONEncoder(TestCase): # lint-amnesty, pylint: disable=missin to_json = json.dumps(obj, cls=DateTimeJSONEncoder) from_json = json.loads(to_json) - self.assertEqual(from_json['number'], 100) - self.assertEqual(from_json['string'], 'hello') - self.assertEqual(from_json['object'], {'a': 1}) + assert from_json['number'] == 100 + assert from_json['string'] == 'hello' + assert from_json['object'] == {'a': 1} - self.assertEqual(from_json['a_datetime'], an_iso_datetime) - self.assertEqual(from_json['a_tz_datetime'], an_iso_datetime) - self.assertEqual(from_json['a_date'], an_iso_date) + assert from_json['a_datetime'] == an_iso_datetime + assert from_json['a_tz_datetime'] == an_iso_datetime + assert from_json['a_date'] == an_iso_date diff --git a/common/djangoapps/track/views/tests/test_segmentio.py b/common/djangoapps/track/views/tests/test_segmentio.py index c07b70b936..9ddf176988 100644 --- a/common/djangoapps/track/views/tests/test_segmentio.py +++ b/common/djangoapps/track/views/tests/test_segmentio.py @@ -39,7 +39,7 @@ class SegmentIOTrackingTestCase(SegmentIOTrackingTestCaseBase): def test_get_request(self): request = self.request_factory.get(SEGMENTIO_TEST_ENDPOINT) response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 405) + assert response.status_code == 405 self.assert_no_events_emitted() @override_settings( @@ -48,19 +48,19 @@ class SegmentIOTrackingTestCase(SegmentIOTrackingTestCaseBase): def test_no_secret_config(self): request = self.request_factory.post(SEGMENTIO_TEST_ENDPOINT) response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 401) + assert response.status_code == 401 self.assert_no_events_emitted() def test_no_secret_provided(self): request = self.request_factory.post(SEGMENTIO_TEST_ENDPOINT) response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 401) + assert response.status_code == 401 self.assert_no_events_emitted() def test_secret_mismatch(self): request = self.create_request(key='y') response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 401) + assert response.status_code == 401 self.assert_no_events_emitted() @data('identify', 'Group', 'Alias', 'Page', 'identify', 'screen') @@ -130,7 +130,7 @@ class SegmentIOTrackingTestCase(SegmentIOTrackingTestCaseBase): self.assert_no_events_emitted() try: response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 expected_event = { 'accept_language': '', @@ -252,7 +252,7 @@ class SegmentIOTrackingTestCase(SegmentIOTrackingTestCaseBase): content_type='application/json' ) response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 self.assert_events_emitted() def test_hiding_failure(self): @@ -263,7 +263,7 @@ class SegmentIOTrackingTestCase(SegmentIOTrackingTestCaseBase): ) response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 self.assert_no_events_emitted() @data( @@ -310,7 +310,7 @@ class SegmentIOTrackingTestCase(SegmentIOTrackingTestCaseBase): middleware.process_request(request) try: response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 expected_event = { 'accept_language': '', @@ -443,7 +443,7 @@ class SegmentIOTrackingTestCase(SegmentIOTrackingTestCaseBase): middleware.process_request(request) try: response = segmentio.segmentio_event(request) - self.assertEqual(response.status_code, 200) + assert response.status_code == 200 expected_event = { 'accept_language': '', diff --git a/common/djangoapps/util/testing.py b/common/djangoapps/util/testing.py index 645deb96be..9a3b0f52b0 100644 --- a/common/djangoapps/util/testing.py +++ b/common/djangoapps/util/testing.py @@ -89,7 +89,7 @@ class EventTestMixin(object): """ Ensures no events were emitted since the last event related assertion. """ - self.assertFalse(self.mock_tracker.emit.called) + assert not self.mock_tracker.emit.called def assert_event_emitted(self, event_name, **kwargs): """ @@ -109,7 +109,7 @@ class EventTestMixin(object): for call_args in self.mock_tracker.emit.call_args_list: if call_args[0][0] == event_name: actual_count += 1 - self.assertEqual(actual_count, expected_count) + assert actual_count == expected_count def reset_tracker(self): """ @@ -134,7 +134,7 @@ class PatchMediaTypeMixin(object): json.dumps({}), content_type=self.unsupported_media_type ) - self.assertEqual(response.status_code, 415) + assert response.status_code == 415 def patch_testcase(): diff --git a/common/djangoapps/util/tests/test_course.py b/common/djangoapps/util/tests/test_course.py index c43a34997f..6ef70a16bd 100644 --- a/common/djangoapps/util/tests/test_course.py +++ b/common/djangoapps/util/tests/test_course.py @@ -81,7 +81,7 @@ class TestCourseSharingLinks(ModuleStoreTestCase): enable_social_sharing=enable_social_sharing, enable_mktg_site=enable_mktg_site, ) - self.assertEqual(actual_course_sharing_link, expected_course_sharing_link) + assert actual_course_sharing_link == expected_course_sharing_link @ddt.data( (['social_sharing_url'], 'test_marketing_url'), @@ -106,7 +106,7 @@ class TestCourseSharingLinks(ModuleStoreTestCase): enable_social_sharing=True, enable_mktg_site=True, ) - self.assertEqual(actual_course_sharing_link, expected_course_sharing_link) + assert actual_course_sharing_link == expected_course_sharing_link @ddt.data( (True, 'test_social_sharing_url'), @@ -126,4 +126,4 @@ class TestCourseSharingLinks(ModuleStoreTestCase): enable_mktg_site=True, use_overview=False, ) - self.assertEqual(actual_course_sharing_link, expected_course_sharing_link) + assert actual_course_sharing_link == expected_course_sharing_link diff --git a/common/djangoapps/util/tests/test_date_utils.py b/common/djangoapps/util/tests/test_date_utils.py index 3a7ea1c6f5..13a69deb14 100644 --- a/common/djangoapps/util/tests/test_date_utils.py +++ b/common/djangoapps/util/tests/test_date_utils.py @@ -6,7 +6,7 @@ Tests for util.date_utils import unittest from datetime import datetime, timedelta, tzinfo - +import pytest import ddt from markupsafe import Markup from mock import patch @@ -135,8 +135,8 @@ class StrftimeLocalizedTest(unittest.TestCase): def test_usual_strftime_behavior(self, fmt_expected): (fmt, expected) = fmt_expected dtime = datetime(2013, 2, 14, 16, 41, 17) - self.assertEqual(expected, strftime_localized(dtime, fmt)) - self.assertEqual(expected, dtime.strftime(fmt)) + assert expected == strftime_localized(dtime, fmt) + assert expected == dtime.strftime(fmt) @ddt.data( ("SHORT_DATE", "Feb 14, 2013"), @@ -148,7 +148,7 @@ class StrftimeLocalizedTest(unittest.TestCase): def test_shortcuts(self, fmt_expected): (fmt, expected) = fmt_expected dtime = datetime(2013, 2, 14, 16, 41, 17) - self.assertEqual(expected, strftime_localized(dtime, fmt)) + assert expected == strftime_localized(dtime, fmt) @patch('common.djangoapps.util.date_utils.pgettext', fake_pgettext(translations={ ("abbreviated month name", "Feb"): "XXfebXX", @@ -167,7 +167,7 @@ class StrftimeLocalizedTest(unittest.TestCase): def test_translated_words(self, fmt_expected): (fmt, expected) = fmt_expected dtime = datetime(2013, 2, 14, 16, 41, 17) - self.assertEqual(expected, strftime_localized(dtime, fmt)) + assert expected == strftime_localized(dtime, fmt) @patch('common.djangoapps.util.date_utils.ugettext', fake_ugettext(translations={ "SHORT_DATE_FORMAT": "date(%Y.%m.%d)", @@ -187,7 +187,7 @@ class StrftimeLocalizedTest(unittest.TestCase): def test_translated_formats(self, fmt_expected): (fmt, expected) = fmt_expected dtime = datetime(2013, 2, 14, 16, 41, 17) - self.assertEqual(expected, strftime_localized(dtime, fmt)) + assert expected == strftime_localized(dtime, fmt) @patch('common.djangoapps.util.date_utils.ugettext', fake_ugettext(translations={ "SHORT_DATE_FORMAT": "oops date(%Y.%x.%d)", @@ -200,7 +200,7 @@ class StrftimeLocalizedTest(unittest.TestCase): def test_recursion_protection(self, fmt_expected): (fmt, expected) = fmt_expected dtime = datetime(2013, 2, 14, 16, 41, 17) - self.assertEqual(expected, strftime_localized(dtime, fmt)) + assert expected == strftime_localized(dtime, fmt) @ddt.data( "%", @@ -209,7 +209,7 @@ class StrftimeLocalizedTest(unittest.TestCase): ) def test_invalid_format_strings(self, fmt): dtime = datetime(2013, 2, 14, 16, 41, 17) - with self.assertRaises(ValueError): + with pytest.raises(ValueError): strftime_localized(dtime, fmt) @@ -227,7 +227,7 @@ class StrftimeLocalizedHtmlTest(unittest.TestCase): with patch('common.djangoapps.util.date_utils.user_timezone_locale_prefs', return_value={'user_timezone': timezone}): html = strftime_localized_html(dtime, 'SHORT_DATE') - self.assertIsInstance(html, Markup) + assert isinstance(html, Markup) self.assertRegex(html, 'Feb 14, 2013') diff --git a/common/djangoapps/util/tests/test_db.py b/common/djangoapps/util/tests/test_db.py index 95c6c52dc8..303be00fd0 100644 --- a/common/djangoapps/util/tests/test_db.py +++ b/common/djangoapps/util/tests/test_db.py @@ -86,11 +86,11 @@ class TransactionManagersTestCase(TransactionTestCase): thread2.join() thread1.join() - self.assertIsInstance(thread1.status.get('exception'), exception_class) - self.assertEqual(thread1.status.get('created'), created_in_1) + assert isinstance(thread1.status.get('exception'), exception_class) + assert thread1.status.get('created') == created_in_1 - self.assertIsNone(thread2.status.get('exception')) - self.assertEqual(thread2.status.get('created'), created_in_2) + assert thread2.status.get('exception') is None + assert thread2.status.get('created') == created_in_2 def test_outer_atomic_nesting(self): """ @@ -176,7 +176,7 @@ class GenerateIntIdTestCase(TestCase): minimum = 1 maximum = times for __ in range(times): - self.assertIn(generate_int_id(minimum, maximum), list(range(minimum, maximum + 1))) + assert generate_int_id(minimum, maximum) in list(range(minimum, (maximum + 1))) @ddt.data(10) def test_used_ids(self, times): @@ -189,7 +189,7 @@ class GenerateIntIdTestCase(TestCase): used_ids = {2, 4, 6, 8} for __ in range(times): int_id = generate_int_id(minimum, maximum, used_ids) - self.assertIn(int_id, list(set(range(minimum, maximum + 1)) - used_ids)) + assert int_id in list((set(range(minimum, (maximum + 1))) - used_ids)) class MigrationTests(TestCase): @@ -213,4 +213,4 @@ class MigrationTests(TestCase): out = StringIO() call_command("makemigrations", dry_run=True, verbosity=3, stdout=out) output = out.getvalue() - self.assertIn("No changes detected", output) + assert 'No changes detected' in output diff --git a/common/djangoapps/util/tests/test_disable_rate_limit.py b/common/djangoapps/util/tests/test_disable_rate_limit.py index 31c117b1b8..35d7b122b2 100644 --- a/common/djangoapps/util/tests/test_disable_rate_limit.py +++ b/common/djangoapps/util/tests/test_disable_rate_limit.py @@ -3,6 +3,7 @@ import unittest +import pytest import mock from django.conf import settings from django.core.cache import cache @@ -44,7 +45,7 @@ class DisableRateLimitTest(TestCase): # Since our fake throttle always rejects requests, # we should expect the request to be rejected. request = mock.Mock() - with self.assertRaises(Throttled): + with pytest.raises(Throttled): self.view.check_throttles(request) def test_disable_rate_limit(self): diff --git a/common/djangoapps/util/tests/test_django_utils.py b/common/djangoapps/util/tests/test_django_utils.py index bb31baec70..22f0d7d945 100644 --- a/common/djangoapps/util/tests/test_django_utils.py +++ b/common/djangoapps/util/tests/test_django_utils.py @@ -19,7 +19,7 @@ class CacheCheckMixin(object): def check_caches(self, key): """Check that caches are empty, and add values.""" for cache in caches.all(): - self.assertIsNone(cache.get(key)) + assert cache.get(key) is None cache.set(key, "Not None") diff --git a/common/djangoapps/util/tests/test_file.py b/common/djangoapps/util/tests/test_file.py index 1f7e0a2ea7..9e4f1bc467 100644 --- a/common/djangoapps/util/tests/test_file.py +++ b/common/djangoapps/util/tests/test_file.py @@ -7,7 +7,7 @@ Tests for file.py import os from datetime import datetime from io import StringIO - +import pytest import ddt import six from django.core import exceptions @@ -37,11 +37,11 @@ class FilenamePrefixGeneratorTestCase(TestCase): """ @ddt.data(CourseLocator(org='foo', course='bar', run='baz'), CourseKey.from_string('foo/bar/baz')) def test_locators(self, course_key): - self.assertEqual(course_filename_prefix_generator(course_key), u'foo_bar_baz') + assert course_filename_prefix_generator(course_key) == u'foo_bar_baz' @ddt.data(CourseLocator(org='foo', course='bar', run='baz'), CourseKey.from_string('foo/bar/baz')) def test_custom_separator(self, course_key): - self.assertEqual(course_filename_prefix_generator(course_key, separator='-'), u'foo-bar-baz') + assert course_filename_prefix_generator(course_key, separator='-') == u'foo-bar-baz' @ddt.ddt @@ -66,15 +66,10 @@ class FilenameGeneratorTestCase(TestCase): """ Tests that the generator creates names based on course_id, base name, and date. """ - self.assertEqual( - u'foo_bar_baz_file_1974-06-22-010203', - course_and_time_based_filename_generator(course_key, 'file') - ) + assert u'foo_bar_baz_file_1974-06-22-010203' == course_and_time_based_filename_generator(course_key, 'file') - self.assertEqual( - u'foo_bar_baz_base_name_ø_1974-06-22-010203', - course_and_time_based_filename_generator(course_key, ' base` name ø ') - ) + assert u'foo_bar_baz_base_name_ø_1974-06-22-010203' ==\ + course_and_time_based_filename_generator(course_key, ' base` name ø ') class StoreUploadedFileTestCase(TestCase): @@ -99,33 +94,33 @@ class StoreUploadedFileTestCase(TestCase): """ Helper method to verify exception text. """ - self.assertEqual(expected_message, text_type(error.exception)) + assert expected_message == text_type(error.value) def test_error_conditions(self): """ Verifies that exceptions are thrown in the expected cases. """ - with self.assertRaises(ValueError) as error: + with pytest.raises(ValueError) as error: self.request.FILES = {"uploaded_file": SimpleUploadedFile("tempfile.csv", self.file_content)} store_uploaded_file(self.request, "wrong_key", [".txt", ".csv"], "stored_file", self.default_max_size) self.verify_exception("No file uploaded with key 'wrong_key'.", error) - with self.assertRaises(exceptions.PermissionDenied) as error: + with pytest.raises(exceptions.PermissionDenied) as error: self.request.FILES = {"uploaded_file": SimpleUploadedFile("tempfile.csv", self.file_content)} store_uploaded_file(self.request, "uploaded_file", [], "stored_file", self.default_max_size) self.verify_exception("The file must end with one of the following extensions: ''.", error) - with self.assertRaises(exceptions.PermissionDenied) as error: + with pytest.raises(exceptions.PermissionDenied) as error: self.request.FILES = {"uploaded_file": SimpleUploadedFile("tempfile.csv", self.file_content)} store_uploaded_file(self.request, "uploaded_file", [".bar"], "stored_file", self.default_max_size) self.verify_exception("The file must end with the extension '.bar'.", error) - with self.assertRaises(exceptions.PermissionDenied) as error: + with pytest.raises(exceptions.PermissionDenied) as error: self.request.FILES = {"uploaded_file": SimpleUploadedFile("tempfile.csv", self.file_content)} store_uploaded_file(self.request, "uploaded_file", [".xxx", ".bar"], "stored_file", self.default_max_size) self.verify_exception("The file must end with one of the following extensions: '.xxx', '.bar'.", error) - with self.assertRaises(exceptions.PermissionDenied) as error: + with pytest.raises(exceptions.PermissionDenied) as error: self.request.FILES = {"uploaded_file": SimpleUploadedFile("tempfile.csv", self.file_content)} store_uploaded_file(self.request, "uploaded_file", [".csv"], "stored_file", 2) self.verify_exception("Maximum upload file size is 2 bytes.", error) @@ -138,7 +133,7 @@ class StoreUploadedFileTestCase(TestCase): def verify_file_presence(should_exist): """ Verify whether or not the stored file, passed to the validator, exists. """ - self.assertEqual(should_exist, validator_data["storage"].exists(validator_data["filename"])) + assert should_exist == validator_data['storage'].exists(validator_data['filename']) def store_file_data(storage, filename): """ Stores file validator data for testing after validation is complete. """ @@ -148,18 +143,18 @@ class StoreUploadedFileTestCase(TestCase): def exception_validator(storage, filename): """ Validation test function that throws an exception """ - self.assertEqual("error_file.csv", os.path.basename(filename)) + assert 'error_file.csv' == os.path.basename(filename) with storage.open(filename, 'rb') as f: - self.assertEqual(self.file_content, f.read()) + assert self.file_content == f.read() store_file_data(storage, filename) raise FileValidationException("validation failed") def success_validator(storage, filename): """ Validation test function that is a no-op """ - self.assertIn("success_file", os.path.basename(filename)) + assert 'success_file' in os.path.basename(filename) store_file_data(storage, filename) - with self.assertRaises(FileValidationException) as error: + with pytest.raises(FileValidationException) as error: self.request.FILES = {"uploaded_file": SimpleUploadedFile("tempfile.csv", self.file_content)} store_uploaded_file( self.request, "uploaded_file", [".csv"], "error_file", @@ -213,15 +208,15 @@ class StoreUploadedFileTestCase(TestCase): file_storage, second_stored_file_name = store_uploaded_file( self.request, "nonunique_file", [".txt"], requested_file_name, self.default_max_size ) - self.assertNotEqual(first_stored_file_name, second_stored_file_name) - self.assertIn(requested_file_name, second_stored_file_name) + assert first_stored_file_name != second_stored_file_name + assert requested_file_name in second_stored_file_name self._verify_successful_upload(file_storage, second_stored_file_name, file_content) def _verify_successful_upload(self, storage, file_name, expected_content): """ Helper method that checks that the stored version of the uploaded file has the correct content """ - self.assertTrue(storage.exists(file_name)) + assert storage.exists(file_name) with storage.open(file_name, 'rb') as f: - self.assertEqual(expected_content, f.read()) + assert expected_content == f.read() @ddt.ddt @@ -231,54 +226,40 @@ class TestUniversalNewlineIterator(TestCase): """ @ddt.data(1, 2, 999) def test_line_feeds(self, buffer_size): - self.assertEqual( - [thing.decode('utf-8') for thing in UniversalNewlineIterator(StringIO(u'foo\nbar\n'), buffer_size=buffer_size)], # lint-amnesty, pylint: disable=line-too-long - ['foo\n', 'bar\n'] - ) + assert [thing.decode('utf-8') for thing + in UniversalNewlineIterator(StringIO(u'foo\nbar\n'), buffer_size=buffer_size)] == ['foo\n', 'bar\n'] @ddt.data(1, 2, 999) def test_carriage_returns(self, buffer_size): - self.assertEqual( - [thing.decode('utf-8') for thing in UniversalNewlineIterator(StringIO(u'foo\rbar\r'), buffer_size=buffer_size)], # lint-amnesty, pylint: disable=line-too-long - ['foo\n', 'bar\n'] - ) + assert [thing.decode('utf-8') for thing in + UniversalNewlineIterator(StringIO(u'foo\rbar\r'), buffer_size=buffer_size)] == ['foo\n', 'bar\n'] @ddt.data(1, 2, 999) def test_carriage_returns_and_line_feeds(self, buffer_size): - self.assertEqual( - [thing.decode('utf-8') for thing in UniversalNewlineIterator(StringIO(u'foo\r\nbar\r\n'), buffer_size=buffer_size)], # lint-amnesty, pylint: disable=line-too-long - ['foo\n', 'bar\n'] - ) + assert [thing.decode('utf-8') for thing in + UniversalNewlineIterator(StringIO(u'foo\r\nbar\r\n'), buffer_size=buffer_size)] == ['foo\n', 'bar\n'] @ddt.data(1, 2, 999) def test_no_trailing_newline(self, buffer_size): - self.assertEqual( - [thing.decode('utf-8') for thing in UniversalNewlineIterator(StringIO(u'foo\nbar'), buffer_size=buffer_size)], # lint-amnesty, pylint: disable=line-too-long - ['foo\n', 'bar'] - ) + assert [thing.decode('utf-8') for thing in + UniversalNewlineIterator(StringIO(u'foo\nbar'), buffer_size=buffer_size)] == ['foo\n', 'bar'] @ddt.data(1, 2, 999) def test_only_one_line(self, buffer_size): - self.assertEqual( - [thing.decode('utf-8') for thing in UniversalNewlineIterator(StringIO(u'foo\n'), buffer_size=buffer_size)], - ['foo\n'] - ) + assert [thing.decode('utf-8') for thing in + UniversalNewlineIterator(StringIO(u'foo\n'), buffer_size=buffer_size)] == ['foo\n'] @ddt.data(1, 2, 999) def test_only_one_line_no_trailing_newline(self, buffer_size): - self.assertEqual( - [thing.decode('utf-8') for thing in UniversalNewlineIterator(StringIO(u'foo'), buffer_size=buffer_size)], - ['foo'] - ) + assert [thing.decode('utf-8') for thing in + UniversalNewlineIterator(StringIO(u'foo'), buffer_size=buffer_size)] == ['foo'] @ddt.data(1, 2, 999) def test_empty_file(self, buffer_size): - self.assertEqual( - [thing.decode('utf-8') for thing in UniversalNewlineIterator(StringIO(u''), buffer_size=buffer_size)], - [] - ) + assert [thing.decode('utf-8') for thing in + UniversalNewlineIterator(StringIO(u''), buffer_size=buffer_size)] == [] @ddt.data(1, 2, 999) def test_unicode_data(self, buffer_size): - self.assertEqual([thing.decode('utf-8') if six.PY3 else thing for thing in - UniversalNewlineIterator(StringIO(u'héllø wo®ld'), buffer_size=buffer_size)], [u'héllø wo®ld']) # lint-amnesty, pylint: disable=line-too-long + assert [thing.decode('utf-8') for thing + in UniversalNewlineIterator(StringIO(u'héllø wo®ld'), buffer_size=buffer_size)] == [u'héllø wo®ld'] diff --git a/common/djangoapps/util/tests/test_json_request.py b/common/djangoapps/util/tests/test_json_request.py index 08590fbdb8..ec0563f412 100644 --- a/common/djangoapps/util/tests/test_json_request.py +++ b/common/djangoapps/util/tests/test_json_request.py @@ -18,58 +18,58 @@ class JsonResponseTestCase(unittest.TestCase): """ def test_empty(self): resp = JsonResponse() - self.assertIsInstance(resp, HttpResponse) - self.assertEqual(resp.content.decode('utf-8'), "") - self.assertEqual(resp.status_code, 204) - self.assertEqual(resp["content-type"], "application/json") + assert isinstance(resp, HttpResponse) + assert resp.content.decode('utf-8') == '' + assert resp.status_code == 204 + assert resp['content-type'] == 'application/json' def test_empty_string(self): resp = JsonResponse("") - self.assertIsInstance(resp, HttpResponse) - self.assertEqual(resp.content.decode('utf-8'), "") - self.assertEqual(resp.status_code, 204) - self.assertEqual(resp["content-type"], "application/json") + assert isinstance(resp, HttpResponse) + assert resp.content.decode('utf-8') == '' + assert resp.status_code == 204 + assert resp['content-type'] == 'application/json' def test_string(self): resp = JsonResponse("foo") - self.assertEqual(resp.content.decode('utf-8'), '"foo"') - self.assertEqual(resp.status_code, 200) - self.assertEqual(resp["content-type"], "application/json") + assert resp.content.decode('utf-8') == '"foo"' + assert resp.status_code == 200 + assert resp['content-type'] == 'application/json' def test_dict(self): obj = {"foo": "bar"} resp = JsonResponse(obj) compare = json.loads(resp.content.decode('utf-8')) - self.assertEqual(obj, compare) - self.assertEqual(resp.status_code, 200) - self.assertEqual(resp["content-type"], "application/json") + assert obj == compare + assert resp.status_code == 200 + assert resp['content-type'] == 'application/json' def test_set_status_kwarg(self): obj = {"error": "resource not found"} resp = JsonResponse(obj, status=404) compare = json.loads(resp.content.decode('utf-8')) - self.assertEqual(obj, compare) - self.assertEqual(resp.status_code, 404) - self.assertEqual(resp["content-type"], "application/json") + assert obj == compare + assert resp.status_code == 404 + assert resp['content-type'] == 'application/json' def test_set_status_arg(self): obj = {"error": "resource not found"} resp = JsonResponse(obj, 404) compare = json.loads(resp.content.decode('utf-8')) - self.assertEqual(obj, compare) - self.assertEqual(resp.status_code, 404) - self.assertEqual(resp["content-type"], "application/json") + assert obj == compare + assert resp.status_code == 404 + assert resp['content-type'] == 'application/json' def test_encoder(self): obj = [1, 2, 3] encoder = object() with mock.patch.object(json, "dumps", return_value="[1,2,3]") as dumps: resp = JsonResponse(obj, encoder=encoder) - self.assertEqual(resp.status_code, 200) + assert resp.status_code == 200 compare = json.loads(resp.content.decode('utf-8')) - self.assertEqual(obj, compare) + assert obj == compare kwargs = dumps.call_args[1] - self.assertIs(kwargs["cls"], encoder) + assert kwargs['cls'] is encoder class JsonResponseBadRequestTestCase(unittest.TestCase): @@ -80,49 +80,49 @@ class JsonResponseBadRequestTestCase(unittest.TestCase): def test_empty(self): resp = JsonResponseBadRequest() - self.assertIsInstance(resp, HttpResponseBadRequest) - self.assertEqual(resp.content.decode("utf-8"), "") - self.assertEqual(resp.status_code, 400) - self.assertEqual(resp["content-type"], "application/json") + assert isinstance(resp, HttpResponseBadRequest) + assert resp.content.decode('utf-8') == '' + assert resp.status_code == 400 + assert resp['content-type'] == 'application/json' def test_empty_string(self): resp = JsonResponseBadRequest("") - self.assertIsInstance(resp, HttpResponse) - self.assertEqual(resp.content.decode('utf-8'), "") - self.assertEqual(resp.status_code, 400) - self.assertEqual(resp["content-type"], "application/json") + assert isinstance(resp, HttpResponse) + assert resp.content.decode('utf-8') == '' + assert resp.status_code == 400 + assert resp['content-type'] == 'application/json' def test_dict(self): obj = {"foo": "bar"} resp = JsonResponseBadRequest(obj) compare = json.loads(resp.content.decode('utf-8')) - self.assertEqual(obj, compare) - self.assertEqual(resp.status_code, 400) - self.assertEqual(resp["content-type"], "application/json") + assert obj == compare + assert resp.status_code == 400 + assert resp['content-type'] == 'application/json' def test_set_status_kwarg(self): obj = {"error": "resource not found"} resp = JsonResponseBadRequest(obj, status=404) compare = json.loads(resp.content.decode('utf-8')) - self.assertEqual(obj, compare) - self.assertEqual(resp.status_code, 404) - self.assertEqual(resp["content-type"], "application/json") + assert obj == compare + assert resp.status_code == 404 + assert resp['content-type'] == 'application/json' def test_set_status_arg(self): obj = {"error": "resource not found"} resp = JsonResponseBadRequest(obj, 404) compare = json.loads(resp.content.decode('utf-8')) - self.assertEqual(obj, compare) - self.assertEqual(resp.status_code, 404) - self.assertEqual(resp["content-type"], "application/json") + assert obj == compare + assert resp.status_code == 404 + assert resp['content-type'] == 'application/json' def test_encoder(self): obj = [1, 2, 3] encoder = object() with mock.patch.object(json, "dumps", return_value="[1,2,3]") as dumps: resp = JsonResponseBadRequest(obj, encoder=encoder) - self.assertEqual(resp.status_code, 400) + assert resp.status_code == 400 compare = json.loads(resp.content.decode('utf-8')) - self.assertEqual(obj, compare) + assert obj == compare kwargs = dumps.call_args[1] - self.assertIs(kwargs["cls"], encoder) + assert kwargs['cls'] is encoder diff --git a/common/djangoapps/util/tests/test_keyword_sub_utils.py b/common/djangoapps/util/tests/test_keyword_sub_utils.py index 8938a49f35..2010a74b93 100644 --- a/common/djangoapps/util/tests/test_keyword_sub_utils.py +++ b/common/djangoapps/util/tests/test_keyword_sub_utils.py @@ -48,8 +48,8 @@ class KeywordSubTest(ModuleStoreTestCase): test_string, self.context, ) - self.assertIn(course_name, result) - self.assertEqual(result, expected) + assert course_name in result + assert result == expected def test_anonymous_id_sub(self): """ @@ -60,8 +60,8 @@ class KeywordSubTest(ModuleStoreTestCase): result = Ks.substitute_keywords_with_data( test_string, self.context, ) - self.assertNotIn('%%USER_ID%%', result) - self.assertIn(anonymous_id, result) + assert '%%USER_ID%%' not in result + assert anonymous_id in result def test_name_sub(self): """ @@ -73,8 +73,8 @@ class KeywordSubTest(ModuleStoreTestCase): test_string, self.context, ) - self.assertNotIn('%%USER_FULLNAME%%', result) - self.assertIn(user_name, result) + assert '%%USER_FULLNAME%%' not in result + assert user_name in result def test_illegal_subtag(self): """ @@ -85,7 +85,7 @@ class KeywordSubTest(ModuleStoreTestCase): test_string, self.context, ) - self.assertEqual(test_string, result) + assert test_string == result def test_should_not_sub(self): """ @@ -96,7 +96,7 @@ class KeywordSubTest(ModuleStoreTestCase): test_string, self.context, ) - self.assertEqual(test_string, result) + assert test_string == result @file_data('fixtures/test_keywordsub_multiple_tags.json') def test_sub_multiple_tags(self, test_string, expected): @@ -107,7 +107,7 @@ class KeywordSubTest(ModuleStoreTestCase): result = Ks.substitute_keywords_with_data( test_string, self.context, ) - self.assertEqual(result, expected) + assert result == expected def test_subbing_no_userid_or_courseid(self): """ @@ -119,10 +119,10 @@ class KeywordSubTest(ModuleStoreTestCase): (key, value) for key, value in six.iteritems(self.context) if key != 'course_title' ) result = Ks.substitute_keywords_with_data(test_string, no_course_context) - self.assertEqual(test_string, result) + assert test_string == result no_user_id_context = dict( (key, value) for key, value in six.iteritems(self.context) if key != 'user_id' ) result = Ks.substitute_keywords_with_data(test_string, no_user_id_context) - self.assertEqual(test_string, result) + assert test_string == result diff --git a/common/djangoapps/util/tests/test_memcache.py b/common/djangoapps/util/tests/test_memcache.py index 1b2e8f1a2c..c636d43e57 100644 --- a/common/djangoapps/util/tests/test_memcache.py +++ b/common/djangoapps/util/tests/test_memcache.py @@ -26,18 +26,18 @@ class MemcacheTest(TestCase): def test_safe_key(self): key = safe_key('test', 'prefix', 'version') - self.assertEqual(key, 'prefix:version:test') + assert key == 'prefix:version:test' def test_numeric_inputs(self): # Numeric key - self.assertEqual(safe_key(1, 'prefix', 'version'), 'prefix:version:1') + assert safe_key(1, 'prefix', 'version') == 'prefix:version:1' # Numeric prefix - self.assertEqual(safe_key('test', 5, 'version'), '5:version:test') + assert safe_key('test', 5, 'version') == '5:version:test' # Numeric version - self.assertEqual(safe_key('test', 'prefix', 5), 'prefix:5:test') + assert safe_key('test', 'prefix', 5) == 'prefix:5:test' def test_safe_key_long(self): @@ -51,22 +51,21 @@ class MemcacheTest(TestCase): key = safe_key(key, '', '') # The key should now be valid - self.assertTrue(self._is_valid_key(key), - msg="Failed for key length {0}".format(length)) + assert self._is_valid_key(key), 'Failed for key length {0}'.format(length) def test_long_key_prefix_version(self): # Long key key = safe_key('a' * 300, 'prefix', 'version') - self.assertTrue(self._is_valid_key(key)) + assert self._is_valid_key(key) # Long prefix key = safe_key('key', 'a' * 300, 'version') - self.assertTrue(self._is_valid_key(key)) + assert self._is_valid_key(key) # Long version key = safe_key('key', 'prefix', 'a' * 300) - self.assertTrue(self._is_valid_key(key)) + assert self._is_valid_key(key) def test_safe_key_unicode(self): @@ -79,8 +78,7 @@ class MemcacheTest(TestCase): key = safe_key(key, '', '') # The key should now be valid - self.assertTrue(self._is_valid_key(key), - msg="Failed for unicode character {0}".format(unicode_char)) + assert self._is_valid_key(key), 'Failed for unicode character {0}'.format(unicode_char) def test_safe_key_prefix_unicode(self): @@ -93,8 +91,7 @@ class MemcacheTest(TestCase): key = safe_key('test', prefix, '') # The key should now be valid - self.assertTrue(self._is_valid_key(key), - msg="Failed for unicode character {0}".format(unicode_char)) + assert self._is_valid_key(key), 'Failed for unicode character {0}'.format(unicode_char) def test_safe_key_version_unicode(self): @@ -107,8 +104,7 @@ class MemcacheTest(TestCase): key = safe_key('test', '', version) # The key should now be valid - self.assertTrue(self._is_valid_key(key), - msg="Failed for unicode character {0}".format(unicode_char)) + assert self._is_valid_key(key), 'Failed for unicode character {0}'.format(unicode_char) def _is_valid_key(self, key): """ diff --git a/common/djangoapps/util/tests/test_milestones_helpers.py b/common/djangoapps/util/tests/test_milestones_helpers.py index f8dcc922cc..1758cb55dd 100644 --- a/common/djangoapps/util/tests/test_milestones_helpers.py +++ b/common/djangoapps/util/tests/test_milestones_helpers.py @@ -65,27 +65,27 @@ class MilestonesHelpersTestCase(ModuleStoreTestCase): 'ENABLE_PREREQUISITE_COURSES': feature_flags[0], 'MILESTONES_APP': feature_flags[1] }): - self.assertEqual(feature_flags[2], milestones_helpers.is_prerequisite_courses_enabled()) + assert feature_flags[2] == milestones_helpers.is_prerequisite_courses_enabled() def test_add_milestone_returns_none_when_app_disabled(self): response = milestones_helpers.add_milestone(milestone_data=self.milestone) - self.assertIsNone(response) + assert response is None def test_get_milestones_returns_none_when_app_disabled(self): response = milestones_helpers.get_milestones(namespace="whatever") - self.assertEqual(len(response), 0) + assert len(response) == 0 def test_get_milestone_relationship_types_returns_none_when_app_disabled(self): response = milestones_helpers.get_milestone_relationship_types() - self.assertEqual(len(response), 0) + assert len(response) == 0 def test_add_course_milestone_returns_none_when_app_disabled(self): response = milestones_helpers.add_course_milestone(six.text_type(self.course.id), 'requires', self.milestone) - self.assertIsNone(response) + assert response is None def test_get_course_milestones_returns_none_when_app_disabled(self): response = milestones_helpers.get_course_milestones(six.text_type(self.course.id)) - self.assertEqual(len(response), 0) + assert len(response) == 0 def test_add_course_content_milestone_returns_none_when_app_disabled(self): response = milestones_helpers.add_course_content_milestone( @@ -94,7 +94,7 @@ class MilestonesHelpersTestCase(ModuleStoreTestCase): 'requires', self.milestone ) - self.assertIsNone(response) + assert response is None def test_get_course_content_milestones_returns_none_when_app_disabled(self): response = milestones_helpers.get_course_content_milestones( @@ -102,28 +102,28 @@ class MilestonesHelpersTestCase(ModuleStoreTestCase): 'i4x://doesnt/matter/for/this/test', 'requires' ) - self.assertEqual(len(response), 0) + assert len(response) == 0 def test_remove_content_references_returns_none_when_app_disabled(self): response = milestones_helpers.remove_content_references("i4x://any/content/id/will/do") - self.assertIsNone(response) + assert response is None def test_get_namespace_choices_returns_values_when_app_disabled(self): response = milestones_helpers.get_namespace_choices() - self.assertIn('ENTRANCE_EXAM', response) + assert 'ENTRANCE_EXAM' in response def test_get_course_milestones_fulfillment_paths_returns_none_when_app_disabled(self): response = milestones_helpers.get_course_milestones_fulfillment_paths(six.text_type(self.course.id), self.user) - self.assertIsNone(response) + assert response is None def test_add_user_milestone_returns_none_when_app_disabled(self): response = milestones_helpers.add_user_milestone(self.user, self.milestone) - self.assertIsNone(response) + assert response is None def test_get_service_returns_none_when_app_disabled(self): """MilestonesService is None when app disabled""" response = milestones_helpers.get_service() - self.assertIsNone(response) + assert response is None @patch.dict(settings.FEATURES, {'MILESTONES_APP': True}) def test_any_unfulfilled_milestones(self): @@ -134,9 +134,9 @@ class MilestonesHelpersTestCase(ModuleStoreTestCase): # Should not raise any exceptions milestones_helpers.any_unfulfilled_milestones(self.course.id, self.user['id']) - with self.assertRaises(InvalidCourseKeyException): + with pytest.raises(InvalidCourseKeyException): milestones_helpers.any_unfulfilled_milestones(None, self.user['id']) - with self.assertRaises(InvalidUserException): + with pytest.raises(InvalidUserException): milestones_helpers.any_unfulfilled_milestones(self.course.id, None) @patch.dict(settings.FEATURES, {'MILESTONES_APP': True}) diff --git a/common/djangoapps/util/tests/test_password_policy_validators.py b/common/djangoapps/util/tests/test_password_policy_validators.py index f7f4eb3da3..3b11fadde5 100644 --- a/common/djangoapps/util/tests/test_password_policy_validators.py +++ b/common/djangoapps/util/tests/test_password_policy_validators.py @@ -4,6 +4,7 @@ import unittest +import pytest from ddt import data, ddt, unpack from django.contrib.auth.models import User # lint-amnesty, pylint: disable=imported-auth-user from django.core.exceptions import ValidationError @@ -41,9 +42,9 @@ class PasswordPolicyValidatorsTestCase(unittest.TestCase): if msg is None: validate_password(password, user) else: - with self.assertRaises(ValidationError) as cm: + with pytest.raises(ValidationError) as cm: validate_password(password, user) - self.assertIn(msg, ' '.join(cm.exception.messages)) + assert msg in ' '.join(cm.value.messages) def test_unicode_password(self): """ Tests that validate_password enforces unicode """ @@ -51,8 +52,8 @@ class PasswordPolicyValidatorsTestCase(unittest.TestCase): byte_str = unicode_str.encode('utf-8') # Sanity checks and demonstration of why this test is useful - self.assertEqual(len(byte_str), 4) - self.assertEqual(len(unicode_str), 1) + assert len(byte_str) == 4 + assert len(unicode_str) == 1 # Test length check self.validation_errors_checker(byte_str, 'This password is too short. It must contain at least 2 characters.') @@ -65,7 +66,7 @@ class PasswordPolicyValidatorsTestCase(unittest.TestCase): """ Tests that validate_password normalizes passwords """ # s ̣ ̇ (s with combining dot below and combining dot above) not_normalized_password = u'\u0073\u0323\u0307' - self.assertEqual(len(not_normalized_password), 3) + assert len(not_normalized_password) == 3 # When we normalize we expect the not_normalized password to fail # because it should be normalized to u'\u1E69' -> ṩ @@ -98,7 +99,7 @@ class PasswordPolicyValidatorsTestCase(unittest.TestCase): def test_password_instructions(self, config, msg): """ Tests password instructions """ with override_settings(AUTH_PASSWORD_VALIDATORS=config): - self.assertIn(msg, password_validators_instruction_texts()) + assert msg in password_validators_instruction_texts() @data( (u'userna', u'username', 'test@example.com', 'The password is too similar to the username.'), diff --git a/common/djangoapps/util/tests/test_sandboxing.py b/common/djangoapps/util/tests/test_sandboxing.py index d3b3240eff..c96ee3c5b2 100644 --- a/common/djangoapps/util/tests/test_sandboxing.py +++ b/common/djangoapps/util/tests/test_sandboxing.py @@ -20,22 +20,22 @@ class SandboxingTest(TestCase): """ Test to make sure that a non-match returns false """ - self.assertFalse(can_execute_unsafe_code(CourseLocator('edX', 'notful', 'empty'))) - self.assertFalse(can_execute_unsafe_code(LibraryLocator('edY', 'test_bank'))) + assert not can_execute_unsafe_code(CourseLocator('edX', 'notful', 'empty')) + assert not can_execute_unsafe_code(LibraryLocator('edY', 'test_bank')) @override_settings(COURSES_WITH_UNSAFE_CODE=['edX/full/.*']) def test_sandbox_inclusion(self): """ Test to make sure that a match works across course runs """ - self.assertTrue(can_execute_unsafe_code(CourseKey.from_string('edX/full/2012_Fall'))) - self.assertTrue(can_execute_unsafe_code(CourseKey.from_string('edX/full/2013_Spring'))) - self.assertFalse(can_execute_unsafe_code(LibraryLocator('edX', 'test_bank'))) + assert can_execute_unsafe_code(CourseKey.from_string('edX/full/2012_Fall')) + assert can_execute_unsafe_code(CourseKey.from_string('edX/full/2013_Spring')) + assert not can_execute_unsafe_code(LibraryLocator('edX', 'test_bank')) def test_courselikes_with_unsafe_code_default(self): """ Test that the default setting for COURSES_WITH_UNSAFE_CODE is an empty setting, e.g. we don't use @override_settings in these tests # lint-amnesty, pylint: disable=line-too-long """ - self.assertFalse(can_execute_unsafe_code(CourseLocator('edX', 'full', '2012_Fall'))) - self.assertFalse(can_execute_unsafe_code(CourseLocator('edX', 'full', '2013_Spring'))) - self.assertFalse(can_execute_unsafe_code(LibraryLocator('edX', 'test_bank'))) + assert not can_execute_unsafe_code(CourseLocator('edX', 'full', '2012_Fall')) + assert not can_execute_unsafe_code(CourseLocator('edX', 'full', '2013_Spring')) + assert not can_execute_unsafe_code(LibraryLocator('edX', 'test_bank')) diff --git a/common/djangoapps/util/tests/test_string_utils.py b/common/djangoapps/util/tests/test_string_utils.py index 49872c54c6..23943473c8 100644 --- a/common/djangoapps/util/tests/test_string_utils.py +++ b/common/djangoapps/util/tests/test_string_utils.py @@ -2,7 +2,7 @@ Tests for string_utils.py """ - +import pytest from django.test import TestCase from common.djangoapps.util.string_utils import str_to_bool @@ -13,22 +13,22 @@ class StringUtilsTest(TestCase): Tests for str_to_bool. """ def test_str_to_bool_true(self): - self.assertTrue(str_to_bool('True')) - self.assertTrue(str_to_bool('true')) - self.assertTrue(str_to_bool('trUe')) + assert str_to_bool('True') + assert str_to_bool('true') + assert str_to_bool('trUe') def test_str_to_bool_false(self): - self.assertFalse(str_to_bool('Tru')) - self.assertFalse(str_to_bool('False')) - self.assertFalse(str_to_bool('false')) - self.assertFalse(str_to_bool('')) - self.assertFalse(str_to_bool(None)) - self.assertFalse(str_to_bool('anything')) + assert not str_to_bool('Tru') + assert not str_to_bool('False') + assert not str_to_bool('false') + assert not str_to_bool('') + assert not str_to_bool(None) + assert not str_to_bool('anything') def test_str_to_bool_errors(self): def test_raises_error(val): - with self.assertRaises(AttributeError): - self.assertFalse(str_to_bool(val)) + with pytest.raises(AttributeError): + assert not str_to_bool(val) test_raises_error({}) test_raises_error([]) diff --git a/common/djangoapps/xblock_django/tests/test_api.py b/common/djangoapps/xblock_django/tests/test_api.py index cffb70be41..9852d0c135 100644 --- a/common/djangoapps/xblock_django/tests/test_api.py +++ b/common/djangoapps/xblock_django/tests/test_api.py @@ -71,10 +71,10 @@ class XBlockSupportTestCase(CacheIsolationTestCase): of whether or not XBlockStudioConfigurationFlag is enabled. """ XBlockStudioConfiguration.objects.all().delete() - self.assertFalse(XBlockStudioConfigurationFlag.is_enabled()) - self.assertEqual(0, len(authorable_xblocks(allow_unsupported=True))) + assert not XBlockStudioConfigurationFlag.is_enabled() + assert 0 == len(authorable_xblocks(allow_unsupported=True)) XBlockStudioConfigurationFlag(enabled=True).save() - self.assertEqual(0, len(authorable_xblocks(allow_unsupported=True))) + assert 0 == len(authorable_xblocks(allow_unsupported=True)) def test_authorable_blocks(self): """ @@ -103,21 +103,21 @@ class XBlockSupportTestCase(CacheIsolationTestCase): """ Verifies the returned xblock state. """ - self.assertEqual(name, block.name) - self.assertEqual(template, block.template) - self.assertEqual(support_level, block.support_level) + assert name == block.name + assert template == block.template + assert support_level == block.support_level # There are no xblocks with name video. authorable_blocks = authorable_xblocks(name="video") - self.assertEqual(0, len(authorable_blocks)) + assert 0 == len(authorable_blocks) # There is only a single html xblock. authorable_blocks = authorable_xblocks(name="html") - self.assertEqual(1, len(authorable_blocks)) + assert 1 == len(authorable_blocks) verify_xblock_fields("html", "zoom", XBlockStudioConfiguration.PROVISIONAL_SUPPORT, authorable_blocks[0]) authorable_blocks = authorable_xblocks(name="problem", allow_unsupported=True) - self.assertEqual(3, len(authorable_blocks)) + assert 3 == len(authorable_blocks) no_template = None circuit = None multiple_choice = None diff --git a/common/djangoapps/xblock_django/tests/test_user_service.py b/common/djangoapps/xblock_django/tests/test_user_service.py index 871ec11477..8edf900093 100644 --- a/common/djangoapps/xblock_django/tests/test_user_service.py +++ b/common/djangoapps/xblock_django/tests/test_user_service.py @@ -2,7 +2,7 @@ Tests for the DjangoXBlockUserService. """ - +import pytest from django.test import TestCase from opaque_keys.edx.keys import CourseKey @@ -38,26 +38,21 @@ class UserServiceTestCase(TestCase): """ A set of assertions for an anonymous XBlockUser. """ - self.assertFalse(xb_user.opt_attrs[ATTR_KEY_IS_AUTHENTICATED]) - self.assertIsNone(xb_user.full_name) + assert not xb_user.opt_attrs[ATTR_KEY_IS_AUTHENTICATED] + assert xb_user.full_name is None self.assertListEqual(xb_user.emails, []) def assert_xblock_user_matches_django(self, xb_user, dj_user): """ A set of assertions for comparing a XBlockUser to a django User """ - self.assertTrue(xb_user.opt_attrs[ATTR_KEY_IS_AUTHENTICATED]) - self.assertEqual(xb_user.emails[0], dj_user.email) - self.assertEqual(xb_user.full_name, dj_user.profile.name) - self.assertEqual(xb_user.opt_attrs[ATTR_KEY_USERNAME], dj_user.username) - self.assertEqual(xb_user.opt_attrs[ATTR_KEY_USER_ID], dj_user.id) - self.assertFalse(xb_user.opt_attrs[ATTR_KEY_USER_IS_STAFF]) - self.assertTrue( - all( - pref in USER_PREFERENCES_WHITE_LIST - for pref in xb_user.opt_attrs[ATTR_KEY_USER_PREFERENCES] - ) - ) + assert xb_user.opt_attrs[ATTR_KEY_IS_AUTHENTICATED] + assert xb_user.emails[0] == dj_user.email + assert xb_user.full_name == dj_user.profile.name + assert xb_user.opt_attrs[ATTR_KEY_USERNAME] == dj_user.username + assert xb_user.opt_attrs[ATTR_KEY_USER_ID] == dj_user.id + assert not xb_user.opt_attrs[ATTR_KEY_USER_IS_STAFF] + assert all(((pref in USER_PREFERENCES_WHITE_LIST) for pref in xb_user.opt_attrs[ATTR_KEY_USER_PREFERENCES])) def test_convert_anon_user(self): """ @@ -65,7 +60,7 @@ class UserServiceTestCase(TestCase): """ django_user_service = DjangoXBlockUserService(self.anon_user) xb_user = django_user_service.get_current_user() - self.assertTrue(xb_user.is_current_user) + assert xb_user.is_current_user self.assert_is_anon_xb_user(xb_user) def test_convert_authenticate_user(self): @@ -74,7 +69,7 @@ class UserServiceTestCase(TestCase): """ django_user_service = DjangoXBlockUserService(self.user) xb_user = django_user_service.get_current_user() - self.assertTrue(xb_user.is_current_user) + assert xb_user.is_current_user self.assert_xblock_user_matches_django(xb_user, self.user) def test_get_anonymous_user_id_returns_none_for_non_staff_users(self): @@ -87,7 +82,7 @@ class UserServiceTestCase(TestCase): username=self.user.username, course_id='edx/toy/2012_Fall' ) - self.assertIsNone(anonymous_user_id) + assert anonymous_user_id is None def test_get_anonymous_user_id_returns_none_for_non_existing_users(self): """ @@ -96,7 +91,7 @@ class UserServiceTestCase(TestCase): django_user_service = DjangoXBlockUserService(self.user, user_is_staff=True) anonymous_user_id = django_user_service.get_anonymous_user_id(username="No User", course_id='edx/toy/2012_Fall') - self.assertIsNone(anonymous_user_id) + assert anonymous_user_id is None def test_get_anonymous_user_id_returns_id_for_existing_users(self): """ @@ -114,7 +109,7 @@ class UserServiceTestCase(TestCase): course_id='edX/toy/2012_Fall' ) - self.assertEqual(anonymous_user_id, anon_user_id) + assert anonymous_user_id == anon_user_id def test_external_id(self): """ @@ -126,5 +121,5 @@ class UserServiceTestCase(TestCase): ext_id1 = django_user_service.get_external_user_id('test1') ext_id2 = django_user_service.get_external_user_id('test2') assert ext_id1 != ext_id2 - with self.assertRaises(ValueError): + with pytest.raises(ValueError): django_user_service.get_external_user_id('unknown')