@@ -237,6 +237,60 @@ def test_list_enriches_course_progress_from_snowflake(self, mock_source_cls):
237237 self .assertEqual (response .data ['results' ][0 ]['enrollment_id' ], enrollment .enrollment_id )
238238 self .assertEqual (response .data ['results' ][0 ]['course_progress' ], 0.87 )
239239
240+ def test_list_excludes_enrollments_of_unlinked_learners (self ):
241+ """
242+ Test that the enrollment list endpoint excludes enrollments belonging to unlinked learners.
243+ """
244+ linked_learner = EnterpriseLearnerFactory (
245+ enterprise_customer_uuid = self .enterprise_id ,
246+ is_linked = True ,
247+ )
248+ unlinked_learner = EnterpriseLearnerFactory (
249+ enterprise_customer_uuid = self .enterprise_id ,
250+ is_linked = False ,
251+ )
252+ linked_enrollment = EnterpriseLearnerEnrollmentFactory (
253+ enterprise_customer_uuid = self .enterprise_id ,
254+ is_consent_granted = True ,
255+ enterprise_user_id = linked_learner .enterprise_user_id ,
256+ )
257+ EnterpriseLearnerEnrollmentFactory (
258+ enterprise_customer_uuid = self .enterprise_id ,
259+ is_consent_granted = True ,
260+ enterprise_user_id = unlinked_learner .enterprise_user_id ,
261+ )
262+
263+ url = reverse ('v1:enterprise-learner-enrollment-list' , kwargs = {'enterprise_id' : self .enterprise_id })
264+ response = self .client .get (url )
265+
266+ self .assertEqual (response .status_code , status .HTTP_200_OK )
267+ results = response .json ()['results' ]
268+ self .assertEqual (len (results ), 1 )
269+ self .assertEqual (results [0 ]['enrollment_id' ], linked_enrollment .enrollment_id )
270+
271+ def test_overview_number_of_users_excludes_unlinked_learners (self ):
272+ """
273+ Test that `number_of_users` in the overview response only counts learners with `is_linked=True`.
274+ """
275+ EnterpriseLearnerFactory (
276+ enterprise_customer_uuid = self .enterprise_id ,
277+ is_linked = True ,
278+ )
279+ EnterpriseLearnerFactory (
280+ enterprise_customer_uuid = self .enterprise_id ,
281+ is_linked = True ,
282+ )
283+ EnterpriseLearnerFactory (
284+ enterprise_customer_uuid = self .enterprise_id ,
285+ is_linked = False ,
286+ )
287+
288+ url = reverse ('v1:enterprise-learner-enrollment-overview' , kwargs = {'enterprise_id' : self .enterprise_id })
289+ response = self .client .get (url )
290+
291+ self .assertEqual (response .status_code , status .HTTP_200_OK )
292+ self .assertEqual (response .json ()['number_of_users' ], 2 )
293+
240294 @mock .patch ('enterprise_data.api.v1.views.enterprise_learner.SnowflakeCourseProgressSource' )
241295 def test_list_returns_200_when_snowflake_enrichment_fails (self , mock_source_cls ):
242296 enterprise_learner = EnterpriseLearnerFactory (
@@ -334,7 +388,10 @@ def test_get_queryset_adds_placeholder_metadata_columns(self, mock_apply_filters
334388
335389 result = viewset .get_queryset ()
336390
337- mock_filter .assert_called_once_with (enterprise_customer_uuid = self .enterprise_id )
391+ mock_filter .assert_called_once_with (
392+ enterprise_customer_uuid = self .enterprise_id ,
393+ enterprise_user__is_linked = True ,
394+ )
338395 enrollments .extra .assert_called_once_with (select = {
339396 'course_progress' : 'NULL' ,
340397 'course_passing_grade' : 'NULL' ,
0 commit comments