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