From b2ccb8d4d21a34191c7cd2c729ccd40d22577306 Mon Sep 17 00:00:00 2001 From: Santhosh Kumar Date: Wed, 2 Sep 2026 10:23:32 +0000 Subject: [PATCH 1/6] feat: route Insights enrollment APIs through Snowflake --- .../insights_snowflake/mappers/enrollment.py | 149 +++++++++ .../insights_snowflake/queries/enrollment.py | 117 +++++++ .../insights_snowflake/service.py | 44 +++ .../tests/test_insights_snowflake.py | 291 +++++++++++++++++- .../v0/tests/views/test_courses.py | 189 ++++++++++++ analytics_data_api/v0/views/courses.py | 69 ++++- 6 files changed, 845 insertions(+), 14 deletions(-) create mode 100644 analytics_data_api/insights_snowflake/mappers/enrollment.py create mode 100644 analytics_data_api/insights_snowflake/queries/enrollment.py diff --git a/analytics_data_api/insights_snowflake/mappers/enrollment.py b/analytics_data_api/insights_snowflake/mappers/enrollment.py new file mode 100644 index 00000000..c9e9e062 --- /dev/null +++ b/analytics_data_api/insights_snowflake/mappers/enrollment.py @@ -0,0 +1,149 @@ +"""Map Snowflake enrollment rows into the existing API response shapes.""" + +from itertools import groupby + +from analytics_data_api.constants import enrollment_modes, genders +from analytics_data_api.v0 import models + +GENDER_FIELD_MAP = { + 'f': genders.FEMALE, + genders.FEMALE: genders.FEMALE, + 'm': genders.MALE, + genders.MALE: genders.MALE, + 'o': genders.OTHER, + genders.OTHER: genders.OTHER, +} + + +def _row_value(row, name): + """Return a row value from dictionary rows produced by the Snowflake client.""" + if name in row: + return row[name] + return row[name.upper()] + + +def _copy_fields(row, field_names): + """Return an API-shaped dictionary with selected fields from a Snowflake row.""" + return { + field_name: _row_value(row, field_name) + for field_name in field_names + } + + +def _gender_field(gender): + """Return the existing API gender field for a Snowflake gender value.""" + if gender is None: + return genders.UNKNOWN + return GENDER_FIELD_MAP.get(gender.lower(), genders.UNKNOWN) + + +def map_course_enrollment_daily_rows(rows): + """Map course enrollment count rows into the existing API shape.""" + return [ + _copy_fields(row, ['course_id', 'date', 'count', 'created']) + for row in rows or [] + ] + + +def map_course_enrollment_education_rows(rows): + """Map enrollment education rows into the existing API shape.""" + return [ + _copy_fields(row, ['course_id', 'date', 'education_level', 'count', 'created']) + for row in rows or [] + ] + + +def map_course_enrollment_mode_rows(rows): + """Pivot enrollment mode rows into the existing API shape.""" + rows = sorted(rows or [], key=lambda row: (_row_value(row, 'course_id'), _row_value(row, 'date'))) + formatted_data = [] + + for key, group in groupby(rows, lambda row: (_row_value(row, 'course_id'), _row_value(row, 'date'))): + item = { + 'course_id': key[0], + 'date': key[1], + 'created': None, + } + total = 0 + cumulative_total = 0 + + for row in group: + mode = _row_value(row, 'mode') + count = int(_row_value(row, 'count')) + cumulative_count = int(_row_value(row, 'cumulative_count')) + created = _row_value(row, 'created') + item[mode] = item.get(mode, 0) + count + item['created'] = max(created, item['created']) if item['created'] else created + total += count + cumulative_total += cumulative_count + + item[enrollment_modes.PROFESSIONAL] = item.get(enrollment_modes.PROFESSIONAL, 0) + item.pop( + enrollment_modes.PROFESSIONAL_NO_ID, + 0, + ) + item['count'] = total + item['cumulative_count'] = cumulative_total + formatted_data.append(item) + + return formatted_data + + +def map_course_enrollment_gender_rows(rows): + """Pivot enrollment gender rows into the existing API shape.""" + rows = sorted(rows or [], key=lambda row: (_row_value(row, 'course_id'), _row_value(row, 'date'))) + formatted_data = [] + + for key, group in groupby(rows, lambda row: (_row_value(row, 'course_id'), _row_value(row, 'date'))): + item = { + 'course_id': key[0], + 'date': key[1], + 'created': None, + genders.MALE: 0, + genders.FEMALE: 0, + genders.OTHER: 0, + genders.UNKNOWN: 0, + } + + for row in group: + gender = _gender_field(_row_value(row, 'gender')) + created = _row_value(row, 'created') + item[gender] += int(_row_value(row, 'count')) + item['created'] = max(created, item['created']) if item['created'] else created + + formatted_data.append(item) + + return formatted_data + + +def map_course_enrollment_location_rows(rows): + """Map enrollment location rows into model instances used by the current serializer.""" + items = [ + models.CourseEnrollmentByCountry( + course_id=_row_value(row, 'course_id'), + date=_row_value(row, 'date'), + country_code=_row_value(row, 'country_code'), + count=int(_row_value(row, 'count')), + created=_row_value(row, 'created'), + ) + for row in rows or [] + ] + items = sorted(items, key=lambda item: '' if item.country.alpha2 is None else item.country.alpha2) + returned_items = [] + + for key, group in groupby(items, lambda item: (item.date, item.country.alpha2, item.course_id)): + count = 0 + created = None + + for item in group: + created = max(created, item.created) if created else item.created + count += item.count + + returned_items.append(models.CourseEnrollmentByCountry( + course_id=key[2], + date=key[0], + country_code=key[1], + count=count, + created=created, + )) + + return returned_items diff --git a/analytics_data_api/insights_snowflake/queries/enrollment.py b/analytics_data_api/insights_snowflake/queries/enrollment.py new file mode 100644 index 00000000..2dc21064 --- /dev/null +++ b/analytics_data_api/insights_snowflake/queries/enrollment.py @@ -0,0 +1,117 @@ +"""Snowflake queries for course enrollment metrics.""" + +import datetime + +from analytics_data_api.insights_snowflake.client import fetch_all, get_qualified_table_name + +COURSE_ENROLLMENT_DAILY_TABLE = 'COURSE_ENROLLMENT_DAILY' +COURSE_ENROLLMENT_MODE_DAILY_TABLE = 'COURSE_ENROLLMENT_MODE_DAILY' +COURSE_ENROLLMENT_EDUCATION_LEVEL_CURRENT_TABLE = 'COURSE_ENROLLMENT_EDUCATION_LEVEL_CURRENT' +COURSE_ENROLLMENT_GENDER_DAILY_TABLE = 'COURSE_ENROLLMENT_GENDER_DAILY' +COURSE_ENROLLMENT_LOCATION_CURRENT_TABLE = 'COURSE_ENROLLMENT_LOCATION_CURRENT' + + +def _date_value(value): + """Return a date value for date-filtered Snowflake enrollment queries.""" + if isinstance(value, datetime.datetime): + return value.date() + return value + + +def _get_course_enrollment_rows(table, columns, order_by, course_id, start_date=None, end_date=None): + """Return Snowflake enrollment rows for one controlled table.""" + table_name = get_qualified_table_name(table) + select_columns = ',\n '.join(columns) + params = { + 'course_id': course_id, + } + + if start_date or end_date: + params.update({ + 'start_date': _date_value(start_date), + 'end_date': _date_value(end_date), + }) + sql = """ +SELECT + {select_columns} +FROM {table_name} +WHERE course_id = %(course_id)s + AND (%(start_date)s IS NULL OR "DATE" >= %(start_date)s) + AND (%(end_date)s IS NULL OR "DATE" < %(end_date)s) +ORDER BY {order_by} +""".format(select_columns=select_columns, table_name=table_name, order_by=order_by) + else: + sql = """ +SELECT + {select_columns} +FROM {table_name} +WHERE course_id = %(course_id)s + AND "DATE" = ( + SELECT MAX("DATE") + FROM {table_name} + WHERE course_id = %(course_id)s + ) +ORDER BY {order_by} +""".format(select_columns=select_columns, table_name=table_name, order_by=order_by) + + return fetch_all(sql, params) + + +def get_course_enrollment_daily_rows(course_id, start_date=None, end_date=None): + """Return Snowflake rows for course enrollment counts.""" + return _get_course_enrollment_rows( + COURSE_ENROLLMENT_DAILY_TABLE, + ['course_id', '"DATE" AS date', '"COUNT" AS count', 'created'], + 'course_id, date', + course_id, + start_date=start_date, + end_date=end_date, + ) + + +def get_course_enrollment_mode_rows(course_id, start_date=None, end_date=None): + """Return Snowflake rows for enrollment mode counts.""" + return _get_course_enrollment_rows( + COURSE_ENROLLMENT_MODE_DAILY_TABLE, + ['course_id', '"DATE" AS date', 'mode', '"COUNT" AS count', 'cumulative_count', 'created'], + 'course_id, date, mode', + course_id, + start_date=start_date, + end_date=end_date, + ) + + +def get_course_enrollment_education_rows(course_id, start_date=None, end_date=None): + """Return Snowflake rows for enrollment education counts.""" + return _get_course_enrollment_rows( + COURSE_ENROLLMENT_EDUCATION_LEVEL_CURRENT_TABLE, + ['course_id', '"DATE" AS date', 'education_level', '"COUNT" AS count', 'created'], + 'course_id, date, education_level', + course_id, + start_date=start_date, + end_date=end_date, + ) + + +def get_course_enrollment_gender_rows(course_id, start_date=None, end_date=None): + """Return Snowflake rows for enrollment gender counts.""" + return _get_course_enrollment_rows( + COURSE_ENROLLMENT_GENDER_DAILY_TABLE, + ['course_id', '"DATE" AS date', 'gender', '"COUNT" AS count', 'created'], + 'course_id, date, gender', + course_id, + start_date=start_date, + end_date=end_date, + ) + + +def get_course_enrollment_location_rows(course_id, start_date=None, end_date=None): + """Return Snowflake rows for enrollment location counts.""" + return _get_course_enrollment_rows( + COURSE_ENROLLMENT_LOCATION_CURRENT_TABLE, + ['course_id', '"DATE" AS date', 'country_code', '"COUNT" AS count', 'created'], + 'course_id, date, country_code', + course_id, + start_date=start_date, + end_date=end_date, + ) diff --git a/analytics_data_api/insights_snowflake/service.py b/analytics_data_api/insights_snowflake/service.py index 5f34b038..4c07dd98 100644 --- a/analytics_data_api/insights_snowflake/service.py +++ b/analytics_data_api/insights_snowflake/service.py @@ -1,10 +1,54 @@ """Service functions for Snowflake-backed Insights endpoints.""" from analytics_data_api.insights_snowflake.mappers.activity import map_course_activity_weekly_rows +from analytics_data_api.insights_snowflake.mappers.enrollment import ( + map_course_enrollment_daily_rows, + map_course_enrollment_education_rows, + map_course_enrollment_gender_rows, + map_course_enrollment_location_rows, + map_course_enrollment_mode_rows, +) from analytics_data_api.insights_snowflake.queries.activity import get_course_activity_weekly_rows +from analytics_data_api.insights_snowflake.queries.enrollment import ( + get_course_enrollment_daily_rows, + get_course_enrollment_education_rows, + get_course_enrollment_gender_rows, + get_course_enrollment_location_rows, + get_course_enrollment_mode_rows, +) def get_course_activity_weekly(course_id, start_date=None, end_date=None): """Return course activity in the existing API response shape.""" rows = get_course_activity_weekly_rows(course_id, start_date=start_date, end_date=end_date) return map_course_activity_weekly_rows(rows) + + +def get_course_enrollment(course_id, start_date=None, end_date=None): + """Return course enrollment counts in the existing API response shape.""" + rows = get_course_enrollment_daily_rows(course_id, start_date=start_date, end_date=end_date) + return map_course_enrollment_daily_rows(rows) + + +def get_course_enrollment_mode(course_id, start_date=None, end_date=None): + """Return course enrollment mode counts in the existing API response shape.""" + rows = get_course_enrollment_mode_rows(course_id, start_date=start_date, end_date=end_date) + return map_course_enrollment_mode_rows(rows) + + +def get_course_enrollment_education(course_id, start_date=None, end_date=None): + """Return course enrollment education counts in the existing API response shape.""" + rows = get_course_enrollment_education_rows(course_id, start_date=start_date, end_date=end_date) + return map_course_enrollment_education_rows(rows) + + +def get_course_enrollment_gender(course_id, start_date=None, end_date=None): + """Return course enrollment gender counts in the existing API response shape.""" + rows = get_course_enrollment_gender_rows(course_id, start_date=start_date, end_date=end_date) + return map_course_enrollment_gender_rows(rows) + + +def get_course_enrollment_location(course_id, start_date=None, end_date=None): + """Return course enrollment location counts in the existing API response shape.""" + rows = get_course_enrollment_location_rows(course_id, start_date=start_date, end_date=end_date) + return map_course_enrollment_location_rows(rows) diff --git a/analytics_data_api/tests/test_insights_snowflake.py b/analytics_data_api/tests/test_insights_snowflake.py index dd956d20..dd7b06da 100644 --- a/analytics_data_api/tests/test_insights_snowflake.py +++ b/analytics_data_api/tests/test_insights_snowflake.py @@ -6,15 +6,40 @@ from django.test import SimpleTestCase, override_settings from rest_framework.response import Response +from analytics_data_api.constants import country, enrollment_modes, genders from analytics_data_api.insights_snowflake.client import fetch_all, get_qualified_table_name from analytics_data_api.insights_snowflake.mappers.activity import map_course_activity_weekly_rows +from analytics_data_api.insights_snowflake.mappers.enrollment import ( + map_course_enrollment_daily_rows, + map_course_enrollment_education_rows, + map_course_enrollment_gender_rows, + map_course_enrollment_location_rows, + map_course_enrollment_mode_rows, +) from analytics_data_api.insights_snowflake.queries.activity import get_course_activity_weekly_rows +from analytics_data_api.insights_snowflake.queries.enrollment import ( + COURSE_ENROLLMENT_EDUCATION_LEVEL_CURRENT_TABLE, + COURSE_ENROLLMENT_GENDER_DAILY_TABLE, + COURSE_ENROLLMENT_LOCATION_CURRENT_TABLE, + get_course_enrollment_daily_rows, + get_course_enrollment_education_rows, + get_course_enrollment_gender_rows, + get_course_enrollment_location_rows, + get_course_enrollment_mode_rows, +) from analytics_data_api.insights_snowflake.response_headers import ( DATA_SOURCE_HEADER, DATA_SOURCE_SNOWFLAKE, InsightsDataSourceResponseMixin, ) -from analytics_data_api.insights_snowflake.service import get_course_activity_weekly +from analytics_data_api.insights_snowflake.service import ( + get_course_activity_weekly, + get_course_enrollment, + get_course_enrollment_education, + get_course_enrollment_gender, + get_course_enrollment_location, + get_course_enrollment_mode, +) from analytics_data_api.insights_snowflake.toggles import ( COURSE_ACTIVITY_SNOWFLAKE_FLAG, INSIGHTS_SNOWFLAKE_FLAG, @@ -110,6 +135,72 @@ def test_get_course_activity_weekly_rows_uses_date_range_query_with_dates(self, }) +class InsightsSnowflakeEnrollmentQueryTests(SimpleTestCase): + """Cover enrollment query construction with mocked Snowflake execution.""" + + @patch('analytics_data_api.insights_snowflake.queries.enrollment.fetch_all') + @patch( + 'analytics_data_api.insights_snowflake.queries.enrollment.get_qualified_table_name', + Mock(return_value='PROD.INSIGHTS.COURSE_ENROLLMENT_DAILY') + ) + def test_get_course_enrollment_daily_rows_uses_latest_date_query_without_dates(self, mock_fetch_all): + mock_fetch_all.return_value = [{'course_id': 'course-v1:edX+DemoX+Demo_Course'}] + + rows = get_course_enrollment_daily_rows('course-v1:edX+DemoX+Demo_Course') + + self.assertEqual(rows, [{'course_id': 'course-v1:edX+DemoX+Demo_Course'}]) + sql, params = mock_fetch_all.call_args[0] + self.assertIn('SELECT MAX("DATE")', sql) + self.assertEqual(params, {'course_id': 'course-v1:edX+DemoX+Demo_Course'}) + + @patch('analytics_data_api.insights_snowflake.queries.enrollment.fetch_all') + @patch( + 'analytics_data_api.insights_snowflake.queries.enrollment.get_qualified_table_name', + Mock(return_value='PROD.INSIGHTS.COURSE_ENROLLMENT_MODE_DAILY') + ) + def test_get_course_enrollment_mode_rows_uses_date_range_query_with_dates(self, mock_fetch_all): + start_date = datetime.datetime(2014, 1, 1, tzinfo=datetime.timezone.utc) + end_date = datetime.datetime(2014, 1, 8, tzinfo=datetime.timezone.utc) + + get_course_enrollment_mode_rows( + 'course-v1:edX+DemoX+Demo_Course', + start_date=start_date, + end_date=end_date, + ) + + sql, params = mock_fetch_all.call_args[0] + self.assertIn('"DATE" >= %(start_date)s', sql) + self.assertIn('"DATE" < %(end_date)s', sql) + self.assertEqual(params, { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'start_date': start_date.date(), + 'end_date': end_date.date(), + }) + + @patch('analytics_data_api.insights_snowflake.queries.enrollment.fetch_all') + @patch('analytics_data_api.insights_snowflake.queries.enrollment.get_qualified_table_name') + def test_enrollment_query_functions_use_expected_tables(self, mock_get_table_name, _mock_fetch_all): + mock_get_table_name.return_value = 'PROD.INSIGHTS.ENROLLMENT_TABLE' + course_id = 'course-v1:edX+DemoX+Demo_Course' + + query_functions = [ + (get_course_enrollment_education_rows, COURSE_ENROLLMENT_EDUCATION_LEVEL_CURRENT_TABLE, 'education_level'), + (get_course_enrollment_gender_rows, COURSE_ENROLLMENT_GENDER_DAILY_TABLE, 'gender'), + (get_course_enrollment_location_rows, COURSE_ENROLLMENT_LOCATION_CURRENT_TABLE, 'country_code'), + ] + + for query_function, table, expected_column in query_functions: + mock_get_table_name.reset_mock() + _mock_fetch_all.reset_mock() + + query_function(course_id) + + mock_get_table_name.assert_called_once_with(table) + sql, params = _mock_fetch_all.call_args[0] + self.assertIn(expected_column, sql) + self.assertEqual(params, {'course_id': course_id}) + + class InsightsSnowflakeActivityMapperTests(SimpleTestCase): """Cover Snowflake activity row mapping into the existing API shape.""" @@ -170,9 +261,172 @@ def test_map_course_activity_weekly_rows_accepts_uppercase_snowflake_keys(self): }]) +class InsightsSnowflakeEnrollmentMapperTests(SimpleTestCase): + """Cover Snowflake enrollment row mapping into the existing API shapes.""" + + def test_map_course_enrollment_daily_rows_accepts_uppercase_snowflake_keys(self): + date = datetime.date(2014, 1, 1) + created = datetime.datetime(2014, 1, 2, tzinfo=datetime.timezone.utc) + rows = [{ + 'COURSE_ID': 'course-v1:edX+DemoX+Demo_Course', + 'DATE': date, + 'COUNT': 203, + 'CREATED': created, + }] + + self.assertEqual(map_course_enrollment_daily_rows(rows), [{ + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'count': 203, + 'created': created, + }]) + + def test_map_course_enrollment_education_rows(self): + date = datetime.date(2014, 1, 1) + created = datetime.datetime(2014, 1, 2, tzinfo=datetime.timezone.utc) + rows = [{ + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'education_level': 'bachelors', + 'count': 25, + 'created': created, + }] + + self.assertEqual(map_course_enrollment_education_rows(rows), [{ + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'education_level': 'bachelors', + 'count': 25, + 'created': created, + }]) + + def test_map_course_enrollment_mode_rows_pivots_modes(self): + date = datetime.date(2014, 1, 1) + created = datetime.datetime(2014, 1, 2, tzinfo=datetime.timezone.utc) + later_created = datetime.datetime(2014, 1, 3, tzinfo=datetime.timezone.utc) + rows = [ + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'mode': enrollment_modes.PROFESSIONAL_NO_ID, + 'count': 3, + 'cumulative_count': 7, + 'created': created, + }, + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'mode': enrollment_modes.PROFESSIONAL, + 'count': 4, + 'cumulative_count': 8, + 'created': later_created, + }, + ] + + self.assertEqual(map_course_enrollment_mode_rows(rows), [{ + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'created': later_created, + enrollment_modes.PROFESSIONAL: 7, + 'count': 7, + 'cumulative_count': 15, + }]) + + def test_map_course_enrollment_gender_rows_pivots_genders(self): + date = datetime.date(2014, 1, 1) + created = datetime.datetime(2014, 1, 2, tzinfo=datetime.timezone.utc) + rows = [ + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'gender': 'f', + 'count': 3, + 'created': created, + }, + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'gender': None, + 'count': 4, + 'created': created, + }, + ] + + self.assertEqual(map_course_enrollment_gender_rows(rows), [{ + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'created': created, + genders.MALE: 0, + genders.FEMALE: 3, + genders.OTHER: 0, + genders.UNKNOWN: 4, + }]) + + def test_map_course_enrollment_location_rows_groups_unknown_countries(self): + date = datetime.date(2014, 1, 1) + created = datetime.datetime(2014, 1, 2, tzinfo=datetime.timezone.utc) + rows = [ + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'country_code': '', + 'count': 3, + 'created': created, + }, + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'country_code': None, + 'count': 4, + 'created': created, + }, + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'country_code': 'US', + 'count': 5, + 'created': created, + }, + ] + + mapped_rows = map_course_enrollment_location_rows(rows) + + self.assertEqual(len(mapped_rows), 2) + self.assertEqual(mapped_rows[0].country.name, country.UNKNOWN_COUNTRY_CODE) + self.assertEqual(mapped_rows[0].count, 7) + self.assertEqual(mapped_rows[1].country.alpha2, 'US') + self.assertEqual(mapped_rows[1].count, 5) + + class InsightsSnowflakeServiceTests(SimpleTestCase): """Cover service orchestration without real Snowflake calls.""" + def assertServiceCallsQueryAndMapper(self, service_function, query_path, mapper_path): + raw_rows = [{'course_id': 'course-v1:edX+DemoX+Demo_Course'}] + mapped_rows = [{'course_id': 'course-v1:edX+DemoX+Demo_Course', 'count': 203}] + start_date = datetime.datetime(2014, 1, 1, tzinfo=datetime.timezone.utc) + end_date = datetime.datetime(2014, 1, 8, tzinfo=datetime.timezone.utc) + + with patch(query_path) as mock_get_rows, patch(mapper_path) as mock_map_rows: + mock_get_rows.return_value = raw_rows + mock_map_rows.return_value = mapped_rows + + self.assertEqual( + service_function( + 'course-v1:edX+DemoX+Demo_Course', + start_date=start_date, + end_date=end_date, + ), + mapped_rows, + ) + + mock_get_rows.assert_called_once_with( + 'course-v1:edX+DemoX+Demo_Course', + start_date=start_date, + end_date=end_date, + ) + mock_map_rows.assert_called_once_with(raw_rows) + @patch('analytics_data_api.insights_snowflake.service.map_course_activity_weekly_rows') @patch('analytics_data_api.insights_snowflake.service.get_course_activity_weekly_rows') def test_get_course_activity_weekly_calls_query_and_mapper(self, mock_get_rows, mock_map_rows): @@ -199,6 +453,41 @@ def test_get_course_activity_weekly_calls_query_and_mapper(self, mock_get_rows, ) mock_map_rows.assert_called_once_with(raw_rows) + def test_get_course_enrollment_calls_query_and_mapper(self): + self.assertServiceCallsQueryAndMapper( + get_course_enrollment, + 'analytics_data_api.insights_snowflake.service.get_course_enrollment_daily_rows', + 'analytics_data_api.insights_snowflake.service.map_course_enrollment_daily_rows', + ) + + def test_get_course_enrollment_mode_calls_query_and_mapper(self): + self.assertServiceCallsQueryAndMapper( + get_course_enrollment_mode, + 'analytics_data_api.insights_snowflake.service.get_course_enrollment_mode_rows', + 'analytics_data_api.insights_snowflake.service.map_course_enrollment_mode_rows', + ) + + def test_get_course_enrollment_education_calls_query_and_mapper(self): + self.assertServiceCallsQueryAndMapper( + get_course_enrollment_education, + 'analytics_data_api.insights_snowflake.service.get_course_enrollment_education_rows', + 'analytics_data_api.insights_snowflake.service.map_course_enrollment_education_rows', + ) + + def test_get_course_enrollment_gender_calls_query_and_mapper(self): + self.assertServiceCallsQueryAndMapper( + get_course_enrollment_gender, + 'analytics_data_api.insights_snowflake.service.get_course_enrollment_gender_rows', + 'analytics_data_api.insights_snowflake.service.map_course_enrollment_gender_rows', + ) + + def test_get_course_enrollment_location_calls_query_and_mapper(self): + self.assertServiceCallsQueryAndMapper( + get_course_enrollment_location, + 'analytics_data_api.insights_snowflake.service.get_course_enrollment_location_rows', + 'analytics_data_api.insights_snowflake.service.map_course_enrollment_location_rows', + ) + class BaseTestView: """Small base class for testing response mixin behavior.""" diff --git a/analytics_data_api/v0/tests/views/test_courses.py b/analytics_data_api/v0/tests/views/test_courses.py index 661666c6..fff2f1a2 100644 --- a/analytics_data_api/v0/tests/views/test_courses.py +++ b/analytics_data_api/v0/tests/views/test_courses.py @@ -32,6 +32,7 @@ from analytics_data_api.v0 import models from analytics_data_api.v0.tests.utils import create_engagement from analytics_data_api.v0.tests.views import CourseSamples, VerifyCsvResponseMixin +from analytics_data_api.v0.views import courses as course_views from analyticsdataserver.tests.utils import TestCaseWithAuthentication @@ -177,6 +178,18 @@ def test_get_with_intervals(self, course_id): expected = self.format_as_response(*self.model.objects.filter(date=self.date)) self.assertIntervalFilteringWorks(expected, course_id, self.date, self.date + datetime.timedelta(days=1)) + def assertSnowflakeResponse(self, course_id, path, view_class, snowflake_data, expected): + mock_get_data = Mock(return_value=snowflake_data) + + with patch('analytics_data_api.v0.views.courses.is_insights_snowflake_enabled', return_value=True): + with patch.object(view_class, 'snowflake_service_function', staticmethod(mock_get_data)): + response = self.authenticated_get(f'/api/v1/courses/{course_id}{path}') + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.data, expected) + self.assertEqual(response['X-Insights-Data-Source'], 'snowflake') + mock_get_data.assert_called_once_with(course_id, None, None) + @ddt.ddt @set_databases @@ -332,6 +345,44 @@ def format_as_response(self, *args): 'created': ce.created.strftime(settings.DATETIME_FORMAT)} for ce in args ] + def test_get_uses_snowflake_service_when_global_flag_enabled(self): + course_id = CourseSamples.course_ids[0] + created = datetime.datetime(2014, 1, 2, tzinfo=pytz.utc) + snowflake_data = [{ + 'course_id': course_id, + 'date': self.date, + 'education_level': self.el1, + 'count': 25, + 'created': created, + }] + expected = [{ + 'course_id': course_id, + 'date': self.date.strftime(settings.DATE_FORMAT), + 'education_level': self.el1, + 'count': 25, + 'created': created.strftime(settings.DATETIME_FORMAT), + }] + + self.assertSnowflakeResponse( + course_id, + '/enrollment/education/', + course_views.CourseEnrollmentByEducationView, + snowflake_data, + expected, + ) + + def test_get_returns_404_when_global_flag_enabled_and_no_snowflake_data(self): + course_id = CourseSamples.course_ids[0] + mock_get_data = Mock(return_value=[]) + + with patch('analytics_data_api.v0.views.courses.is_insights_snowflake_enabled', return_value=True): + with patch.object(course_views.CourseEnrollmentView, 'snowflake_service_function', + staticmethod(mock_get_data)): + response = self.authenticated_get(f'/api/v1/courses/{course_id}/enrollment/') + + self.assertEqual(response.status_code, 404) + mock_get_data.assert_called_once_with(course_id, None, None) + @ddt.ddt @set_databases @@ -396,6 +447,36 @@ def test_default_fill(self, course_id): self.assertViewReturnsExpectedData([expected], course_id) + def test_get_uses_snowflake_service_when_global_flag_enabled(self): + course_id = CourseSamples.course_ids[0] + created = datetime.datetime(2014, 1, 2, tzinfo=pytz.utc) + snowflake_data = [{ + 'course_id': course_id, + 'date': self.date, + 'female': 3, + 'male': 4, + 'other': 5, + 'unknown': 6, + 'created': created, + }] + expected = [{ + 'course_id': course_id, + 'date': self.date.strftime(settings.DATE_FORMAT), + 'female': 3, + 'male': 4, + 'other': 5, + 'unknown': 6, + 'created': created.strftime(settings.DATETIME_FORMAT), + }] + + self.assertSnowflakeResponse( + course_id, + '/enrollment/gender/', + course_views.CourseEnrollmentByGenderView, + snowflake_data, + expected, + ) + @set_databases class CourseEnrollmentViewTests(CourseEnrollmentViewTestCaseMixin, TestCaseWithAuthentication): @@ -414,6 +495,46 @@ def format_as_response(self, *args): for ce in args ] + def test_get_uses_aurora_when_global_snowflake_flag_disabled(self): + course_id = CourseSamples.course_ids[0] + self.generate_data(course_id) + expected = self.format_as_response(*self.get_latest_data(course_id)) + mock_get_data = Mock() + + with patch('analytics_data_api.v0.views.courses.is_insights_snowflake_enabled', return_value=False): + with patch.object(course_views.CourseEnrollmentView, 'snowflake_service_function', + staticmethod(mock_get_data)): + response = self.authenticated_get(f'/api/v1/courses/{course_id}/enrollment/') + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.data, expected) + self.assertEqual(response['X-Insights-Data-Source'], 'aurora') + mock_get_data.assert_not_called() + + def test_get_uses_snowflake_service_when_global_flag_enabled(self): + course_id = CourseSamples.course_ids[0] + created = datetime.datetime(2014, 1, 2, tzinfo=pytz.utc) + snowflake_data = [{ + 'course_id': course_id, + 'date': self.date, + 'count': 203, + 'created': created, + }] + expected = [{ + 'course_id': course_id, + 'date': self.date.strftime(settings.DATE_FORMAT), + 'count': 203, + 'created': created.strftime(settings.DATETIME_FORMAT), + }] + + self.assertSnowflakeResponse( + course_id, + '/enrollment/', + course_views.CourseEnrollmentView, + snowflake_data, + expected, + ) + @ddt.ddt @set_databases @@ -476,6 +597,44 @@ def test_default_fill(self, course_id): self.assertViewReturnsExpectedData([expected], course_id) + def test_get_uses_snowflake_service_when_global_flag_enabled(self): + course_id = CourseSamples.course_ids[0] + created = datetime.datetime(2014, 1, 2, tzinfo=pytz.utc) + snowflake_data = [{ + 'course_id': course_id, + 'date': self.date, + 'count': 25, + 'cumulative_count': 100, + 'created': created, + 'audit': 5, + 'credit': 4, + 'honor': 3, + 'professional': 2, + 'verified': 1, + 'masters': 10, + }] + expected = [{ + 'course_id': course_id, + 'date': self.date.strftime(settings.DATE_FORMAT), + 'count': 25, + 'cumulative_count': 100, + 'created': created.strftime(settings.DATETIME_FORMAT), + 'audit': 5, + 'credit': 4, + 'honor': 3, + 'professional': 2, + 'verified': 1, + 'masters': 10, + }] + + self.assertSnowflakeResponse( + course_id, + '/enrollment/mode/', + course_views.CourseEnrollmentModeView, + snowflake_data, + expected, + ) + @set_databases class CourseEnrollmentByLocationViewTests(CourseEnrollmentViewTestCaseMixin, TestCaseWithAuthentication): @@ -523,6 +682,36 @@ def setUpClass(cls): super().setUpClass() cls.country = get_country('US') + def test_get_uses_snowflake_service_when_global_flag_enabled(self): + course_id = CourseSamples.course_ids[0] + created = datetime.datetime(2014, 1, 2, tzinfo=pytz.utc) + snowflake_data = [models.CourseEnrollmentByCountry( + course_id=course_id, + date=self.date, + country_code='US', + count=455, + created=created, + )] + expected = [{ + 'course_id': course_id, + 'date': self.date.strftime(settings.DATE_FORMAT), + 'country': { + 'alpha2': self.country.alpha2, + 'alpha3': self.country.alpha3, + 'name': self.country.name, + }, + 'count': 455, + 'created': created.strftime(settings.DATETIME_FORMAT), + }] + + self.assertSnowflakeResponse( + course_id, + '/enrollment/location/', + course_views.CourseEnrollmentByLocationView, + snowflake_data, + expected, + ) + @ddt.ddt @set_databases diff --git a/analytics_data_api/v0/views/courses.py b/analytics_data_api/v0/views/courses.py index ba1e254e..cab15631 100644 --- a/analytics_data_api/v0/views/courses.py +++ b/analytics_data_api/v0/views/courses.py @@ -15,8 +15,18 @@ from analytics_data_api.constants import enrollment_modes from analytics_data_api.insights_snowflake.response_headers import InsightsDataSourceResponseMixin -from analytics_data_api.insights_snowflake.service import get_course_activity_weekly -from analytics_data_api.insights_snowflake.toggles import is_course_activity_snowflake_enabled +from analytics_data_api.insights_snowflake.service import ( + get_course_activity_weekly, + get_course_enrollment, + get_course_enrollment_education, + get_course_enrollment_gender, + get_course_enrollment_location, + get_course_enrollment_mode, +) +from analytics_data_api.insights_snowflake.toggles import ( + is_course_activity_snowflake_enabled, + is_insights_snowflake_enabled, +) from analytics_data_api.utils import dictfetchall, get_course_report_download_details from analytics_data_api.v0 import models, serializers from analytics_data_api.v0.exceptions import ReportFileNotFoundError @@ -289,6 +299,34 @@ def apply_date_filtering(self, queryset): return queryset +class SnowflakeCourseEnrollmentMixin(InsightsDataSourceResponseMixin): + """Route migrated enrollment endpoints through Snowflake when the global flag is enabled.""" + + snowflake_service_function = None + + def get_snowflake_queryset(self): + """Return the Snowflake-backed queryset replacement for this endpoint.""" + if self.snowflake_service_function is None: + raise NotImplementedError + + data = self.snowflake_service_function(self.course_id, self.start_date, self.end_date) + if data: + return data + raise Http404 + + def get_aurora_queryset(self): + """Return the existing Aurora queryset for this endpoint.""" + return super().get_queryset() + + def get_queryset(self): + if is_insights_snowflake_enabled(self.request): + self.set_insights_data_source_snowflake() + return self.get_snowflake_queryset() + + self.set_insights_data_source_aurora() + return self.get_aurora_queryset() + + class CourseEnrollmentByBirthYearView(BaseCourseEnrollmentView): """ Get the number of enrolled users by birth year. @@ -329,7 +367,7 @@ class CourseEnrollmentByBirthYearView(BaseCourseEnrollmentView): model = models.CourseEnrollmentByBirthYear -class CourseEnrollmentByEducationView(BaseCourseEnrollmentView): +class CourseEnrollmentByEducationView(SnowflakeCourseEnrollmentMixin, BaseCourseEnrollmentView): """ Get the number of enrolled users by education level. @@ -368,9 +406,10 @@ class CourseEnrollmentByEducationView(BaseCourseEnrollmentView): slug = 'enrollment-education' serializer_class = serializers.CourseEnrollmentByEducationSerializer model = models.CourseEnrollmentByEducation + snowflake_service_function = staticmethod(get_course_enrollment_education) -class CourseEnrollmentByGenderView(BaseCourseEnrollmentView): +class CourseEnrollmentByGenderView(SnowflakeCourseEnrollmentMixin, BaseCourseEnrollmentView): """ Get the number of enrolled users by gender. @@ -408,9 +447,10 @@ class CourseEnrollmentByGenderView(BaseCourseEnrollmentView): slug = 'enrollment-gender' serializer_class = serializers.CourseEnrollmentByGenderSerializer model = models.CourseEnrollmentByGender + snowflake_service_function = staticmethod(get_course_enrollment_gender) - def get_queryset(self): - queryset = super().get_queryset() + def get_aurora_queryset(self): + queryset = super().get_aurora_queryset() formatted_data = [] items = queryset.all() @@ -439,7 +479,7 @@ def get_queryset(self): return formatted_data -class CourseEnrollmentView(BaseCourseEnrollmentView): +class CourseEnrollmentView(SnowflakeCourseEnrollmentMixin, BaseCourseEnrollmentView): """ Get the number of enrolled users. @@ -474,9 +514,10 @@ class CourseEnrollmentView(BaseCourseEnrollmentView): slug = 'enrollment' serializer_class = serializers.CourseEnrollmentDailySerializer model = models.CourseEnrollmentDaily + snowflake_service_function = staticmethod(get_course_enrollment) -class CourseEnrollmentModeView(BaseCourseEnrollmentView): +class CourseEnrollmentModeView(SnowflakeCourseEnrollmentMixin, BaseCourseEnrollmentView): """ Get the number of enrolled users by enrollment mode. @@ -516,9 +557,10 @@ class CourseEnrollmentModeView(BaseCourseEnrollmentView): slug = 'enrollment_mode' serializer_class = serializers.CourseEnrollmentModeDailySerializer model = models.CourseEnrollmentModeDaily + snowflake_service_function = staticmethod(get_course_enrollment_mode) - def get_queryset(self): - queryset = super().get_queryset() + def get_aurora_queryset(self): + queryset = super().get_aurora_queryset() formatted_data = [] items = queryset.all() @@ -553,7 +595,7 @@ def get_queryset(self): # pylint: disable=line-too-long -class CourseEnrollmentByLocationView(BaseCourseEnrollmentView): +class CourseEnrollmentByLocationView(SnowflakeCourseEnrollmentMixin, BaseCourseEnrollmentView): """ Get the number of enrolled users by location. @@ -599,10 +641,11 @@ class CourseEnrollmentByLocationView(BaseCourseEnrollmentView): slug = 'enrollment-location' serializer_class = serializers.CourseEnrollmentByCountrySerializer model = models.CourseEnrollmentByCountry + snowflake_service_function = staticmethod(get_course_enrollment_location) - def get_queryset(self): + def get_aurora_queryset(self): # Get all of the data from the database - queryset = super().get_queryset() + queryset = super().get_aurora_queryset() items = queryset.all() # Data must be sorted in order for groupby to work properly From b32322389e605a9ec85eeef954e168b0a0ae002a Mon Sep 17 00:00:00 2001 From: Santhosh Kumar Date: Wed, 2 Sep 2026 10:33:51 +0000 Subject: [PATCH 2/6] test: reset analytics DB state in enrollment tests --- analytics_data_api/v0/tests/views/test_courses.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/analytics_data_api/v0/tests/views/test_courses.py b/analytics_data_api/v0/tests/views/test_courses.py index fff2f1a2..43a8222e 100644 --- a/analytics_data_api/v0/tests/views/test_courses.py +++ b/analytics_data_api/v0/tests/views/test_courses.py @@ -169,6 +169,11 @@ def setUpClass(cls): super().setUpClass() cls.date = datetime.date(2014, 1, 1) + def tearDown(self): + if hasattr(thread_data, 'analyticsapi_database'): + del thread_data.analyticsapi_database + super().tearDown() + def get_latest_data(self, course_id): return self.model.objects.filter(course_id=course_id, date=self.date).order_by('date', *self.order_by) @@ -407,6 +412,7 @@ def generate_data(self, course_id): def tearDown(self): self.destroy_data() + super().tearDown() def serialize_enrollment(self, enrollment): return { From 7fe3e6f40227fde426d160fff8bf3af1335594ff Mon Sep 17 00:00:00 2001 From: Santhosh Kumar Date: Wed, 2 Sep 2026 10:41:20 +0000 Subject: [PATCH 3/6] test: create enrollment fallback data in v1 analytics DB --- analytics_data_api/v0/tests/views/test_courses.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/analytics_data_api/v0/tests/views/test_courses.py b/analytics_data_api/v0/tests/views/test_courses.py index 43a8222e..f813f7ca 100644 --- a/analytics_data_api/v0/tests/views/test_courses.py +++ b/analytics_data_api/v0/tests/views/test_courses.py @@ -503,8 +503,17 @@ def format_as_response(self, *args): def test_get_uses_aurora_when_global_snowflake_flag_disabled(self): course_id = CourseSamples.course_ids[0] - self.generate_data(course_id) - expected = self.format_as_response(*self.get_latest_data(course_id)) + latest_enrollment = self.model.objects.using(settings.ANALYTICS_DATABASE_V1).create( + course_id=course_id, + date=self.date, + count=203, + ) + self.model.objects.using(settings.ANALYTICS_DATABASE_V1).create( + course_id=course_id, + date=self.date - datetime.timedelta(days=5), + count=203, + ) + expected = self.format_as_response(latest_enrollment) mock_get_data = Mock() with patch('analytics_data_api.v0.views.courses.is_insights_snowflake_enabled', return_value=False): From c69f4eecf5f4022e2e019d9386a867e1d3f44c28 Mon Sep 17 00:00:00 2001 From: Santhosh Kumar Date: Wed, 2 Sep 2026 11:00:27 +0000 Subject: [PATCH 4/6] test: cover enrollment date object query params --- .../tests/test_insights_snowflake.py | 22 +++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/analytics_data_api/tests/test_insights_snowflake.py b/analytics_data_api/tests/test_insights_snowflake.py index dd7b06da..e36334f9 100644 --- a/analytics_data_api/tests/test_insights_snowflake.py +++ b/analytics_data_api/tests/test_insights_snowflake.py @@ -177,6 +177,28 @@ def test_get_course_enrollment_mode_rows_uses_date_range_query_with_dates(self, 'end_date': end_date.date(), }) + @patch('analytics_data_api.insights_snowflake.queries.enrollment.fetch_all') + @patch( + 'analytics_data_api.insights_snowflake.queries.enrollment.get_qualified_table_name', + Mock(return_value='PROD.INSIGHTS.COURSE_ENROLLMENT_DAILY') + ) + def test_get_course_enrollment_daily_rows_accepts_date_objects(self, mock_fetch_all): + start_date = datetime.date(2014, 1, 1) + end_date = datetime.date(2014, 1, 8) + + get_course_enrollment_daily_rows( + 'course-v1:edX+DemoX+Demo_Course', + start_date=start_date, + end_date=end_date, + ) + + _sql, params = mock_fetch_all.call_args[0] + self.assertEqual(params, { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'start_date': start_date, + 'end_date': end_date, + }) + @patch('analytics_data_api.insights_snowflake.queries.enrollment.fetch_all') @patch('analytics_data_api.insights_snowflake.queries.enrollment.get_qualified_table_name') def test_enrollment_query_functions_use_expected_tables(self, mock_get_table_name, _mock_fetch_all): From 4fbd945ff02d49299967bd1b7bc3046ba079caaa Mon Sep 17 00:00:00 2001 From: Santhosh Kumar Date: Wed, 2 Sep 2026 12:23:54 +0000 Subject: [PATCH 5/6] fix: sort enrollment location rows by grouping key --- .../insights_snowflake/mappers/enrollment.py | 2 +- .../tests/test_insights_snowflake.py | 36 +++++++++++++++++++ 2 files changed, 37 insertions(+), 1 deletion(-) diff --git a/analytics_data_api/insights_snowflake/mappers/enrollment.py b/analytics_data_api/insights_snowflake/mappers/enrollment.py index c9e9e062..e8ce4bb2 100644 --- a/analytics_data_api/insights_snowflake/mappers/enrollment.py +++ b/analytics_data_api/insights_snowflake/mappers/enrollment.py @@ -127,7 +127,7 @@ def map_course_enrollment_location_rows(rows): ) for row in rows or [] ] - items = sorted(items, key=lambda item: '' if item.country.alpha2 is None else item.country.alpha2) + items = sorted(items, key=lambda item: (item.date, item.country.alpha2 or '', item.course_id)) returned_items = [] for key, group in groupby(items, lambda item: (item.date, item.country.alpha2, item.course_id)): diff --git a/analytics_data_api/tests/test_insights_snowflake.py b/analytics_data_api/tests/test_insights_snowflake.py index e36334f9..21d36ee2 100644 --- a/analytics_data_api/tests/test_insights_snowflake.py +++ b/analytics_data_api/tests/test_insights_snowflake.py @@ -419,6 +419,42 @@ def test_map_course_enrollment_location_rows_groups_unknown_countries(self): self.assertEqual(mapped_rows[1].country.alpha2, 'US') self.assertEqual(mapped_rows[1].count, 5) + def test_map_course_enrollment_location_rows_groups_unsorted_dates(self): + date = datetime.date(2014, 1, 1) + next_date = datetime.date(2014, 1, 2) + created = datetime.datetime(2014, 1, 3, tzinfo=datetime.timezone.utc) + rows = [ + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'country_code': 'US', + 'count': 3, + 'created': created, + }, + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': next_date, + 'country_code': 'US', + 'count': 4, + 'created': created, + }, + { + 'course_id': 'course-v1:edX+DemoX+Demo_Course', + 'date': date, + 'country_code': 'US', + 'count': 5, + 'created': created, + }, + ] + + mapped_rows = map_course_enrollment_location_rows(rows) + + self.assertEqual(len(mapped_rows), 2) + self.assertEqual(mapped_rows[0].date, date) + self.assertEqual(mapped_rows[0].count, 8) + self.assertEqual(mapped_rows[1].date, next_date) + self.assertEqual(mapped_rows[1].count, 4) + class InsightsSnowflakeServiceTests(SimpleTestCase): """Cover service orchestration without real Snowflake calls.""" From 861cc0cbe108f4216067b6be7241ec972c1c52bb Mon Sep 17 00:00:00 2001 From: Santhosh Kumar Date: Wed, 2 Sep 2026 12:34:34 +0000 Subject: [PATCH 6/6] test: scope enrollment interval expectation to course --- analytics_data_api/v0/tests/views/test_courses.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/analytics_data_api/v0/tests/views/test_courses.py b/analytics_data_api/v0/tests/views/test_courses.py index f813f7ca..c2aeefdc 100644 --- a/analytics_data_api/v0/tests/views/test_courses.py +++ b/analytics_data_api/v0/tests/views/test_courses.py @@ -180,7 +180,7 @@ def get_latest_data(self, course_id): @ddt.data(*CourseSamples.course_ids) def test_get_with_intervals(self, course_id): self.generate_data(course_id) - expected = self.format_as_response(*self.model.objects.filter(date=self.date)) + expected = self.format_as_response(*self.model.objects.filter(course_id=course_id, date=self.date)) self.assertIntervalFilteringWorks(expected, course_id, self.date, self.date + datetime.timedelta(days=1)) def assertSnowflakeResponse(self, course_id, path, view_class, snowflake_data, expected):