diff --git a/compliance/__init__.py b/compliance/__init__.py new file mode 100644 index 0000000000..e73626d479 --- /dev/null +++ b/compliance/__init__.py @@ -0,0 +1 @@ +"""Compliance utilities for MITx Online.""" diff --git a/compliance/api.py b/compliance/api.py new file mode 100644 index 0000000000..e98c4a440c --- /dev/null +++ b/compliance/api.py @@ -0,0 +1,256 @@ +"""CyberSource export compliance helpers.""" + +from __future__ import annotations + +import json +import logging +from dataclasses import dataclass +from typing import Any +from uuid import uuid4 + +from CyberSource.api.verification_api import VerificationApi +from CyberSource.models.riskv1exportcomplianceinquiries_order_information import ( + Riskv1exportcomplianceinquiriesOrderInformation, +) +from CyberSource.models.riskv1exportcomplianceinquiries_order_information_bill_to import ( + Riskv1exportcomplianceinquiriesOrderInformationBillTo, +) +from CyberSource.models.riskv1liststypeentries_client_reference_information import ( + Riskv1liststypeentriesClientReferenceInformation, +) +from CyberSource.models.validate_export_compliance_request import ( + ValidateExportComplianceRequest, +) +from django.conf import settings +from django.core.exceptions import ImproperlyConfigured, ObjectDoesNotExist + +from compliance.exceptions import ExportComplianceDataError + +log = logging.getLogger(__name__) + +ISO_3166_2_PART_COUNT = 2 + + +@dataclass(frozen=True) +class ExportComplianceResult: + """Normalized export compliance response.""" + + decision: str | None + reason_code: str | int | None + request_id: str | None + raw: Any + + @property + def accepted(self) -> bool: + """Return True when CyberSource accepted the export check.""" + return self.decision in {"ACCEPT", "COMPLETED"} + + +def _require_setting(name: str) -> str: + """Return a non-empty setting value or raise an error.""" + value = getattr(settings, name, None) + if not value: + message = f"{name} must be configured for export checks" + raise ImproperlyConfigured(message) + return value + + +def _get_cybersource_configuration() -> dict[str, str | int]: + """Return REST client configuration for CyberSource export checks.""" + return { + "authentication_type": "HTTP_SIGNATURE", + "merchantid": _require_setting("MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_ID"), + "merchant_keyid": _require_setting( + "MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_SECRET_KEY_ID" + ), + "merchant_secretkey": _require_setting( + "MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_SECRET" + ), + "run_environment": _require_setting( + "MITOL_PAYMENT_GATEWAY_CYBERSOURCE_REST_API_ENVIRONMENT" + ), + "timeout": 1000, + } + + +def _split_user_name(user) -> tuple[str, str]: + """Split a user's display name into first/last values.""" + full_name = (user.name or "").strip() + if not full_name: + return ("", "") + + name_parts = full_name.split(maxsplit=1) + if len(name_parts) == 1: + return (name_parts[0], "") + + return (name_parts[0], name_parts[1]) + + +def _normalize_administrative_area( + country: str | None, state: str | None +) -> str | None: + """Normalize ISO-3166-2 style subdivision values for CyberSource bill-to data.""" + if not state: + return None + + normalized_country = (country or "").strip().upper() + normalized_state = state.strip() + subdivision_parts = normalized_state.split("-", maxsplit=1) + + if ( + normalized_country + and len(subdivision_parts) == ISO_3166_2_PART_COUNT + and subdivision_parts[0].upper() == normalized_country + and subdivision_parts[1] + ): + return subdivision_parts[1] + + return normalized_state + + +def _validate_bill_to_fields(user, bill_to: dict[str, str]) -> None: + """Raise a clear error when required CyberSource bill-to fields are missing.""" + missing_fields = [] + + if not bill_to.get("first_name"): + missing_fields.append("first_name") + if not bill_to.get("last_name"): + missing_fields.append("last_name") + + required_fields = ["address1", "locality", "country", "email"] + if bill_to.get("country") in {"US", "CA"}: + required_fields.extend(["administrative_area", "postal_code"]) + + missing_fields.extend(field for field in required_fields if not bill_to.get(field)) + + if missing_fields: + raise ExportComplianceDataError(user, missing_fields) + + +def _build_export_payload(user) -> Any: + """Build the CyberSource export compliance REST request payload.""" + try: + legal_address = user.legal_address + except ObjectDoesNotExist: + legal_address = None + first_name, last_name = _split_user_name(user) + + bill_to = { + "first_name": first_name, + "last_name": last_name, + "email": user.email, + } + + if legal_address and legal_address.country: + bill_to["country"] = legal_address.country + if legal_address and legal_address.street_address_1: + bill_to["address1"] = legal_address.street_address_1 + if legal_address and legal_address.street_address_2: + bill_to["address2"] = legal_address.street_address_2 + if legal_address and legal_address.city: + bill_to["locality"] = legal_address.city + if legal_address and legal_address.state: + bill_to["administrative_area"] = _normalize_administrative_area( + legal_address.country, + legal_address.state, + ) + if legal_address and legal_address.postal_code: + bill_to["postal_code"] = legal_address.postal_code + + _validate_bill_to_fields(user, bill_to) + + return ValidateExportComplianceRequest( + client_reference_information=Riskv1liststypeentriesClientReferenceInformation( + code=str(uuid4()) + ), + order_information=Riskv1exportcomplianceinquiriesOrderInformation( + bill_to=Riskv1exportcomplianceinquiriesOrderInformationBillTo( + **{ + key: value + for key, value in bill_to.items() + if value not in [None, ""] + } + ) + ), + ) + + +def _remove_none_values(value: Any) -> Any: + """Recursively remove None values from SDK payload data.""" + if isinstance(value, dict): + return { + key: _remove_none_values(item) + for key, item in value.items() + if item is not None + } + if isinstance(value, list): + return [_remove_none_values(item) for item in value if item is not None] + return value + + +def _serialize_export_payload(payload: Any) -> str: + """Serialize a CyberSource payload to the JSON string expected by this SDK build.""" + return json.dumps(_remove_none_values(payload.to_dict())) + + +def _get_response_payload(response: Any) -> Any: + """Return the response body object from SDK return values.""" + if isinstance(response, tuple) and response: + return response[0] + return response + + +def _get_response_value(response: Any, *names: str) -> Any: + """Read a value from an SDK response object or dict using any provided name.""" + payload = _get_response_payload(response) + + if isinstance(payload, dict): + for name in names: + if name in payload: + return payload[name] + return None + + for name in names: + value = getattr(payload, name, None) + if value is not None: + return value + + return None + + +def get_cybersource_client(): + """Create an authenticated REST client for CyberSource export checks.""" + return VerificationApi(_get_cybersource_configuration()) + + +def _get_reason_code(response) -> str | None: + """Extract the most useful reason code from a REST response.""" + export_info = _get_response_value( + response, + "export_compliance_information", + "exportComplianceInformation", + ) + info_codes = _get_response_value(export_info, "info_codes", "infoCodes") or [] + if info_codes: + return ",".join(info_codes) + + error_info = _get_response_value(response, "error_information", "errorInformation") + return _get_response_value(error_info, "reason") or _get_response_value( + response, "message" + ) + + +def verify_user_with_exports(user) -> ExportComplianceResult: + """Verify a user against CyberSource export compliance services.""" + client = get_cybersource_client() + payload = _serialize_export_payload(_build_export_payload(user)) + + log.info("Running CyberSource export compliance check for user=%s", user.id) + response = client.validate_export_compliance(payload) + + return ExportComplianceResult( + decision=_get_response_value(response, "status"), + reason_code=_get_reason_code(response), + request_id=_get_response_value(response, "id"), + raw=response, + ) diff --git a/compliance/api_test.py b/compliance/api_test.py new file mode 100644 index 0000000000..7907e18568 --- /dev/null +++ b/compliance/api_test.py @@ -0,0 +1,201 @@ +"""Tests for compliance API helpers.""" + +import json +import uuid +from types import SimpleNamespace + +import pytest +from django.core.exceptions import ImproperlyConfigured + +from compliance.api import ( + ExportComplianceResult, + _build_export_payload, + _normalize_administrative_area, + get_cybersource_client, + verify_user_with_exports, +) +from compliance.exceptions import ExportComplianceDataError +from users.factories import UserFactory +from users.models import User + +pytestmark = [pytest.mark.django_db] + + +@pytest.fixture +def export_settings(settings): + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_ID = "merchant-id" + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_SECRET_KEY_ID = uuid.uuid4().hex + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_SECRET = uuid.uuid4().hex + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_REST_API_ENVIRONMENT = ( + "apitest.cybersource.com" + ) + return settings + + +def test_build_export_payload_uses_user_and_legal_address(export_settings): + """Payload should include user identifying fields and address values.""" + user = UserFactory.create(name="Ada Lovelace", email="ada@example.com") + user.legal_address.country = "US" + user.legal_address.street_address_1 = "77 Massachusetts Ave" + user.legal_address.street_address_2 = "Building 1" + user.legal_address.city = "Cambridge" + user.legal_address.state = "US-MA" + user.legal_address.postal_code = "02139" + user.legal_address.save() + + payload = _build_export_payload(user) + + assert payload.client_reference_information.code + assert payload.order_information.bill_to.first_name == "Ada" + assert payload.order_information.bill_to.last_name == "Lovelace" + assert payload.order_information.bill_to.email == "ada@example.com" + assert payload.order_information.bill_to.address1 == "77 Massachusetts Ave" + assert payload.order_information.bill_to.address2 == "Building 1" + assert payload.order_information.bill_to.locality == "Cambridge" + assert payload.order_information.bill_to.country == "US" + assert payload.order_information.bill_to.administrative_area == "MA" + assert payload.order_information.bill_to.postal_code == "02139" + + +def test_build_export_payload_requires_cybersource_bill_to_fields(export_settings): + """Payload creation should fail fast when required CyberSource address fields are missing.""" + user = UserFactory.create(name="Ada Lovelace", email="ada@example.com") + user.legal_address.country = "US" + user.legal_address.street_address_1 = "" + user.legal_address.city = "" + user.legal_address.state = "US-MA" + user.legal_address.postal_code = "" + user.legal_address.save() + + with pytest.raises(ExportComplianceDataError) as exc_info: + _build_export_payload(user) + + assert exc_info.value.missing_fields == ["address1", "locality", "postal_code"] + assert exc_info.value.to_error_detail() == { + "detail": str(exc_info.value), + "missing_fields": ["address1", "locality", "postal_code"], + } + + +def test_build_export_payload_requires_legal_address(export_settings): + """Payload creation should surface a clear, actionable error when a user has no legal address at all.""" + user = UserFactory.create(name="Ada Lovelace", email="ada@example.com") + user.legal_address.delete() + user = User.objects.get(pk=user.pk) + + with pytest.raises(ExportComplianceDataError) as exc_info: + _build_export_payload(user) + + assert "contact support" in str(exc_info.value) + assert exc_info.value.missing_fields == ["address1", "country", "locality"] + + +def test_normalize_administrative_area_strips_country_prefix(): + """ISO-3166-2 values should be reduced to the region code for CyberSource.""" + assert _normalize_administrative_area("US", "US-MA") == "MA" + assert _normalize_administrative_area("ca", "CA-ON") == "ON" + + +def test_normalize_administrative_area_preserves_non_prefixed_values(): + """Plain state values should pass through unchanged.""" + assert _normalize_administrative_area("US", "MA") == "MA" + assert _normalize_administrative_area("US", "Massachusetts") == "Massachusetts" + + +def test_get_cybersource_client_requires_configuration(settings): + """Client creation should fail if required settings are missing.""" + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_ID = "" + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_SECRET_KEY_ID = "" + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_SECRET = "" + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_REST_API_ENVIRONMENT = "" + with pytest.raises( + ImproperlyConfigured, + match="MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_ID must be configured", + ): + get_cybersource_client() + + +def test_get_cybersource_client_uses_official_rest_sdk_configuration(export_settings): + """Client creation should use the official CyberSource REST SDK config keys.""" + client = get_cybersource_client() + + assert client.api_client.mconfig.authentication_type == "HTTP_SIGNATURE" + assert client.api_client.mconfig.merchant_id == "merchant-id" + assert client.api_client.mconfig.run_environment == "apitest.cybersource.com" + + +def test_verify_user_with_exports_calls_validate_export_compliance( + mocker, export_settings +): + """Verification should call CyberSource and normalize the response.""" + user = UserFactory.create(name="Ada Lovelace", email="ada@example.com") + user.legal_address.country = "US" + user.legal_address.street_address_1 = "77 Massachusetts Ave" + user.legal_address.street_address_2 = "Building 1" + user.legal_address.city = "Cambridge" + user.legal_address.state = "US-MA" + user.legal_address.postal_code = "02139" + user.legal_address.save() + + response = SimpleNamespace( + status="COMPLETED", + id="abc123", + export_compliance_information=SimpleNamespace(info_codes=["MATCH-BCO"]), + error_information=None, + message=None, + ) + mock_client = mocker.Mock() + mock_client.validate_export_compliance.return_value = response + mocker.patch("compliance.api.get_cybersource_client", return_value=mock_client) + + result = verify_user_with_exports(user) + + assert isinstance(result, ExportComplianceResult) + assert result.accepted is True + assert result.decision == "COMPLETED" + assert result.reason_code == "MATCH-BCO" + assert result.request_id == "abc123" + mock_client.validate_export_compliance.assert_called_once() + payload = json.loads(mock_client.validate_export_compliance.call_args.args[0]) + assert payload["order_information"]["bill_to"]["email"] == "ada@example.com" + assert payload["order_information"]["bill_to"]["address1"] == "77 Massachusetts Ave" + assert payload["order_information"]["bill_to"]["address2"] == "Building 1" + assert payload["order_information"]["bill_to"]["locality"] == "Cambridge" + assert payload["order_information"]["bill_to"]["country"] == "US" + assert payload["order_information"]["bill_to"]["postal_code"] == "02139" + assert payload["client_reference_information"].get("partner") is None + + +def test_verify_user_with_exports_normalizes_tuple_response(mocker, export_settings): + """Verification should handle SDK responses returned as (body, status, raw_json).""" + user = UserFactory.create(name="Ada Lovelace", email="ada@example.com") + user.legal_address.country = "US" + user.legal_address.street_address_1 = "77 Massachusetts Ave" + user.legal_address.city = "Cambridge" + user.legal_address.state = "US-MA" + user.legal_address.postal_code = "02139" + user.legal_address.save() + + response = ( + { + "status": "COMPLETED", + "id": "abc123", + "export_compliance_information": {"info_codes": ["MATCH-BCO"]}, + "error_information": None, + "message": None, + }, + 201, + '{"status":"COMPLETED","id":"abc123"}', + ) + mock_client = mocker.Mock() + mock_client.validate_export_compliance.return_value = response + mocker.patch("compliance.api.get_cybersource_client", return_value=mock_client) + + result = verify_user_with_exports(user) + + assert isinstance(result, ExportComplianceResult) + assert result.accepted is True + assert result.decision == "COMPLETED" + assert result.reason_code == "MATCH-BCO" + assert result.request_id == "abc123" + assert result.raw == response diff --git a/compliance/exceptions.py b/compliance/exceptions.py new file mode 100644 index 0000000000..44e84faca4 --- /dev/null +++ b/compliance/exceptions.py @@ -0,0 +1,69 @@ +"""Exceptions for the compliance app""" + + +class ExportComplianceCheckError(Exception): + """Base class for export compliance verification failures""" + + def to_error_detail(self) -> dict: + """Return a JSON-serializable representation of this error for API responses.""" + return {"detail": str(self)} + + +class ExportComplianceError(ExportComplianceCheckError): + """A user failed a CyberSource export compliance check""" + + def __init__(self, user, decision, reason_code, msg=None): + """ + Sets exception properties and adds a default message + + Args: + user (users.models.User): The user who failed the export compliance check + decision (str): The decision returned by CyberSource (e.g. "REJECT") + reason_code (str or int): The reason code returned by CyberSource + """ + self.user = user + self.decision = decision + self.reason_code = reason_code + if msg is None: + msg = ( + "Export compliance check did not accept enrollment for " + f"user={user.id}: decision={decision!r}, reason_code={reason_code!r}" + ) + super().__init__(msg) + + def to_error_detail(self) -> dict: + """Return a JSON-serializable representation of this error for API responses.""" + return { + "detail": str(self), + "decision": self.decision, + "reason_code": self.reason_code, + } + + +class ExportComplianceDataError(ExportComplianceCheckError): + """A user is missing the profile data required to run an export compliance check""" + + def __init__(self, user, missing_fields, msg=None): + """ + Sets exception properties and adds a default message + + Args: + user (users.models.User): The user missing required profile data + missing_fields (list of str): The billTo fields that could not be populated + """ + self.user = user + self.missing_fields = sorted(set(missing_fields)) + if msg is None: + msg = ( + f"Unable to verify export compliance for user={user.id}: missing " + f"required profile information ({', '.join(self.missing_fields)}). " + "Please update your profile or contact support." + ) + super().__init__(msg) + + def to_error_detail(self) -> dict: + """Return a JSON-serializable representation of this error for API responses.""" + return { + "detail": str(self), + "missing_fields": self.missing_fields, + } diff --git a/conftest.py b/conftest.py index 384b65b667..20e29aba25 100644 --- a/conftest.py +++ b/conftest.py @@ -2,6 +2,7 @@ import uuid from pathlib import Path +from types import SimpleNamespace import pytest from faker import Faker @@ -24,6 +25,7 @@ def default_settings(monkeypatch, settings): settings.FEATURES[features.IGNORE_EDX_FAILURES] = False settings.FEATURES[features.SYNC_ON_DASHBOARD_LOAD] = False + settings.FEATURES[features.EXPORT_COMPLIANCE_CHECK_ENABLED] = True @pytest.fixture(autouse=True) @@ -40,9 +42,30 @@ def mocked_flexibleprice_signal(mocker): @pytest.fixture(autouse=True) def payment_gateway_settings(settings): + """Set default CyberSource settings for tests.""" settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_SECURITY_KEY = "Test Security Key" settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_ACCESS_KEY = "Test Access Key" settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_PROFILE_ID = uuid.uuid4() + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_ID = "merchant-id" + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_SECRET_KEY_ID = uuid.uuid4().hex + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_MERCHANT_SECRET = uuid.uuid4().hex + settings.MITOL_PAYMENT_GATEWAY_CYBERSOURCE_REST_API_ENVIRONMENT = ( + "apitest.cybersource.com" + ) + + +@pytest.fixture(autouse=True) +def mocked_export_compliance(mocker): + """Mock export compliance checks in shared enrollment helpers by default.""" + return mocker.patch( + "courses.api.verify_user_with_exports", + return_value=SimpleNamespace( + accepted=True, + decision="ACCEPT", + reason_code=100, + request_id="test-request-id", + ), + ) @pytest.fixture(autouse=True) diff --git a/courses/api.py b/courses/api.py index 8d3cfd946c..80b7d5f599 100644 --- a/courses/api.py +++ b/courses/api.py @@ -31,6 +31,8 @@ from b2b.api import process_add_org_membership from cms.api import create_default_courseware_page +from compliance.api import verify_user_with_exports +from compliance.exceptions import ExportComplianceError from courses import mail_api from courses.constants import ( COURSE_KEY_PATTERN, @@ -202,6 +204,8 @@ def create_run_enrollments( # noqa: C901 created in mitxonline, paired with a boolean indicating whether or not the edX enrollment API call was successful for all of the given course runs """ + _verify_exports_compliance_for_enrollment(user) + if keep_failed_enrollments is None: keep_failed_enrollments = settings.FEATURES.get( features.IGNORE_EDX_FAILURES, False @@ -318,6 +322,8 @@ def create_program_enrollments( Returns: list of ProgramEnrollment: A list of enrollment objects that were successfully created """ + _verify_exports_compliance_for_enrollment(user) + successful_enrollments = [] for program in programs: try: @@ -396,6 +402,25 @@ def upgrade_audit_run_enrollments_for_program_purchase(user, program): return upgraded_enrollments +def _verify_exports_compliance_for_enrollment(user) -> None: + """Verify users with CyberSource before creating enrollments.""" + if not settings.FEATURES.get(features.EXPORT_COMPLIANCE_CHECK_ENABLED, False): + return + + result = verify_user_with_exports(user) + if result.accepted: + return + + log.warning( + "Export compliance check did not accept enrollment for user=%s: " + "decision=%r, reason_code=%r", + user.id, + result.decision, + result.reason_code, + ) + raise ExportComplianceError(user, result.decision, result.reason_code) + + def downgrade_learner(enrollment): """ Downgrades given enrollment from verified to audit. diff --git a/courses/api_test.py b/courses/api_test.py index bd1f977d07..bf34d35472 100644 --- a/courses/api_test.py +++ b/courses/api_test.py @@ -35,6 +35,8 @@ OrganizationPageFactory, ) from cms.factories import CourseIndexPageFactory +from compliance.api import ExportComplianceResult +from compliance.exceptions import ExportComplianceError from courses.api import ( check_course_modes, create_local_enrollment, @@ -98,6 +100,7 @@ ) from ecommerce.factories import LineFactory, OrderFactory, ProductFactory from ecommerce.models import Basket, OrderStatus +from main import features from main.constants import USER_MSG_TYPE_B2B_ENROLL_SUCCESS from main.test_utils import MockHttpError from openedx.constants import ( @@ -859,6 +862,184 @@ def test_mixed_enrollments_upgrades_only_audit( assert verified_enrollment.enrollment_mode == EDX_ENROLLMENT_VERIFIED_MODE +def test_create_run_enrollments_verifies_exports_for_verified_mode( + mocker, user, django_capture_on_commit_callbacks +): + """Verified course enrollments should require an accepted export check.""" + run = CourseRunFactory.create() + patched_verify = mocker.patch( + "courses.api.verify_user_with_exports", + return_value=ExportComplianceResult( + decision="ACCEPT", + reason_code=100, + request_id="req-123", + raw={}, + ), + ) + patched_edx_enroll = mocker.patch("courses.api.enroll_in_edx_course_runs") + mocker.patch("courses.api.mail_api.send_course_run_enrollment_email") + mocker.patch("courses.tasks.subscribe_edx_course_emails.delay") + + with django_capture_on_commit_callbacks(execute=True): + successful_enrollments, edx_request_success = create_run_enrollments( + user, [run], mode=EDX_ENROLLMENT_VERIFIED_MODE + ) + + patched_verify.assert_called_once_with(user) + patched_edx_enroll.assert_called_once_with( + user, + [run], + mode=EDX_ENROLLMENT_VERIFIED_MODE, + ) + assert edx_request_success is True + assert len(successful_enrollments) == 1 + + +def test_create_run_enrollments_verifies_exports_for_audit_mode(mocker, user): + """Audit course enrollments should also require an accepted export check.""" + run = CourseRunFactory.create() + patched_verify = mocker.patch( + "courses.api.verify_user_with_exports", + return_value=ExportComplianceResult( + decision="ACCEPT", + reason_code=100, + request_id="req-123", + raw={}, + ), + ) + patched_edx_enroll = mocker.patch("courses.api.enroll_in_edx_course_runs") + mocker.patch("courses.api.mail_api.send_course_run_enrollment_email") + mocker.patch("courses.tasks.subscribe_edx_course_emails.delay") + + create_run_enrollments(user, [run], mode=EDX_ENROLLMENT_AUDIT_MODE) + + patched_verify.assert_called_once_with(user) + patched_edx_enroll.assert_called_once() + + +def test_create_run_enrollments_rejects_nonaccepted_exports(mocker, user): + """Verified course enrollments should fail closed when exports are not accepted.""" + run = CourseRunFactory.create() + patched_verify = mocker.patch( + "courses.api.verify_user_with_exports", + return_value=ExportComplianceResult( + decision="REJECT", + reason_code=102, + request_id="req-123", + raw={}, + ), + ) + patched_edx_enroll = mocker.patch("courses.api.enroll_in_edx_course_runs") + + with pytest.raises( + ExportComplianceError, match="Export compliance check did not accept" + ): + create_run_enrollments(user, [run], mode=EDX_ENROLLMENT_VERIFIED_MODE) + + patched_verify.assert_called_once_with(user) + patched_edx_enroll.assert_not_called() + assert not CourseRunEnrollment.objects.filter(user=user, run=run).exists() + + +def test_create_run_enrollments_skips_exports_check_when_feature_disabled( + settings, mocker, user +): + """The export compliance check should be skipped entirely when the feature flag is off.""" + settings.FEATURES[features.EXPORT_COMPLIANCE_CHECK_ENABLED] = False + run = CourseRunFactory.create() + patched_verify = mocker.patch( + "courses.api.verify_user_with_exports", + return_value=ExportComplianceResult( + decision="REJECT", + reason_code=102, + request_id="req-123", + raw={}, + ), + ) + patched_edx_enroll = mocker.patch("courses.api.enroll_in_edx_course_runs") + + successful_enrollments, _ = create_run_enrollments( + user, [run], mode=EDX_ENROLLMENT_VERIFIED_MODE + ) + + patched_verify.assert_not_called() + patched_edx_enroll.assert_called_once() + assert len(successful_enrollments) == 1 + + +def test_create_program_enrollments_verifies_exports_for_verified_mode(mocker, user): + """Verified program enrollments should require an accepted export check.""" + program = ProgramFactory.create() + patched_verify = mocker.patch( + "courses.api.verify_user_with_exports", + return_value=ExportComplianceResult( + decision="ACCEPT", + reason_code=100, + request_id="req-123", + raw={}, + ), + ) + + successful_enrollments = create_program_enrollments( + user, + [program], + enrollment_mode=EDX_ENROLLMENT_VERIFIED_MODE, + ) + + patched_verify.assert_called_once_with(user) + assert len(successful_enrollments) == 1 + assert successful_enrollments[0].program == program + + +def test_create_program_enrollments_verifies_exports_for_default_mode(mocker, user): + """Default program enrollments should also require an accepted export check.""" + program = ProgramFactory.create() + patched_verify = mocker.patch( + "courses.api.verify_user_with_exports", + return_value=ExportComplianceResult( + decision="ACCEPT", + reason_code=100, + request_id="req-123", + raw={}, + ), + ) + + successful_enrollments = create_program_enrollments( + user, + [program], + ) + + patched_verify.assert_called_once_with(user) + assert len(successful_enrollments) == 1 + assert successful_enrollments[0].program == program + + +def test_create_program_enrollments_rejects_nonaccepted_exports(mocker, user): + """Verified program enrollments should fail closed when exports are not accepted.""" + program = ProgramFactory.create() + patched_verify = mocker.patch( + "courses.api.verify_user_with_exports", + return_value=ExportComplianceResult( + decision="REVIEW", + reason_code=480, + request_id="req-123", + raw={}, + ), + ) + + with pytest.raises( + ExportComplianceError, match="Export compliance check did not accept" + ): + create_program_enrollments( + user, + [program], + enrollment_mode=EDX_ENROLLMENT_VERIFIED_MODE, + ) + + patched_verify.assert_called_once_with(user) + assert not ProgramEnrollment.objects.filter(user=user, program=program).exists() + + class TestDeactivateEnrollments: """Test cases for functions that deactivate enrollments""" diff --git a/courses/management/commands/migrate_edx_data.py b/courses/management/commands/migrate_edx_data.py index 8e02400046..63b941e9d8 100644 --- a/courses/management/commands/migrate_edx_data.py +++ b/courses/management/commands/migrate_edx_data.py @@ -298,7 +298,17 @@ def _bulk_create_legal_addresses(created_users, row_lookup_by_id, batch_size): if not user_data: continue country = user_data.get("user_address_country") or "" - legal_addresses.append(LegalAddress(user=user, country=country)) + legal_addresses.append( + LegalAddress( + user=user, + country=country, + state=user_data.get("user_address_state") or None, + postal_code=user_data.get("user_address_postal_code") or "", + street_address_1=user_data.get("user_address_street_1") or "", + street_address_2=user_data.get("user_address_street_2") or "", + city=user_data.get("user_address_city") or "", + ) + ) if legal_addresses: LegalAddress.objects.bulk_create( legal_addresses, batch_size=batch_size, ignore_conflicts=True diff --git a/courses/serializers/v1/courses.py b/courses/serializers/v1/courses.py index b07b8305fc..3ea513a98c 100644 --- a/courses/serializers/v1/courses.py +++ b/courses/serializers/v1/courses.py @@ -10,6 +10,7 @@ from rest_framework.exceptions import ValidationError from cms.serializers import CoursePageSerializer +from compliance.exceptions import ExportComplianceCheckError from courses import models from courses.api import create_run_enrollments from courses.serializers.v1.base import ( @@ -175,13 +176,16 @@ def create(self, validated_data): if run.b2b_contract is not None: raise ValidationError({"run_id": f"Invalid course run id: {run_id}"}) - successful_enrollments, _ = create_run_enrollments( - user, - [run], - keep_failed_enrollments=settings.FEATURES.get( - features.IGNORE_EDX_FAILURES, False - ), - ) + try: + successful_enrollments, _ = create_run_enrollments( + user, + [run], + keep_failed_enrollments=settings.FEATURES.get( + features.IGNORE_EDX_FAILURES, False + ), + ) + except ExportComplianceCheckError as exc: + raise ValidationError(exc.to_error_detail()) from exc return successful_enrollments[0] if successful_enrollments else None diff --git a/courses/serializers/v2/courses.py b/courses/serializers/v2/courses.py index c6bf688fcf..e43a465a69 100644 --- a/courses/serializers/v2/courses.py +++ b/courses/serializers/v2/courses.py @@ -12,6 +12,7 @@ from rest_framework.exceptions import ValidationError from cms.serializers import CoursePageSerializer +from compliance.exceptions import ExportComplianceCheckError from courses import models from courses.api import create_run_enrollments from courses.serializers.utils import get_topics_from_page @@ -340,13 +341,18 @@ def create(self, validated_data): if run.b2b_contract is not None: raise ValidationError({"run_id": f"Invalid course run id: {run_id}"}) - successful_enrollments, _ = create_run_enrollments( - user, - [run], - keep_failed_enrollments=settings.FEATURES.get( - features.IGNORE_EDX_FAILURES, False - ), - ) + + try: + successful_enrollments, _ = create_run_enrollments( + user, + [run], + keep_failed_enrollments=settings.FEATURES.get( + features.IGNORE_EDX_FAILURES, False + ), + ) + except ExportComplianceCheckError as exc: + raise ValidationError(exc.to_error_detail()) from exc + return successful_enrollments[0] if successful_enrollments else None @extend_schema_field(serializers.IntegerField(allow_null=True)) diff --git a/courses/serializers/v3/courses.py b/courses/serializers/v3/courses.py index cf6a32e761..ff38929fb7 100644 --- a/courses/serializers/v3/courses.py +++ b/courses/serializers/v3/courses.py @@ -9,6 +9,7 @@ from rest_framework import serializers from rest_framework.exceptions import ValidationError +from compliance.exceptions import ExportComplianceCheckError from courses import models from courses.api import create_run_enrollments from courses.serializers.v1.base import ( @@ -113,13 +114,17 @@ def create(self, validated_data): if run is None or run.b2b_contract_id is not None: raise ValidationError({"run_id": f"Invalid course run id: {run_id}"}) - successful_enrollments, _ = create_run_enrollments( - user, - [run], - keep_failed_enrollments=settings.FEATURES.get( - features.IGNORE_EDX_FAILURES, False - ), - ) + try: + successful_enrollments, _ = create_run_enrollments( + user, + [run], + keep_failed_enrollments=settings.FEATURES.get( + features.IGNORE_EDX_FAILURES, False + ), + ) + except ExportComplianceCheckError as exc: + raise ValidationError(exc.to_error_detail()) from exc + if not successful_enrollments: msg = "Unable to create course run enrollment" raise ValueError(msg) diff --git a/courses/views/v1/__init__.py b/courses/views/v1/__init__.py index 91b622ad01..f7e408040f 100644 --- a/courses/views/v1/__init__.py +++ b/courses/views/v1/__init__.py @@ -27,6 +27,7 @@ from rest_framework.views import APIView from reversion.models import Version +from compliance.exceptions import ExportComplianceCheckError from courses.api import ( create_run_enrollments, deactivate_run_enrollment, @@ -366,13 +367,19 @@ def post(self, request): create_user(user) user.refresh_from_db() - _, edx_request_success = create_run_enrollments( - user=user, - runs=[run], - keep_failed_enrollments=settings.FEATURES.get( - features.IGNORE_EDX_FAILURES, False - ), - ) + try: + _, edx_request_success = create_run_enrollments( + user=user, + runs=[run], + keep_failed_enrollments=settings.FEATURES.get( + features.IGNORE_EDX_FAILURES, False + ), + ) + except ExportComplianceCheckError as exc: + return redirect_with_user_message( + reverse("user-dashboard"), + {"type": USER_MSG_TYPE_ENROLL_BLOCKED, **exc.to_error_detail()}, + ) def respond(data, status=True): # noqa: FBT002 """ diff --git a/courses/views/v2/__init__.py b/courses/views/v2/__init__.py index 24b494c953..37963c9b86 100644 --- a/courses/views/v2/__init__.py +++ b/courses/views/v2/__init__.py @@ -32,6 +32,7 @@ from rest_framework.response import Response from cms.models import CoursePage +from compliance.exceptions import ExportComplianceCheckError from courses.api import ( create_program_enrollments, create_run_enrollments, @@ -784,12 +785,15 @@ def _create_course_enrollment_from_program(request, courserun_id, program_enroll if should_create_audit_enrollment: # Audit enrollments just get created, regardless of whether or not # the course is an elective. - enrollments, _ = create_run_enrollments( - request.user, - [run], - mode=EDX_ENROLLMENT_AUDIT_MODE, - keep_failed_enrollments=True, - ) + try: + enrollments, _ = create_run_enrollments( + request.user, + [run], + mode=EDX_ENROLLMENT_AUDIT_MODE, + keep_failed_enrollments=True, + ) + except ExportComplianceCheckError as exc: + return Response(exc.to_error_detail(), status=status.HTTP_400_BAD_REQUEST) if len(enrollments) == 0: raise EnrollmentCreationFailedError return Response( @@ -810,6 +814,51 @@ def _create_course_enrollment_from_program(request, courserun_id, program_enroll ) +def _reconcile_verified_program_enrollments( + request, courserun_id, root_program, verified_program_enrollments, programs +): + """ + Create audit/verified program enrollments as needed to reconcile the + learner's verified enrollment state before enrolling them in the course run. + + Returns: + Response: an error response if reconciliation failed, otherwise None + """ + try: + if len(verified_program_enrollments) == 0: + # No verified enrollments, so it doesn't matter - the user will get an + # audit one. (But make the audit enrollment to not confuse the course run + # process later.) + create_program_enrollments( + request.user, programs, enrollment_mode=EDX_ENROLLMENT_AUDIT_MODE + ) + elif ( + len(verified_program_enrollments) == 1 + and verified_program_enrollments[0].program == root_program + ): + # The verified enrollment that's here is for the root program, so we can + # create a verified enrollment for the other program. + create_program_enrollments( + request.user, programs, enrollment_mode=EDX_ENROLLMENT_VERIFIED_MODE + ) + elif ( + len(verified_program_enrollments) == 1 + and verified_program_enrollments[0].program != root_program + ): + # The verified enrollment that's here is _not_ for the root program, so + # we should stop. + log.error( + "add_verified_program_course_enrollment: user %s enrolling in %s has no verified enrollment in %s", + request.user, + courserun_id, + root_program, + ) + return Response(status=status.HTTP_400_BAD_REQUEST) + except ExportComplianceCheckError as exc: + return Response(exc.to_error_detail(), status=status.HTTP_400_BAD_REQUEST) + return None + + @extend_schema( request=list[str], responses={ @@ -910,35 +959,11 @@ def add_verified_program_course_enrollment(request, courserun_id: str): if enrollment.enrollment_mode == EDX_ENROLLMENT_VERIFIED_MODE ] - if len(verified_program_enrollments) == 0: - # No verified enrollments, so it doesn't matter - the user will get an - # audit one. (But make the audit enrollment to not confuse the course run - # process later.) - create_program_enrollments( - request.user, programs, enrollment_mode=EDX_ENROLLMENT_AUDIT_MODE - ) - elif ( - len(verified_program_enrollments) == 1 - and verified_program_enrollments[0].program == root_program - ): - # The verified enrollment that's here is for the root program, so we can - # create a verified enrollment for the other program. - create_program_enrollments( - request.user, programs, enrollment_mode=EDX_ENROLLMENT_VERIFIED_MODE - ) - elif ( - len(verified_program_enrollments) == 1 - and verified_program_enrollments[0].program != root_program - ): - # The verified enrollment that's here is _not_ for the root program, so - # we should stop. - log.error( - "add_verified_program_course_enrollment: user %s enrolling in %s has no verified enrollment in %s", - request.user, - courserun_id, - root_program, - ) - return Response(status=status.HTTP_400_BAD_REQUEST) + error_response = _reconcile_verified_program_enrollments( + request, courserun_id, root_program, verified_program_enrollments, programs + ) + if error_response is not None: + return error_response # If we fell out the bottom, we have all verified enrollments, or we've made # sufficient enrollments to fill the gaps. diff --git a/courses/views/v3/__init__.py b/courses/views/v3/__init__.py index e72ae9eacd..d01d6ee3ba 100644 --- a/courses/views/v3/__init__.py +++ b/courses/views/v3/__init__.py @@ -25,6 +25,7 @@ from rest_framework.response import Response from b2b.models import ContractPage +from compliance.exceptions import ExportComplianceCheckError from courses.api import create_program_enrollments, deactivate_run_enrollment from courses.constants import COURSE_KEY_PATTERN, ENROLL_CHANGE_STATUS_UNENROLLED from courses.models import ( @@ -235,7 +236,11 @@ def create(self, request, *args, **kwargs): # noqa: ARG002 return Response(response_serializer.data, status=status.HTTP_200_OK) # Create the enrollment using default enrollment mode (audit) - enrollments = create_program_enrollments(request.user, [program]) + try: + enrollments = create_program_enrollments(request.user, [program]) + except ExportComplianceCheckError as exc: + raise serializers.ValidationError(exc.to_error_detail()) from exc + if not enrollments: raise ValueError("Failed to create program enrollment.") # noqa: EM101 response_serializer = ProgramEnrollmentSerializer(enrollments[0]) diff --git a/drf_lint_baseline.json b/drf_lint_baseline.json index baad7c7f54..a7c792c838 100644 --- a/drf_lint_baseline.json +++ b/drf_lint_baseline.json @@ -9,7 +9,7 @@ "courses/serializers/base.py:53:16:ORM001", "courses/serializers/v1/base.py:75:20:ORM002", "courses/serializers/v1/courses.py:171:18:ORM001", - "courses/serializers/v1/courses.py:57:16:ORM001", + "courses/serializers/v1/courses.py:58:16:ORM001", "courses/serializers/v1/programs.py:181:12:ORM001", "courses/serializers/v1/programs.py:196:12:ORM001", "courses/serializers/v1/programs.py:208:12:ORM001", @@ -19,14 +19,14 @@ "courses/serializers/v1/programs.py:332:16:ORM001", "courses/serializers/v1/programs.py:346:17:ORM001", "courses/serializers/v1/programs.py:368:16:ORM001", - "courses/serializers/v2/courses.py:272:17:ORM002", + "courses/serializers/v2/courses.py:273:17:ORM002", "courses/serializers/v2/departments.py:35:40:ORM002", "courses/serializers/v2/departments.py:49:42:ORM002", "courses/serializers/v2/programs.py:385:53:ORM002", "courses/serializers/v2/programs.py:495:12:ORM002", "courses/serializers/v2/programs.py:596:12:ORM002", "courses/serializers/v3/courses.py:111:14:ORM001", - "courses/serializers/v3/courses.py:55:12:ORM002", + "courses/serializers/v3/courses.py:56:12:ORM002", "courses/serializers/v3/programs.py:55:22:ORM001", "ecommerce/serializers/__init__.py:199:17:ORM001", "ecommerce/serializers/__init__.py:201:18:ORM001", diff --git a/fixtures/common.py b/fixtures/common.py index c33324837f..31ec9e9bcb 100644 --- a/fixtures/common.py +++ b/fixtures/common.py @@ -98,6 +98,7 @@ def valid_address_dict(): return dict( # noqa: C408 country="US", state="US-MA", + postal_code="02139", ) @@ -107,6 +108,7 @@ def invalid_address_dict(): return dict( # noqa: C408 country="US", state="XX", + postal_code="02139", ) @@ -116,6 +118,7 @@ def address_no_state_dict(): return dict( # noqa: C408 country="US", state=None, + postal_code="02139", ) @@ -126,6 +129,7 @@ def intl_address_dict(): return dict( # noqa: C408 last_name="User", country="JP", + postal_code="100-0001", ) diff --git a/main/features.py b/main/features.py index f5be3154a1..6485093b7e 100644 --- a/main/features.py +++ b/main/features.py @@ -8,3 +8,5 @@ ENABLE_GOOGLE_ANALYTICS_DATA_PUSH = "mitxonline-4099-dedp-google-analytics" REDIRECT_LEARN_DASHBOARD = "redirect-to-learn-dashboard" + +EXPORT_COMPLIANCE_CHECK_ENABLED = "EXPORT_COMPLIANCE_CHECK_ENABLED" diff --git a/main/settings.py b/main/settings.py index 5ba9f80cbb..c4e078a1f4 100644 --- a/main/settings.py +++ b/main/settings.py @@ -265,7 +265,7 @@ "cms.apps.CustomWagtailUsersAppConfig", "cms", "sheets", - # "compliance", + "compliance", "openedx", # must be after "users" to pick up custom user model "ecommerce", diff --git a/openapi/specs/v0.yaml b/openapi/specs/v0.yaml index cb1efb9a4a..a5fa4f16e6 100644 --- a/openapi/specs/v0.yaml +++ b/openapi/specs/v0.yaml @@ -6540,10 +6540,22 @@ components: country: type: string maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 email: type: string description: Get email from the linked user object @@ -6930,10 +6942,22 @@ components: country: type: string maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 required: - country LegalAddressRequest: @@ -6944,10 +6968,22 @@ components: type: string minLength: 1 maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 required: - country Line: diff --git a/openapi/specs/v1.yaml b/openapi/specs/v1.yaml index ea37999834..8425bc9b91 100644 --- a/openapi/specs/v1.yaml +++ b/openapi/specs/v1.yaml @@ -6540,10 +6540,22 @@ components: country: type: string maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 email: type: string description: Get email from the linked user object @@ -6930,10 +6942,22 @@ components: country: type: string maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 required: - country LegalAddressRequest: @@ -6944,10 +6968,22 @@ components: type: string minLength: 1 maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 required: - country Line: diff --git a/openapi/specs/v2.yaml b/openapi/specs/v2.yaml index 9f7511d1f3..e93ad88c0c 100644 --- a/openapi/specs/v2.yaml +++ b/openapi/specs/v2.yaml @@ -6540,10 +6540,22 @@ components: country: type: string maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 email: type: string description: Get email from the linked user object @@ -6930,10 +6942,22 @@ components: country: type: string maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 required: - country LegalAddressRequest: @@ -6944,10 +6968,22 @@ components: type: string minLength: 1 maxLength: 2 + street_address_1: + type: string + maxLength: 255 + street_address_2: + type: string + maxLength: 255 + city: + type: string + maxLength: 255 state: type: string nullable: true maxLength: 10 + postal_code: + type: string + maxLength: 20 required: - country Line: diff --git a/users/factories.py b/users/factories.py index 5a63ccda0a..a2a6603884 100644 --- a/users/factories.py +++ b/users/factories.py @@ -54,6 +54,10 @@ class LegalAddressFactory(DjangoModelFactory): user = SubFactory("users.factories.UserFactory") country = Faker("country_code", representation="alpha-2") + street_address_1 = Faker("street_address") + street_address_2 = Faker("secondary_address") + city = Faker("city") + postal_code = Faker("postcode") class Meta: model = LegalAddress diff --git a/users/migrations/0041_legaladdress_postal_code.py b/users/migrations/0041_legaladdress_postal_code.py new file mode 100644 index 0000000000..5ab1f6f71b --- /dev/null +++ b/users/migrations/0041_legaladdress_postal_code.py @@ -0,0 +1,16 @@ +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("users", "0040_add_is_etl_flag"), + ] + + operations = [ + migrations.AddField( + model_name="legaladdress", + name="postal_code", + field=models.CharField(blank=True, default="", max_length=20), + preserve_default=False, + ), + ] diff --git a/users/migrations/0042_legaladdress_street_and_city.py b/users/migrations/0042_legaladdress_street_and_city.py new file mode 100644 index 0000000000..b134853196 --- /dev/null +++ b/users/migrations/0042_legaladdress_street_and_city.py @@ -0,0 +1,28 @@ +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("users", "0041_legaladdress_postal_code"), + ] + + operations = [ + migrations.AddField( + model_name="legaladdress", + name="street_address_1", + field=models.CharField(blank=True, default="", max_length=255), + preserve_default=False, + ), + migrations.AddField( + model_name="legaladdress", + name="street_address_2", + field=models.CharField(blank=True, default="", max_length=255), + preserve_default=False, + ), + migrations.AddField( + model_name="legaladdress", + name="city", + field=models.CharField(blank=True, default="", max_length=255), + preserve_default=False, + ), + ] diff --git a/users/models.py b/users/models.py index 120dbb35b8..c4b92ad0a8 100644 --- a/users/models.py +++ b/users/models.py @@ -457,6 +457,10 @@ class LegalAddress(TimestampedModel): max_length=2, blank=True, validators=[validate_iso_3166_1_code] ) # ISO-3166-1 state = models.CharField(max_length=255, blank=True, null=True) # noqa: DJ001 + postal_code = models.CharField(max_length=20, blank=True) + street_address_1 = models.CharField(max_length=255, blank=True) + street_address_2 = models.CharField(max_length=255, blank=True) + city = models.CharField(max_length=255, blank=True) @property def us_state(self): diff --git a/users/serializers.py b/users/serializers.py index 2268b5aafd..86c3e3472b 100644 --- a/users/serializers.py +++ b/users/serializers.py @@ -99,9 +99,17 @@ class LegalAddressSerializer(serializers.ModelSerializer): # NOTE: the model defines these as allowing empty values for backwards compatibility # so we override them here to require them for new writes country = serializers.CharField(max_length=2) + street_address_1 = serializers.CharField( + max_length=255, required=False, allow_blank=True + ) + street_address_2 = serializers.CharField( + max_length=255, required=False, allow_blank=True + ) + city = serializers.CharField(max_length=255, required=False, allow_blank=True) state = serializers.CharField( max_length=10, required=False, allow_blank=True, allow_null=True ) + postal_code = serializers.CharField(max_length=20, required=False, allow_blank=True) def validate(self, data): """Validate the legal address data""" @@ -131,7 +139,11 @@ class Meta: model = LegalAddress fields = ( "country", + "street_address_1", + "street_address_2", + "city", "state", + "postal_code", ) diff --git a/users/views_test.py b/users/views_test.py index ec758e6ee5..a39125525c 100644 --- a/users/views_test.py +++ b/users/views_test.py @@ -17,6 +17,18 @@ from variants.serializers import SupportedVariantSerializer +def expected_legal_address(user): + """Return the serialized legal address shape for a user.""" + return { + "country": user.legal_address.country, + "street_address_1": user.legal_address.street_address_1, + "street_address_2": user.legal_address.street_address_2, + "city": user.legal_address.city, + "state": user.legal_address.state, + "postal_code": user.legal_address.postal_code, + } + + @pytest.mark.django_db def test_cannot_create_user(client): """Verify the api to create a user is nonexistent""" @@ -118,10 +130,7 @@ def test_get_user_by_me(mocker, client, user, is_anonymous, has_orgs): "email": user.email, "name": user.name, "global_id": user.global_id, - "legal_address": { - "country": user.legal_address.country, - "state": user.legal_address.state, - }, + "legal_address": expected_legal_address(user), "user_profile": { "gender": user.user_profile.gender, "year_of_birth": user.user_profile.year_of_birth, @@ -268,10 +277,7 @@ def test_get_userinfo(client, user, is_anonymous, has_openedx_user, has_edx_user "email": user.email, "name": user.name, "global_id": user.global_id, - "legal_address": { - "country": user.legal_address.country, - "state": user.legal_address.state, - }, + "legal_address": expected_legal_address(user), "user_profile": { "gender": user.user_profile.gender, "year_of_birth": user.user_profile.year_of_birth,