Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
127 changes: 125 additions & 2 deletions enterprise_access/apps/api/serializers/customer_billing.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from urllib.parse import urljoin

from django.conf import settings
from django_countries.serializer_fields import CountryField
from django_countries.serializers import CountryFieldMixin
from drf_spectacular.utils import extend_schema_field
from rest_framework import serializers
Expand Down Expand Up @@ -33,6 +34,53 @@ class RecordConflictError(APIException):
default_code = 'record_conflict_error'


BILLING_ADDRESS_FIELDS = (
'billing_address_country',
'billing_address_line_1',
'billing_address_line_2',
'billing_address_city',
'billing_address_state',
'billing_address_postal_code',
)
BILLING_ADDRESS_REQUIRED_FIELDS = (
'billing_address_country',
'billing_address_line_1',
'billing_address_city',
'billing_address_state',
'billing_address_postal_code',
)


def _is_billing_address_value_present(value):
return value not in (None, '')


def validate_billing_address_fields(serializer, attrs):
"""
Require a complete billing address whenever any billing address field is supplied.
"""
final_values = {}
for field_name in BILLING_ADDRESS_FIELDS:
if field_name in attrs:
final_values[field_name] = attrs[field_name]
elif serializer.instance is not None:
final_values[field_name] = getattr(serializer.instance, field_name)
else:
final_values[field_name] = None

if not any(_is_billing_address_value_present(value) for value in final_values.values()):
return attrs

errors = {
field_name: 'This field is required when billing address details are provided.'
for field_name in BILLING_ADDRESS_REQUIRED_FIELDS
if not _is_billing_address_value_present(final_values[field_name])
}
if errors:
raise serializers.ValidationError(errors)
return attrs


# pylint: disable=abstract-method
class CustomerBillingCreateCheckoutSessionRequestSerializer(serializers.Serializer):
"""
Expand Down Expand Up @@ -66,6 +114,52 @@ class CustomerBillingCreateCheckoutSessionRequestSerializer(serializers.Serializ
required=False,
help_text='The slug of the SSP product representing the plan selection.',
)
billing_address_country = CountryField(
required=False,
allow_null=True,
help_text='Two-letter ISO country code for the billing address.',
)
billing_address_line_1 = serializers.CharField(
required=False,
allow_null=True,
allow_blank=True,
max_length=255,
help_text='First line of the billing street address.',
)
billing_address_line_2 = serializers.CharField(
required=False,
allow_null=True,
allow_blank=True,
max_length=255,
help_text='Second line of the billing street address (optional).',
)
billing_address_city = serializers.CharField(
required=False,
allow_null=True,
allow_blank=True,
max_length=255,
help_text='Billing address city.',
)
billing_address_state = serializers.CharField(
required=False,
allow_null=True,
allow_blank=True,
max_length=255,
help_text='Billing address state or province.',
)
billing_address_postal_code = serializers.CharField(
required=False,
allow_null=True,
allow_blank=True,
max_length=20,
help_text='Billing address postal code.',
)

def validate(self, attrs):
"""
Require a complete billing address when any billing address field is supplied.
"""
return validate_billing_address_fields(self, attrs)


# pylint: disable=abstract-method
Expand Down Expand Up @@ -160,7 +254,17 @@ class Meta:
fields = '__all__'
read_only_fields = [
field.name for field in CheckoutIntent._meta.get_fields()
if field.name not in ('state', 'country', 'terms_metadata')
if field.name not in (
'state',
'country',
'billing_address_country',
'billing_address_line_1',
'billing_address_line_2',
'billing_address_city',
'billing_address_state',
'billing_address_postal_code',
'terms_metadata',
)
]

def validate_state(self, value):
Expand Down Expand Up @@ -199,6 +303,13 @@ def validate_terms_metadata(self, value):
)
return value

def validate(self, attrs):
"""
Perform cross-field validation, including optional billing address completeness.
"""
attrs = super().validate(attrs)
return validate_billing_address_fields(self, attrs)


class CheckoutIntentCreateRequestSerializer(CountryFieldMixin, serializers.ModelSerializer):
"""
Expand All @@ -221,6 +332,12 @@ class Meta:
'enterprise_name',
'quantity',
'country',
'billing_address_country',
'billing_address_line_1',
'billing_address_line_2',
'billing_address_city',
'billing_address_state',
'billing_address_postal_code',
'terms_metadata',
'ssp_product'
]
Expand Down Expand Up @@ -252,7 +369,7 @@ def validate(self, attrs):
raise serializers.ValidationError(
{'enterprise_slug': 'enterprise_slug is required when enterprise_name is provided.'}
)
return attrs
return validate_billing_address_fields(self, attrs)

def create(self, validated_data):
"""
Expand All @@ -266,6 +383,12 @@ def create(self, validated_data):
slug=validated_data.get('enterprise_slug'),
name=validated_data.get('enterprise_name'),
country=validated_data.get('country'),
billing_address_country=validated_data.get('billing_address_country'),
billing_address_line_1=validated_data.get('billing_address_line_1'),
billing_address_line_2=validated_data.get('billing_address_line_2'),
billing_address_city=validated_data.get('billing_address_city'),
billing_address_state=validated_data.get('billing_address_state'),
billing_address_postal_code=validated_data.get('billing_address_postal_code'),
terms_metadata=validated_data.get('terms_metadata'),
ssp_product=ssp_product,
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -360,6 +360,12 @@ def test_create_checkout_intent_success(self):
'enterprise_name': 'Test Enterprise post',
'quantity': 13,
'country': 'NZ',
'billing_address_country': 'US',
'billing_address_line_1': '123 Main St',
'billing_address_line_2': 'Suite 200',
'billing_address_city': 'New York',
'billing_address_state': 'NY',
'billing_address_postal_code': '10001',
'terms_metadata': {'version': '1.0', 'accepted_at': '2024-01-15T10:30:00Z'},
'ssp_product': 'teams-yearly',
}
Expand All @@ -377,6 +383,12 @@ def test_create_checkout_intent_success(self):
self.assertEqual(response_data['quantity'], 13)
self.assertEqual(response_data['state'], CheckoutIntentState.CREATED)
self.assertEqual(response_data['country'], 'NZ')
self.assertEqual(response_data['billing_address_country'], 'US')
self.assertEqual(response_data['billing_address_line_1'], '123 Main St')
self.assertEqual(response_data['billing_address_line_2'], 'Suite 200')
self.assertEqual(response_data['billing_address_city'], 'New York')
self.assertEqual(response_data['billing_address_state'], 'NY')
self.assertEqual(response_data['billing_address_postal_code'], '10001')
self.assertEqual(response_data['terms_metadata'], {'version': '1.0', 'accepted_at': '2024-01-15T10:30:00Z'})

def test_create_or_update_checkout_intent_success(self):
Expand All @@ -391,6 +403,11 @@ def test_create_or_update_checkout_intent_success(self):
'enterprise_name': self.checkout_intent_1.enterprise_name,
'quantity': 33,
'country': 'IT',
'billing_address_country': 'IT',
'billing_address_line_1': 'Via Roma 1',
'billing_address_city': 'Rome',
'billing_address_state': 'RM',
'billing_address_postal_code': '00100',
'terms_metadata': {'version': '2.0', 'updated': True},
'ssp_product': 'teams-yearly',
}
Expand All @@ -408,12 +425,48 @@ def test_create_or_update_checkout_intent_success(self):
self.assertEqual(response_data['quantity'], 33)
self.assertEqual(response_data['state'], CheckoutIntentState.CREATED)
self.assertEqual(response_data['country'], 'IT')
self.assertEqual(response_data['billing_address_country'], 'IT')
self.assertEqual(response_data['billing_address_line_1'], 'Via Roma 1')
self.assertEqual(response_data['billing_address_city'], 'Rome')
self.assertEqual(response_data['billing_address_state'], 'RM')
self.assertEqual(response_data['billing_address_postal_code'], '00100')
self.assertEqual(response_data['terms_metadata'], {'version': '2.0', 'test_mode': True, 'updated': True})
self.checkout_intent_1.refresh_from_db()
self.assertEqual(self.checkout_intent_1.quantity, 33)
self.assertEqual(self.checkout_intent_1.country, 'IT')
self.assertEqual(self.checkout_intent_1.billing_address_country, 'IT')
self.assertEqual(self.checkout_intent_1.billing_address_line_1, 'Via Roma 1')
self.assertEqual(self.checkout_intent_1.billing_address_city, 'Rome')
self.assertEqual(self.checkout_intent_1.billing_address_state, 'RM')
self.assertEqual(self.checkout_intent_1.billing_address_postal_code, '00100')
self.assertEqual(self.checkout_intent_1.terms_metadata, {'version': '2.0', 'test_mode': True, 'updated': True})

def test_create_checkout_intent_requires_complete_billing_address(self):
"""Test creation fails when only part of the billing address is provided."""
other_user = UserFactory()
self.set_jwt_cookie([{
'system_wide_role': SYSTEM_ENTERPRISE_LEARNER_ROLE,
'context': str(uuid.uuid4()),
}], user=other_user)

response = self.client.post(
self.list_url,
{
'enterprise_slug': 'test-enterprise-billing',
'enterprise_name': 'Test Enterprise Billing',
'quantity': 5,
'billing_address_line_1': '123 Main St',
'ssp_product': 'teams-yearly',
},
format='json'
)

self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn('billing_address_country', response.data)
self.assertIn('billing_address_city', response.data)
self.assertIn('billing_address_state', response.data)
self.assertIn('billing_address_postal_code', response.data)

@ddt.data(
# Invalid quantity cases:
{'quantity': -1, 'enterprise_slug': 'valid', 'enterprise_name': 'Valid'},
Expand Down Expand Up @@ -503,6 +556,39 @@ def test_update_terms_metadata_and_country(self):
self.assertEqual(self.checkout_intent_1.terms_metadata, new_terms)
self.assertEqual(self.checkout_intent_1.country, 'AU')

def test_update_billing_address(self):
"""Test updating billing address fields via PATCH."""
self.set_jwt_cookie([{
'system_wide_role': SYSTEM_ENTERPRISE_LEARNER_ROLE,
'context': str(uuid.uuid4()),
}])

response = self.client.patch(
self.detail_url_1,
{
'billing_address_country': 'US',
'billing_address_line_1': '77 Massachusetts Ave',
'billing_address_city': 'Cambridge',
'billing_address_state': 'MA',
'billing_address_postal_code': '02139',
},
format='json'
)

self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data['billing_address_country'], 'US')
self.assertEqual(response.data['billing_address_line_1'], '77 Massachusetts Ave')
self.assertEqual(response.data['billing_address_city'], 'Cambridge')
self.assertEqual(response.data['billing_address_state'], 'MA')
self.assertEqual(response.data['billing_address_postal_code'], '02139')

self.checkout_intent_1.refresh_from_db()
self.assertEqual(self.checkout_intent_1.billing_address_country, 'US')
self.assertEqual(self.checkout_intent_1.billing_address_line_1, '77 Massachusetts Ave')
self.assertEqual(self.checkout_intent_1.billing_address_city, 'Cambridge')
self.assertEqual(self.checkout_intent_1.billing_address_state, 'MA')
self.assertEqual(self.checkout_intent_1.billing_address_postal_code, '02139')

@ddt.data(
# Test that strings are rejected
{'terms_metadata': 'invalid_string'},
Expand Down
58 changes: 56 additions & 2 deletions enterprise_access/apps/api/v1/tests/test_customer_billing.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,12 @@
)
from enterprise_access.apps.core.tests.factories import UserFactory
from enterprise_access.apps.customer_billing.constants import CheckoutIntentState
from enterprise_access.apps.customer_billing.models import CheckoutIntent, StripeEventData, StripeEventSummary
from enterprise_access.apps.customer_billing.models import (
CheckoutIntent,
SspProduct,
StripeEventData,
StripeEventSummary
)
from enterprise_access.apps.customer_billing.tests.utils import AttrDict
from test_utils import APITest

Expand Down Expand Up @@ -3873,6 +3878,14 @@ class CreateCheckoutSessionViewTests(APITest):
def setUp(self):
super().setUp()
self.url = reverse('api:v1:customer-billing-create-checkout-session')
SspProduct.objects.get_or_create(
slug='quarterly_license_plan',
defaults={
'stripe_price_lookup_key': 'price_quarterly_0002',
'is_active': True,
'catalog_query_uuid': uuid.uuid4(),
},
)

def tearDown(self):
CheckoutIntent.objects.all().delete()
Expand Down Expand Up @@ -3922,7 +3935,12 @@ def test_create_checkout_session_returns_client_secret_from_dict(
'company_name': 'Test Co',
'quantity': 5,
'stripe_price_id': 'price_abc123',
'ssp_product': 'quarterly_license_plan',
'ssp_product_slug': 'quarterly_license_plan',
'billing_address_country': 'US',
'billing_address_line_1': '123 Main St',
'billing_address_city': 'Boston',
'billing_address_state': 'MA',
'billing_address_postal_code': '02110',
},
format='json',
)
Expand All @@ -3932,4 +3950,40 @@ def test_create_checkout_session_returns_client_secret_from_dict(
response.data['checkout_session_client_secret'],
'cs_test_abc_secret_xyz',
)
mock_create_intent.assert_called_once_with(
user=self.user,
quantity=5,
slug='test-slug',
name='Test Co',
billing_address_country='US',
billing_address_line_1='123 Main St',
billing_address_line_2=None,
billing_address_city='Boston',
billing_address_state='MA',
billing_address_postal_code='02110',
ssp_product=mock.ANY,
)
mock_enterprise_catalog_metadata.assert_called()

def test_create_checkout_session_requires_complete_billing_address(self):
"""Billing address fields must be complete when any billing address detail is supplied."""
self.set_jwt_cookie()

response = self.client.post(
self.url,
data={
'admin_email': self.user.email,
'enterprise_slug': 'test-slug',
'company_name': 'Test Co',
'quantity': 5,
'stripe_price_id': 'price_abc123',
'billing_address_line_1': '123 Main St',
},
format='json',
)

self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertIn('billing_address_country', response.data)
self.assertIn('billing_address_city', response.data)
self.assertIn('billing_address_state', response.data)
self.assertIn('billing_address_postal_code', response.data)
Loading