diff --git a/common/djangoapps/student/models.py b/common/djangoapps/student/models.py index da7ce8d4b1..6da10a300f 100644 --- a/common/djangoapps/student/models.py +++ b/common/djangoapps/student/models.py @@ -2269,9 +2269,11 @@ def get_user_by_username_or_email(username_or_email): retired, or the user is in the process of being retired. """ username_or_email = strip_if_string(username_or_email) + user = None if '@' in username_or_email: - user = User.objects.get(email=username_or_email) - else: + user = User.objects.filter(email=username_or_email).first() + + if user is None: user = User.objects.get(username=username_or_email) UserRetirementRequest = apps.get_model('user_api', 'UserRetirementRequest') if UserRetirementRequest.has_user_requested_retirement(user): diff --git a/lms/djangoapps/instructor/tests/test_tools.py b/lms/djangoapps/instructor/tests/test_tools.py index 7d05da99dd..22236e3266 100644 --- a/lms/djangoapps/instructor/tests/test_tools.py +++ b/lms/djangoapps/instructor/tests/test_tools.py @@ -8,6 +8,7 @@ import unittest import mock import six +from django.contrib.auth.models import User from django.test import TestCase from django.test.utils import override_settings from pytz import UTC @@ -357,3 +358,31 @@ def msk_from_problem_urlname(course_id, urlname, block_type='problem'): urlname = urlname[:-4] return course_id.make_usage_key(block_type, urlname) + + +@attr(shard=1) +class TestStudentFromIdentifier(TestCase): + """ + Test get_student_from_identifier() + """ + def setUp(self): + """ + Fixtures + """ + super(TestStudentFromIdentifier, self).setUp() + self.students = [ + UserFactory.create(username='foo@touchstone'), # a student with character `@` in user name + UserFactory.create() + ] + + def test_valid_student_id(self): + for student in self.students: + assert student == tools.get_student_from_identifier(student.username) + + def test_valid_student_email(self): + for student in self.students: + assert student == tools.get_student_from_identifier(student.email) + + def test_invalid_student_id(self): + with self.assertRaises(User.DoesNotExist): + assert tools.get_student_from_identifier("invalid")