diff --git a/common/djangoapps/student/admin.py b/common/djangoapps/student/admin.py index a650ac28fc..c6a0888779 100644 --- a/common/djangoapps/student/admin.py +++ b/common/djangoapps/student/admin.py @@ -159,7 +159,6 @@ class LinkedInAddToProfileConfigurationAdmin(admin.ModelAdmin): class CourseEnrollmentForm(forms.ModelForm): - def __init__(self, *args, **kwargs): # If args is a QueryDict, then the ModelForm addition request came in as a POST with a course ID string. # Change the course ID string to a CourseLocator object by copying the QueryDict to make it mutable. @@ -194,6 +193,15 @@ class CourseEnrollmentForm(forms.ModelForm): return course_key + def save(self, *args, **kwargs): + course_enrollment = super(CourseEnrollmentForm, self).save(commit=False) + user = self.cleaned_data['user'] + course_overview = self.cleaned_data['course'] + enrollment = CourseEnrollment.get_or_create_enrollment(user, course_overview.id) + course_enrollment.id = enrollment.id + course_enrollment.created = enrollment.created + return course_enrollment + class Meta: model = CourseEnrollment fields = '__all__' @@ -204,7 +212,7 @@ class CourseEnrollmentAdmin(admin.ModelAdmin): """ Admin interface for the CourseEnrollment model. """ list_display = ('id', 'course_id', 'mode', 'user', 'is_active',) list_filter = ('mode', 'is_active',) - raw_id_fields = ('user',) + raw_id_fields = ('user', 'course') search_fields = ('course__id', 'mode', 'user__username',) form = CourseEnrollmentForm diff --git a/common/djangoapps/student/tests/test_admin_views.py b/common/djangoapps/student/tests/test_admin_views.py index 6d99b8c9a2..c3b57e6d6d 100644 --- a/common/djangoapps/student/tests/test_admin_views.py +++ b/common/djangoapps/student/tests/test_admin_views.py @@ -13,13 +13,15 @@ from django.contrib.auth.models import User from django.forms import ValidationError from django.test import TestCase, override_settings from django.urls import reverse +from django.utils.timezone import now from mock import Mock -from student.admin import COURSE_ENROLLMENT_ADMIN_SWITCH, UserAdmin -from student.models import LoginFailures +from student.admin import COURSE_ENROLLMENT_ADMIN_SWITCH, UserAdmin, CourseEnrollmentForm +from student.models import CourseEnrollment, LoginFailures from student.tests.factories import CourseEnrollmentFactory, UserFactory from xmodule.modulestore.tests.django_utils import SharedModuleStoreTestCase from xmodule.modulestore.tests.factories import CourseFactory +from openedx.core.djangoapps.content.course_overviews.tests.factories import CourseOverviewFactory class AdminCourseRolesPageTest(SharedModuleStoreTestCase): @@ -370,3 +372,53 @@ class LoginFailuresAdminTest(TestCase): ) count = LoginFailures.objects.count() self.assertEqual(count, start_count - 1) + + +class CourseEnrollmentAdminFormTest(SharedModuleStoreTestCase): + """ + Unit test for CourseEnrollment admin ModelForm. + """ + @classmethod + def setUpClass(cls): + super(CourseEnrollmentAdminFormTest, cls).setUpClass() + cls.course = CourseOverviewFactory(start=now()) + + def setUp(self): + super(CourseEnrollmentAdminFormTest, self).setUp() + self.user = UserFactory.create() + + def test_admin_model_form_create(self): + """ + Test CourseEnrollmentAdminForm creation. + """ + self.assertEqual(CourseEnrollment.objects.count(), 0) + + form = CourseEnrollmentForm({ + 'user': self.user.id, + 'course': six.text_type(self.course.id), + 'is_active': True, + 'mode': 'audit', + }) + self.assertTrue(form.is_valid()) + enrollment = form.save() + self.assertEqual(CourseEnrollment.objects.count(), 1) + self.assertEqual(CourseEnrollment.objects.first(), enrollment) + + def test_admin_model_form_update(self): + """ + Test CourseEnrollmentAdminForm update. + """ + enrollment = CourseEnrollment.get_or_create_enrollment(self.user, self.course.id) + count = CourseEnrollment.objects.count() + form = CourseEnrollmentForm({ + 'user': self.user.id, + 'course': six.text_type(self.course.id), + 'is_active': False, + 'mode': 'audit'}, + instance=enrollment + ) + self.assertTrue(form.is_valid()) + course_enrollment = form.save() + self.assertEqual(count, CourseEnrollment.objects.count()) + self.assertFalse(course_enrollment.is_active) + self.assertEqual(enrollment.id, course_enrollment.id)