diff --git a/.gitignore b/.gitignore
index 6379115442..02f7688a47 100644
--- a/.gitignore
+++ b/.gitignore
@@ -132,3 +132,5 @@ docker-compose.*.yml
certs/*
.claude/
+
+.drf_lint_cache.json
diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml
index e67de1005f..51bd82210f 100644
--- a/.pre-commit-config.yaml
+++ b/.pre-commit-config.yaml
@@ -80,5 +80,5 @@ repos:
language: python
files: '(serializers\.py$|serializers/.*\.py$)'
additional_dependencies:
- - mitol-drf-lint
+ - "mitol-drf-lint==2026.8.28"
- "setuptools<82"
diff --git a/RELEASE.rst b/RELEASE.rst
index 2536944d5c..9269990d96 100644
--- a/RELEASE.rst
+++ b/RELEASE.rst
@@ -1,6 +1,19 @@
Release Notes
=============
+Version 1.166.0
+---------------
+
+- feat (hq11846): Complete your Purchase, paid-amount-off discount behavior (#3926)
+- Retire b2b_contract create --create, demote the org sync to a reconciler (C1 5/5) (#3932)
+- Expose the provisioning API under /api/v0/b2b/provisioning/ (C1 4/5) (#3931)
+- Provision Keycloak organizations and IdPs at runtime (C1 3/5) (#3930)
+- Add the B2B onboarding and identity provider records (C1 2/5) (#3929)
+- Give the Keycloak admin client the calls provisioning needs (C1 1/5) (#3928)
+- Harden test for locals (#3950)
+- chore: pin mitol-drf-lint in the drf-serializer-orm-check hook (#3943)
+- Reuse prefetched course runs in the v2 course API again (#3948)
+
Version 1.165.4
---------------
diff --git a/b2b/api.py b/b2b/api.py
index 3f45c672dd..b8da4bb9aa 100644
--- a/b2b/api.py
+++ b/b2b/api.py
@@ -16,7 +16,7 @@
from django.contrib.contenttypes.models import ContentType
from django.core.cache import caches
from django.core.exceptions import ValidationError
-from django.db import transaction
+from django.db import IntegrityError, transaction
from django.db.models import Count, Manager, Prefetch, Q
from mitol.common.utils import now_in_utc
from opaque_keys.edx.keys import CourseKey
@@ -32,6 +32,7 @@
MAILGUN_LOGS_DESC,
MAILGUN_LOGS_PAGE_LIMIT,
MAILGUN_LOGS_RETENTION_DAYS,
+ ONBOARDING_STATE_ORG_CREATED,
ORG_KEY_MAX_LENGTH,
RETIREMENT_CONTRACT_NAME,
RETIREMENT_ORG_KEY,
@@ -49,6 +50,7 @@
ContractProgramItem,
DiscountContractAttachmentRedemption,
OrganizationIndexPage,
+ OrganizationOnboarding,
OrganizationPage,
UserOrganization,
)
@@ -1810,31 +1812,65 @@ def reconcile_keycloak_orgs():
create or update corresponding records in MITx Online. This does not manage
memberships, just base org info.
+ Since the provisioning API (capability C1) writes both systems together,
+ this is a drift reconciler rather than the primary create path: it adopts
+ the organizations Pulumi still owns, ones made in the console, and ones left
+ behind by a provisioning saga whose compensating delete also failed. That
+ last case is why it has to see the whole realm, not a first page of it.
+
Returns
- tuple (created, updated): number of orgs created and updated
"""
org_model = get_keycloak_model(*KCAM_ORGANIZATIONS)
- orgs = org_model.list()
+ orgs = org_model.list_all()
parent_org_page = OrganizationIndexPage.objects.first()
created_count = 0
updated_count = 0
for org in orgs:
try:
- page, created = reconcile_single_keycloak_org(org)
+ # Each org gets its own savepoint. Postgres aborts the whole
+ # transaction on any failed statement, so catching a database error
+ # and carrying on with the loop only works if that error was
+ # contained - otherwise every later query raises
+ # TransactionManagementError and skipping one org still loses the
+ # rest of the pass, just less legibly.
+ with transaction.atomic():
+ page, created = reconcile_single_keycloak_org(org)
+
+ if created:
+ parent_org_page.add_child(instance=page)
+ page.save()
+ parent_org_page.save()
+ else:
+ page.save()
+
+ # An adopted organization needs an onboarding record too, so
+ # that orgs that arrived this way show up in the same place as
+ # the ones the provisioning API made.
+ OrganizationOnboarding.objects.get_or_create(
+ organization=page,
+ defaults={
+ "state": ONBOARDING_STATE_ORG_CREATED,
+ "state_changed_at": now_in_utc(),
+ },
+ )
+ # Counted after the savepoint commits, so a rolled-back org is not
+ # reported as reconciled.
if created:
created_count += 1
- parent_org_page.add_child(instance=page)
- page.save()
- parent_org_page.save()
else:
updated_count += 1
- page.save()
- except ValidationError: # noqa: PERF203
+ except (ValidationError, IntegrityError): # noqa: PERF203
+ # IntegrityError because OrganizationOnboarding.organization is a
+ # OneToOneField: a concurrent provisioning saga or a second
+ # reconcile run can insert the row between this one's check and its
+ # insert. The per-org catch is the point - one org losing that race
+ # must not abandon the rest of the pass.
log.exception(
- "Validation error: could not create or update organization for Keycloak org %s",
+ "Could not create or update organization for Keycloak org %s",
org.id,
)
diff --git a/b2b/api_test.py b/b2b/api_test.py
index 815ef83e19..3dc8870201 100644
--- a/b2b/api_test.py
+++ b/b2b/api_test.py
@@ -44,6 +44,7 @@
B2B_RUN_TAG_FORMAT,
CONTRACT_MEMBERSHIP_CODE,
CONTRACT_MEMBERSHIP_MANAGED,
+ ONBOARDING_STATE_ORG_CREATED,
)
from b2b.exceptions import SourceCourseIncompleteError
from b2b.factories import ContractPageFactory, OrganizationPageFactory
@@ -929,6 +930,18 @@ def list(self):
return self.orgs
+ def list_all(self, page_size=None, **kwargs): # noqa: ARG002
+ """
+ Return every fake org.
+
+ reconcile_keycloak_orgs pages rather than calling list, because
+ Keycloak's collection endpoints answer with 10 results when no max
+ is given and a drift reconciler that sees a first page is not a
+ reconciler.
+ """
+
+ return self.orgs
+
org_model = MockedOrgModel()
org_model.orgs = factories.OrganizationRepresentationFactory.create_batch(3)
@@ -982,6 +995,13 @@ def list(self):
assert found_count == (3 if not update_an_org else 4)
+ # An adopted org needs an onboarding record too, so orgs that arrived this
+ # way show up in the same place as the ones the provisioning API made.
+ assert all(
+ org_page.onboarding.state == ONBOARDING_STATE_ORG_CREATED
+ for org_page in org_pages
+ )
+
def test_reconcile_bad_keycloak_org(mocker):
"""Test that reconciliation works when there's bad data"""
diff --git a/b2b/constants.py b/b2b/constants.py
index 6e3c92cb4f..d99e5d351c 100644
--- a/b2b/constants.py
+++ b/b2b/constants.py
@@ -42,6 +42,73 @@
RETIREMENT_ORG_NAME = "Retired Runs"
RETIREMENT_CONTRACT_NAME = "Retired Runs Holding Contract"
+# Onboarding states for a B2B organization, in order. `blocked` is reachable
+# from anywhere, with the reason in OrganizationOnboarding.notes.
+#
+# The state is descriptive, not enforcing: it records what has been observed to
+# be true so an operator can answer "what is left for this customer" without
+# reading four systems. Nothing in the provisioning API gates on it.
+ONBOARDING_STATE_REQUESTED = "requested"
+ONBOARDING_STATE_ORG_CREATED = "org_created"
+ONBOARDING_STATE_IDP_CONFIGURED = "idp_configured"
+ONBOARDING_STATE_IDP_VALIDATED = "idp_validated"
+ONBOARDING_STATE_CONTRACT_READY = "contract_ready"
+ONBOARDING_STATE_LIVE = "live"
+ONBOARDING_STATE_BLOCKED = "blocked"
+
+ONBOARDING_STATE_CHOICES = [
+ (ONBOARDING_STATE_REQUESTED, "Requested"),
+ (ONBOARDING_STATE_ORG_CREATED, "Organization created"),
+ (ONBOARDING_STATE_IDP_CONFIGURED, "Identity provider configured"),
+ (ONBOARDING_STATE_IDP_VALIDATED, "Identity provider validated"),
+ (ONBOARDING_STATE_CONTRACT_READY, "Contract ready"),
+ (ONBOARDING_STATE_LIVE, "Live"),
+ (ONBOARDING_STATE_BLOCKED, "Blocked"),
+]
+
+IDP_PROTOCOL_SAML = "saml"
+IDP_PROTOCOL_OIDC = "oidc"
+IDP_PROTOCOL_CHOICES = [
+ (IDP_PROTOCOL_SAML, "SAML"),
+ (IDP_PROTOCOL_OIDC, "OIDC"),
+]
+
+IDP_STATE_DRAFT = "draft"
+IDP_STATE_TESTING = "testing"
+IDP_STATE_ACTIVE = "active"
+IDP_STATE_DISABLED = "disabled"
+
+IDP_LIFECYCLE_CHOICES = [
+ (IDP_STATE_DRAFT, "Draft"),
+ (IDP_STATE_TESTING, "Testing"),
+ (IDP_STATE_ACTIVE, "Active"),
+ (IDP_STATE_DISABLED, "Disabled"),
+]
+
+# The lifecycle state is written to Keycloak's own `enabled`/`hideOnLogin` as
+# well as our row, so the two cannot drift. `hideOnLogin` stays true
+# throughout, matching what the Pulumi resources set today: partner IdPs are
+# reached by organization/domain routing or an explicit kc_idp_hint, never by a
+# button on the shared login page. `testing` and `active` therefore carry the
+# same Keycloak flags -- they differ in whether the organization has an
+# email-domain redirect pointing at the IdP, which is org-level config.
+IDP_STATE_KEYCLOAK_FLAGS = {
+ IDP_STATE_DRAFT: {"enabled": False, "hideOnLogin": True},
+ IDP_STATE_TESTING: {"enabled": True, "hideOnLogin": True},
+ IDP_STATE_ACTIVE: {"enabled": True, "hideOnLogin": True},
+ IDP_STATE_DISABLED: {"enabled": False, "hideOnLogin": True},
+}
+
+# An IdP goes live only after somebody has actually logged in through it, so
+# there is no draft -> active edge. An IdP that has already been through
+# testing can be re-enabled directly.
+IDP_ALLOWED_TRANSITIONS = {
+ IDP_STATE_DRAFT: [IDP_STATE_TESTING],
+ IDP_STATE_TESTING: [IDP_STATE_DRAFT, IDP_STATE_ACTIVE, IDP_STATE_DISABLED],
+ IDP_STATE_ACTIVE: [IDP_STATE_TESTING, IDP_STATE_DISABLED],
+ IDP_STATE_DISABLED: [IDP_STATE_TESTING, IDP_STATE_ACTIVE],
+}
+
MAILGUN_LOGS_API_URL = "https://api.mailgun.net/v1/analytics/logs"
MAILGUN_LOGS_PAGE_LIMIT = 100
MAILGUN_LOGS_DESC = "timestamp:desc"
diff --git a/b2b/exceptions.py b/b2b/exceptions.py
index 778ef30e78..0921646a78 100644
--- a/b2b/exceptions.py
+++ b/b2b/exceptions.py
@@ -16,3 +16,41 @@ class TargetCourseRunExistsError(Exception):
class KeycloakAdminImproperlyConfiguredError(Exception):
"""Raised if Keycloak admin client is improperly configured."""
+
+
+class AliasCollisionError(Exception):
+ """
+ Raised when a Keycloak alias is already taken.
+
+ Organization and identity provider aliases are realm-wide, and the realm is
+ shared with the resources Pulumi still declares, so an alias that is free in
+ our own tables can still collide. Creating one anyway would break the next
+ pulumi up that declares the same name.
+ """
+
+
+class InvalidLifecycleTransitionError(Exception):
+ """Raised when an identity provider is asked to skip a lifecycle state."""
+
+
+class OrphanedKeycloakOrganizationError(Exception):
+ """
+ Raised when a Keycloak organization is left behind by a failed create.
+
+ The organization creation saga compensates for a failed MITx Online write by
+ deleting the Keycloak organization it just made. When that compensating
+ delete also fails, the organization is orphaned and this is raised with its
+ ID so the caller can surface it.
+ """
+
+
+class OrganizationNotProvisionedError(Exception):
+ """
+ Raised when an organization has no Keycloak counterpart to act on.
+
+ An OrganizationPage with a null sso_organization_id cannot be updated in
+ Keycloak, because there is nothing there to update. Roughly 24 production
+ organizations are in this state, inherited from mitodl/hq#10552 and from
+ the b2b_contract create --create path that made them; they need backfilling
+ through this API rather than patching.
+ """
diff --git a/b2b/keycloak_admin_api.py b/b2b/keycloak_admin_api.py
index 05997f3139..35b0eba26c 100644
--- a/b2b/keycloak_admin_api.py
+++ b/b2b/keycloak_admin_api.py
@@ -14,6 +14,7 @@
from b2b.exceptions import KeycloakAdminImproperlyConfiguredError
from b2b.keycloak_admin_dataclasses import (
+ IdentityProviderRepresentation,
OrganizationRepresentation,
RealmRepresentation,
UserRepresentation,
@@ -21,6 +22,16 @@
KCAM_ORGANIZATIONS = (OrganizationRepresentation, "organizations")
KCAM_USERS = (UserRepresentation, "users")
+KCAM_IDENTITY_PROVIDERS = (
+ IdentityProviderRepresentation,
+ "identity-provider/instances",
+)
+
+IDENTITY_PROVIDER_IMPORT_CONFIG_ENDPOINT = "identity-provider/import-config"
+
+# Keycloak's collection endpoints default to returning 10 items when `max` is
+# not supplied, so anything that needs the whole collection has to page.
+KEYCLOAK_LIST_PAGE_SIZE = 100
class KeycloakAdminClient:
@@ -298,11 +309,93 @@ def disassociate(self, endpoint):
- requests.HTTPError if the request fails.
"""
+ return self.delete(endpoint)
+
+ def delete(self, endpoint):
+ """
+ Delete the object at the endpoint in the realm.
+
+ Args:
+ - endpoint: The endpoint to use (e.g., "identity-provider/instances/{alias}").
+ Returns:
+ - True if successful.
+ Raises:
+ - requests.HTTPError if the request fails.
+ """
+
response = self.realm_request("DELETE", endpoint)
response.raise_for_status()
return True
+ def create_returning_id(self, endpoint, data):
+ """
+ Create an object at the endpoint in the realm and return its ID.
+
+ Several Keycloak creation endpoints - organizations and identity
+ provider instances among them - answer 201 with an empty body and the
+ new resource's location in the `Location` header, so `create` cannot be
+ used for them: it has no JSON to parse into a representation.
+
+ Args:
+ - endpoint: The endpoint to use (e.g., "organizations").
+ - data: The data to send.
+
+ Returns:
+ - The last path segment of the Location header, or None if Keycloak did
+ not send one.
+ """
+
+ response = self.realm_request("POST", endpoint, json=data)
+ response.raise_for_status()
+
+ location = response.headers.get("Location")
+
+ return location.rstrip("/").rsplit("/", 1)[-1] if location else None
+
+ def post_raw(self, endpoint, data):
+ """
+ POST to the endpoint in the realm and return the decoded JSON body.
+
+ For endpoints whose response is not one of the generated representation
+ classes - `identity-provider/import-config` answers with a flat config
+ map - so `create` cannot coerce it.
+
+ Args:
+ - endpoint: The endpoint to use.
+ - data: The data to send.
+
+ Returns:
+ - The decoded response body.
+ """
+
+ response = self.realm_request("POST", endpoint, json=data)
+ response.raise_for_status()
+
+ return response.json()
+
+ def post_file(self, endpoint, data, files):
+ """
+ POST a multipart form to the endpoint in the realm.
+
+ `identity-provider/import-config` takes either a JSON body naming a URL
+ for Keycloak to fetch, or a multipart upload of the metadata document
+ itself. This is the second form.
+
+ Args:
+ - endpoint: The endpoint to use.
+ - data: The form fields to send alongside the file.
+ - files: The requests-style files mapping.
+
+ Returns:
+ - The decoded response body.
+ """
+
+ response = self.realm_request("POST", endpoint, data=data, files=files)
+ response.raise_for_status()
+
+ return response.json()
+
class KeycloakAdminModel:
"""Middleware class to help with working with Keycloak data."""
@@ -344,6 +437,31 @@ def list(self, **kwargs):
self.endpoint, self.representation_class, **kwargs
)
+ def list_all(self, page_size=KEYCLOAK_LIST_PAGE_SIZE, **kwargs):
+ """
+ List every object in the realm, paging until the results run out.
+
+ `list` passes no `max`, and Keycloak's collection endpoints default to
+ 10 results, so it silently truncates any collection larger than that.
+ Use this wherever the whole collection is the point - a realm-wide
+ alias check, a reconciliation pass.
+
+ Args:
+ - page_size: How many results to ask for per request.
+ Returns:
+ - A list of representation instances.
+ """
+
+ items = []
+ first = 0
+
+ while True:
+ page = self.list(first=first, max=page_size, **kwargs)
+ items.extend(page)
+ if len(page) < page_size:
+ return items
+ first += page_size
+
def get(self, item_id):
"""
Get a single object by its ID.
@@ -360,6 +478,54 @@ def get(self, item_id):
self.representation_class,
)
+ def create(self, data):
+ """
+ Create an object at the endpoint in the realm.
+
+ Args:
+ - data: The data to send.
+
+ Returns:
+ - The new object's ID (its alias, for endpoints keyed by one), or None
+ if Keycloak did not report a location.
+ Raises:
+ - requests.HTTPError if the request fails.
+ """
+
+ return self.admin_client.create_returning_id(self.endpoint, data)
+
+ def update(self, item_id, data):
+ """
+ Update the object with the given ID.
+
+ Args:
+ - item_id: The ID of the object to update.
+ - data: The full representation to write. Keycloak's PUT endpoints
+ replace rather than merge, so send everything.
+
+ Returns:
+ - True if successful.
+ Raises:
+ - requests.HTTPError if the request fails.
+ """
+
+ return self.admin_client.save(f"{self.endpoint}/{item_id}", data)
+
+ def delete(self, item_id):
+ """
+ Delete the object with the given ID.
+
+ Args:
+ - item_id: The ID of the object to delete.
+
+ Returns:
+ - True if successful.
+ Raises:
+ - requests.HTTPError if the request fails.
+ """
+
+ return self.admin_client.delete(f"{self.endpoint}/{item_id}")
+
def associate(self, association_type, parent_id, child_id):
"""
Associate the object with the given ID with the target ID.
@@ -414,6 +580,55 @@ def bootstrap_client(*, verify_realm=False):
return client
+def import_identity_provider_config(
+ provider_id, *, from_url=None, metadata=None, client=None
+):
+ """
+ Ask Keycloak to parse identity provider metadata into an IdP config map.
+
+ Keycloak fetches and parses the metadata itself, which is why we do not
+ port ol-infrastructure's saml_helpers into this codebase. Note that the URL
+ form makes Keycloak fetch a caller-supplied address, so the endpoints that
+ reach this stay staff-only.
+
+ This is not a KeycloakAdminModel method because import-config is a sibling
+ of identity-provider/instances rather than an operation on one.
+
+ Args:
+ - provider_id: The Keycloak provider ID ("saml" or "oidc").
+ - from_url: A metadata or discovery URL for Keycloak to fetch.
+ - metadata: The metadata document itself, uploaded instead of fetched.
+ - client: An optional KeycloakAdminClient instance.
+ Returns:
+ - The flat config dict Keycloak parsed out of the metadata.
+ Raises:
+ - ValueError if neither from_url nor metadata was given.
+ - requests.HTTPError if the request fails.
+ """
+
+ if not from_url and not metadata:
+ # Both are keyword arguments defaulting to None, and callers can reach
+ # this from a shell as well as through the serializers. Say which
+ # argument is missing rather than letting requests raise a TypeError on
+ # a None file body several frames down.
+ msg = "Supply either from_url or metadata to parse."
+ raise ValueError(msg)
+
+ client = client or bootstrap_client()
+
+ if from_url:
+ return client.post_raw(
+ IDENTITY_PROVIDER_IMPORT_CONFIG_ENDPOINT,
+ {"providerId": provider_id, "fromUrl": from_url},
+ )
+
+ return client.post_file(
+ IDENTITY_PROVIDER_IMPORT_CONFIG_ENDPOINT,
+ {"providerId": provider_id},
+ {"file": ("metadata.xml", metadata, "application/xml")},
+ )
+
+
def get_keycloak_model(representation, endpoint, *, client=None):
"""
Get a KeycloakAdminModel instance for the given model type.
diff --git a/b2b/keycloak_admin_api_test.py b/b2b/keycloak_admin_api_test.py
index a86e01391c..a00d65019c 100644
--- a/b2b/keycloak_admin_api_test.py
+++ b/b2b/keycloak_admin_api_test.py
@@ -14,8 +14,12 @@
KeycloakAdminClient,
KeycloakAdminModel,
bootstrap_client,
+ import_identity_provider_config,
+)
+from b2b.keycloak_admin_dataclasses import (
+ IdentityProviderRepresentation,
+ RealmRepresentation,
)
-from b2b.keycloak_admin_dataclasses import RealmRepresentation
pytestmark = [pytest.mark.django_db]
FAKE = faker.Faker()
@@ -657,3 +661,154 @@ def test_realm_representation_ignores_extra_fields():
assert realm.id == fake_realm.id
assert realm.realm == fake_realm.realm
+
+
+def test_create_returning_id_reads_the_location_header(settings, mocker):
+ """Keycloak answers a create with 201 and an empty body plus a Location."""
+
+ client, _, _, _ = _mocked_admin_client(settings, mocker)
+
+ new_id = FAKE.uuid4()
+ response = requests.Response()
+ response.status_code = 201
+ response.headers["Location"] = (
+ f"https://keycloak.example.com/admin/realms/olapps/organizations/{new_id}"
+ )
+ mocker.patch.object(client, "realm_request", return_value=response)
+
+ assert client.create_returning_id("organizations", {"alias": "acme"}) == new_id
+
+
+def test_create_returning_id_without_a_location(settings, mocker):
+ """A create with no Location header returns None rather than guessing."""
+
+ client, _, _, _ = _mocked_admin_client(settings, mocker)
+
+ response = requests.Response()
+ response.status_code = 201
+ mocker.patch.object(client, "realm_request", return_value=response)
+
+ assert client.create_returning_id("organizations", {"alias": "acme"}) is None
+
+
+def test_post_raw_returns_the_body_uncoerced(settings, mocker):
+ """import-config answers with a flat config map, not a representation."""
+
+ client, _, _, _ = _mocked_admin_client(settings, mocker)
+
+ config = {"singleSignOnServiceUrl": FAKE.url(), "idpEntityId": FAKE.url()}
+ mocker.patch.object(client, "realm_request", return_value=_faked_response(config))
+
+ assert client.post_raw("identity-provider/import-config", {}) == config
+
+
+def test_model_list_all_pages_until_exhausted(settings, mocker):
+ """
+ list_all keeps asking until a short page comes back.
+
+ Keycloak's collection endpoints default to 10 results when `max` is absent,
+ so anything that needs the whole collection has to page for it.
+ """
+
+ client, _, _, _ = _mocked_admin_client(settings, mocker)
+
+ first_page = RealmRepresentationFactory.create_batch(2)
+ second_page = RealmRepresentationFactory.create_batch(1)
+ mocked_list = mocker.patch.object(
+ client, "list", side_effect=[first_page, second_page]
+ )
+
+ realm_model = KeycloakAdminModel(client, RealmRepresentation, "realms")
+
+ assert realm_model.list_all(page_size=2) == [*first_page, *second_page]
+ assert mocked_list.call_args_list == [
+ mocker.call("realms", RealmRepresentation, first=0, max=2),
+ mocker.call("realms", RealmRepresentation, first=2, max=2),
+ ]
+
+
+def test_model_create_update_delete(settings, mocker):
+ """The three model-level methods the provisioning API needs."""
+
+ client, _, _, _ = _mocked_admin_client(settings, mocker)
+
+ fake_alias = FAKE.word()
+ mocked_create = mocker.patch.object(
+ client, "create_returning_id", return_value=fake_alias
+ )
+ mocked_save = mocker.patch.object(client, "save", return_value=True)
+ mocked_delete = mocker.patch.object(client, "delete", return_value=True)
+
+ idp_model = KeycloakAdminModel(
+ client, IdentityProviderRepresentation, "identity-provider/instances"
+ )
+
+ assert idp_model.create({"alias": fake_alias}) == fake_alias
+ assert idp_model.update(fake_alias, {"enabled": True}) is True
+ assert idp_model.delete(fake_alias) is True
+
+ mocked_create.assert_called_once_with(
+ "identity-provider/instances", {"alias": fake_alias}
+ )
+ mocked_save.assert_called_once_with(
+ f"identity-provider/instances/{fake_alias}", {"enabled": True}
+ )
+ mocked_delete.assert_called_once_with(f"identity-provider/instances/{fake_alias}")
+
+
+def test_import_identity_provider_config_from_url(settings, mocker):
+ """The URL form asks Keycloak to fetch and parse the metadata itself."""
+
+ client, _, _, _ = _mocked_admin_client(settings, mocker)
+
+ config = {"idpEntityId": FAKE.url()}
+ mocked_post = mocker.patch.object(client, "post_raw", return_value=config)
+ metadata_url = FAKE.url()
+
+ result = import_identity_provider_config(
+ "saml", from_url=metadata_url, client=client
+ )
+
+ assert result == config
+ mocked_post.assert_called_once_with(
+ "identity-provider/import-config",
+ {"providerId": "saml", "fromUrl": metadata_url},
+ )
+
+
+def test_import_identity_provider_config_from_document(settings, mocker):
+ """Metadata supplied inline is uploaded rather than fetched."""
+
+ client, _, _, _ = _mocked_admin_client(settings, mocker)
+
+ config = {"idpEntityId": FAKE.url()}
+ mocked_post = mocker.patch.object(client, "post_file", return_value=config)
+ metadata = ""
+
+ result = import_identity_provider_config("saml", metadata=metadata, client=client)
+
+ assert result == config
+ mocked_post.assert_called_once_with(
+ "identity-provider/import-config",
+ {"providerId": "saml"},
+ {"file": ("metadata.xml", metadata, "application/xml")},
+ )
+
+
+def test_import_identity_provider_config_needs_a_source(settings, mocker):
+ """
+ With neither source, say so rather than posting a None file body.
+
+ Both sources are keyword arguments defaulting to None, so this is reachable
+ from a shell even though the serializers require one of them.
+ """
+
+ client, _, _, _ = _mocked_admin_client(settings, mocker)
+ mocked_post_raw = mocker.patch.object(client, "post_raw")
+ mocked_post_file = mocker.patch.object(client, "post_file")
+
+ with pytest.raises(ValueError, match="Supply either from_url or metadata"):
+ import_identity_provider_config("saml", client=client)
+
+ mocked_post_raw.assert_not_called()
+ mocked_post_file.assert_not_called()
diff --git a/b2b/management/commands/b2b_contract.py b/b2b/management/commands/b2b_contract.py
index af2e571ba5..8690369d7d 100644
--- a/b2b/management/commands/b2b_contract.py
+++ b/b2b/management/commands/b2b_contract.py
@@ -19,7 +19,6 @@
from b2b.models import (
ContractPage,
ContractProgramItem,
- OrganizationIndexPage,
OrganizationPage,
)
from courses.models import (
@@ -131,16 +130,12 @@ def add_arguments(self, parser):
type=str,
help="The end date of the contract.",
)
- create_parser.add_argument(
- "--create",
- action="store_true",
- help="Create an organization if it does not exist.",
- )
- create_parser.add_argument(
- "--org-key",
- type=str,
- help="The org key to use for the new organization.",
- )
+ # No --create: this used to build an OrganizationPage with no
+ # sso_organization_id, and orgs in that state are silently broken -
+ # attach_user() returns False without doing anything, so every
+ # membership write is a no-op (mitodl/hq#10552). Provision the
+ # organization through /api/v0/b2b/provisioning/organizations/ instead,
+ # which writes Keycloak and MITx Online together.
create_parser.add_argument(
"--max-learners",
type=int,
@@ -272,10 +267,8 @@ def handle_create(self, *args, **kwargs): # noqa: ARG002
description = kwargs.pop("description")
start_date = kwargs.pop("start")
end_date = kwargs.pop("end")
- create_organization = kwargs.pop("create")
max_learners = kwargs.pop("max_learners")
price = kwargs.pop("price")
- new_org_key = kwargs.pop("org_key")
self.stdout.write(
f"Creating contract '{contract_name}' for organization '{organization_name}'"
@@ -287,20 +280,12 @@ def handle_create(self, *args, **kwargs): # noqa: ARG002
log.info("Got organization %s", org)
- if not org and create_organization:
- if not new_org_key:
- msg = f"To create '{organization_name}', you must supply an org key."
- raise CommandError(msg)
-
- parent = OrganizationIndexPage.objects.first()
- org = OrganizationPage(name=organization_name, org_key=new_org_key)
- parent.add_child(instance=org)
- org.save()
- parent.save()
- org.refresh_from_db()
- self.stdout.write(f"Created organization '{organization_name}'")
- elif not org:
- msg = f"Organization '{organization_name}' does not exist. Use --create to create it."
+ if not org:
+ msg = (
+ f"Organization '{organization_name}' does not exist. Create it "
+ "through POST /api/v0/b2b/provisioning/organizations/, which "
+ "creates the Keycloak organization alongside it."
+ )
raise CommandError(msg)
contract = ContractPage(
diff --git a/b2b/migrations/0028_organizationidentityprovider_organizationonboarding.py b/b2b/migrations/0028_organizationidentityprovider_organizationonboarding.py
new file mode 100644
index 0000000000..e02ca11906
--- /dev/null
+++ b/b2b/migrations/0028_organizationidentityprovider_organizationonboarding.py
@@ -0,0 +1,161 @@
+# Generated by Django 5.2.15 on 2026-09-04 17:58
+
+import django.db.models.deletion
+from django.db import migrations, models
+
+# Imported by name rather than as `mitol.common.utils.datetime.now_in_utc`,
+# which is what makemigrations writes: mitol.common.utils re-exports the
+# datetime *class*, so that dotted path resolves to it and the attribute
+# lookup fails at import time.
+from mitol.common.utils.datetime import now_in_utc
+
+# Imported by name rather than as `mitol.common.utils.datetime.now_in_utc`,
+# which is what makemigrations writes: mitol.common.utils re-exports the
+# datetime *class*, so the dotted path resolves to that and the attribute
+# lookup fails at import time.
+
+
+class Migration(migrations.Migration):
+ dependencies = [
+ ("b2b", "0027_discountcontractattachmentredemption_email_message_id_and_more"),
+ ]
+
+ operations = [
+ migrations.CreateModel(
+ name="OrganizationIdentityProvider",
+ fields=[
+ (
+ "id",
+ models.BigAutoField(
+ auto_created=True,
+ primary_key=True,
+ serialize=False,
+ verbose_name="ID",
+ ),
+ ),
+ ("created_on", models.DateTimeField(auto_now_add=True)),
+ ("updated_on", models.DateTimeField(auto_now=True)),
+ (
+ "alias",
+ models.CharField(
+ help_text="The Keycloak IdP alias. Realm-wide, not per-organization.",
+ max_length=255,
+ unique=True,
+ ),
+ ),
+ (
+ "protocol",
+ models.CharField(
+ choices=[("saml", "SAML"), ("oidc", "OIDC")], max_length=8
+ ),
+ ),
+ (
+ "lifecycle_state",
+ models.CharField(
+ choices=[
+ ("draft", "Draft"),
+ ("testing", "Testing"),
+ ("active", "Active"),
+ ("disabled", "Disabled"),
+ ],
+ default="draft",
+ max_length=16,
+ ),
+ ),
+ (
+ "display_name",
+ models.CharField(blank=True, default="", max_length=255),
+ ),
+ (
+ "internal_id",
+ models.CharField(
+ blank=True,
+ default="",
+ help_text="Keycloak's internalId for the IdP instance.",
+ max_length=255,
+ ),
+ ),
+ (
+ "metadata_source",
+ models.TextField(
+ help_text="The metadata URL, or the inline XML, the config was parsed from. Not blankable: refreshing an IdP re-reads this, so a row without one cannot be refreshed."
+ ),
+ ),
+ (
+ "metadata_artifact",
+ models.JSONField(
+ blank=True,
+ help_text="The config map Keycloak parsed out of the metadata. Persisted so a partner's metadata endpoint going away can neither destroy config nor block a deploy.",
+ null=True,
+ ),
+ ),
+ ("metadata_fetched_at", models.DateTimeField(blank=True, null=True)),
+ (
+ "organization",
+ models.ForeignKey(
+ on_delete=django.db.models.deletion.CASCADE,
+ related_name="identity_providers",
+ to="b2b.organizationpage",
+ ),
+ ),
+ ],
+ options={
+ "abstract": False,
+ },
+ ),
+ migrations.CreateModel(
+ name="OrganizationOnboarding",
+ fields=[
+ (
+ "id",
+ models.BigAutoField(
+ auto_created=True,
+ primary_key=True,
+ serialize=False,
+ verbose_name="ID",
+ ),
+ ),
+ ("created_on", models.DateTimeField(auto_now_add=True)),
+ ("updated_on", models.DateTimeField(auto_now=True)),
+ (
+ "state",
+ models.CharField(
+ choices=[
+ ("requested", "Requested"),
+ ("org_created", "Organization created"),
+ ("idp_configured", "Identity provider configured"),
+ ("idp_validated", "Identity provider validated"),
+ ("contract_ready", "Contract ready"),
+ ("live", "Live"),
+ ("blocked", "Blocked"),
+ ],
+ default="requested",
+ max_length=32,
+ ),
+ ),
+ (
+ "state_changed_at",
+ models.DateTimeField(default=now_in_utc),
+ ),
+ (
+ "notes",
+ models.TextField(
+ blank=True,
+ default="",
+ help_text="Free-form operator notes; holds the reason when state is blocked.",
+ ),
+ ),
+ (
+ "organization",
+ models.OneToOneField(
+ on_delete=django.db.models.deletion.CASCADE,
+ related_name="onboarding",
+ to="b2b.organizationpage",
+ ),
+ ),
+ ],
+ options={
+ "abstract": False,
+ },
+ ),
+ ]
diff --git a/b2b/models.py b/b2b/models.py
index d4530e6a03..cd16a282d5 100644
--- a/b2b/models.py
+++ b/b2b/models.py
@@ -26,10 +26,16 @@
CONTRACT_MEMBERSHIP_AUTOS,
CONTRACT_MEMBERSHIP_MANAGED,
CONTRACT_MEMBERSHIP_TYPE_CHOICES,
+ IDP_LIFECYCLE_CHOICES,
+ IDP_PROTOCOL_CHOICES,
+ IDP_STATE_DRAFT,
+ ONBOARDING_STATE_CHOICES,
+ ONBOARDING_STATE_REQUESTED,
ORG_INDEX_SLUG,
)
from courses.constants import UAI_COURSEWARE_ID_PREFIX
from courses.models import Program
+from main.models import ValidateOnSaveMixin
from variants.models import SupportedVariant
log = logging.getLogger(__name__)
@@ -859,6 +865,111 @@ def __str__(self):
return f"UserOrganization: {self.user} in {self.organization}"
+class OrganizationOnboarding(TimestampedModel, ValidateOnSaveMixin):
+ """
+ Where an organization is in the B2B onboarding sequence.
+
+ The system of record that did not exist before: onboarding state lived in
+ people's heads and in four separate systems. The state is descriptive, not
+ enforcing - nothing gates on it. A state machine that blocks operators
+ before the operators trust it is a state machine they work around.
+ """
+
+ organization = models.OneToOneField(
+ "b2b.OrganizationPage",
+ on_delete=models.CASCADE,
+ related_name="onboarding",
+ )
+ state = models.CharField(
+ max_length=32,
+ choices=ONBOARDING_STATE_CHOICES,
+ default=ONBOARDING_STATE_REQUESTED,
+ )
+ state_changed_at = models.DateTimeField(default=now_in_utc)
+ notes = models.TextField(
+ blank=True,
+ default="",
+ help_text="Free-form operator notes; holds the reason when state is blocked.",
+ )
+
+ def set_state(self, state, notes=None):
+ """
+ Move to the given onboarding state and stamp when it happened.
+
+ Args:
+ - state (str): the state to move to
+ - notes (str): replacement notes, if any
+ """
+
+ self.state = state
+ self.state_changed_at = now_in_utc()
+ if notes is not None:
+ self.notes = notes
+ self.save()
+
+ def __str__(self):
+ """Return a reasonable representation of the object as a string."""
+
+ return f"OrganizationOnboarding: {self.organization} is {self.state}"
+
+
+class OrganizationIdentityProvider(TimestampedModel, ValidateOnSaveMixin):
+ """
+ An identity provider we provisioned in Keycloak for an organization.
+
+ Keycloak has no field for where an IdP is in its rollout, so the lifecycle
+ state lives here - but it is written through to Keycloak's `enabled` and
+ `hideOnLogin` on every transition (see IDP_STATE_KEYCLOAK_FLAGS) so the two
+ representations cannot drift apart silently.
+ """
+
+ organization = models.ForeignKey(
+ "b2b.OrganizationPage",
+ on_delete=models.CASCADE,
+ related_name="identity_providers",
+ )
+ alias = models.CharField(
+ max_length=255,
+ unique=True,
+ help_text="The Keycloak IdP alias. Realm-wide, not per-organization.",
+ )
+ protocol = models.CharField(max_length=8, choices=IDP_PROTOCOL_CHOICES)
+ lifecycle_state = models.CharField(
+ max_length=16,
+ choices=IDP_LIFECYCLE_CHOICES,
+ default=IDP_STATE_DRAFT,
+ )
+ display_name = models.CharField(max_length=255, blank=True, default="")
+ internal_id = models.CharField(
+ max_length=255,
+ blank=True,
+ default="",
+ help_text="Keycloak's internalId for the IdP instance.",
+ )
+ metadata_source = models.TextField(
+ help_text=(
+ "The metadata URL, or the inline XML, the config was parsed from. "
+ "Not blankable: refreshing an IdP re-reads this, so a row without "
+ "one cannot be refreshed."
+ ),
+ )
+ metadata_artifact = models.JSONField(
+ null=True,
+ blank=True,
+ help_text=(
+ "The config map Keycloak parsed out of the metadata. Persisted so a "
+ "partner's metadata endpoint going away can neither destroy config "
+ "nor block a deploy."
+ ),
+ )
+ metadata_fetched_at = models.DateTimeField(null=True, blank=True)
+
+ def __str__(self):
+ """Return a reasonable representation of the object as a string."""
+
+ return f"OrganizationIdentityProvider: {self.alias} ({self.lifecycle_state})"
+
+
def is_organization_manager(user, org_id):
"""
Check if a user is a manager of the specified organization.
diff --git a/b2b/provisioning.py b/b2b/provisioning.py
new file mode 100644
index 0000000000..d7532b79bb
--- /dev/null
+++ b/b2b/provisioning.py
@@ -0,0 +1,668 @@
+"""
+Staff-only provisioning of per-customer Keycloak resources (capability C1).
+
+The ownership split this implements: Pulumi keeps the realm, its authentication
+flows, client scopes, clients and service-account grants; this module owns
+Keycloak organizations, their domains, identity providers, IdP attribute
+mappers and org<->IdP links, alongside the MITx Online records that go with
+them.
+
+Pulumi only deletes resources that are in its own state, so an organization
+created here is invisible to it and safe. The cross-system hazard is alias
+collision, not deletion: organization and IdP aliases are realm-wide, so an
+alias created here will break a later pulumi up that declares the same name.
+Every create in this module checks the realm before writing.
+
+See docs/source/b2b/provisioning_api.md.
+"""
+
+import logging
+
+from django.core.exceptions import ImproperlyConfigured
+from django.db import transaction
+from mitol.common.utils import now_in_utc
+
+from b2b.constants import (
+ IDP_ALLOWED_TRANSITIONS,
+ IDP_PROTOCOL_OIDC,
+ IDP_PROTOCOL_SAML,
+ IDP_STATE_DRAFT,
+ IDP_STATE_KEYCLOAK_FLAGS,
+ ONBOARDING_STATE_ORG_CREATED,
+)
+from b2b.exceptions import (
+ AliasCollisionError,
+ InvalidLifecycleTransitionError,
+ OrganizationNotProvisionedError,
+ OrphanedKeycloakOrganizationError,
+)
+from b2b.keycloak_admin_api import (
+ KCAM_IDENTITY_PROVIDERS,
+ KCAM_ORGANIZATIONS,
+ bootstrap_client,
+ get_keycloak_model,
+ import_identity_provider_config,
+)
+from b2b.models import (
+ OrganizationIdentityProvider,
+ OrganizationIndexPage,
+ OrganizationOnboarding,
+ OrganizationPage,
+)
+
+log = logging.getLogger(__name__)
+
+ORG_IDP_ASSOCIATION = "identity-providers"
+
+# Keycloak's attribute-importer mapper, per protocol, and the config key each
+# one reads the source attribute from. These are the same mappers
+# ol-infrastructure's AttributeImporterIdentityProviderMapper resources produce,
+# so an IdP built here is shaped like the ones Pulumi still declares.
+IDP_ATTRIBUTE_MAPPERS = {
+ IDP_PROTOCOL_SAML: "saml-user-attribute-idp-mapper",
+ IDP_PROTOCOL_OIDC: "oidc-user-attribute-idp-mapper",
+}
+
+
+class KeycloakConnection:
+ """
+ A Keycloak admin client and the models built on it.
+
+ Bootstrapping the client fetches OIDC discovery and a token, so the
+ provisioning calls share one rather than each making their own.
+ """
+
+ def __init__(self, client=None):
+ self.client = client or bootstrap_client()
+ self.organizations = get_keycloak_model(*KCAM_ORGANIZATIONS, client=self.client)
+ self.identity_providers = get_keycloak_model(
+ *KCAM_IDENTITY_PROVIDERS, client=self.client
+ )
+
+
+def realm_organization_aliases(connection):
+ """
+ Return every organization alias in the realm, lowercased.
+
+ Includes the organizations Pulumi still owns - that is the point.
+
+ Args:
+ - connection (KeycloakConnection): the Keycloak connection to use.
+ Returns:
+ - set[str]: the aliases in use.
+ """
+
+ return {
+ org.alias.lower()
+ for org in connection.organizations.list_all()
+ if org.alias is not None
+ }
+
+
+def realm_identity_provider_aliases(connection):
+ """
+ Return every identity provider alias in the realm, lowercased.
+
+ Args:
+ - connection (KeycloakConnection): the Keycloak connection to use.
+ Returns:
+ - set[str]: the aliases in use.
+ """
+
+ return {
+ idp.alias.lower()
+ for idp in connection.identity_providers.list_all()
+ if idp.alias is not None
+ }
+
+
+def _require_provisioned(organization):
+ """
+ Refuse to act on an organization that has no Keycloak counterpart.
+
+ Without this the admin call goes to organizations/None, Keycloak answers
+ 404, and the operator is told the Keycloak API failed - a 502 they would
+ retry forever. Keycloak is fine; our record is the incomplete one.
+
+ Args:
+ - organization (OrganizationPage): the organization to check.
+ Raises:
+ - OrganizationNotProvisionedError: it has no sso_organization_id.
+ """
+
+ if not organization.sso_organization_id:
+ msg = (
+ f"Organization '{organization.org_key}' has no Keycloak "
+ "organization. It predates this API and needs backfilling before "
+ "it can be managed here."
+ )
+ raise OrganizationNotProvisionedError(msg)
+
+
+def _find_organization_by_alias(connection, alias):
+ """Return the realm's organization with this alias, or None."""
+
+ return next(
+ (
+ org
+ for org in connection.organizations.list_all()
+ if org.alias is not None and org.alias.lower() == alias.lower()
+ ),
+ None,
+ )
+
+
+def create_organization( # noqa: PLR0913
+ *,
+ name,
+ org_key,
+ org_key_prefix=None,
+ domains=(),
+ description="",
+ redirect_url="",
+ connection=None,
+):
+ """
+ Create a Keycloak organization and the MITx Online records beside it.
+
+ The ordering is what makes this safe. Keycloak has no transactions, so the
+ Keycloak write goes first - it is the step most likely to fail for reasons
+ outside our control - and the MITx Online writes go into one atomic block
+ afterwards. A failure there is compensated by deleting the Keycloak
+ organization we just made, because a Keycloak org with no MITx Online
+ counterpart is exactly the orphan shape this exists to prevent.
+
+ The Keycloak alias is the org_key verbatim. That is not cosmetic:
+ reconcile_single_keycloak_org derives org_key from the alias when it adopts
+ an organization, so any other choice would give an adopted org a different
+ org_key from the one it was created with - and org_key is baked into every
+ B2B courseware ID.
+
+ Args:
+ - name (str): the organization's display name
+ - org_key (str): the immutable short key, used as the Keycloak alias
+ - org_key_prefix (str): courseware ID prefix; defaults to the model default
+ - domains (list[str]): email domains to assert for the organization
+ - description (str): free-form description
+ - redirect_url (str): where Keycloak sends members after login
+ - connection (KeycloakConnection): an existing connection, if any
+ Returns:
+ - OrganizationPage: the new organization
+ Raises:
+ - AliasCollisionError: org_key is taken here or in the realm
+ - OrphanedKeycloakOrganizationError: the MITx Online write and its
+ compensating delete both failed
+ - requests.HTTPError: Keycloak rejected the create
+ """
+
+ connection = connection or KeycloakConnection()
+
+ if OrganizationPage.objects.filter(org_key=org_key).exists():
+ msg = f"An organization with org key '{org_key}' already exists."
+ raise AliasCollisionError(msg)
+
+ if org_key.lower() in realm_organization_aliases(connection):
+ msg = (
+ f"The alias '{org_key}' is already in use by an organization in the "
+ "Keycloak realm."
+ )
+ raise AliasCollisionError(msg)
+
+ # Resolve the parent page here rather than inside the transaction below. It
+ # is a precondition we can check before touching Keycloak, and a
+ # precondition checked after the irreversible external write is one the
+ # compensation has to clean up for no reason.
+ organization_index = OrganizationIndexPage.objects.first()
+ if organization_index is None:
+ msg = (
+ "No OrganizationIndexPage exists; the CMS is not set up to hold "
+ "organizations."
+ )
+ raise ImproperlyConfigured(msg)
+
+ # Domains are written verified, with no verification having occurred: staff
+ # are asserting them. That is defensible only while the asserting party is
+ # MIT staff, and stops being so the moment the partner-facing wizard (C2)
+ # lets a customer assert their own domain.
+ keycloak_payload = {
+ "name": name,
+ "alias": org_key,
+ "enabled": True,
+ "description": description,
+ "redirectUrl": redirect_url,
+ "domains": [{"name": domain, "verified": True} for domain in domains],
+ }
+
+ sso_organization_id = connection.organizations.create(keycloak_payload)
+
+ if not sso_organization_id:
+ # Keycloak answers the create with 201 and an empty body, so the ID
+ # normally comes from the Location header. Fall back to looking the
+ # organization up by the alias we just claimed.
+ created = _find_organization_by_alias(connection, org_key)
+ sso_organization_id = created.id if created else None
+
+ if not sso_organization_id:
+ # Keycloak accepted the create but we cannot find what it made. Writing
+ # our row anyway would produce an OrganizationPage with a null
+ # sso_organization_id, which is the silently-broken shape this saga
+ # exists to prevent - attach_user() would no-op for every member. We
+ # cannot compensate either, because the delete needs the ID we do not
+ # have. Leave it for reconcile_keycloak_orgs to adopt by alias.
+ log.error(
+ "Created a Keycloak organization with alias %s but could not "
+ "resolve its ID, from the Location header or by lookup",
+ org_key,
+ )
+ msg = (
+ f"Keycloak accepted the organization '{org_key}' but did not "
+ "report its ID, so no MITx Online record could be written."
+ )
+ raise OrphanedKeycloakOrganizationError(msg)
+
+ try:
+ with transaction.atomic():
+ organization = OrganizationPage(
+ name=name,
+ org_key=org_key,
+ description=description,
+ sso_organization_id=sso_organization_id,
+ )
+ if org_key_prefix:
+ organization.org_key_prefix = org_key_prefix
+
+ organization_index.add_child(instance=organization)
+ organization.save()
+
+ OrganizationOnboarding.objects.create(
+ organization=organization,
+ state=ONBOARDING_STATE_ORG_CREATED,
+ state_changed_at=now_in_utc(),
+ )
+ except Exception as write_error:
+ try:
+ connection.organizations.delete(sso_organization_id)
+ except Exception as compensation_error:
+ # Both writes failed. Do not retry against a system that just
+ # failed - log the orphan loudly and let reconcile_keycloak_orgs
+ # adopt it on its next run, which is the correct outcome anyway.
+ log.error( # noqa: TRY400
+ "Orphaned Keycloak organization %s (alias %s): the MITx Online "
+ "write failed and the compensating delete failed too",
+ sso_organization_id,
+ org_key,
+ )
+ msg = (
+ f"Keycloak organization {sso_organization_id} was created but "
+ "MITx Online records could not be written and the organization "
+ "could not be removed."
+ )
+ raise OrphanedKeycloakOrganizationError(msg) from compensation_error
+ raise write_error # noqa: TRY201
+
+ return organization
+
+
+def update_organization( # noqa: PLR0913
+ organization,
+ *,
+ name=None,
+ description=None,
+ redirect_url=None,
+ domains=None,
+ connection=None,
+):
+ """
+ Update an organization in both systems.
+
+ org_key is deliberately not updatable: it is in every B2B courseware ID via
+ create_contract_run_key, which is also why reconcile_single_keycloak_org
+ refuses to change it on adoption.
+
+ Keycloak's organization PUT replaces rather than merges, so this reads the
+ current representation and writes it back with the changes applied.
+
+ Args:
+ - organization (OrganizationPage): the organization to update
+ - name (str): new display name, if changing
+ - description (str): new description, if changing
+ - redirect_url (str): new post-login redirect, if changing
+ - domains (list[str]): the complete new domain list, if changing
+ - connection (KeycloakConnection): an existing connection, if any
+ Returns:
+ - OrganizationPage: the updated organization
+ Raises:
+ - OrganizationNotProvisionedError: the organization has no Keycloak record
+ """
+
+ _require_provisioned(organization)
+
+ connection = connection or KeycloakConnection()
+
+ keycloak_org = connection.organizations.get(organization.sso_organization_id)
+ payload = keycloak_org.model_dump(by_alias=True, exclude_none=True)
+
+ if name is not None:
+ payload["name"] = name
+ organization.name = name
+ if description is not None:
+ payload["description"] = description
+ organization.description = description
+ if redirect_url is not None:
+ payload["redirectUrl"] = redirect_url
+ if domains is not None:
+ payload["domains"] = [{"name": domain, "verified": True} for domain in domains]
+
+ connection.organizations.update(organization.sso_organization_id, payload)
+ organization.save()
+
+ return organization
+
+
+def parse_identity_provider_metadata(
+ protocol, *, metadata_url=None, metadata_xml=None, connection=None
+):
+ """
+ Parse IdP metadata without creating anything.
+
+ Keycloak does the parsing. This is the cheapest useful call in the API and
+ the one the eventual wizard leans on hardest: paste a metadata URL, see what
+ Keycloak makes of it, before committing to a resource.
+
+ It is also the endpoint most likely to be abused as an SSRF probe, since the
+ URL form makes Keycloak fetch a caller-supplied address and that egress is
+ confirmed working. Keep it staff-only.
+
+ Args:
+ - protocol (str): "saml" or "oidc"
+ - metadata_url (str): a metadata or discovery URL for Keycloak to fetch
+ - metadata_xml (str): the metadata document, uploaded instead of fetched
+ - connection (KeycloakConnection): an existing connection, if any
+ Returns:
+ - dict: the config map Keycloak parsed out of the metadata
+ """
+
+ connection = connection or KeycloakConnection()
+
+ return import_identity_provider_config(
+ protocol,
+ from_url=metadata_url,
+ metadata=metadata_xml,
+ client=connection.client,
+ )
+
+
+def _attribute_mapper_payload(protocol, alias, user_attribute, source, *, friendly):
+ """Build one attribute-importer mapper representation."""
+
+ if protocol == IDP_PROTOCOL_SAML:
+ source_key = "attribute.friendly.name" if friendly else "attribute.name"
+ config = {
+ source_key: source,
+ "user.attribute": user_attribute,
+ "syncMode": "INHERIT",
+ "attribute.name.format": "ATTRIBUTE_FORMAT_URI",
+ }
+ else:
+ config = {
+ "claim": source,
+ "user.attribute": user_attribute,
+ "syncMode": "INHERIT",
+ }
+
+ return {
+ "name": f"{alias}-{user_attribute}-mapper",
+ "identityProviderAlias": alias,
+ "identityProviderMapper": IDP_ATTRIBUTE_MAPPERS[protocol],
+ "config": config,
+ }
+
+
+def _create_attribute_mappers(
+ connection, alias, protocol, attribute_map, attribute_name_map
+):
+ """Create the IdP's attribute-importer mappers in Keycloak."""
+
+ for user_attribute, source in (attribute_map or {}).items():
+ connection.client.create_returning_id(
+ f"identity-provider/instances/{alias}/mappers",
+ _attribute_mapper_payload(
+ protocol, alias, user_attribute, source, friendly=True
+ ),
+ )
+
+ for user_attribute, source in (attribute_name_map or {}).items():
+ connection.client.create_returning_id(
+ f"identity-provider/instances/{alias}/mappers",
+ _attribute_mapper_payload(
+ protocol, alias, user_attribute, source, friendly=False
+ ),
+ )
+
+
+def create_identity_provider( # noqa: PLR0913
+ organization,
+ *,
+ alias,
+ protocol,
+ display_name="",
+ metadata_url=None,
+ metadata_xml=None,
+ client_id=None,
+ client_secret=None,
+ attribute_map=None,
+ attribute_name_map=None,
+ connection=None,
+):
+ """
+ Create an identity provider for an organization and link the two.
+
+ The IdP starts in `draft`, which is disabled in Keycloak. Nobody can reach
+ it until somebody transitions it to `testing`.
+
+ Same ordering rule as the organization saga: Keycloak first, our row last,
+ compensating delete if our row cannot be written.
+
+ Args:
+ - organization (OrganizationPage): the organization to attach the IdP to
+ - alias (str): the realm-wide IdP alias
+ - protocol (str): "saml" or "oidc"
+ - display_name (str): the IdP's display name
+ - metadata_url (str): SAML metadata or OIDC discovery URL
+ - metadata_xml (str): SAML metadata document, instead of a URL
+ - client_id (str): OIDC client ID
+ - client_secret (str): OIDC client secret
+ - attribute_map (dict): user attribute -> SAML friendly name / OIDC claim
+ - attribute_name_map (dict): user attribute -> SAML attribute name
+ - connection (KeycloakConnection): an existing connection, if any
+ Returns:
+ - OrganizationIdentityProvider: the new record
+ Raises:
+ - AliasCollisionError: the alias is taken here or in the realm
+ - OrganizationNotProvisionedError: the organization has no Keycloak record
+ """
+
+ # Before anything is written. Without it the IdP is created in Keycloak and
+ # only the org<->IdP link fails, so the compensation deletes an IdP we
+ # should never have made - a wasted round trip reported as 502.
+ _require_provisioned(organization)
+
+ connection = connection or KeycloakConnection()
+
+ if OrganizationIdentityProvider.objects.filter(alias=alias).exists():
+ msg = f"An identity provider with alias '{alias}' already exists."
+ raise AliasCollisionError(msg)
+
+ if alias.lower() in realm_identity_provider_aliases(connection):
+ msg = (
+ f"The alias '{alias}' is already in use by an identity provider in "
+ "the Keycloak realm."
+ )
+ raise AliasCollisionError(msg)
+
+ # The artifact is what Keycloak parsed out of the partner's metadata, and
+ # nothing else: the OIDC client secret we add below goes to Keycloak but is
+ # never persisted here, because this field is served back over the API.
+ metadata_artifact = parse_identity_provider_metadata(
+ protocol,
+ metadata_url=metadata_url,
+ metadata_xml=metadata_xml,
+ connection=connection,
+ )
+ fetched_at = now_in_utc()
+
+ config = dict(metadata_artifact)
+ if protocol == IDP_PROTOCOL_OIDC:
+ config.update({"clientId": client_id, "clientSecret": client_secret})
+ elif metadata_url:
+ # Let Keycloak re-read the descriptor itself, matching what the Pulumi
+ # SAML resources set.
+ config.update(
+ {
+ "metadataDescriptorUrl": metadata_url,
+ "useMetadataDescriptorUrl": "true",
+ }
+ )
+
+ draft_flags = IDP_STATE_KEYCLOAK_FLAGS[IDP_STATE_DRAFT]
+ internal_id = connection.identity_providers.create(
+ {
+ "alias": alias,
+ "displayName": display_name,
+ "providerId": protocol,
+ "enabled": draft_flags["enabled"],
+ "hideOnLogin": draft_flags["hideOnLogin"],
+ "config": config,
+ }
+ )
+
+ try:
+ _create_attribute_mappers(
+ connection, alias, protocol, attribute_map, attribute_name_map
+ )
+ connection.organizations.associate(
+ ORG_IDP_ASSOCIATION, organization.sso_organization_id, alias
+ )
+ identity_provider = OrganizationIdentityProvider.objects.create(
+ organization=organization,
+ alias=alias,
+ protocol=protocol,
+ display_name=display_name,
+ internal_id=internal_id or "",
+ metadata_source=metadata_url or metadata_xml or "",
+ metadata_artifact=metadata_artifact,
+ metadata_fetched_at=fetched_at,
+ )
+ except Exception:
+ connection.identity_providers.delete(alias)
+ raise
+
+ return identity_provider
+
+
+def refresh_identity_provider_metadata(identity_provider, *, connection=None):
+ """
+ Re-fetch the IdP's metadata and store what came back.
+
+ An explicit operation, never a side effect of an unrelated change. On
+ failure the stored artifact is left exactly as it was - that is the whole
+ reason it is stored.
+
+ Args:
+ - identity_provider (OrganizationIdentityProvider): the IdP to refresh
+ - connection (KeycloakConnection): an existing connection, if any
+ Returns:
+ - OrganizationIdentityProvider: the refreshed record
+ """
+
+ connection = connection or KeycloakConnection()
+ source = identity_provider.metadata_source
+
+ config = parse_identity_provider_metadata(
+ identity_provider.protocol,
+ metadata_url=source if not source.lstrip().startswith("<") else None,
+ metadata_xml=source if source.lstrip().startswith("<") else None,
+ connection=connection,
+ )
+
+ keycloak_idp = connection.identity_providers.get(identity_provider.alias)
+ payload = keycloak_idp.model_dump(by_alias=True, exclude_none=True)
+ payload["config"] = {**(payload.get("config") or {}), **config}
+ connection.identity_providers.update(identity_provider.alias, payload)
+
+ identity_provider.metadata_artifact = config
+ identity_provider.metadata_fetched_at = now_in_utc()
+ identity_provider.save()
+
+ return identity_provider
+
+
+def transition_identity_provider(identity_provider, state, *, connection=None):
+ """
+ Move an identity provider to a new lifecycle state.
+
+ The only thing that moves the lifecycle, and it writes Keycloak's own
+ enabled/hideOnLogin in the same operation so our record and the realm cannot
+ drift. Keycloak goes first: our row lagging the realm is recoverable, a
+ partner integration disabled without our knowing is not.
+
+ Args:
+ - identity_provider (OrganizationIdentityProvider): the IdP to move
+ - state (str): the lifecycle state to move to
+ - connection (KeycloakConnection): an existing connection, if any
+ Returns:
+ - OrganizationIdentityProvider: the updated record
+ Raises:
+ - InvalidLifecycleTransitionError: the transition is not allowed
+ """
+
+ current = identity_provider.lifecycle_state
+
+ if state not in IDP_ALLOWED_TRANSITIONS[current]:
+ msg = (
+ f"Cannot move identity provider '{identity_provider.alias}' from "
+ f"'{current}' to '{state}'."
+ )
+ raise InvalidLifecycleTransitionError(msg)
+
+ connection = connection or KeycloakConnection()
+ flags = IDP_STATE_KEYCLOAK_FLAGS[state]
+
+ keycloak_idp = connection.identity_providers.get(identity_provider.alias)
+ payload = keycloak_idp.model_dump(by_alias=True, exclude_none=True)
+ payload.update(flags)
+ connection.identity_providers.update(identity_provider.alias, payload)
+
+ identity_provider.lifecycle_state = state
+ identity_provider.save()
+
+ return identity_provider
+
+
+def delete_identity_provider(identity_provider, *, connection=None):
+ """
+ Unlink and delete an identity provider.
+
+ Args:
+ - identity_provider (OrganizationIdentityProvider): the IdP to delete
+ - connection (KeycloakConnection): an existing connection, if any
+ Raises:
+ - OrganizationNotProvisionedError: the organization has no Keycloak record
+ """
+
+ # Creation requires a provisioned organization, so this looks unreachable -
+ # but sso_organization_id is an editable panel on OrganizationPage
+ # (content_panels) and is nullable, so staff can blank it in the Wagtail
+ # admin after the IdP exists. Without the guard the unlink 502s and the IdP
+ # can never be deleted, with a message blaming Keycloak.
+ _require_provisioned(identity_provider.organization)
+
+ connection = connection or KeycloakConnection()
+
+ connection.organizations.disassociate(
+ ORG_IDP_ASSOCIATION,
+ identity_provider.organization.sso_organization_id,
+ identity_provider.alias,
+ )
+ connection.identity_providers.delete(identity_provider.alias)
+ identity_provider.delete()
diff --git a/b2b/provisioning_test.py b/b2b/provisioning_test.py
new file mode 100644
index 0000000000..5d56f8945e
--- /dev/null
+++ b/b2b/provisioning_test.py
@@ -0,0 +1,653 @@
+"""Tests for the staff-only B2B provisioning API's business logic."""
+
+import faker
+import pytest
+from django.core.exceptions import ImproperlyConfigured
+
+from b2b.constants import (
+ IDP_ALLOWED_TRANSITIONS,
+ IDP_LIFECYCLE_CHOICES,
+ IDP_PROTOCOL_OIDC,
+ IDP_PROTOCOL_SAML,
+ IDP_STATE_ACTIVE,
+ IDP_STATE_DISABLED,
+ IDP_STATE_DRAFT,
+ IDP_STATE_TESTING,
+ ONBOARDING_STATE_ORG_CREATED,
+)
+from b2b.exceptions import (
+ AliasCollisionError,
+ InvalidLifecycleTransitionError,
+ OrganizationNotProvisionedError,
+ OrphanedKeycloakOrganizationError,
+)
+from b2b.factories import OrganizationIndexPageFactory, OrganizationPageFactory
+from b2b.keycloak_admin_dataclasses import (
+ IdentityProviderRepresentation,
+ OrganizationRepresentation,
+)
+from b2b.models import (
+ OrganizationIdentityProvider,
+ OrganizationOnboarding,
+ OrganizationPage,
+)
+from b2b.provisioning import (
+ create_identity_provider,
+ create_organization,
+ delete_identity_provider,
+ refresh_identity_provider_metadata,
+ transition_identity_provider,
+ update_organization,
+)
+
+pytestmark = [pytest.mark.django_db]
+FAKE = faker.Faker()
+
+PARSED_METADATA = {
+ "singleSignOnServiceUrl": "https://idp.example.edu/sso",
+ "idpEntityId": "https://idp.example.edu/entity",
+}
+
+
+@pytest.fixture(autouse=True)
+def organization_index():
+ """
+ The index page new OrganizationPages are added under.
+
+ create_organization adds under the existing index rather than calling
+ ensure_b2b_organization_index, which would also move every existing
+ organization page as a side effect of creating one.
+ """
+
+ return OrganizationIndexPageFactory.create()
+
+
+@pytest.fixture
+def connection(mocker):
+ """
+ A KeycloakConnection whose models are mocks.
+
+ Every provisioning function takes a connection, so the tests never touch
+ the network and never need the client bootstrapped.
+ """
+
+ fake_connection = mocker.Mock()
+ fake_connection.organizations.list_all.return_value = []
+ fake_connection.identity_providers.list_all.return_value = []
+ fake_connection.organizations.create.return_value = str(FAKE.uuid4())
+ fake_connection.identity_providers.create.return_value = str(FAKE.uuid4())
+ return fake_connection
+
+
+@pytest.fixture
+def mocked_import_config(mocker):
+ """Stand in for Keycloak's metadata parsing."""
+
+ return mocker.patch(
+ "b2b.provisioning.import_identity_provider_config",
+ return_value=dict(PARSED_METADATA),
+ )
+
+
+def _organization_kwargs(**overrides):
+ return {
+ "name": "Example University",
+ "org_key": "EXAMPLEU",
+ "domains": ["example.edu"],
+ "description": "",
+ "redirect_url": "https://learn.mit.edu/dashboard/organization/exampleu",
+ **overrides,
+ }
+
+
+def test_create_organization_writes_both_systems(connection):
+ """The happy path: Keycloak first, then our records, alias from org_key."""
+
+ organization = create_organization(connection=connection, **_organization_kwargs())
+
+ keycloak_payload = connection.organizations.create.call_args.args[0]
+ assert keycloak_payload["alias"] == "EXAMPLEU"
+ assert keycloak_payload["domains"] == [{"name": "example.edu", "verified": True}]
+
+ assert organization.pk is not None
+ assert str(organization.sso_organization_id) == str(
+ connection.organizations.create.return_value
+ )
+ assert organization.onboarding.state == ONBOARDING_STATE_ORG_CREATED
+
+
+def test_create_organization_rejects_a_duplicate_org_key(connection):
+ """An org_key already used here is a 409, not a second organization."""
+
+ existing = OrganizationPageFactory.create()
+
+ with pytest.raises(AliasCollisionError):
+ create_organization(
+ connection=connection, **_organization_kwargs(org_key=existing.org_key)
+ )
+
+ connection.organizations.create.assert_not_called()
+
+
+def test_create_organization_rejects_an_alias_taken_in_the_realm(connection):
+ """
+ An alias free in our tables can still be taken in the realm.
+
+ The realm is shared with the organizations Pulumi still declares, so
+ creating this anyway would break the next pulumi up that declares the
+ same name.
+ """
+
+ connection.organizations.list_all.return_value = [
+ OrganizationRepresentation(id=str(FAKE.uuid4()), alias="exampleu")
+ ]
+
+ with pytest.raises(AliasCollisionError):
+ create_organization(connection=connection, **_organization_kwargs())
+
+ connection.organizations.create.assert_not_called()
+
+
+def test_create_organization_compensates_a_failed_local_write(connection, mocker):
+ """
+ A failed MITx Online write deletes the Keycloak organization it made.
+
+ A Keycloak org with no MITx Online counterpart is exactly the orphan shape
+ the compensation exists to prevent.
+ """
+
+ mocker.patch(
+ "b2b.provisioning.OrganizationOnboarding.objects.create",
+ side_effect=ValueError("no"),
+ )
+
+ with pytest.raises(ValueError, match="no"):
+ create_organization(connection=connection, **_organization_kwargs())
+
+ connection.organizations.delete.assert_called_once_with(
+ connection.organizations.create.return_value
+ )
+ assert not OrganizationPage.objects.filter(org_key="EXAMPLEU").exists()
+
+
+def test_create_organization_reports_an_orphan_when_compensation_fails(
+ connection, mocker
+):
+ """
+ When the compensating delete also fails, say so and stop.
+
+ Retrying against a system that just failed is worse than logging the orphan
+ and letting reconcile_keycloak_orgs adopt it on its next run.
+ """
+
+ mocker.patch(
+ "b2b.provisioning.OrganizationOnboarding.objects.create",
+ side_effect=ValueError("no"),
+ )
+ connection.organizations.delete.side_effect = ValueError("also no")
+
+ with pytest.raises(OrphanedKeycloakOrganizationError) as failure:
+ create_organization(connection=connection, **_organization_kwargs())
+
+ assert str(connection.organizations.create.return_value) in str(failure.value)
+ assert not OrganizationPage.objects.filter(org_key="EXAMPLEU").exists()
+
+
+def test_create_organization_never_leaves_a_null_sso_organization_id(connection):
+ """
+ An OrganizationPage without sso_organization_id is silently broken.
+
+ attach_user() returns False without doing anything for those, so every
+ membership write is a no-op. This is the hq#10552 shape.
+ """
+
+ create_organization(connection=connection, **_organization_kwargs())
+
+ assert not OrganizationPage.objects.filter(
+ sso_organization_id__isnull=True
+ ).exists()
+
+
+def test_create_organization_refuses_to_write_a_row_without_an_id(connection):
+ """
+ An unresolvable ID stops the saga rather than writing a null.
+
+ Keycloak sent no Location and the alias lookup found nothing, so there is
+ no ID to store and nothing to compensate with either - the delete needs the
+ ID we do not have. Writing the row anyway would mint exactly the
+ hq#10552 orphan this saga exists to prevent.
+ """
+
+ connection.organizations.create.return_value = None
+ connection.organizations.list_all.return_value = []
+
+ with pytest.raises(OrphanedKeycloakOrganizationError):
+ create_organization(connection=connection, **_organization_kwargs())
+
+ assert not OrganizationPage.objects.filter(org_key="EXAMPLEU").exists()
+
+
+def test_create_organization_looks_up_the_id_when_keycloak_sends_no_location(
+ connection,
+):
+ """Keycloak's create answers 201 with an empty body; fall back to the alias."""
+
+ known_id = str(FAKE.uuid4())
+ connection.organizations.create.return_value = None
+ connection.organizations.list_all.return_value = [
+ OrganizationRepresentation(id=known_id, alias="EXAMPLEU")
+ ]
+
+ # The pre-flight collision check runs against the same list, so start from
+ # an alias that is not yet in the realm and add it on the second call.
+ connection.organizations.list_all.side_effect = [
+ [],
+ [OrganizationRepresentation(id=known_id, alias="EXAMPLEU")],
+ ]
+
+ organization = create_organization(connection=connection, **_organization_kwargs())
+
+ assert str(organization.sso_organization_id) == known_id
+
+
+def test_update_organization_replaces_the_keycloak_representation(connection):
+ """Keycloak's organization PUT replaces, so read-modify-write."""
+
+ organization = OrganizationPageFactory.create()
+ connection.organizations.get.return_value = OrganizationRepresentation(
+ id=str(organization.sso_organization_id),
+ name=organization.name,
+ alias=organization.org_key,
+ )
+
+ update_organization(
+ organization,
+ connection=connection,
+ name="Renamed",
+ domains=["renamed.edu"],
+ )
+
+ _, payload = connection.organizations.update.call_args.args
+ assert payload["name"] == "Renamed"
+ assert payload["domains"] == [{"name": "renamed.edu", "verified": True}]
+ assert payload["alias"] == organization.org_key
+
+ organization.refresh_from_db()
+ assert organization.name == "Renamed"
+
+
+def test_update_organization_refuses_an_unprovisioned_organization(connection):
+ """
+ An org with no Keycloak record cannot be updated in Keycloak.
+
+ Roughly 24 production organizations are in this state (hq#10552). Without
+ the check the GET goes to organizations/None and the operator is told the
+ Keycloak API failed, which is a 502 they would retry forever.
+ """
+
+ organization = OrganizationPageFactory.create(sso_organization_id=None)
+
+ with pytest.raises(OrganizationNotProvisionedError):
+ update_organization(organization, connection=connection, name="Renamed")
+
+ connection.organizations.get.assert_not_called()
+ connection.organizations.update.assert_not_called()
+
+
+def test_create_identity_provider_refuses_an_unprovisioned_organization(
+ connection, mocked_import_config
+):
+ """
+ Same guard as update_organization, checked before anything is written.
+
+ Without it the IdP is created in Keycloak and only the org link fails, so
+ the compensation deletes an IdP that should never have been made.
+ """
+
+ organization = OrganizationPageFactory.create(sso_organization_id=None)
+
+ with pytest.raises(OrganizationNotProvisionedError):
+ create_identity_provider(
+ organization,
+ connection=connection,
+ alias="exampleu",
+ protocol=IDP_PROTOCOL_SAML,
+ metadata_url="https://idp.example.edu/metadata.xml",
+ attribute_map={"email": "E-Mail Address"},
+ )
+
+ mocked_import_config.assert_not_called()
+ connection.identity_providers.create.assert_not_called()
+ connection.identity_providers.delete.assert_not_called()
+
+
+def test_create_identity_provider_starts_in_draft(connection, mocked_import_config):
+ """A new IdP is disabled in Keycloak until somebody moves it to testing."""
+
+ organization = OrganizationPageFactory.create()
+ metadata_url = "https://idp.example.edu/metadata.xml"
+
+ identity_provider = create_identity_provider(
+ organization,
+ connection=connection,
+ alias="exampleu",
+ protocol=IDP_PROTOCOL_SAML,
+ display_name="Example University",
+ metadata_url=metadata_url,
+ attribute_map={"email": "E-Mail Address"},
+ )
+
+ payload = connection.identity_providers.create.call_args.args[0]
+ assert payload["enabled"] is False
+ assert payload["hideOnLogin"] is True
+ assert payload["providerId"] == IDP_PROTOCOL_SAML
+ assert payload["config"]["metadataDescriptorUrl"] == metadata_url
+
+ connection.organizations.associate.assert_called_once_with(
+ "identity-providers", organization.sso_organization_id, "exampleu"
+ )
+
+ assert identity_provider.lifecycle_state == IDP_STATE_DRAFT
+ assert identity_provider.metadata_artifact == PARSED_METADATA
+ assert identity_provider.metadata_source == metadata_url
+ mocked_import_config.assert_called_once()
+
+
+def test_create_identity_provider_creates_attribute_mappers(
+ connection,
+ mocked_import_config,
+):
+ """Attribute mappers are created the same shape ol-infrastructure makes."""
+
+ organization = OrganizationPageFactory.create()
+
+ create_identity_provider(
+ organization,
+ connection=connection,
+ alias="exampleu",
+ protocol=IDP_PROTOCOL_SAML,
+ metadata_url="https://idp.example.edu/metadata.xml",
+ attribute_map={"email": "E-Mail Address"},
+ )
+
+ endpoint, payload = connection.client.create_returning_id.call_args.args
+ assert endpoint == "identity-provider/instances/exampleu/mappers"
+ assert payload["identityProviderMapper"] == "saml-user-attribute-idp-mapper"
+ assert payload["config"]["attribute.friendly.name"] == "E-Mail Address"
+ assert payload["config"]["user.attribute"] == "email"
+
+
+def test_create_identity_provider_keeps_the_client_secret_out_of_the_artifact(
+ connection,
+ mocked_import_config,
+):
+ """
+ The OIDC client secret reaches Keycloak but is never persisted.
+
+ metadata_artifact is served back over the API, so it holds what Keycloak
+ parsed out of the partner's metadata and nothing else.
+ """
+
+ organization = OrganizationPageFactory.create()
+ secret = FAKE.password()
+
+ identity_provider = create_identity_provider(
+ organization,
+ connection=connection,
+ alias="exampleu-oidc",
+ protocol=IDP_PROTOCOL_OIDC,
+ metadata_url="https://idp.example.edu/.well-known/openid-configuration",
+ client_id="mitxonline",
+ client_secret=secret,
+ attribute_map={"email": "email"},
+ )
+
+ payload = connection.identity_providers.create.call_args.args[0]
+ assert payload["config"]["clientSecret"] == secret
+ assert secret not in str(identity_provider.metadata_artifact)
+
+
+def test_create_identity_provider_rejects_a_realm_alias_collision(
+ connection, mocked_import_config
+):
+ """IdP aliases are realm-wide, so a collision across customers is real."""
+
+ organization = OrganizationPageFactory.create()
+ connection.identity_providers.list_all.return_value = [
+ IdentityProviderRepresentation(alias="exampleu")
+ ]
+
+ with pytest.raises(AliasCollisionError):
+ create_identity_provider(
+ organization,
+ connection=connection,
+ alias="exampleu",
+ protocol=IDP_PROTOCOL_SAML,
+ metadata_url="https://idp.example.edu/metadata.xml",
+ attribute_map={"email": "E-Mail Address"},
+ )
+
+ connection.identity_providers.create.assert_not_called()
+ mocked_import_config.assert_not_called()
+
+
+def test_create_identity_provider_compensates_a_failed_local_write(
+ connection,
+ mocked_import_config,
+):
+ """A failed link or row write deletes the IdP Keycloak just made."""
+
+ organization = OrganizationPageFactory.create()
+ connection.organizations.associate.side_effect = ValueError("no")
+
+ with pytest.raises(ValueError, match="no"):
+ create_identity_provider(
+ organization,
+ connection=connection,
+ alias="exampleu",
+ protocol=IDP_PROTOCOL_SAML,
+ metadata_url="https://idp.example.edu/metadata.xml",
+ attribute_map={"email": "E-Mail Address"},
+ )
+
+ connection.identity_providers.delete.assert_called_once_with("exampleu")
+ assert not OrganizationIdentityProvider.objects.filter(alias="exampleu").exists()
+
+
+def _identity_provider(organization, state=IDP_STATE_DRAFT):
+ return OrganizationIdentityProvider.objects.create(
+ organization=organization,
+ alias="exampleu",
+ protocol=IDP_PROTOCOL_SAML,
+ lifecycle_state=state,
+ metadata_source="https://idp.example.edu/metadata.xml",
+ metadata_artifact=dict(PARSED_METADATA),
+ )
+
+
+def test_transition_writes_keycloak_and_our_row_together(connection):
+ """The lifecycle cannot drift from the realm because one call does both."""
+
+ identity_provider = _identity_provider(OrganizationPageFactory.create())
+ connection.identity_providers.get.return_value = IdentityProviderRepresentation(
+ alias="exampleu", enabled=False, hide_on_login=True
+ )
+
+ transition_identity_provider(
+ identity_provider, IDP_STATE_TESTING, connection=connection
+ )
+
+ _, payload = connection.identity_providers.update.call_args.args
+ assert payload["enabled"] is True
+ assert payload["hideOnLogin"] is True
+
+ identity_provider.refresh_from_db()
+ assert identity_provider.lifecycle_state == IDP_STATE_TESTING
+
+
+def test_transition_refuses_to_skip_testing(connection):
+ """An IdP goes live only after somebody has logged in through it."""
+
+ identity_provider = _identity_provider(OrganizationPageFactory.create())
+
+ with pytest.raises(InvalidLifecycleTransitionError):
+ transition_identity_provider(
+ identity_provider, IDP_STATE_ACTIVE, connection=connection
+ )
+
+ connection.identity_providers.update.assert_not_called()
+ identity_provider.refresh_from_db()
+ assert identity_provider.lifecycle_state == IDP_STATE_DRAFT
+
+
+def test_transition_allows_re_enabling_a_disabled_provider(connection):
+ """An IdP that already went through testing can be turned back on."""
+
+ identity_provider = _identity_provider(
+ OrganizationPageFactory.create(), state=IDP_STATE_DISABLED
+ )
+ connection.identity_providers.get.return_value = IdentityProviderRepresentation(
+ alias="exampleu", enabled=False, hide_on_login=True
+ )
+
+ transition_identity_provider(
+ identity_provider, IDP_STATE_ACTIVE, connection=connection
+ )
+
+ identity_provider.refresh_from_db()
+ assert identity_provider.lifecycle_state == IDP_STATE_ACTIVE
+
+
+def test_refresh_metadata_leaves_the_artifact_alone_when_the_fetch_fails(
+ connection, mocker
+):
+ """
+ An unreachable partner endpoint never clears the stored artifact.
+
+ Storing it is what stops a partner's metadata going away from destroying
+ config, so a failed refresh must not undo that.
+ """
+
+ identity_provider = _identity_provider(OrganizationPageFactory.create())
+ mocker.patch(
+ "b2b.provisioning.import_identity_provider_config",
+ side_effect=ValueError("unreachable"),
+ )
+
+ with pytest.raises(ValueError, match="unreachable"):
+ refresh_identity_provider_metadata(identity_provider, connection=connection)
+
+ identity_provider.refresh_from_db()
+ assert identity_provider.metadata_artifact == PARSED_METADATA
+ connection.identity_providers.update.assert_not_called()
+
+
+def test_refresh_metadata_stores_what_came_back(connection, mocker):
+ """A successful refresh replaces the artifact and stamps when it happened."""
+
+ identity_provider = _identity_provider(OrganizationPageFactory.create())
+ refreshed = {"idpEntityId": "https://idp.example.edu/rotated"}
+ mocker.patch(
+ "b2b.provisioning.import_identity_provider_config", return_value=refreshed
+ )
+ connection.identity_providers.get.return_value = IdentityProviderRepresentation(
+ alias="exampleu", enabled=True, config=dict(PARSED_METADATA)
+ )
+
+ refresh_identity_provider_metadata(identity_provider, connection=connection)
+
+ identity_provider.refresh_from_db()
+ assert identity_provider.metadata_artifact == refreshed
+ assert identity_provider.metadata_fetched_at is not None
+
+
+def test_delete_identity_provider_unlinks_before_deleting(connection):
+ """Unlink from the organization, then remove the instance, then our row."""
+
+ organization = OrganizationPageFactory.create()
+ identity_provider = _identity_provider(organization)
+
+ delete_identity_provider(identity_provider, connection=connection)
+
+ connection.organizations.disassociate.assert_called_once_with(
+ "identity-providers", organization.sso_organization_id, "exampleu"
+ )
+ connection.identity_providers.delete.assert_called_once_with("exampleu")
+ assert not OrganizationIdentityProvider.objects.filter(alias="exampleu").exists()
+
+
+def test_delete_identity_provider_refuses_an_unprovisioned_organization(connection):
+ """
+ Reachable because sso_organization_id is editable in the Wagtail admin.
+
+ Creation requires a provisioned organization, but staff can blank the field
+ afterwards. Without the guard the unlink 502s and the IdP can never be
+ deleted.
+ """
+
+ organization = OrganizationPageFactory.create()
+ identity_provider = _identity_provider(organization)
+ organization.sso_organization_id = None
+ organization.save()
+ identity_provider.refresh_from_db()
+
+ with pytest.raises(OrganizationNotProvisionedError):
+ delete_identity_provider(identity_provider, connection=connection)
+
+ connection.organizations.disassociate.assert_not_called()
+ connection.identity_providers.delete.assert_not_called()
+ assert OrganizationIdentityProvider.objects.filter(alias="exampleu").exists()
+
+
+def test_onboarding_set_state_stamps_the_change():
+ """state_changed_at tracks state changes, not every save."""
+
+ organization = OrganizationPageFactory.create()
+ onboarding = OrganizationOnboarding.objects.create(organization=organization)
+ original = onboarding.state_changed_at
+
+ onboarding.set_state(ONBOARDING_STATE_ORG_CREATED, notes="created by hand")
+
+ onboarding.refresh_from_db()
+ assert onboarding.state == ONBOARDING_STATE_ORG_CREATED
+ assert onboarding.notes == "created by hand"
+ assert onboarding.state_changed_at > original
+
+
+def test_every_lifecycle_state_has_transitions():
+ """
+ IDP_ALLOWED_TRANSITIONS covers exactly the lifecycle choices.
+
+ transition_identity_provider indexes the dict by the IdP's current state.
+ That is a KeyError waiting to happen only if the two constants drift apart,
+ so guard the drift rather than the lookup: a runtime .get() would turn
+ "somebody added a fifth state and forgot the transitions" into a silently
+ untransitionable IdP, where this fails the build.
+ """
+
+ assert set(IDP_ALLOWED_TRANSITIONS) == {state for state, _ in IDP_LIFECYCLE_CHOICES}
+
+ valid_states = set(IDP_ALLOWED_TRANSITIONS)
+ for state, destinations in IDP_ALLOWED_TRANSITIONS.items():
+ assert set(destinations) <= valid_states, state
+ assert state not in destinations, state
+
+
+def test_create_organization_checks_the_index_page_before_touching_keycloak(
+ connection, organization_index
+):
+ """
+ A missing index page fails before the Keycloak write, not after.
+
+ It is a precondition we can check for nothing; checking it after the
+ irreversible external write would mean compensating for a failure we could
+ have seen coming.
+ """
+
+ organization_index.delete()
+
+ with pytest.raises(ImproperlyConfigured):
+ create_organization(connection=connection, **_organization_kwargs())
+
+ connection.organizations.create.assert_not_called()
+ connection.organizations.delete.assert_not_called()
diff --git a/b2b/serializers/v0/provisioning.py b/b2b/serializers/v0/provisioning.py
new file mode 100644
index 0000000000..8dfdf84ee6
--- /dev/null
+++ b/b2b/serializers/v0/provisioning.py
@@ -0,0 +1,239 @@
+"""Serializers for the staff-only B2B provisioning API (v0)."""
+
+from rest_framework import serializers
+
+from b2b.constants import (
+ IDP_LIFECYCLE_CHOICES,
+ IDP_PROTOCOL_CHOICES,
+ IDP_PROTOCOL_OIDC,
+ IDP_PROTOCOL_SAML,
+ ONBOARDING_STATE_CHOICES,
+)
+from b2b.models import (
+ OrganizationIdentityProvider,
+ OrganizationOnboarding,
+ OrganizationPage,
+)
+
+
+class OrganizationOnboardingSerializer(serializers.ModelSerializer):
+ """Where an organization is in the onboarding sequence."""
+
+ class Meta:
+ model = OrganizationOnboarding
+ fields = ["state", "state_changed_at", "notes"]
+ read_only_fields = ["state_changed_at"]
+
+
+class OrganizationIdentityProviderSerializer(serializers.ModelSerializer):
+ """
+ An identity provider we provisioned for an organization.
+
+ metadata_artifact is what Keycloak parsed out of the partner's metadata.
+ Credentials are never written to it, so this is safe to serve.
+ """
+
+ class Meta:
+ model = OrganizationIdentityProvider
+ fields = [
+ "id",
+ "alias",
+ "protocol",
+ "display_name",
+ "lifecycle_state",
+ "internal_id",
+ "metadata_source",
+ "metadata_artifact",
+ "metadata_fetched_at",
+ "created_on",
+ "updated_on",
+ ]
+ read_only_fields = fields
+
+
+class ProvisionedOrganizationSerializer(serializers.ModelSerializer):
+ """
+ An organization as the provisioning API sees it.
+
+ `domains` and `redirect_url` live only in Keycloak, so they are populated
+ from the representation the view fetched rather than from our database.
+ """
+
+ onboarding = OrganizationOnboardingSerializer(read_only=True)
+ identity_providers = OrganizationIdentityProviderSerializer(
+ many=True, read_only=True
+ )
+ domains = serializers.SerializerMethodField()
+ redirect_url = serializers.SerializerMethodField()
+
+ def get_domains(self, instance) -> list[str] | None:
+ """Return the organization's Keycloak domains, if they were fetched."""
+
+ keycloak_org = getattr(instance, "keycloak_organization", None)
+ if keycloak_org is None:
+ return None
+ return [domain.name for domain in keycloak_org.domains or []]
+
+ def get_redirect_url(self, instance) -> str | None:
+ """Return the organization's Keycloak redirect URL, if it was fetched."""
+
+ keycloak_org = getattr(instance, "keycloak_organization", None)
+ return keycloak_org.redirect_url if keycloak_org else None
+
+ class Meta:
+ model = OrganizationPage
+ fields = [
+ "id",
+ "name",
+ "org_key",
+ "org_key_prefix",
+ "description",
+ "slug",
+ "sso_organization_id",
+ "domains",
+ "redirect_url",
+ "onboarding",
+ "identity_providers",
+ ]
+ read_only_fields = fields
+
+
+class CreateOrganizationSerializer(serializers.Serializer):
+ """Request body for provisioning a new organization."""
+
+ name = serializers.CharField(max_length=255)
+ org_key = serializers.CharField(max_length=30)
+ org_key_prefix = serializers.CharField(
+ max_length=30, required=False, allow_blank=True
+ )
+ domains = serializers.ListField(
+ child=serializers.CharField(), required=False, default=list
+ )
+ description = serializers.CharField(required=False, allow_blank=True, default="")
+ redirect_url = serializers.CharField(required=False, allow_blank=True, default="")
+
+
+class UpdateOrganizationSerializer(serializers.Serializer):
+ """
+ Request body for updating an organization.
+
+ org_key is rejected rather than silently ignored. It is in every B2B
+ courseware ID via create_contract_run_key, so a caller who thinks they
+ changed it and did not is worse off than one who got an error.
+ """
+
+ name = serializers.CharField(max_length=255, required=False)
+ description = serializers.CharField(required=False, allow_blank=True)
+ redirect_url = serializers.CharField(required=False, allow_blank=True)
+ domains = serializers.ListField(child=serializers.CharField(), required=False)
+
+ def validate(self, attrs):
+ """Reject any attempt to change the immutable org key."""
+
+ if "org_key" in self.initial_data:
+ msg = (
+ "org_key is immutable: it is part of every courseware ID for "
+ "this organization."
+ )
+ raise serializers.ValidationError({"org_key": msg})
+ return attrs
+
+
+class ParseMetadataSerializer(serializers.Serializer):
+ """
+ Request body for parsing IdP metadata without creating anything.
+
+ Exactly one of metadata_url or metadata_xml.
+ """
+
+ protocol = serializers.ChoiceField(choices=IDP_PROTOCOL_CHOICES)
+ metadata_url = serializers.URLField(required=False, allow_blank=True)
+ metadata_xml = serializers.CharField(required=False, allow_blank=True)
+
+ def validate(self, attrs):
+ """Require exactly one metadata source."""
+
+ has_url = bool(attrs.get("metadata_url"))
+ has_xml = bool(attrs.get("metadata_xml"))
+
+ if has_url == has_xml:
+ msg = "Supply exactly one of metadata_url or metadata_xml."
+ raise serializers.ValidationError(msg)
+
+ return attrs
+
+
+class CreateIdentityProviderSerializer(serializers.Serializer):
+ """Request body for provisioning an identity provider."""
+
+ alias = serializers.CharField(max_length=255)
+ protocol = serializers.ChoiceField(choices=IDP_PROTOCOL_CHOICES)
+ display_name = serializers.CharField(
+ max_length=255, required=False, allow_blank=True, default=""
+ )
+ metadata_url = serializers.URLField(required=False, allow_blank=True)
+ metadata_xml = serializers.CharField(required=False, allow_blank=True)
+ discovery_url = serializers.URLField(required=False, allow_blank=True)
+ client_id = serializers.CharField(required=False, allow_blank=True)
+ client_secret = serializers.CharField(required=False, allow_blank=True)
+ attribute_map = serializers.DictField(
+ child=serializers.CharField(), required=False, default=dict
+ )
+ attribute_name_map = serializers.DictField(
+ child=serializers.CharField(), required=False, default=dict
+ )
+
+ def validate(self, attrs):
+ """Check that the protocol has the inputs it needs, and only those."""
+
+ protocol = attrs["protocol"]
+
+ if protocol == IDP_PROTOCOL_OIDC:
+ if not attrs.get("discovery_url"):
+ msg = "discovery_url is required for an OIDC identity provider."
+ raise serializers.ValidationError({"discovery_url": msg})
+ if not attrs.get("client_id") or not attrs.get("client_secret"):
+ msg = (
+ "client_id and client_secret are required for an OIDC "
+ "identity provider."
+ )
+ raise serializers.ValidationError(msg)
+ # The saga takes one metadata source regardless of protocol; for
+ # OIDC that source is the discovery document.
+ attrs["metadata_url"] = attrs["discovery_url"]
+ elif protocol == IDP_PROTOCOL_SAML:
+ if bool(attrs.get("metadata_url")) == bool(attrs.get("metadata_xml")):
+ msg = (
+ "Supply exactly one of metadata_url or metadata_xml for a "
+ "SAML identity provider."
+ )
+ raise serializers.ValidationError(msg)
+
+ # SAML only. Nothing maps a SAML assertion's attributes onto the
+ # user unless we say so, so a SAML IdP with no mappers brokers
+ # users with no email or name. ol-infrastructure treats them as
+ # required too: onboard_saml_org falls back to scanning the
+ # metadata for common friendly names when none are given, while
+ # onboard_oidc_org creates no mappers at all - the OIDC IdPs
+ # Pulumi runs in production today have none.
+ if not attrs.get("attribute_map") and not attrs.get("attribute_name_map"):
+ msg = (
+ "Supply attribute_map or attribute_name_map for a SAML "
+ "identity provider."
+ )
+ raise serializers.ValidationError(msg)
+
+ return attrs
+
+
+class IdentityProviderTransitionSerializer(serializers.Serializer):
+ """Request body for moving an identity provider's lifecycle state."""
+
+ state = serializers.ChoiceField(choices=IDP_LIFECYCLE_CHOICES)
+
+
+class SetOnboardingStateSerializer(serializers.Serializer):
+ """Request body for recording an organization's onboarding state."""
+
+ state = serializers.ChoiceField(choices=ONBOARDING_STATE_CHOICES)
+ notes = serializers.CharField(required=False, allow_blank=True)
diff --git a/b2b/views/v0/provisioning.py b/b2b/views/v0/provisioning.py
new file mode 100644
index 0000000000..07c16f2eb4
--- /dev/null
+++ b/b2b/views/v0/provisioning.py
@@ -0,0 +1,401 @@
+"""
+Staff-only B2B provisioning API (v0).
+
+The runtime replacement for onboarding a customer by opening a Pulumi PR
+against ol-infrastructure's olapps.py, and for waiting up to
+KEYCLOAK_ORG_SYNC_FREQUENCY seconds for the organization to appear in MITx
+Online. See docs/source/b2b/provisioning_api.md.
+
+Everything here is staff-write via IsAdminOrReadOnly. b2b.permissions'
+IsOrganizationManager is deliberately not used: an org manager is a
+customer-side role and must not be able to provision. The partner-facing
+wizard (C2) gets its own permission class scoped by invite token.
+"""
+
+import logging
+
+from django.shortcuts import get_object_or_404
+from drf_spectacular.types import OpenApiTypes
+from drf_spectacular.utils import OpenApiParameter, extend_schema, inline_serializer
+from requests.exceptions import HTTPError
+from rest_framework import serializers, status, viewsets
+from rest_framework.decorators import action
+from rest_framework.response import Response
+from rest_framework_extensions.mixins import NestedViewSetMixin
+
+from b2b.exceptions import (
+ AliasCollisionError,
+ InvalidLifecycleTransitionError,
+ OrganizationNotProvisionedError,
+ OrphanedKeycloakOrganizationError,
+)
+from b2b.models import (
+ OrganizationIdentityProvider,
+ OrganizationOnboarding,
+ OrganizationPage,
+)
+from b2b.provisioning import (
+ KeycloakConnection,
+ create_identity_provider,
+ create_organization,
+ delete_identity_provider,
+ parse_identity_provider_metadata,
+ refresh_identity_provider_metadata,
+ transition_identity_provider,
+ update_organization,
+)
+from b2b.serializers.v0.provisioning import (
+ CreateIdentityProviderSerializer,
+ CreateOrganizationSerializer,
+ IdentityProviderTransitionSerializer,
+ OrganizationIdentityProviderSerializer,
+ OrganizationOnboardingSerializer,
+ ParseMetadataSerializer,
+ ProvisionedOrganizationSerializer,
+ SetOnboardingStateSerializer,
+ UpdateOrganizationSerializer,
+)
+from main.permissions import IsAdminOrReadOnly
+
+log = logging.getLogger(__name__)
+
+DetailSerializer = inline_serializer(
+ name="ProvisioningDetailSerializer",
+ fields={"detail": serializers.CharField()},
+)
+
+
+class ProvisioningExceptionMixin:
+ """
+ Turn the provisioning failure modes into the responses the spec calls for.
+
+ A Keycloak call that fails becomes a 502 rather than a 500: our records are
+ intact and the operator's next move is to retry, not to open a ticket
+ against MITx Online.
+ """
+
+ def handle_exception(self, exc):
+ """Map provisioning exceptions onto HTTP responses."""
+
+ if isinstance(exc, AliasCollisionError | OrganizationNotProvisionedError):
+ return Response({"detail": str(exc)}, status=status.HTTP_409_CONFLICT)
+ if isinstance(exc, InvalidLifecycleTransitionError):
+ return Response({"detail": str(exc)}, status=status.HTTP_400_BAD_REQUEST)
+ if isinstance(exc, OrphanedKeycloakOrganizationError):
+ return Response(
+ {"detail": str(exc)},
+ status=status.HTTP_500_INTERNAL_SERVER_ERROR,
+ )
+ if isinstance(exc, HTTPError):
+ log.exception("Keycloak call failed during provisioning")
+ return Response(
+ {"detail": "The Keycloak admin API call failed."},
+ status=status.HTTP_502_BAD_GATEWAY,
+ )
+
+ return super().handle_exception(exc)
+
+
+class OrganizationProvisioningViewSet(
+ ProvisioningExceptionMixin,
+ viewsets.GenericViewSet,
+):
+ """Provision and inspect B2B organizations."""
+
+ permission_classes = [IsAdminOrReadOnly]
+ serializer_class = ProvisionedOrganizationSerializer
+ lookup_field = "org_key"
+ lookup_url_kwarg = "org_key"
+ queryset = OrganizationPage.objects.select_related("onboarding").prefetch_related(
+ "identity_providers"
+ )
+
+ def _with_keycloak(self, organization, connection=None):
+ """
+ Attach the organization's live Keycloak representation to the instance.
+
+ Domains and the post-login redirect live only in Keycloak, so a
+ response that did not read them back would be reporting what we asked
+ for rather than what is there.
+ """
+
+ if organization.sso_organization_id:
+ connection = connection or KeycloakConnection()
+ organization.keycloak_organization = connection.organizations.get(
+ organization.sso_organization_id
+ )
+
+ return organization
+
+ @extend_schema(
+ request=CreateOrganizationSerializer,
+ responses={
+ 201: ProvisionedOrganizationSerializer,
+ 409: DetailSerializer,
+ 502: DetailSerializer,
+ },
+ )
+ def create(self, request):
+ """
+ Provision a new organization.
+
+ Writes the Keycloak organization first, then the OrganizationPage and
+ its onboarding record in one transaction, compensating by deleting the
+ Keycloak organization if that transaction fails.
+ """
+
+ request_serializer = CreateOrganizationSerializer(data=request.data)
+ request_serializer.is_valid(raise_exception=True)
+
+ connection = KeycloakConnection()
+ organization = create_organization(
+ connection=connection, **request_serializer.validated_data
+ )
+
+ return Response(
+ self.get_serializer(self._with_keycloak(organization, connection)).data,
+ status=status.HTTP_201_CREATED,
+ )
+
+ @extend_schema(
+ responses={200: ProvisionedOrganizationSerializer, 502: DetailSerializer}
+ )
+ def retrieve(self, request, org_key=None): # noqa: ARG002
+ """Return an organization, including what Keycloak currently holds."""
+
+ organization = self.get_object()
+
+ return Response(self.get_serializer(self._with_keycloak(organization)).data)
+
+ @extend_schema(
+ request=UpdateOrganizationSerializer,
+ responses={
+ 200: ProvisionedOrganizationSerializer,
+ 400: DetailSerializer,
+ 502: DetailSerializer,
+ },
+ )
+ def partial_update(self, request, org_key=None): # noqa: ARG002
+ """Update an organization in both systems. org_key cannot change."""
+
+ organization = self.get_object()
+
+ request_serializer = UpdateOrganizationSerializer(data=request.data)
+ request_serializer.is_valid(raise_exception=True)
+
+ connection = KeycloakConnection()
+ organization = update_organization(
+ organization, connection=connection, **request_serializer.validated_data
+ )
+
+ return Response(
+ self.get_serializer(self._with_keycloak(organization, connection)).data
+ )
+
+ @extend_schema(
+ request=SetOnboardingStateSerializer,
+ responses={200: OrganizationOnboardingSerializer},
+ )
+ @action(detail=True, methods=["post"])
+ def onboarding(self, request, org_key=None): # noqa: ARG002
+ """
+ Record where this organization is in the onboarding sequence.
+
+ Descriptive only. Nothing in this API gates on the state; it exists so
+ a human can answer "what is left for this customer" without reading
+ four systems.
+ """
+
+ organization = self.get_object()
+
+ request_serializer = SetOnboardingStateSerializer(data=request.data)
+ request_serializer.is_valid(raise_exception=True)
+
+ onboarding, _ = OrganizationOnboarding.objects.get_or_create(
+ organization=organization
+ )
+ onboarding.set_state(
+ request_serializer.validated_data["state"],
+ notes=request_serializer.validated_data.get("notes"),
+ )
+
+ return Response(OrganizationOnboardingSerializer(onboarding).data)
+
+
+@extend_schema(
+ parameters=[
+ # The nested router names the parent lookup after the ORM path it
+ # filters on, which is not a field on OrganizationIdentityProvider, so
+ # spectacular cannot infer its type.
+ OpenApiParameter(
+ name="parent_lookup_organization__org_key",
+ type=OpenApiTypes.STR,
+ location=OpenApiParameter.PATH,
+ description="The organization's org key.",
+ )
+ ]
+)
+class IdentityProviderProvisioningViewSet(
+ ProvisioningExceptionMixin,
+ NestedViewSetMixin,
+ viewsets.GenericViewSet,
+):
+ """Provision and manage an organization's identity providers."""
+
+ permission_classes = [IsAdminOrReadOnly]
+ serializer_class = OrganizationIdentityProviderSerializer
+ lookup_field = "alias"
+ lookup_url_kwarg = "alias"
+ queryset = OrganizationIdentityProvider.objects.select_related("organization")
+
+ def _organization(self):
+ """Return the organization this route is nested under."""
+
+ return get_object_or_404(
+ OrganizationPage,
+ org_key=self.kwargs["parent_lookup_organization__org_key"],
+ )
+
+ @extend_schema(responses={200: OrganizationIdentityProviderSerializer(many=True)})
+ def list(self, request, **kwargs): # noqa: ARG002
+ """
+ List the organization's identity providers.
+
+ Resolve the parent first so an unknown org_key is a 404 rather than an
+ empty list. A mistyped key otherwise reads as "this organization has no
+ identity providers", which is the wrong answer to act on.
+ """
+
+ self._organization()
+
+ return Response(self.get_serializer(self.get_queryset(), many=True).data)
+
+ @extend_schema(responses={200: OrganizationIdentityProviderSerializer})
+ def retrieve(self, request, alias=None, **kwargs): # noqa: ARG002
+ """Return a single identity provider."""
+
+ return Response(self.get_serializer(self.get_object()).data)
+
+ @extend_schema(
+ request=CreateIdentityProviderSerializer,
+ responses={
+ 201: OrganizationIdentityProviderSerializer,
+ 409: DetailSerializer,
+ 502: DetailSerializer,
+ },
+ )
+ def create(self, request, **kwargs): # noqa: ARG002
+ """
+ Provision an identity provider and link it to the organization.
+
+ The IdP lands in `draft`, which is disabled in Keycloak. Nobody can
+ reach it until it is transitioned to `testing`.
+ """
+
+ request_serializer = CreateIdentityProviderSerializer(data=request.data)
+ request_serializer.is_valid(raise_exception=True)
+
+ payload = dict(request_serializer.validated_data)
+ payload.pop("discovery_url", None)
+
+ identity_provider = create_identity_provider(self._organization(), **payload)
+
+ return Response(
+ self.get_serializer(identity_provider).data,
+ status=status.HTTP_201_CREATED,
+ )
+
+ @extend_schema(responses={204: None, 502: DetailSerializer})
+ def destroy(self, request, alias=None, **kwargs): # noqa: ARG002
+ """Unlink and delete an identity provider."""
+
+ delete_identity_provider(self.get_object())
+
+ return Response(status=status.HTTP_204_NO_CONTENT)
+
+ @extend_schema(
+ request=None,
+ responses={
+ 200: OrganizationIdentityProviderSerializer,
+ 502: DetailSerializer,
+ },
+ )
+ @action(detail=True, methods=["post"], url_path="refresh-metadata")
+ def refresh_metadata(self, request, alias=None, **kwargs): # noqa: ARG002
+ """
+ Re-fetch the partner's metadata and store what came back.
+
+ On failure the stored artifact is left untouched, which is the whole
+ reason it is stored.
+ """
+
+ return Response(
+ self.get_serializer(
+ refresh_identity_provider_metadata(self.get_object())
+ ).data
+ )
+
+ @extend_schema(
+ request=IdentityProviderTransitionSerializer,
+ responses={
+ 200: OrganizationIdentityProviderSerializer,
+ 400: DetailSerializer,
+ 502: DetailSerializer,
+ },
+ )
+ @action(detail=True, methods=["post"])
+ def transition(self, request, alias=None, **kwargs): # noqa: ARG002
+ """
+ Move the identity provider's lifecycle state.
+
+ The only mover, and it writes Keycloak's enabled flag in the same
+ operation so the two cannot drift. draft -> active is rejected: an IdP
+ goes live only after somebody has logged in through it.
+ """
+
+ request_serializer = IdentityProviderTransitionSerializer(data=request.data)
+ request_serializer.is_valid(raise_exception=True)
+
+ identity_provider = transition_identity_provider(
+ self.get_object(), request_serializer.validated_data["state"]
+ )
+
+ return Response(self.get_serializer(identity_provider).data)
+
+
+class ParseMetadataView(ProvisioningExceptionMixin, viewsets.ViewSet):
+ """
+ Parse IdP metadata without creating anything.
+
+ Deliberately unnested: an operator (later, a wizard) pastes a metadata URL
+ or document and sees what Keycloak makes of it before committing to a
+ resource. Staff-only, and it stays that way in phase 1 - the URL form makes
+ Keycloak fetch a caller-supplied address, which is an SSRF surface that
+ needs an allowlist and a rate limit before it goes anywhere near a partner.
+ """
+
+ permission_classes = [IsAdminOrReadOnly]
+
+ @extend_schema(
+ request=ParseMetadataSerializer,
+ responses={
+ 200: inline_serializer(
+ name="ParsedIdentityProviderConfigSerializer",
+ fields={"config": serializers.DictField(child=serializers.CharField())},
+ ),
+ 502: DetailSerializer,
+ },
+ )
+ def create(self, request):
+ """Return the config map Keycloak parses out of the given metadata."""
+
+ request_serializer = ParseMetadataSerializer(data=request.data)
+ request_serializer.is_valid(raise_exception=True)
+
+ config = parse_identity_provider_metadata(
+ request_serializer.validated_data["protocol"],
+ metadata_url=request_serializer.validated_data.get("metadata_url"),
+ metadata_xml=request_serializer.validated_data.get("metadata_xml"),
+ )
+
+ return Response({"config": config})
diff --git a/b2b/views/v0/provisioning_test.py b/b2b/views/v0/provisioning_test.py
new file mode 100644
index 0000000000..b3ba79250f
--- /dev/null
+++ b/b2b/views/v0/provisioning_test.py
@@ -0,0 +1,472 @@
+"""Tests for the staff-only B2B provisioning API's HTTP surface."""
+
+import faker
+import pytest
+from django.urls import reverse
+from requests.exceptions import HTTPError
+from rest_framework import status
+
+from b2b.constants import (
+ IDP_PROTOCOL_OIDC,
+ IDP_PROTOCOL_SAML,
+ IDP_STATE_ACTIVE,
+ IDP_STATE_DRAFT,
+ IDP_STATE_TESTING,
+ ONBOARDING_STATE_LIVE,
+)
+from b2b.exceptions import AliasCollisionError
+from b2b.factories import OrganizationIndexPageFactory, OrganizationPageFactory
+from b2b.keycloak_admin_dataclasses import (
+ OrganizationDomainRepresentation,
+ OrganizationRepresentation,
+)
+from b2b.models import OrganizationIdentityProvider, OrganizationPage
+
+pytestmark = [pytest.mark.django_db]
+FAKE = faker.Faker()
+
+
+@pytest.fixture(autouse=True)
+def organization_index():
+ """The index page organizations are added under."""
+
+ return OrganizationIndexPageFactory.create()
+
+
+@pytest.fixture(autouse=True)
+def mocked_connection(mocker):
+ """
+ Stop the views bootstrapping a real Keycloak client.
+
+ The views construct a KeycloakConnection and hand it to the provisioning
+ functions, so patching the class covers both.
+ """
+
+ connection = mocker.Mock()
+ connection.organizations.get.return_value = OrganizationRepresentation(
+ id=str(FAKE.uuid4()),
+ name="Example University",
+ alias="EXAMPLEU",
+ redirect_url="https://learn.mit.edu/dashboard/organization/exampleu",
+ domains=[OrganizationDomainRepresentation(name="example.edu", verified=True)],
+ )
+ for target in (
+ "b2b.views.v0.provisioning.KeycloakConnection",
+ "b2b.provisioning.KeycloakConnection",
+ ):
+ mocker.patch(target, return_value=connection)
+ return connection
+
+
+def _organizations_url():
+ return reverse("b2b:b2b-provisioning-organization-list")
+
+
+def _organization_url(org_key):
+ return reverse(
+ "b2b:b2b-provisioning-organization-detail", kwargs={"org_key": org_key}
+ )
+
+
+def _identity_providers_url(org_key):
+ return reverse(
+ "b2b:b2b-provisioning-organization-idp-list",
+ kwargs={"parent_lookup_organization__org_key": org_key},
+ )
+
+
+def _identity_provider_url(org_key, alias, suffix="detail"):
+ return reverse(
+ f"b2b:b2b-provisioning-organization-idp-{suffix}",
+ kwargs={
+ "parent_lookup_organization__org_key": org_key,
+ "alias": alias,
+ },
+ )
+
+
+CREATE_BODY = {
+ "name": "Example University",
+ "org_key": "EXAMPLEU",
+ "domains": ["example.edu"],
+ "redirect_url": "https://learn.mit.edu/dashboard/organization/exampleu",
+}
+
+
+def test_create_organization_requires_staff(user_drf_client):
+ """
+ An org manager is a customer-side role and must not provision.
+
+ IsAdminOrReadOnly grants write to is_staff only; a plain authenticated user
+ gets read access and nothing else.
+ """
+
+ response = user_drf_client.post(_organizations_url(), CREATE_BODY, format="json")
+
+ assert response.status_code == status.HTTP_403_FORBIDDEN
+ assert not OrganizationPage.objects.filter(org_key="EXAMPLEU").exists()
+
+
+def test_create_organization(admin_drf_client, mocker):
+ """A staff create returns 201 and the organization it made."""
+
+ organization = OrganizationPageFactory.build(org_key="EXAMPLEU")
+ mocker.patch(
+ "b2b.views.v0.provisioning.create_organization",
+ return_value=OrganizationPageFactory.create(org_key="EXAMPLEU"),
+ )
+
+ response = admin_drf_client.post(_organizations_url(), CREATE_BODY, format="json")
+
+ assert response.status_code == status.HTTP_201_CREATED
+ assert response.json()["org_key"] == organization.org_key
+
+
+def test_create_organization_alias_collision_is_a_conflict(admin_drf_client, mocker):
+ """A taken alias is 409 with the reason, not a 500."""
+
+ mocker.patch(
+ "b2b.views.v0.provisioning.create_organization",
+ side_effect=AliasCollisionError("taken"),
+ )
+
+ response = admin_drf_client.post(_organizations_url(), CREATE_BODY, format="json")
+
+ assert response.status_code == status.HTTP_409_CONFLICT
+ assert response.json()["detail"] == "taken"
+
+
+def test_keycloak_failure_is_a_bad_gateway(admin_drf_client, mocker):
+ """
+ A failed Keycloak call is 502, not 500.
+
+ Our records are intact and the operator's next move is to retry, not to
+ open a ticket against MITx Online.
+ """
+
+ mocker.patch(
+ "b2b.views.v0.provisioning.create_organization",
+ side_effect=HTTPError("keycloak said no"),
+ )
+
+ response = admin_drf_client.post(_organizations_url(), CREATE_BODY, format="json")
+
+ assert response.status_code == status.HTTP_502_BAD_GATEWAY
+
+
+def test_retrieve_organization_includes_what_keycloak_holds(admin_drf_client):
+ """Domains and the redirect URL live only in Keycloak, so read them back."""
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+
+ response = admin_drf_client.get(_organization_url(organization.org_key))
+
+ assert response.status_code == status.HTTP_200_OK
+ assert response.json()["domains"] == ["example.edu"]
+ assert (
+ response.json()["redirect_url"]
+ == "https://learn.mit.edu/dashboard/organization/exampleu"
+ )
+
+
+def test_patch_rejects_an_org_key_change(admin_drf_client):
+ """
+ org_key is immutable, and saying so beats accepting and ignoring it.
+
+ It is part of every B2B courseware ID via create_contract_run_key, which is
+ also why reconcile_single_keycloak_org refuses to update it.
+ """
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+
+ response = admin_drf_client.patch(
+ _organization_url(organization.org_key),
+ {"name": "Renamed", "org_key": "SOMETHINGELSE"},
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_400_BAD_REQUEST
+ assert "org_key" in response.json()["errors"]
+
+ organization.refresh_from_db()
+ assert organization.org_key == "EXAMPLEU"
+ assert organization.name != "Renamed"
+
+
+def test_patch_updates_the_mutable_fields(admin_drf_client, mocker):
+ """name, description, redirect_url and domains are all updatable."""
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+ mocked_update = mocker.patch(
+ "b2b.views.v0.provisioning.update_organization", return_value=organization
+ )
+
+ response = admin_drf_client.patch(
+ _organization_url(organization.org_key),
+ {"name": "Renamed", "domains": ["renamed.edu"]},
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_200_OK
+ assert mocked_update.call_args.kwargs["name"] == "Renamed"
+ assert mocked_update.call_args.kwargs["domains"] == ["renamed.edu"]
+
+
+def test_patching_an_unprovisioned_organization_is_a_conflict(admin_drf_client):
+ """
+ An org with no Keycloak record is 409, not 502.
+
+ 502 would say the Keycloak API failed. It did not - our record is the
+ incomplete one, and retrying will never fix it. Roughly 24 production
+ organizations are in this state (hq#10552).
+ """
+
+ organization = OrganizationPageFactory.create(
+ org_key="LEGACYU", sso_organization_id=None
+ )
+
+ response = admin_drf_client.patch(
+ _organization_url(organization.org_key),
+ {"name": "Renamed"},
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_409_CONFLICT
+ assert "LEGACYU" in response.json()["detail"]
+
+
+def test_set_onboarding_state(admin_drf_client):
+ """The onboarding record is how an operator says where a customer is."""
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+
+ response = admin_drf_client.post(
+ reverse(
+ "b2b:b2b-provisioning-organization-onboarding",
+ kwargs={"org_key": organization.org_key},
+ ),
+ {"state": ONBOARDING_STATE_LIVE, "notes": "first cohort enrolled"},
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_200_OK
+ assert response.json()["state"] == ONBOARDING_STATE_LIVE
+
+ organization.refresh_from_db()
+ assert organization.onboarding.notes == "first cohort enrolled"
+
+
+def _identity_provider(organization, alias="exampleu", state=IDP_STATE_DRAFT):
+ return OrganizationIdentityProvider.objects.create(
+ organization=organization,
+ alias=alias,
+ protocol=IDP_PROTOCOL_SAML,
+ lifecycle_state=state,
+ metadata_source="https://idp.example.edu/metadata.xml",
+ metadata_artifact={"idpEntityId": "https://idp.example.edu/entity"},
+ )
+
+
+def test_identity_providers_are_scoped_to_their_organization(admin_drf_client):
+ """The nested route lists only that organization's providers."""
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+ _identity_provider(organization)
+ _identity_provider(OrganizationPageFactory.create(org_key="OTHERU"), alias="otheru")
+
+ response = admin_drf_client.get(_identity_providers_url(organization.org_key))
+
+ assert response.status_code == status.HTTP_200_OK
+ assert [idp["alias"] for idp in response.json()] == ["exampleu"]
+
+
+def test_listing_providers_for_an_unknown_organization_is_a_404(admin_drf_client):
+ """
+ An unknown org_key is a 404, not an empty list.
+
+ A mistyped key otherwise reads as "this organization has no identity
+ providers", which is a much worse thing for an operator to act on.
+ """
+
+ response = admin_drf_client.get(_identity_providers_url("NOSUCHORG"))
+
+ assert response.status_code == status.HTTP_404_NOT_FOUND
+
+
+def test_create_identity_provider_requires_a_metadata_source(admin_drf_client):
+ """A SAML IdP takes exactly one of metadata_url or metadata_xml."""
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+
+ response = admin_drf_client.post(
+ _identity_providers_url(organization.org_key),
+ {
+ "alias": "exampleu",
+ "protocol": IDP_PROTOCOL_SAML,
+ "attribute_map": {"email": "E-Mail Address"},
+ },
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_400_BAD_REQUEST
+
+
+def test_saml_identity_provider_requires_attribute_mappers(admin_drf_client):
+ """
+ Nothing maps a SAML assertion's attributes onto the user unless we say so.
+
+ A SAML IdP with no mappers brokers users with no email or name, which is a
+ support ticket rather than a working integration.
+ """
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+
+ response = admin_drf_client.post(
+ _identity_providers_url(organization.org_key),
+ {
+ "alias": "exampleu",
+ "protocol": IDP_PROTOCOL_SAML,
+ "metadata_url": "https://idp.example.edu/metadata.xml",
+ },
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_400_BAD_REQUEST
+
+
+def test_oidc_identity_provider_does_not_require_attribute_mappers(
+ admin_drf_client, mocker
+):
+ """
+ OIDC providers are valid with no mappers, and ours run that way.
+
+ ol-infrastructure's onboard_oidc_org creates no attribute-importer mappers
+ at all, so every OIDC IdP Pulumi manages in production today has none.
+ Rejecting that configuration would refuse what we already run.
+ """
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+ mocked_create = mocker.patch(
+ "b2b.views.v0.provisioning.create_identity_provider",
+ return_value=_identity_provider(organization),
+ )
+
+ response = admin_drf_client.post(
+ _identity_providers_url(organization.org_key),
+ {
+ "alias": "exampleu",
+ "protocol": IDP_PROTOCOL_OIDC,
+ "discovery_url": "https://idp.example.edu/.well-known/openid-configuration",
+ "client_id": "mitxonline",
+ "client_secret": "shh",
+ },
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_201_CREATED
+ mocked_create.assert_called_once()
+
+
+def test_create_identity_provider(admin_drf_client, mocker):
+ """A staff create returns 201 and the provider record."""
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+ mocker.patch(
+ "b2b.views.v0.provisioning.create_identity_provider",
+ return_value=_identity_provider(organization),
+ )
+
+ response = admin_drf_client.post(
+ _identity_providers_url(organization.org_key),
+ {
+ "alias": "exampleu",
+ "protocol": IDP_PROTOCOL_SAML,
+ "metadata_url": "https://idp.example.edu/metadata.xml",
+ "attribute_map": {"email": "E-Mail Address"},
+ },
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_201_CREATED
+ assert response.json()["lifecycle_state"] == IDP_STATE_DRAFT
+
+
+def test_transition_rejects_skipping_testing(admin_drf_client):
+ """Draft -> active is a 400: an IdP goes live only after a real login."""
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+ identity_provider = _identity_provider(organization)
+
+ response = admin_drf_client.post(
+ _identity_provider_url(organization.org_key, "exampleu", suffix="transition"),
+ {"state": IDP_STATE_ACTIVE},
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_400_BAD_REQUEST
+
+ identity_provider.refresh_from_db()
+ assert identity_provider.lifecycle_state == IDP_STATE_DRAFT
+
+
+def test_transition_to_testing(admin_drf_client, mocker):
+ """The allowed move writes both systems through the provisioning layer."""
+
+ organization = OrganizationPageFactory.create(org_key="EXAMPLEU")
+ identity_provider = _identity_provider(organization)
+ mocked_transition = mocker.patch(
+ "b2b.views.v0.provisioning.transition_identity_provider",
+ return_value=identity_provider,
+ )
+
+ response = admin_drf_client.post(
+ _identity_provider_url(organization.org_key, "exampleu", suffix="transition"),
+ {"state": IDP_STATE_TESTING},
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_200_OK
+ assert mocked_transition.call_args.args[1] == IDP_STATE_TESTING
+
+
+def test_parse_metadata_creates_nothing(admin_drf_client, mocker):
+ """The cheapest useful call: see what Keycloak makes of the metadata."""
+
+ config = {"idpEntityId": "https://idp.example.edu/entity"}
+ mocker.patch(
+ "b2b.views.v0.provisioning.parse_identity_provider_metadata",
+ return_value=config,
+ )
+
+ response = admin_drf_client.post(
+ reverse("b2b:b2b-provisioning-parse-metadata-list"),
+ {
+ "protocol": IDP_PROTOCOL_SAML,
+ "metadata_url": "https://idp.example.edu/metadata.xml",
+ },
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_200_OK
+ assert response.json()["config"] == config
+ assert not OrganizationIdentityProvider.objects.exists()
+
+
+def test_parse_metadata_is_staff_only(user_drf_client):
+ """
+ Keycloak fetches a caller-supplied URL here, so this stays staff-only.
+
+ Exposing it to partners needs an allowlist or deny-private-ranges policy
+ and a rate limit, which belongs to C2's threat model.
+ """
+
+ response = user_drf_client.post(
+ reverse("b2b:b2b-provisioning-parse-metadata-list"),
+ {
+ "protocol": IDP_PROTOCOL_SAML,
+ "metadata_url": "http://169.254.169.254/latest/meta-data/",
+ },
+ format="json",
+ )
+
+ assert response.status_code == status.HTTP_403_FORBIDDEN
diff --git a/b2b/views/v0/urls.py b/b2b/views/v0/urls.py
index 4c31885951..1ca8b795ee 100644
--- a/b2b/views/v0/urls.py
+++ b/b2b/views/v0/urls.py
@@ -13,6 +13,11 @@
ManagerOrganizationViewSet,
ProcessMailgunWebhook,
)
+from b2b.views.v0.provisioning import (
+ IdentityProviderProvisioningViewSet,
+ OrganizationProvisioningViewSet,
+ ParseMetadataView,
+)
from b2b.views.v0.service import OrganizationManagerCheckView
from main.routers import SimpleRouterWithNesting
@@ -45,6 +50,28 @@
],
)
+# Staff-only provisioning routes (capability C1). These take ownership of
+# per-customer Keycloak resources from Pulumi; see
+# docs/source/b2b/provisioning_api.md.
+provisioning_org = v0_router.register(
+ r"provisioning/organizations",
+ OrganizationProvisioningViewSet,
+ basename="b2b-provisioning-organization",
+)
+provisioning_org.register(
+ r"identity-providers",
+ IdentityProviderProvisioningViewSet,
+ basename="b2b-provisioning-organization-idp",
+ parents_query_lookups=[
+ "organization__org_key",
+ ],
+)
+v0_router.register(
+ r"provisioning/parse-metadata",
+ ParseMetadataView,
+ basename="b2b-provisioning-parse-metadata",
+)
+
urlpatterns = [
path("", include(v0_router.urls)),
path(r"enroll//", Enroll.as_view(), name="enroll-user"),
diff --git a/courses/api_test.py b/courses/api_test.py
index 91e654a1c7..7851b2c900 100644
--- a/courses/api_test.py
+++ b/courses/api_test.py
@@ -18,7 +18,7 @@
from django.core.exceptions import ValidationError
from django.db import connection
from django.db.models import Prefetch
-from django.test import RequestFactory
+from django.test import RequestFactory, override_settings
from django.test.utils import CaptureQueriesContext
from edx_api.course_detail import CourseDetail, CourseMode, CourseModes
from mitol.common.utils.datetime import now_in_utc
@@ -3863,6 +3863,7 @@ def test_course_run_certificate_verifiable_credentials_feature_flag_disabled(
assert certificate.verifiable_credential is None
+@override_settings(MIT_LEARN_BASE_URL="https://learn.mit.edu")
@patch("courses.api.CourseRunEnrollment.all_objects.get")
@patch("courses.api.get_thumbnail_url")
@patch("courses.signals.upsert_custom_properties")
@@ -3974,6 +3975,7 @@ def test_course_run_certificate_verifiable_credentials_signing_payload(
assert payload == expected_payload
+@override_settings(MIT_LEARN_BASE_URL="https://learn.mit.edu")
@patch("courses.api.ProgramEnrollment.all_objects.get")
@patch("courses.api.get_thumbnail_url")
def test_program_certificate_verifiable_credentials_signing_payload(
diff --git a/courses/models.py b/courses/models.py
index 58ce1960b7..2e007f0fdb 100644
--- a/courses/models.py
+++ b/courses/models.py
@@ -1199,7 +1199,7 @@ def get_filtered_runs(
courseruns = (
self.prefetched_courseruns
if hasattr(self, "prefetched_courseruns")
- else list(self.courseruns.prefetch_related("b2b_contracts").all())
+ else list(self.courseruns.all())
)
courseruns = sorted(courseruns, key=lambda r: r.id)
diff --git a/courses/models_test.py b/courses/models_test.py
index dd8c27a8f7..699e7af342 100644
--- a/courses/models_test.py
+++ b/courses/models_test.py
@@ -1499,6 +1499,18 @@ def test_get_filtered_runs_includes_runs_from_b2b_contracts(filter_name):
)
+def test_get_filtered_runs_reuses_prefetched_courseruns(django_assert_num_queries):
+ """A courseruns prefetch is reused rather than re-queried."""
+ course = CourseFactory.create()
+ course_run = CourseRunFactory.create(course=course)
+ course = Course.objects.prefetch_related("courseruns").get(pk=course.pk)
+
+ with django_assert_num_queries(0):
+ runs = course.get_filtered_runs(courserun_is_enrollable=None)
+
+ assert runs == [course_run]
+
+
# Test for course run constraints
# As a default we expect uniqueness on course, courseware_id, and run_tag
# We also shouldn't allow:
diff --git a/courses/views/v2/__init__.py b/courses/views/v2/__init__.py
index 6110a47896..4fd43c893e 100644
--- a/courses/views/v2/__init__.py
+++ b/courses/views/v2/__init__.py
@@ -471,7 +471,7 @@ def get_queryset(self):
"courseruns",
queryset=CourseRun.objects.order_by("id")
.select_related("b2b_contract")
- .prefetch_related("b2b_contracts", modes_prefetch, products_prefetch),
+ .prefetch_related(modes_prefetch, products_prefetch),
)
queryset = queryset.prefetch_related(
"departments", "in_programs", course_runs_prefetch
diff --git a/courses/views/v2/views_test.py b/courses/views/v2/views_test.py
index 6e8b58b28a..c51a6165e1 100644
--- a/courses/views/v2/views_test.py
+++ b/courses/views/v2/views_test.py
@@ -2823,8 +2823,7 @@ def test_course_run_and_product_prefetch_optimized(
ProductFactory(
purchasable_object=run,
)
- # increased below from 21 to 27 - the M2M for b2b_contracts adds some queries
- max_expected_queries = 27
+ max_expected_queries = 21
num_queries_before = len(connection.queries)
with django_assert_max_num_queries(max_expected_queries):
resp = user_drf_client.get(reverse("v2:courses_api-list"))
@@ -2839,9 +2838,8 @@ def test_course_run_and_product_prefetch_optimized(
product_queries = [
q for q in queries_after if 'FROM "ecommerce_product"' in q.get("sql", "")
]
- # increased below from 2 to 3 - the M2M for b2b_contracts adds some queries
- assert len(product_queries) == 3, (
- f"Expected 3 product query, got {len(product_queries)}: {[q['sql'] for q in product_queries]}"
+ assert len(product_queries) == 2, (
+ f"Expected 2 product queries, got {len(product_queries)}: {[q['sql'] for q in product_queries]}"
)
diff --git a/drf_lint_baseline.json b/drf_lint_baseline.json
index 4dc7d4c2e4..65c1ba54e7 100644
--- a/drf_lint_baseline.json
+++ b/drf_lint_baseline.json
@@ -1,14 +1,35 @@
[
+ "b2b/serializers/v0/__init__.py:56:50:ORM003",
+ "b2b/serializers/v0/__init__.py:89:31:ORM003",
"cms/serializers.py:115:16:ORM001",
"cms/serializers.py:154:41:ORM001",
+ "cms/serializers.py:185:19:ORM004",
+ "cms/serializers.py:194:29:ORM004",
+ "cms/serializers.py:197:21:ORM004",
+ "cms/serializers.py:202:15:ORM004",
+ "cms/serializers.py:277:39:ORM006",
+ "cms/serializers.py:28:39:ORM006",
+ "cms/serializers.py:307:11:ORM006",
+ "cms/serializers.py:309:18:ORM006",
"cms/serializers.py:321:12:ORM001",
+ "cms/serializers.py:322:36:ORM006",
+ "cms/serializers.py:350:31:ORM006",
"cms/serializers.py:356:20:ORM001",
"cms/serializers.py:367:39:ORM001",
+ "cms/serializers.py:407:39:ORM006",
"cms/serializers.py:84:12:ORM001",
"cms/serializers.py:97:12:ORM001",
"courses/serializers/base.py:53:16:ORM001",
+ "courses/serializers/v1/base.py:195:15:ORM004",
"courses/serializers/v1/base.py:75:20:ORM002",
+ "courses/serializers/v1/base.py:84:15:ORM003",
+ "courses/serializers/v1/courses.py:109:15:ORM004",
+ "courses/serializers/v1/courses.py:38:14:ORM003",
"courses/serializers/v1/courses.py:59:16:ORM001",
+ "courses/serializers/v1/courses.py:60:58:ORM003",
+ "courses/serializers/v1/courses.py:77:12:ORM007",
+ "courses/serializers/v1/programs.py:128:49:ORM003",
+ "courses/serializers/v1/programs.py:129:50:ORM003",
"courses/serializers/v1/programs.py:181:12:ORM001",
"courses/serializers/v1/programs.py:196:12:ORM001",
"courses/serializers/v1/programs.py:208:12:ORM001",
@@ -16,6 +37,10 @@
"courses/serializers/v1/programs.py:303:27:ORM001",
"courses/serializers/v1/programs.py:318:16:ORM001",
"courses/serializers/v1/programs.py:335:17:ORM001",
+ "courses/serializers/v1/programs.py:360:16:ORM004",
+ "courses/serializers/v1/programs.py:96:37:ORM003",
+ "courses/serializers/v2/certificates.py:174:11:ORM006",
+ "courses/serializers/v2/certificates.py:175:19:ORM006",
"courses/serializers/v2/courses.py:138:18:ORM004",
"courses/serializers/v2/courses.py:142:18:ORM003",
"courses/serializers/v2/courses.py:150:19:ORM004",
@@ -25,49 +50,113 @@
"courses/serializers/v2/courses.py:219:12:ORM007",
"courses/serializers/v2/courses.py:268:15:ORM004",
"courses/serializers/v2/courses.py:276:17:ORM002",
- "courses/serializers/v2/courses.py:297:19:ORM007",
+ "courses/serializers/v2/courses.py:297:19:ORM004",
"courses/serializers/v2/departments.py:35:40:ORM002",
"courses/serializers/v2/departments.py:49:42:ORM002",
+ "courses/serializers/v2/programs.py:115:24:ORM003",
+ "courses/serializers/v2/programs.py:194:43:ORM003",
+ "courses/serializers/v2/programs.py:197:15:ORM004",
+ "courses/serializers/v2/programs.py:262:27:ORM004",
+ "courses/serializers/v2/programs.py:376:15:ORM004",
"courses/serializers/v2/programs.py:387:50:ORM002",
+ "courses/serializers/v2/programs.py:439:19:ORM003",
"courses/serializers/v2/programs.py:500:12:ORM002",
+ "courses/serializers/v3/courses.py:100:11:ORM006",
+ "courses/serializers/v3/courses.py:101:19:ORM006",
+ "courses/serializers/v3/courses.py:107:15:ORM006",
"courses/serializers/v3/courses.py:57:12:ORM002",
- "ecommerce/serializers/__init__.py:256:17:ORM001",
- "ecommerce/serializers/__init__.py:258:18:ORM001",
- "ecommerce/serializers/__init__.py:259:18:ORM001",
- "ecommerce/serializers/__init__.py:281:26:ORM002",
- "ecommerce/serializers/__init__.py:369:26:ORM002",
- "ecommerce/serializers/__init__.py:377:31:ORM002",
- "ecommerce/serializers/__init__.py:383:20:ORM002",
- "ecommerce/serializers/__init__.py:388:31:ORM002",
- "ecommerce/serializers/__init__.py:404:35:ORM002",
- "ecommerce/serializers/__init__.py:463:24:ORM002",
- "ecommerce/serializers/__init__.py:482:12:ORM001",
- "ecommerce/serializers/__init__.py:506:22:ORM002",
- "ecommerce/serializers/__init__.py:553:22:ORM002",
- "ecommerce/serializers/__init__.py:623:20:ORM002",
- "ecommerce/serializers/__init__.py:798:22:ORM002",
- "ecommerce/serializers/__init__.py:980:28:ORM002",
+ "courses/serializers/v3/courses.py:62:18:ORM004",
+ "courses/serializers/v3/courses.py:69:18:ORM004",
+ "courses/serializers/v3/courses.py:74:18:ORM004",
+ "ecommerce/serializers/__init__.py:243:12:ORM005",
+ "ecommerce/serializers/__init__.py:254:17:ORM001",
+ "ecommerce/serializers/__init__.py:256:18:ORM001",
+ "ecommerce/serializers/__init__.py:257:18:ORM001",
+ "ecommerce/serializers/__init__.py:279:26:ORM002",
+ "ecommerce/serializers/__init__.py:317:42:ORM006",
+ "ecommerce/serializers/__init__.py:367:26:ORM002",
+ "ecommerce/serializers/__init__.py:375:31:ORM002",
+ "ecommerce/serializers/__init__.py:381:20:ORM002",
+ "ecommerce/serializers/__init__.py:383:19:ORM004",
+ "ecommerce/serializers/__init__.py:386:31:ORM002",
+ "ecommerce/serializers/__init__.py:402:35:ORM002",
+ "ecommerce/serializers/__init__.py:422:4:ORM005",
+ "ecommerce/serializers/__init__.py:423:4:ORM005",
+ "ecommerce/serializers/__init__.py:424:4:ORM005",
+ "ecommerce/serializers/__init__.py:431:12:ORM005",
+ "ecommerce/serializers/__init__.py:451:15:ORM003",
+ "ecommerce/serializers/__init__.py:461:24:ORM002",
+ "ecommerce/serializers/__init__.py:480:12:ORM001",
+ "ecommerce/serializers/__init__.py:504:22:ORM002",
+ "ecommerce/serializers/__init__.py:551:22:ORM002",
+ "ecommerce/serializers/__init__.py:588:46:ORM006",
+ "ecommerce/serializers/__init__.py:615:15:ORM003",
+ "ecommerce/serializers/__init__.py:621:20:ORM002",
+ "ecommerce/serializers/__init__.py:658:26:ORM004",
+ "ecommerce/serializers/__init__.py:796:22:ORM002",
+ "ecommerce/serializers/__init__.py:978:28:ORM002",
"ecommerce/serializers/v0/__init__.py:289:17:ORM001",
"ecommerce/serializers/v0/__init__.py:291:18:ORM001",
"ecommerce/serializers/v0/__init__.py:292:18:ORM001",
"ecommerce/serializers/v0/__init__.py:314:26:ORM002",
+ "ecommerce/serializers/v0/__init__.py:352:42:ORM006",
"ecommerce/serializers/v0/__init__.py:402:26:ORM002",
"ecommerce/serializers/v0/__init__.py:411:35:ORM002",
"ecommerce/serializers/v0/__init__.py:418:20:ORM002",
+ "ecommerce/serializers/v0/__init__.py:420:19:ORM004",
"ecommerce/serializers/v0/__init__.py:424:35:ORM002",
"ecommerce/serializers/v0/__init__.py:441:35:ORM002",
+ "ecommerce/serializers/v0/__init__.py:463:4:ORM005",
+ "ecommerce/serializers/v0/__init__.py:464:4:ORM005",
+ "ecommerce/serializers/v0/__init__.py:465:4:ORM005",
+ "ecommerce/serializers/v0/__init__.py:495:15:ORM003",
+ "ecommerce/serializers/v0/__init__.py:499:15:ORM003",
+ "ecommerce/serializers/v0/__init__.py:503:15:ORM003",
+ "ecommerce/serializers/v0/__init__.py:507:18:ORM003",
+ "ecommerce/serializers/v0/__init__.py:512:15:ORM003",
"ecommerce/serializers/v0/__init__.py:522:24:ORM002",
"ecommerce/serializers/v0/__init__.py:541:12:ORM001",
"ecommerce/serializers/v0/__init__.py:565:22:ORM002",
"ecommerce/serializers/v0/__init__.py:612:22:ORM002",
+ "ecommerce/serializers/v0/__init__.py:649:46:ORM006",
+ "ecommerce/serializers/v0/__init__.py:66:12:ORM005",
+ "ecommerce/serializers/v0/__init__.py:681:15:ORM003",
"ecommerce/serializers/v0/__init__.py:687:20:ORM002",
+ "ecommerce/serializers/v0/__init__.py:724:26:ORM004",
"ecommerce/serializers/v0/__init__.py:819:22:ORM002",
"ecommerce/serializers/v0/__init__.py:963:28:ORM002",
"flexiblepricing/serializers.py:147:34:ORM001",
+ "flexiblepricing/serializers.py:163:43:ORM006",
"flexiblepricing/serializers.py:170:34:ORM001",
"flexiblepricing/serializers.py:173:20:ORM001",
"flexiblepricing/serializers.py:204:30:ORM001",
"flexiblepricing/serializers.py:207:31:ORM001",
"flexiblepricing/serializers.py:212:16:ORM001",
- "flexiblepricing/serializers.py:216:16:ORM001"
+ "flexiblepricing/serializers.py:216:16:ORM001",
+ "flexiblepricing/serializers.py:226:15:ORM004",
+ "flexiblepricing/serializers.py:233:15:ORM004",
+ "flexiblepricing/serializers.py:240:15:ORM004",
+ "flexiblepricing/serializers.py:247:15:ORM004",
+ "flexiblepricing/serializers.py:44:11:ORM006",
+ "hubspot_sync/serializers.py:104:21:ORM004",
+ "hubspot_sync/serializers.py:109:21:ORM004",
+ "hubspot_sync/serializers.py:118:11:ORM006",
+ "hubspot_sync/serializers.py:124:15:ORM006",
+ "hubspot_sync/serializers.py:126:15:ORM004",
+ "hubspot_sync/serializers.py:130:15:ORM006",
+ "hubspot_sync/serializers.py:184:21:ORM004",
+ "hubspot_sync/serializers.py:193:21:ORM004",
+ "hubspot_sync/serializers.py:196:30:ORM006",
+ "hubspot_sync/serializers.py:222:19:ORM004",
+ "hubspot_sync/serializers.py:242:15:ORM004",
+ "hubspot_sync/serializers.py:243:15:ORM004",
+ "hubspot_sync/serializers.py:250:17:ORM004",
+ "hubspot_sync/serializers.py:258:17:ORM004",
+ "hubspot_sync/serializers.py:270:19:ORM004",
+ "hubspot_sync/serializers.py:98:21:ORM004",
+ "users/serializers.py:163:15:ORM006",
+ "users/serializers.py:174:15:ORM005",
+ "users/serializers.py:210:15:ORM005",
+ "users/serializers.py:250:12:ORM004",
+ "users/serializers.py:413:12:ORM005"
]
diff --git a/ecommerce/admin.py b/ecommerce/admin.py
index d0c6ac4e70..70ff5e3cc0 100644
--- a/ecommerce/admin.py
+++ b/ecommerce/admin.py
@@ -16,6 +16,7 @@
from viewflow import fsm
from ecommerce.api import refund_order
+from ecommerce.discount_sources import fulfilled_redemptions_funded_by
from ecommerce.forms import AdminRefundOrderForm
from ecommerce.models import (
Basket,
@@ -402,6 +403,11 @@ def get_queryset(self, request):
return super().get_queryset(request).filter(state=OrderStatus.REFUNDED)
+def _used_source_redemptions(order):
+ """Paid-amount-off credits a line of this order funded, which a refund leaves in place."""
+ return fulfilled_redemptions_funded_by(order)
+
+
class AdminRefundOrderView(LoginRequiredMixin, PermissionRequiredMixin, TemplateView):
template_name = "refund_order_confirm.html"
permission_required = "is_superuser"
@@ -476,6 +482,7 @@ def post(self, request):
"form_valid": refund_form.is_valid(),
"errors": errors,
"error_messages": error_messages,
+ "used_source_redemptions": _used_source_redemptions(order),
},
)
except NotImplementedError:
@@ -521,6 +528,7 @@ def get(self, request):
"order": order,
"form_valid": True,
"errors": {},
+ "used_source_redemptions": _used_source_redemptions(order),
},
)
diff --git a/ecommerce/admin_test.py b/ecommerce/admin_test.py
index ab1009ba31..e7681449c9 100644
--- a/ecommerce/admin_test.py
+++ b/ecommerce/admin_test.py
@@ -7,7 +7,10 @@
from reversion.models import Version
from courses.factories import CourseRunFactory
-from ecommerce.factories import OrderFactory
+from ecommerce.factories import (
+ DiscountRedemptionFactory,
+ OrderFactory,
+)
from ecommerce.models import OrderStatus, Product
pytestmark = [pytest.mark.django_db]
@@ -291,3 +294,45 @@ def test_admin_product_create_generates_reversion(client, admin_user):
product = Product.all_objects.get(description="Admin created versioned product")
assert Version.objects.get_for_object(product).count() == 1
+
+
+def test_admin_refund_view_lists_the_credits_a_refund_would_leave_behind(
+ client, admin_user, paid_amount_off_source
+):
+ """Support sees which order the credit went to before knowingly refunding."""
+ _login_admin(client, admin_user)
+ order = paid_amount_off_source.source_line.order
+ redemption = DiscountRedemptionFactory.create(
+ redeemed_discount=paid_amount_off_source.discount,
+ source_line=paid_amount_off_source.source_line,
+ redeemed_order=OrderFactory.create(state=OrderStatus.FULFILLED),
+ )
+
+ response = client.get(f"{reverse('refund-order')}?order={order.id}")
+
+ assert response.status_code == 200
+ assert list(response.context["used_source_redemptions"]) == [redemption]
+ assert redemption.redeemed_order.reference_number in response.content.decode()
+
+
+def test_admin_refund_view_keeps_the_credit_warning_on_a_rejected_form(
+ client, admin_user, paid_amount_off_source
+):
+ """The re-rendered form after a validation error still carries the warning."""
+ _login_admin(client, admin_user)
+ order = paid_amount_off_source.source_line.order
+ redemption = DiscountRedemptionFactory.create(
+ redeemed_discount=paid_amount_off_source.discount,
+ source_line=paid_amount_off_source.source_line,
+ redeemed_order=OrderFactory.create(state=OrderStatus.FULFILLED),
+ )
+
+ response = client.post(
+ reverse("refund-order"),
+ data={"order": str(order.id), "_selected_action": str(order.id)},
+ )
+
+ assert response.status_code == 200
+ assert response.context["form_valid"] is False
+ assert list(response.context["used_source_redemptions"]) == [redemption]
+ assert redemption.redeemed_order.reference_number in response.content.decode()
diff --git a/ecommerce/api.py b/ecommerce/api.py
index fd68b5d926..954e86bfd3 100644
--- a/ecommerce/api.py
+++ b/ecommerce/api.py
@@ -63,6 +63,10 @@
STRIPE_TRANSACTION_REASON_INITIAL_CHECKOUTSESSION,
ZERO_PAYMENT_DATA,
)
+from ecommerce.discount_sources import (
+ double_spent_source_line_ids,
+ fulfilled_paid_amount_off_redemptions,
+)
from ecommerce.exceptions import (
VerifiedProgramInvalidBasketError,
VerifiedProgramInvalidOrderError,
@@ -275,7 +279,7 @@ def generate_checkout_payload( # noqa: PLR0911, C901
return payload
-def check_discount_for_products(discount, basket):
+def check_discount_for_products(discount, basket, products=None):
"""
Checks the validity of the discount against what's in the basket.
@@ -286,13 +290,14 @@ def check_discount_for_products(discount, basket):
Args:
- basket (Basket): the current basket
- discount (Discount|string: the discount to apply (if a string, loads the discount code specified)
+ - products (list or None): basket.get_products(), for a caller that already has it
Returns:
boolean
"""
if not isinstance(discount, Discount):
discount = Discount.objects.filter(discount_code=discount).first()
- basket_products = basket.get_products()
+ basket_products = basket.get_products() if products is None else products
return discount.check_validity_with_products(basket_products)
@@ -307,10 +312,14 @@ def check_basket_discounts_for_validity(request):
"""
basket = establish_basket(request)
+ basket_products = basket.get_products()
+
for basket_discount in basket.discounts.all():
if not basket_discount.redeemed_discount.is_redeemable_by(
- basket.user
- ) or not check_discount_for_products(basket_discount.redeemed_discount, basket):
+ basket.user, basket_products
+ ) or not check_discount_for_products(
+ basket_discount.redeemed_discount, basket, basket_products
+ ):
return False
return True
@@ -358,10 +367,11 @@ def apply_user_discounts(request):
discount = user_discount.discount
if discount:
+ basket_products = basket.get_products()
# check for product specificity in the discount
if not check_discount_for_products(
- discount, basket
- ) or not discount.is_redeemable_by(user):
+ discount, basket, basket_products
+ ) or not discount.is_redeemable_by(user, basket_products):
return
bd = BasketDiscount(
@@ -955,8 +965,8 @@ def check_and_process_pending_orders_for_resolution(
def check_for_duplicate_discount_redemptions():
"""
- Checks for multiple redemptions for discount codes, and makes noise if there
- are any.
+ Checks for multiple redemptions for discount codes, and makes noise if
+ there are any.
For discounts that are one-time or one-time-per-user redemptions, there's a
possibility that the code can be redeemed more than once. This will check
@@ -1027,6 +1037,34 @@ def check_for_duplicate_discount_redemptions():
return seen
+def check_for_double_spent_sources():
+ """
+ The safety net behind OrderFlow.fulfill's source check: log every source
+ line funding a fulfilled paid-amount-off redemption on more than one order,
+ naming the orders to review.
+
+ Returns:
+ - List of the double-spent source line IDs
+ """
+ double_spent = double_spent_source_line_ids()
+ for source_line_id in double_spent:
+ reference_numbers = (
+ fulfilled_paid_amount_off_redemptions()
+ .filter(source_line_id=source_line_id)
+ .order_by("redeemed_order__reference_number")
+ .values_list("redeemed_order__reference_number", flat=True)
+ .distinct()
+ )
+ log.error(
+ "Line %s funds fulfilled paid-amount-off redemptions on more than one "
+ "order (%s); review manually.",
+ source_line_id,
+ ", ".join(reference_numbers),
+ )
+
+ return double_spent
+
+
def _coerce_supplied_date(value):
"""
Normalize a date reaching generate_discount_code either as a management
diff --git a/ecommerce/api_test.py b/ecommerce/api_test.py
index 9b2c8a1d66..5d706da5d7 100644
--- a/ecommerce/api_test.py
+++ b/ecommerce/api_test.py
@@ -1,5 +1,6 @@
"""Tests for Ecommerce api"""
+import logging
import random
import uuid
from datetime import datetime, timedelta
@@ -32,7 +33,10 @@
ANONYMOUS_BASKET_SESSION_KEY,
_retrieve_pending_cybersource_orders,
apply_discount_to_basket,
+ apply_user_discounts,
check_and_process_pending_orders_for_resolution,
+ check_basket_discounts_for_validity,
+ check_for_double_spent_sources,
check_for_duplicate_discount_redemptions,
claim_anonymous_basket,
create_verified_program_course_run_enrollment,
@@ -41,6 +45,7 @@
downgrade_learner_from_order,
establish_basket,
establish_basket_for_request,
+ fulfill_completed_order,
generate_checkout_payload,
get_anonymous_basket_id,
get_auto_apply_discounts_for_basket,
@@ -72,6 +77,7 @@
STRIPE_PAYMENT_STATUS_UNPAID,
TRANSACTION_TYPE_PAYMENT,
TRANSACTION_TYPE_REFUND,
+ ZERO_PAYMENT_DATA,
)
from ecommerce.exceptions import (
VerifiedProgramNoEnrollmentError,
@@ -82,9 +88,11 @@
OneTimeDiscountFactory,
OneTimePerUserDiscountFactory,
OrderFactory,
+ PaidAmountOffDiscountFactory,
ProductFactory,
TransactionFactory,
UnlimitedUseDiscountFactory,
+ make_purchase,
)
from ecommerce.fixtures import (
stripe_checkout_session,
@@ -810,6 +818,79 @@ def test_check_and_process_pending_orders_options(mocker):
mocked_create_enrollments.assert_not_called()
+def _pending_credit_order(user, source_line):
+ """
+ A pending program order carrying a paid-amount-off redemption funded by
+ source_line: the shape checkout leaves behind before payment. The line's
+ price is irrelevant to these tests; only the redemption's FK matters.
+ """
+ line = make_purchase(
+ user, ProgramFactory.create(), Decimal("500.00"), state=OrderStatus.PENDING
+ )
+ DiscountRedemption.objects.create(
+ redemption_date=now_in_utc(),
+ redeemed_by=user,
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_order=line.order,
+ source_line=source_line,
+ )
+ return line.order
+
+
+def test_fulfillment_logs_a_double_spent_source_and_proceeds(
+ paid_amount_off_source, mocker, caplog
+):
+ """Two pending orders share one source; the second still fulfills, loudly."""
+ mocker.patch("ecommerce.api.sync_hubspot_deal")
+ # Enrollment side effects are another test's concern.
+ mocker.patch("ecommerce.models.OrderFlow.create_enrollments")
+ source_line = paid_amount_off_source.source_line
+ first = _pending_credit_order(paid_amount_off_source.user, source_line)
+ second = _pending_credit_order(paid_amount_off_source.user, source_line)
+
+ with caplog.at_level(logging.ERROR, logger="ecommerce.discount_sources"):
+ # Factory setup logs at INFO before at_level narrows the capture
+ # handler, so the records already collected are not fulfillment's.
+ caplog.clear()
+ fulfill_completed_order(first, ZERO_PAYMENT_DATA)
+ assert [
+ r for r in caplog.records if r.name == "ecommerce.discount_sources"
+ ] == []
+
+ fulfill_completed_order(second, ZERO_PAYMENT_DATA)
+
+ first.refresh_from_db()
+ second.refresh_from_db()
+ assert first.state == OrderStatus.FULFILLED
+ assert second.state == OrderStatus.FULFILLED
+ assert first.reference_number in caplog.text
+ assert second.reference_number in caplog.text
+ assert str(source_line.id) in caplog.text
+
+
+def test_fulfillment_logs_a_source_refunded_after_pricing(
+ paid_amount_off_source, mocker, caplog
+):
+ """The credit is already baked into the price, so log the released source
+ and let the payment settle.
+ """
+ mocker.patch("ecommerce.api.sync_hubspot_deal")
+ mocker.patch("ecommerce.models.OrderFlow.create_enrollments")
+ source_line = paid_amount_off_source.source_line
+ order = _pending_credit_order(paid_amount_off_source.user, source_line)
+ Order.objects.filter(pk=source_line.order_id).update(state=OrderStatus.REFUNDED)
+
+ with caplog.at_level(logging.ERROR, logger="ecommerce.discount_sources"):
+ caplog.clear()
+ fulfill_completed_order(order, ZERO_PAYMENT_DATA)
+
+ order.refresh_from_db()
+ assert order.state == OrderStatus.FULFILLED
+ assert order.reference_number in caplog.text
+ assert source_line.order.reference_number in caplog.text
+ assert str(source_line.id) in caplog.text
+
+
@pytest.mark.parametrize("peruser", [True, False])
def test_duplicate_redemption_check(peruser):
"""
@@ -845,6 +926,32 @@ def make_stuff(user, discount):
assert discount.id in seen_ids
+def test_duplicate_redemption_monitor_flags_shared_source_lines(
+ paid_amount_off_source, caplog
+):
+ """The safety net reports a source funding two fulfilled redemptions and
+ names the orders to review.
+ """
+ source_line = paid_amount_off_source.source_line
+ orders = OrderFactory.create_batch(2, state=OrderStatus.FULFILLED)
+ for order in orders:
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_order=order,
+ source_line=source_line,
+ )
+
+ with caplog.at_level(logging.ERROR, logger="ecommerce.api"):
+ caplog.clear()
+ double_spent_ids = check_for_double_spent_sources()
+
+ assert double_spent_ids == [source_line.id]
+ assert str(source_line.id) in caplog.text
+ for order in orders:
+ assert order.reference_number in caplog.text
+
+
def test_create_verified_program_discount():
"""Test that creating a special discount for programs works"""
@@ -1888,3 +1995,50 @@ def test_retrieve_pending_cs_orders(mocker, test_type):
mocked_cs_gateway.assert_called()
assert len(completed.keys()) == (0 if test_type == "cancelled" else 1)
assert len(cancelled.keys()) == (0 if test_type == "completed" else 1)
+
+
+def test_apply_user_discounts_validates_against_the_whole_basket(
+ paid_amount_off_source,
+):
+ """A user discount linked to the second basket item is applied, not refused against the first."""
+ request = RequestFactory().get("/")
+ request.user = paid_amount_off_source.user
+ basket = Basket.objects.create(user=paid_amount_off_source.user)
+ BasketItem.objects.create(
+ basket=basket, product=ProductFactory.create(), quantity=1
+ )
+ BasketItem.objects.create(
+ basket=basket, product=paid_amount_off_source.program_product, quantity=1
+ )
+ UserDiscount.objects.create(
+ user=paid_amount_off_source.user, discount=paid_amount_off_source.discount
+ )
+
+ apply_user_discounts(request)
+
+ assert basket.discounts.get().redeemed_discount == paid_amount_off_source.discount
+
+
+def test_revalidation_passes_a_resolvable_program_child_purchase_discount(
+ paid_amount_off_source,
+):
+ """
+ Both revalidation call sites hand the basket's products to is_redeemable_by.
+ Without them a program-child-purchase discount fails closed, and a False from
+ check_basket_discounts_for_validity wipes every basket discount and blocks
+ checkout.
+ """
+ request = RequestFactory().get("/")
+ request.user = paid_amount_off_source.user
+ basket = Basket.objects.create(user=paid_amount_off_source.user)
+ BasketItem.objects.create(
+ basket=basket, product=paid_amount_off_source.program_product, quantity=1
+ )
+ UserDiscount.objects.create(
+ user=paid_amount_off_source.user, discount=paid_amount_off_source.discount
+ )
+
+ apply_user_discounts(request)
+
+ assert basket.discounts.get().redeemed_discount == paid_amount_off_source.discount
+ assert check_basket_discounts_for_validity(request) is True
diff --git a/ecommerce/conftest.py b/ecommerce/conftest.py
index b1ab575b12..422c3f5e72 100644
--- a/ecommerce/conftest.py
+++ b/ecommerce/conftest.py
@@ -1,8 +1,48 @@
"""Common fixtures for ecommerce tests"""
+from decimal import Decimal
+from types import SimpleNamespace
+
import pytest
+import reversion
+from reversion.models import Version
+
+from courses.factories import CourseRunFactory, ProgramFactory
+from ecommerce.factories import (
+ PaidAmountOffDiscountFactory,
+ ProgramProductFactory,
+ make_purchase,
+)
+from ecommerce.models import DiscountProduct
@pytest.fixture(autouse=True)
def mocked_hubspot_deal_sync(mocker):
return mocker.patch("hubspot_sync.task_helpers.sync_hubspot_deal")
+
+
+@pytest.fixture
+def paid_amount_off_source(user):
+ """
+ One learner holding exactly one qualifying source: a $999 program product,
+ a $100 paid run of a direct child course, and a paid-amount-off discount
+ linked to the program product. Resolving it returns a 100.00 credit, so the
+ program prices at 899.00.
+ """
+ program = ProgramFactory.create()
+ run = CourseRunFactory.create()
+ program.add_requirement(run.course)
+ source_line = make_purchase(user, run, Decimal("100.00"))
+ with reversion.create_revision():
+ program_product = ProgramProductFactory.create(
+ purchasable_object=program, price=Decimal("999.00")
+ )
+ discount = PaidAmountOffDiscountFactory.create()
+ DiscountProduct.objects.create(discount=discount, product=program_product)
+ return SimpleNamespace(
+ user=user,
+ program_product=program_product,
+ program_product_version=Version.objects.get_for_object(program_product).first(),
+ source_line=source_line,
+ discount=discount,
+ )
diff --git a/ecommerce/discount_sources.py b/ecommerce/discount_sources.py
new file mode 100644
index 0000000000..e6b5e19104
--- /dev/null
+++ b/ecommerce/discount_sources.py
@@ -0,0 +1,414 @@
+"""
+Source resolution for program-child-purchase redemptions (hq#11846, "complete
+your program"): a learner who already paid for a course or sub-program in a
+program's own requirement tree qualifies, and a paid-amount-off discount takes
+what they paid off the program purchase. The redemption type decides
+eligibility; only the paid-amount-off discount type gets a resolved amount, a
+persisted source line, and consumes its source.
+
+This module owns eligibility and value. The calculation layer
+(ecommerce/discounts.py) stays user-blind and receives the resolved amount
+via DiscountType.get_discounted_price(..., resolved_amounts=...). Order paths
+never re-resolve: PendingOrder persists the winning line on
+DiscountRedemption.source_line and resolved_amounts_from_redemptions reads it
+back.
+
+Sources are read from Line + Order.state directly, not PaidCourseRun /
+PaidProgram: those rows are deleted on learner self-unenroll and skipped by
+fulfill(skip_fulfillment=True), while the fulfilled order and line remain.
+
+Two things about the module's shape. It imports ecommerce.models at module
+level, so the model layer reaches it through function-local imports (the same
+arrangement ecommerce.discounts has); the orchestration that would let the
+models stop calling it lives in ecommerce.api and is not moved here. And the
+second half of the module is not eligibility or value at all but ledger
+queries over persisted redemptions (source conflicts, released sources,
+funded credits, double spends); they share the source-line vocabulary and the
+paid-amount-off filter, which is why they are colocated.
+"""
+
+from __future__ import annotations
+
+import logging
+from dataclasses import dataclass
+from decimal import Decimal # noqa: TC003
+
+from django.contrib.contenttypes.models import ContentType
+from django.db.models import (
+ Count,
+ Exists,
+ OuterRef,
+ Q,
+ QuerySet,
+ prefetch_related_objects,
+)
+
+from courses.models import (
+ CourseRun,
+ Program,
+ ProgramRequirement,
+ ProgramRequirementNodeType,
+)
+from ecommerce.constants import DISCOUNT_TYPE_PAID_AMOUNT_OFF
+from ecommerce.models import DiscountRedemption, Line, OrderStatus
+
+log = logging.getLogger(__name__)
+
+
+@dataclass(frozen=True)
+class SourceResolution:
+ """The single prior purchase that qualifies a program-child-purchase redemption."""
+
+ source_line: Line
+ amount: Decimal
+
+
+def spends_source(discount) -> bool:
+ """
+ Whether ``discount`` is funded by, and consumes, a specific source line.
+
+ Only a paid-amount-off discount does; a program-child-purchase redemption
+ on any other discount type is eligibility-only and records no source. The
+ fulfilled-redemptions queryset below is the ORM spelling of the same rule.
+ """
+ return discount.discount_type == DISCOUNT_TYPE_PAID_AMOUNT_OFF
+
+
+def fulfilled_paid_amount_off_redemptions() -> QuerySet[DiscountRedemption]:
+ """
+ Every FULFILLED redemption of a paid-amount-off discount.
+
+ The discount type is filtered here rather than trusted from the FK: tying
+ source_line to the discount type would need a constraint spanning Discount
+ and DiscountRedemption, which Postgres cannot express, so a standard
+ redemption carrying a source_line is legal at the DB level.
+ """
+ return DiscountRedemption.objects.filter(
+ redeemed_discount__discount_type=DISCOUNT_TYPE_PAID_AMOUNT_OFF,
+ redeemed_order__state=OrderStatus.FULFILLED,
+ )
+
+
+def _program_for_product(product) -> Program | None:
+ """The program ``product`` sells, if it is an eligible current purchase."""
+ program = product.purchasable_object
+ if not isinstance(program, Program):
+ return None
+ # B2B programs are priced by contract, not by this offer. b2b_only alone
+ # is incomplete: contract-linked programs are not required to set it.
+ if program.b2b_only or program.contract_memberships.exists():
+ return None
+ return program
+
+
+def resolve_program_child_purchase(
+ user, product, *, exclude_consumed=True
+) -> SourceResolution | None:
+ """
+ The most expensive unconsumed qualifying purchase of a course or
+ sub-program in ``product``'s program's own requirement tree, or None.
+
+ Qualifying: a Line on one of the user's FULFILLED orders whose purchased
+ run or program is not B2B, for a run of a course in the tree or for a
+ sub-program in the tree, at any depth but not through a nested program's
+ own tree, with a recorded price above zero. Consumed: the
+ line already funds a paid-amount-off redemption on a FULFILLED order
+ (pending/abandoned checkouts never burn a source). Only a paid-amount-off
+ discount spends its source, so resolve_for_discount passes
+ exclude_consumed=False for any other discount type with the
+ program-child-purchase redemption.
+ """
+ if user is None or user.is_anonymous:
+ return None
+ program = _program_for_product(product)
+ if program is None:
+ return None
+
+ # Every row of this program's requirement tree carries this program's FK,
+ # whatever its depth, so this matches a course or sub-program anywhere in
+ # the tree. A nested program's own courses are rows of that program's tree
+ # and never match. Both queries span required and elective children, because
+ # the operator node they hang off is the parent, not the row matched here.
+ #
+ # Program.courses and Program.program_nodes are the equivalent accessors;
+ # the ids are queried directly instead because nothing here needs hydrated
+ # objects, and values_list keeps the whole thing a subquery the source-line
+ # filter can join against.
+ child_course_ids = ProgramRequirement.objects.filter(
+ program=program,
+ node_type=ProgramRequirementNodeType.COURSE,
+ course__isnull=False,
+ ).values_list("course_id", flat=True)
+ child_program_ids = ProgramRequirement.objects.filter(
+ program=program,
+ node_type=ProgramRequirementNodeType.PROGRAM,
+ required_program__isnull=False,
+ ).values_list("required_program_id", flat=True)
+
+ # B2B purchases never fund the offer. A run is B2B through its contract;
+ # a program through either marker, the same test _program_for_product
+ # applies to the target.
+ source_run_ids = CourseRun.objects.filter(
+ course_id__in=child_course_ids, b2b_contract__isnull=True
+ ).values_list("id", flat=True)
+ source_program_ids = Program.objects.filter(
+ id__in=child_program_ids, b2b_only=False, contract_memberships__isnull=True
+ ).values_list("id", flat=True)
+
+ run_ct = ContentType.objects.get_for_model(CourseRun)
+ program_ct = ContentType.objects.get_for_model(Program)
+
+ candidates = Line.objects.filter(
+ order__purchaser=user,
+ order__state=OrderStatus.FULFILLED,
+ # discounted_unit_price is the price frozen when the source order
+ # was priced (non-null); a $0 purchase has nothing to credit.
+ discounted_unit_price__gt=0,
+ ).filter(
+ Q(purchased_content_type=run_ct, purchased_object_id__in=source_run_ids)
+ | Q(
+ purchased_content_type=program_ct,
+ purchased_object_id__in=source_program_ids,
+ )
+ )
+ if exclude_consumed:
+ # A multi-condition exclude() over funded_redemptions compiles to one
+ # EXISTS per condition, so the type and the state would be checked on
+ # different redemption rows (Django docs, "Spanning multi-valued
+ # relationships"). A correlated subquery holds both on the same row.
+ candidates = candidates.exclude(
+ Exists(
+ fulfilled_paid_amount_off_redemptions().filter(
+ source_line=OuterRef("pk")
+ )
+ )
+ )
+ # Sibling courses often share a price; the id tie-break keeps the choice
+ # stable across the eligibility and pricing resolves of one request.
+ source_line = candidates.order_by("-discounted_unit_price", "-id").first()
+ if source_line is None:
+ return None
+ return SourceResolution(
+ source_line=source_line,
+ amount=source_line.get_discounted_unit_price(),
+ )
+
+
+def resolve_for_discount(discount, user, products) -> SourceResolution | None:
+ """
+ Resolve for the first of ``products`` that ``discount`` is linked to, or
+ None when it is linked to none of them.
+
+ The eligibility guard, the basket display and order creation all resolve
+ through here, so a discount cannot validate against one line and price
+ another. A program-child-purchase discount names its programs through its
+ links, so an unlinked one has nothing to resolve for.
+
+ The pricing sites apply a discount to every linked line, so a multi-item
+ basket that holds more than one linked program gets this one resolution on
+ each of them; hq#12815 narrows pricing to the same single product.
+ """
+ # Several discounts resolve in one request, and a lazy read of the links per
+ # discount trips the N+1 guard (zeal); an explicit prefetch does not.
+ prefetch_related_objects([discount], "products")
+ linked_product_ids = {link.product_id for link in discount.products.all()}
+ product = next(
+ (product for product in products if product.id in linked_product_ids), None
+ )
+ if product is None:
+ return None
+ return resolve_program_child_purchase(
+ user, product, exclude_consumed=spends_source(discount)
+ )
+
+
+def source_line_for(discount, user, products) -> Line | None:
+ """
+ The prior purchase line that funds ``discount`` for ``user``, to persist on
+ the redemption at order creation; None for a discount that spends no
+ source. Resolving for the same target product the eligibility guard uses
+ keeps the persisted source from disagreeing with is_redeemable_by.
+ """
+ if not spends_source(discount):
+ return None
+ resolution = resolve_for_discount(discount, user, products)
+ return resolution.source_line if resolution else None
+
+
+def has_paid_amount_off(discounts) -> bool:
+ """Whether any of ``discounts`` spends a source, i.e. needs resolving."""
+ return any(spends_source(discount) for discount in discounts)
+
+
+def resolved_amounts_for_user(user, discounts, products) -> dict[int, Decimal]:
+ """
+ {discount_id: amount} for the paid-amount-off discounts in ``discounts``.
+
+ Each discount resolves for its own target line, so two paid-amount-off
+ discounts covering different programs in one basket get different answers. For
+ user/basket pricing paths; order paths must use
+ resolved_amounts_from_redemptions so the price stays frozen.
+ """
+ resolutions = {
+ discount.id: resolve_for_discount(discount, user, products)
+ for discount in discounts
+ if spends_source(discount)
+ }
+ return {
+ discount_id: resolution.amount
+ for discount_id, resolution in resolutions.items()
+ if resolution is not None
+ }
+
+
+def resolved_amounts_from_redemptions(redemptions) -> dict[int, Decimal]:
+ """
+ {discount_id: amount} read back from persisted source_line FKs — zero
+ resolver queries, frozen at order creation.
+
+ Callers should select_related("redeemed_discount", "source_line"):
+ redeemed_discount is read on every row, source_line on the paid-amount-off
+ rows.
+
+ Keying on source_line_id alone would be wrong: a standard redemption may
+ legally carry one, and DiscountType.for_discount raises TypeError when a
+ resolved amount reaches a discount type that has no field for it.
+ """
+ return {
+ redemption.redeemed_discount_id: redemption.source_line.get_discounted_unit_price()
+ for redemption in redemptions
+ if redemption.source_line_id is not None
+ and spends_source(redemption.redeemed_discount)
+ }
+
+
+def paid_amount_off_source_line_ids(order) -> list[int]:
+ """Ids of the source lines ``order``'s paid-amount-off redemptions were priced from."""
+ return [
+ redemption.source_line_id
+ for redemption in order.discounts.select_related("redeemed_discount")
+ if redemption.source_line_id is not None
+ and spends_source(redemption.redeemed_discount)
+ ]
+
+
+def find_source_conflict(order, source_line_ids=None) -> DiscountRedemption | None:
+ """
+ A FULFILLED redemption on another order that already consumed one of the
+ source lines ``order`` was priced from, or None.
+
+ Called at fulfillment time: nothing re-validates between pricing and
+ payment, so one source line can otherwise deterministically fund two
+ different program purchases. ``source_line_ids`` is
+ paid_amount_off_source_line_ids(order), for a caller that already has it.
+ """
+ if source_line_ids is None:
+ source_line_ids = paid_amount_off_source_line_ids(order)
+ if not source_line_ids:
+ return None
+ return (
+ fulfilled_paid_amount_off_redemptions()
+ .filter(source_line_id__in=source_line_ids)
+ .exclude(redeemed_order=order)
+ .select_related("redeemed_order")
+ .first()
+ )
+
+
+def released_source_lines(order) -> QuerySet[Line]:
+ """
+ The source lines ``order`` was priced from whose own order is no longer
+ FULFILLED — refunded between pricing and payment, so the credit baked into
+ this order's price rests on a purchase the learner no longer holds.
+ """
+ return (
+ Line.objects.filter(
+ funded_redemptions__redeemed_order=order,
+ funded_redemptions__redeemed_discount__discount_type=DISCOUNT_TYPE_PAID_AMOUNT_OFF,
+ )
+ .exclude(order__state=OrderStatus.FULFILLED)
+ .select_related("order")
+ .distinct()
+ )
+
+
+def log_source_anomalies(order) -> None:
+ """
+ At fulfillment, log the two states the double-spend window can leave
+ ``order`` in: a source line another FULFILLED order already spent, or a
+ source line whose own order was refunded between pricing and payment.
+
+ Nothing re-validates between pricing and payment, and the single-cart
+ checkout makes either state rare, so the credit is honored and the error is
+ made Sentry-visible for a manual refund rather than blocking a payment that
+ already went through. Each message names both orders and the source line.
+ """
+ source_line_ids = paid_amount_off_source_line_ids(order)
+ if not source_line_ids:
+ return
+ conflict = find_source_conflict(order, source_line_ids)
+ if conflict is not None:
+ log.error(
+ "Order %s fulfilled with paid-amount-off source line %s that "
+ "already funds fulfilled order %s — double credit honored; "
+ "review manually.",
+ order.reference_number,
+ conflict.source_line_id,
+ conflict.redeemed_order.reference_number,
+ )
+ for line in released_source_lines(order):
+ log.error(
+ "Order %s fulfilled with paid-amount-off source line %s whose "
+ "own order %s is no longer fulfilled — credit rests on a "
+ "purchase the learner no longer holds; review manually.",
+ order.reference_number,
+ line.id,
+ line.order.reference_number,
+ )
+
+
+def _fulfilled_redemptions_funded_by(order) -> QuerySet[DiscountRedemption]:
+ """``order`` may be an Order or an OuterRef to one."""
+ return (
+ fulfilled_paid_amount_off_redemptions()
+ .filter(source_line__order=order)
+ .exclude(redeemed_order=order)
+ )
+
+
+def fulfilled_redemptions_funded_by(order) -> QuerySet[DiscountRedemption]:
+ """
+ FULFILLED paid-amount-off redemptions on other orders that a line of
+ ``order`` funded — the credits that survive refunding ``order``, since
+ OrderFlow.refund touches only transactions.
+ """
+ # Callers that read rows name the order the credit went to; an exists()
+ # caller pays nothing for the join.
+ return _fulfilled_redemptions_funded_by(order).select_related("redeemed_order")
+
+
+def funds_fulfilled_redemption_exists() -> Exists:
+ """
+ The question ``Order.funds_fulfilled_redemption`` asks, as an ``Exists``
+ for annotating an Order queryset under that same name, so a page of orders
+ answers it in the page query instead of once per row.
+ """
+ return Exists(_fulfilled_redemptions_funded_by(OuterRef("pk")))
+
+
+def double_spent_source_line_ids() -> list[int]:
+ """
+ Source lines funding a FULFILLED paid-amount-off redemption on more than one
+ order — the state the fulfillment-time check logs and honors.
+
+ Counted over distinct orders, not rows: two paid-amount-off discounts on
+ one order resolving the same line are two rows but one spend.
+ """
+ return list(
+ fulfilled_paid_amount_off_redemptions()
+ .filter(source_line__isnull=False)
+ .values("source_line")
+ .annotate(order_count=Count("redeemed_order", distinct=True))
+ .filter(order_count__gt=1)
+ .values_list("source_line", flat=True)
+ )
diff --git a/ecommerce/discount_sources_test.py b/ecommerce/discount_sources_test.py
new file mode 100644
index 0000000000..cfab934d5e
--- /dev/null
+++ b/ecommerce/discount_sources_test.py
@@ -0,0 +1,528 @@
+from decimal import Decimal
+
+import pytest
+import reversion
+from django.contrib.auth.models import AnonymousUser
+
+from b2b.factories import ContractPageFactory
+from b2b.models import ContractProgramItem
+from courses.factories import CourseRunFactory, ProgramFactory
+from ecommerce.constants import (
+ DISCOUNT_TYPE_PERCENT_OFF,
+ REDEMPTION_TYPE_PROGRAM_CHILD_PURCHASE,
+)
+from ecommerce.discount_sources import (
+ double_spent_source_line_ids,
+ find_source_conflict,
+ funds_fulfilled_redemption_exists,
+ released_source_lines,
+ resolve_for_discount,
+ resolve_program_child_purchase,
+ resolved_amounts_for_user,
+ resolved_amounts_from_redemptions,
+)
+from ecommerce.factories import (
+ DiscountFactory,
+ DiscountRedemptionFactory,
+ OrderFactory,
+ PaidAmountOffDiscountFactory,
+ ProductFactory,
+ ProgramProductFactory,
+ make_purchase,
+)
+from ecommerce.models import DiscountProduct, Order, OrderStatus
+
+pytestmark = [pytest.mark.django_db]
+
+
+@pytest.fixture
+def program():
+ return ProgramFactory.create()
+
+
+@pytest.fixture
+def program_product(program):
+ with reversion.create_revision():
+ return ProgramProductFactory.create(purchasable_object=program)
+
+
+def test_resolves_a_fulfilled_child_course_purchase(user, program, program_product):
+ """A paid run of a direct child course funds the discount at what was paid."""
+ run = CourseRunFactory.create()
+ program.add_requirement(run.course)
+ line = make_purchase(user, run, Decimal("100.00"), paid=Decimal("75.00"))
+
+ resolution = resolve_program_child_purchase(user, program_product)
+
+ assert resolution.source_line == line
+ assert resolution.amount == Decimal("75.00")
+
+
+def test_resolves_a_child_program_purchase(user, program, program_product):
+ """A purchased vertical (child program) is a valid source."""
+ vertical = ProgramFactory.create()
+ program.add_requirement(vertical)
+ line = make_purchase(user, vertical, Decimal("300.00"))
+
+ assert resolve_program_child_purchase(user, program_product).source_line == line
+
+
+def test_an_elective_child_course_qualifies(user, program, program_product):
+ """Electives are children too: eligibility is one tree level, not one operator."""
+ run = CourseRunFactory.create()
+ program.add_elective(run.course)
+ line = make_purchase(user, run, Decimal("100.00"))
+
+ assert resolve_program_child_purchase(user, program_product).source_line == line
+
+
+def test_picks_the_most_expensive_qualifying_source(user, program, program_product):
+ """The credit is one purchase, the most expensive, not the sum."""
+ run_a, run_b = CourseRunFactory.create_batch(2)
+ program.add_requirement(run_a.course)
+ program.add_requirement(run_b.course)
+ make_purchase(user, run_a, Decimal("100.00"))
+ best = make_purchase(user, run_b, Decimal("150.00"))
+
+ assert resolve_program_child_purchase(user, program_product).source_line == best
+
+
+def test_only_fulfilled_source_orders_qualify(user, program, program_product):
+ """A refunded (or never-paid) purchase is not a source."""
+ run = CourseRunFactory.create()
+ program.add_requirement(run.course)
+ make_purchase(user, run, Decimal("100.00"), state=OrderStatus.REFUNDED)
+
+ assert resolve_program_child_purchase(user, program_product) is None
+
+
+def test_equal_prices_break_the_tie_on_the_newest_line(user, program, program_product):
+ """Sibling courses often share a price; the choice must be stable, not arbitrary."""
+ run_a, run_b = CourseRunFactory.create_batch(2)
+ program.add_requirement(run_a.course)
+ program.add_requirement(run_b.course)
+ make_purchase(user, run_a, Decimal("100.00"))
+ newest = make_purchase(user, run_b, Decimal("100.00"))
+
+ assert resolve_program_child_purchase(user, program_product).source_line == newest
+
+
+def test_a_free_source_credits_nothing(user, program, program_product):
+ """A $0 purchase (a 100% code, a free enrollment) is not a source."""
+ run = CourseRunFactory.create()
+ program.add_requirement(run.course)
+ make_purchase(user, run, Decimal("0.00"))
+
+ assert resolve_program_child_purchase(user, program_product) is None
+
+
+def test_a_consumed_source_does_not_resolve_again(user, program, program_product):
+ """A source funds at most one redemption; only FULFILLED consumption counts."""
+ run = CourseRunFactory.create()
+ program.add_requirement(run.course)
+ line = make_purchase(user, run, Decimal("100.00"))
+ redemption = DiscountRedemptionFactory.create(
+ redeemed_by=user,
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_order=OrderFactory.create(purchaser=user, state=OrderStatus.PENDING),
+ source_line=line,
+ )
+ # An abandoned checkout never burns the source...
+ assert resolve_program_child_purchase(user, program_product).source_line == line
+
+ redemption.redeemed_order.state = OrderStatus.FULFILLED
+ redemption.redeemed_order.save()
+
+ # ...but a fulfilled one does.
+ assert resolve_program_child_purchase(user, program_product) is None
+
+
+def test_a_standard_discount_never_consumes_a_source(user, program, program_product):
+ """Nothing at the DB level keeps source_line off a percent-off redemption, so
+ the consumed-source exclusion filters the discount type itself.
+ """
+ run = CourseRunFactory.create()
+ program.add_requirement(run.course)
+ line = make_purchase(user, run, Decimal("100.00"))
+ DiscountRedemptionFactory.create(
+ redeemed_by=user,
+ redeemed_discount=DiscountFactory.create(),
+ redeemed_order=OrderFactory.create(purchaser=user, state=OrderStatus.FULFILLED),
+ source_line=line,
+ )
+
+ assert resolve_program_child_purchase(user, program_product).source_line == line
+
+
+def test_consumption_is_a_single_redemption_row(user, program, program_product):
+ """
+ A pending paid-amount-off hold and a separate standard redemption on a
+ fulfilled order are not, together, a spend: the type and the state must
+ hold on the same redemption row.
+ """
+ run = CourseRunFactory.create()
+ program.add_requirement(run.course)
+ line = make_purchase(user, run, Decimal("100.00"))
+ DiscountRedemptionFactory.create(
+ redeemed_by=user,
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_order=OrderFactory.create(purchaser=user, state=OrderStatus.PENDING),
+ source_line=line,
+ )
+ DiscountRedemptionFactory.create(
+ redeemed_by=user,
+ redeemed_discount=DiscountFactory.create(),
+ redeemed_order=OrderFactory.create(purchaser=user, state=OrderStatus.FULFILLED),
+ source_line=line,
+ )
+
+ assert resolve_program_child_purchase(user, program_product).source_line == line
+
+
+def test_only_paid_amount_off_eligibility_needs_an_unconsumed_source(
+ paid_amount_off_source,
+):
+ """
+ Consumption is a paid-amount-off concept: once a source funds a fulfilled
+ paid-amount-off order no paid-amount-off discount resolves it again, but a
+ percent-off discount with the program-child-purchase redemption type still
+ qualifies on it — that offer takes nothing from the purchase.
+ """
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=paid_amount_off_source.discount,
+ redeemed_order=OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.FULFILLED
+ ),
+ source_line=paid_amount_off_source.source_line,
+ )
+ percent_off = DiscountFactory.create(
+ discount_type=DISCOUNT_TYPE_PERCENT_OFF,
+ redemption_type=REDEMPTION_TYPE_PROGRAM_CHILD_PURCHASE,
+ automatic=True,
+ )
+ DiscountProduct.objects.create(
+ discount=percent_off, product=paid_amount_off_source.program_product
+ )
+ products = [paid_amount_off_source.program_product]
+
+ assert (
+ resolve_for_discount(
+ paid_amount_off_source.discount, paid_amount_off_source.user, products
+ )
+ is None
+ )
+ assert (
+ resolve_for_discount(
+ percent_off, paid_amount_off_source.user, products
+ ).source_line
+ == paid_amount_off_source.source_line
+ )
+
+
+def test_a_credit_resolves_for_the_first_linked_product(paid_amount_off_source):
+ """
+ The credit is resolved for the first product the discount is linked to:
+ unlinked products ahead of it are skipped, and a linked product ahead of it
+ decides the answer even when a later one would have qualified.
+ """
+ source = paid_amount_off_source
+ unlinked_product = ProductFactory.create()
+ sourceless_program_product = ProgramProductFactory.create()
+ DiscountProduct.objects.create(
+ discount=source.discount, product=sourceless_program_product
+ )
+
+ assert (
+ resolve_for_discount(
+ source.discount, source.user, [unlinked_product, source.program_product]
+ ).source_line
+ == source.source_line
+ )
+ assert (
+ resolve_for_discount(
+ source.discount,
+ source.user,
+ [sourceless_program_product, source.program_product],
+ )
+ is None
+ )
+
+
+def test_grandchild_courses_do_not_qualify(user, program, program_product):
+ """A nested program's courses are rows of its own tree, not of the parent's."""
+ vertical = ProgramFactory.create()
+ program.add_requirement(vertical)
+ run = CourseRunFactory.create()
+ vertical.add_requirement(run.course)
+ make_purchase(user, run, Decimal("100.00"))
+
+ assert resolve_program_child_purchase(user, program_product) is None
+
+
+def test_b2b_run_sources_do_not_qualify(user, program, program_product):
+ run = CourseRunFactory.create(b2b_contract=ContractPageFactory.create())
+ program.add_requirement(run.course)
+ make_purchase(user, run, Decimal("100.00"))
+
+ assert resolve_program_child_purchase(user, program_product) is None
+
+
+@pytest.mark.parametrize("field", ["b2b_only", "contract"])
+def test_b2b_child_program_sources_do_not_qualify(
+ user, program, program_product, field
+):
+ """A B2B vertical is bought under contract terms, not this offer."""
+ vertical = ProgramFactory.create(b2b_only=field == "b2b_only")
+ if field == "contract":
+ ContractProgramItem.objects.create(
+ contract=ContractPageFactory.create(), program=vertical
+ )
+ program.add_requirement(vertical)
+ make_purchase(user, vertical, Decimal("300.00"))
+
+ assert resolve_program_child_purchase(user, program_product) is None
+
+
+def test_non_program_products_do_not_resolve(user):
+ run_product = ProductFactory.create()
+
+ assert resolve_program_child_purchase(user, run_product) is None
+
+
+@pytest.mark.parametrize("field", ["b2b_only", "contract"])
+def test_b2b_programs_do_not_resolve(user, field):
+ """A B2B program is priced by contract, so it is never offered the credit."""
+ program = ProgramFactory.create(b2b_only=field == "b2b_only")
+ if field == "contract":
+ ContractProgramItem.objects.create(
+ contract=ContractPageFactory.create(), program=program
+ )
+ with reversion.create_revision():
+ product = ProgramProductFactory.create(purchasable_object=program)
+ run = CourseRunFactory.create()
+ program.add_requirement(run.course)
+ make_purchase(user, run, Decimal("100.00"))
+
+ assert resolve_program_child_purchase(user, product) is None
+
+
+def test_anonymous_users_do_not_resolve(program_product):
+ assert resolve_program_child_purchase(AnonymousUser(), program_product) is None
+ assert resolve_program_child_purchase(None, program_product) is None
+
+
+def test_resolved_amounts_for_user_keys_by_discount(paid_amount_off_source):
+ """Each linked discount resolves over its own links, so a discount that is not
+ linked to the product in hand gets no amount rather than the neighbour's.
+ """
+ unlinked = PaidAmountOffDiscountFactory.create()
+
+ amounts = resolved_amounts_for_user(
+ paid_amount_off_source.user,
+ [paid_amount_off_source.discount, unlinked],
+ [paid_amount_off_source.program_product],
+ )
+
+ assert amounts == {paid_amount_off_source.discount.id: Decimal("100.00")}
+
+
+def test_resolved_amounts_for_user_is_empty_without_paid_amount_off_discounts(
+ user, program_product
+):
+ """A standard discount gets no resolved amount."""
+ assert (
+ resolved_amounts_for_user(user, [DiscountFactory.create()], [program_product])
+ == {}
+ )
+
+
+def test_resolved_amounts_from_redemptions_reads_the_frozen_fk(user):
+ """Order paths price from the persisted line, never re-resolving."""
+ line = make_purchase(user, CourseRunFactory.create(), Decimal("80.00"))
+ discount = PaidAmountOffDiscountFactory.create()
+ redemption = DiscountRedemptionFactory.create(
+ redeemed_discount=discount, redeemed_by=user, source_line=line
+ )
+ sourceless = DiscountRedemptionFactory.create(redeemed_by=user)
+
+ amounts = resolved_amounts_from_redemptions([redemption, sourceless])
+
+ assert amounts == {discount.id: Decimal("80.00")}
+
+
+def test_resolved_amounts_from_redemptions_ignores_standard_discounts(user):
+ """The discount type, not the FK, decides who carries a resolved amount."""
+ line = make_purchase(user, CourseRunFactory.create(), Decimal("80.00"))
+ redemption = DiscountRedemptionFactory.create(
+ redeemed_discount=DiscountFactory.create(), redeemed_by=user, source_line=line
+ )
+
+ assert resolved_amounts_from_redemptions([redemption]) == {}
+
+
+def _redemption_on(order, source_line, discount=None):
+ return DiscountRedemptionFactory.create(
+ redeemed_by=order.purchaser,
+ redeemed_discount=discount or PaidAmountOffDiscountFactory.create(),
+ redeemed_order=order,
+ source_line=source_line,
+ )
+
+
+def test_find_source_conflict_reports_a_fulfilled_competitor(paid_amount_off_source):
+ """Another fulfilled order already spent this source: that redemption is the conflict."""
+ source_line = paid_amount_off_source.source_line
+ order = OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.FULFILLED
+ )
+ _redemption_on(order, source_line, paid_amount_off_source.discount)
+ competitor = _redemption_on(
+ OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.FULFILLED
+ ),
+ source_line,
+ )
+ # A standard redemption on a fulfilled order may carry the same line and is
+ # not a competitor.
+ _redemption_on(
+ OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.FULFILLED
+ ),
+ source_line,
+ DiscountFactory.create(),
+ )
+
+ assert find_source_conflict(order) == competitor
+
+
+def test_released_source_lines_lists_each_refunded_source_once(
+ paid_amount_off_source,
+):
+ """
+ A source whose order was refunded after pricing is reported once, however
+ many paid-amount-off redemptions on this order it funds; a standard
+ redemption's source_line is not a source at all.
+ """
+ order = OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.PENDING
+ )
+ refunded = paid_amount_off_source.source_line
+ refunded.order.state = OrderStatus.REFUNDED
+ refunded.order.save()
+ _redemption_on(order, refunded)
+ _redemption_on(order, refunded)
+ standard_source = make_purchase(
+ paid_amount_off_source.user,
+ CourseRunFactory.create(),
+ Decimal("50.00"),
+ state=OrderStatus.REFUNDED,
+ )
+ _redemption_on(order, standard_source, DiscountFactory.create())
+
+ assert list(released_source_lines(order)) == [refunded]
+
+
+def test_find_source_conflict_ignores_a_pending_competitor(paid_amount_off_source):
+ """An abandoned checkout holding the same source is not a double spend."""
+ order = OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.PENDING
+ )
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=paid_amount_off_source.discount,
+ redeemed_order=order,
+ source_line=paid_amount_off_source.source_line,
+ )
+ competitor = OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.PENDING
+ )
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_order=competitor,
+ source_line=paid_amount_off_source.source_line,
+ )
+
+ assert find_source_conflict(order) is None
+
+
+def test_find_source_conflict_ignores_the_orders_own_redemption(paid_amount_off_source):
+ """The order being fulfilled always holds the source itself — that is not a
+ conflict, and this is the case the check runs against on every fulfillment.
+ """
+ order = OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.FULFILLED
+ )
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=paid_amount_off_source.discount,
+ redeemed_order=order,
+ source_line=paid_amount_off_source.source_line,
+ )
+
+ assert find_source_conflict(order) is None
+
+
+def test_double_spent_source_line_ids_ignores_standard_redemptions(
+ paid_amount_off_source,
+):
+ """A standard redemption may legally carry a source_line, but only a
+ paid-amount-off redemption spends one, so the pair is not a double spend.
+ """
+ source_line = paid_amount_off_source.source_line
+ for discount in (PaidAmountOffDiscountFactory.create(), DiscountFactory.create()):
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=discount,
+ redeemed_order=OrderFactory.create(state=OrderStatus.FULFILLED),
+ source_line=source_line,
+ )
+
+ assert double_spent_source_line_ids() == []
+
+
+def test_double_spent_source_line_ids_counts_orders_not_rows(paid_amount_off_source):
+ """Two redemptions on one fulfilled order are a re-priced order, not two spends."""
+ order = OrderFactory.create(state=OrderStatus.FULFILLED)
+ for _ in range(2):
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_order=order,
+ source_line=paid_amount_off_source.source_line,
+ )
+
+ assert double_spent_source_line_ids() == []
+
+
+def test_funds_fulfilled_redemption_exists_matches_the_property(
+ user, django_assert_num_queries
+):
+ """
+ The annotation answers exactly what Order.funds_fulfilled_redemption does,
+ and an annotated instance reads it without a query of its own.
+ """
+ funder = make_purchase(user, CourseRunFactory.create(), Decimal("100.00")).order
+ bystander = make_purchase(user, CourseRunFactory.create(), Decimal("100.00")).order
+ consumer = OrderFactory.create(state=OrderStatus.FULFILLED)
+ DiscountRedemptionFactory.create(
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ source_line=funder.lines.first(),
+ redeemed_order=consumer,
+ )
+
+ annotated = list(
+ Order.objects.annotate(
+ funds_fulfilled_redemption=funds_fulfilled_redemption_exists()
+ )
+ )
+ with django_assert_num_queries(0):
+ by_id = {order.id: order.funds_fulfilled_redemption for order in annotated}
+
+ assert by_id[funder.id] is True
+ assert by_id[bystander.id] is False
+ assert by_id == {
+ order.id: Order.objects.get(id=order.id).funds_fulfilled_redemption
+ for order in annotated
+ }
diff --git a/ecommerce/factories.py b/ecommerce/factories.py
index fb85e8f148..aa27c4a032 100644
--- a/ecommerce/factories.py
+++ b/ecommerce/factories.py
@@ -1,6 +1,8 @@
import faker
+import reversion
from factory import LazyAttribute, SubFactory, fuzzy
from factory.django import DjangoModelFactory
+from reversion.models import Version
from courses.factories import CourseRunFactory, ProgramFactory
from ecommerce import models
@@ -140,3 +142,32 @@ class DiscountRedemptionFactory(DjangoModelFactory):
class Meta:
model = models.DiscountRedemption
+
+
+def make_purchase(
+ user,
+ purchasable,
+ price,
+ *,
+ paid=None,
+ state=models.OrderStatus.FULFILLED,
+):
+ """
+ An order for ``purchasable`` with one line listed at ``price`` and charged
+ ``paid`` (defaults to ``price``). Returns the Line.
+
+ The Product is created inside a revision because Line.product_version is
+ non-null and reversion only records a Version for objects saved under one.
+ """
+ with reversion.create_revision():
+ product = ProductFactory.create(purchasable_object=purchasable, price=price)
+ charged = price if paid is None else paid
+ order = OrderFactory.create(purchaser=user, state=state, total_price_paid=charged)
+ return models.Line.objects.create(
+ order=order,
+ purchased_object_id=product.object_id,
+ purchased_content_type_id=product.content_type_id,
+ product_version=Version.objects.get_for_object(product).first(),
+ quantity=1,
+ discounted_unit_price=charged,
+ )
diff --git a/ecommerce/models.py b/ecommerce/models.py
index 915f352596..b22fb562be 100644
--- a/ecommerce/models.py
+++ b/ecommerce/models.py
@@ -1,6 +1,7 @@
from __future__ import annotations
import uuid
+from collections.abc import Iterable # noqa: TC003
from datetime import datetime, timedelta
from decimal import Decimal
from typing import List # noqa: UP035
@@ -240,8 +241,13 @@ class BasketItem(TimestampedModel):
@cached_property
def discounted_price(self):
"""Return the price of the product with discounts"""
+ from ecommerce.discount_sources import ( # noqa: PLC0415
+ has_paid_amount_off,
+ resolved_amounts_for_user,
+ )
from ecommerce.discounts import DiscountType # noqa: PLC0415
+ products = self.basket.get_products()
discounts = [
discount_redemption.redeemed_discount
for discount_redemption in self.basket.discounts.prefetch_related(
@@ -253,6 +259,14 @@ def discounted_price(self):
DiscountType.get_discounted_price(
discounts,
self.product,
+ # basket.user is a lazy FK: touch it only when a paid-amount-off
+ # discount is on the basket, so an ordinary cart pays no
+ # resolver cost at all.
+ resolved_amounts=(
+ resolved_amounts_for_user(self.basket.user, discounts, products)
+ if has_paid_amount_off(discounts)
+ else {}
+ ),
).quantize(Decimal("0.01"))
* self.quantity
)
@@ -442,24 +456,40 @@ def is_redeemed(self) -> bool:
"""Returns True if the discount has been redeemed"""
return DiscountRedemption.objects.filter(redeemed_discount=self).exists()
- def is_redeemable_by(self, user: User):
+ def is_redeemable_by(self, user: User, products: Iterable[Product] | None = None):
"""
- Enforces the redemption rules for a given discount.
+ Enforces the redemption rules for a given discount: how often it may be
+ redeemed, whether it is inside its date window, and — for a
+ program-child-purchase redemption — whether this user still holds an
+ unconsumed qualifying purchase for one of the products in hand.
+
+ Independent of check_validity_with_products (product scope and
+ liveness); is_valid_for_basket composes the two.
+
+ ``products`` is context for the program-child-purchase arm alone.
+ Omitting it means no product is in hand (a code redemption, the CMS
+ finaid quote), and a program-child-purchase discount is never
+ redeemable there.
Args:
- user (User): The user requesting the discount.
+ - products (Iterable[Product] or None): the products the discount is
+ being checked against — the basket's, or the one being priced.
Returns:
- boolean
"""
- # A program-child-purchase discount is redeemable only by a learner
- # holding the qualifying purchase; that eligibility is decided per source
- # line — a question this method has no way to answer until the resolver
- # lands (hq#11846). Refuse rather than fall through to the unlimited
- # semantics at the bottom, which would let any code-redemption endpoint
- # attach one.
if self.redemption_type == REDEMPTION_TYPE_PROGRAM_CHILD_PURCHASE:
- return False
+ from ecommerce.discount_sources import resolve_for_discount # noqa: PLC0415
+ if products is None:
+ return False
+ if resolve_for_discount(self, user, products) is None:
+ return False
+
+ return self._within_redemption_limits(user)
+
+ def _within_redemption_limits(self, user: User) -> bool:
+ """The redemption-count and date-window rules, without the source check."""
if (
self.redemption_type == REDEMPTION_TYPE_ONE_TIME
and DiscountRedemption.objects.filter(
@@ -528,8 +558,9 @@ def is_valid_for_basket(self, basket, *, allow_finaid=False) -> bool:
Check if the discount is valid for the basket.
Performs the finaid gate and user-tied-discount checks, then delegates
- product scope to check_validity_with_products and the redemption-limit
- and date-window rules to is_redeemable_by.
+ product scope to check_validity_with_products and the redemption-limit,
+ date-window and program-child-purchase eligibility rules to
+ is_redeemable_by.
Financial assistance discounts are excluded by default, because this
check is used for discount codes that are submitted by the user, and
@@ -559,11 +590,13 @@ def _discount_user_has_discount() -> bool:
or self.user_discount_discount.filter(user=basket.user).count() > 0
)
+ if not allow_finaid and self.payment_type == PAYMENT_TYPE_FINANCIAL_ASSISTANCE:
+ return False
+ products = basket.get_products()
return (
- (allow_finaid or self.payment_type != PAYMENT_TYPE_FINANCIAL_ASSISTANCE)
- and self.check_validity_with_products(basket.get_products())
+ self.check_validity_with_products(products)
and _discount_user_has_discount()
- and self.is_redeemable_by(basket.user)
+ and self.is_redeemable_by(basket.user, products)
)
def friendly_format(self):
@@ -583,22 +616,42 @@ def friendly_format(self):
def discount_product(self, product, user=None):
"""
- Returns the calculated discount amount for a given product.
+ Returns the price of ``product`` after this discount.
Args:
product (Product): the product to discount
user (User or None): the current user
Returns:
- Number; the calculated amount of the discounts
+ Decimal; the discounted price, or None when ``user`` may not redeem
+ this discount for ``product``
"""
+ from ecommerce.discount_sources import ( # noqa: PLC0415
+ resolve_for_discount,
+ spends_source,
+ )
from ecommerce.discounts import DiscountType # noqa: PLC0415
- if (user is None and self.valid_now()) or self.is_redeemable_by(user):
- return DiscountType.get_discounted_price([self], product).quantize(
- Decimal("0.01")
- )
-
- return None
+ resolved_amounts = {}
+ if user is None:
+ # No user bypasses the eligibility arm and resolves no amount, so a
+ # program-child-purchase discount quotes full price here; only
+ # callers that have already gated the discount on a real user may
+ # pass None.
+ if not self.valid_now():
+ return None
+ else:
+ if not self._within_redemption_limits(user):
+ return None
+ if self.redemption_type == REDEMPTION_TYPE_PROGRAM_CHILD_PURCHASE:
+ # One resolve serves both the eligibility check and the amount.
+ resolution = resolve_for_discount(self, user, [product])
+ if resolution is None:
+ return None
+ if spends_source(self):
+ resolved_amounts = {self.id: resolution.amount}
+ return DiscountType.get_discounted_price(
+ [self], product, resolved_amounts=resolved_amounts
+ ).quantize(Decimal("0.01"))
def b2b_contracts(self):
"""Return the applicable B2B contract(s), if any."""
@@ -698,6 +751,7 @@ class OrderRefundStatus(TextChoices):
REQUESTED = "requested"
DENIED = "denied"
ELIGIBLE = "eligible"
+ REVIEW_REQUIRED = "review_required"
WINDOW_CLOSED = "window_closed"
INELIGIBLE = "ineligible"
@@ -870,10 +924,18 @@ def fulfill(
skip_fulfillment=False, # noqa: FBT002
):
"""Fulfill the order - create a transaction, send email, trigger plugins."""
+ from ecommerce.discount_sources import log_source_anomalies # noqa: PLC0415
# record the transaction
self.create_transaction(payment_data)
+ # Monitoring only: it logs and never raises, which is what keeps a
+ # charged order from being stranded (viewflow restores the initial
+ # state on any exception in a transition body, wherever it happens).
+ # Running after the transaction just keeps the Transaction row on a
+ # query failure.
+ log_source_anomalies(self.order)
+
# record all the courseruns in the order (unless we're told not to)
if not skip_fulfillment:
self.create_enrollments()
@@ -968,7 +1030,12 @@ def is_within_refund_window(self):
@property
def is_refund_eligible(self):
- """Return True if the learner could request a refund for this order now."""
+ """
+ True when `refund_status` is `eligible`: the in-window self-service case.
+
+ A `review_required` or `window_closed` order is still submittable
+ through the free-text request path; see `refund_status`.
+ """
return self.refund_status == OrderRefundStatus.ELIGIBLE
@cached_property
@@ -984,6 +1051,23 @@ def is_b2b_order(self):
"""
return any(run.b2b_contract_id for run in self.purchased_runs)
+ @cached_property
+ def funds_fulfilled_redemption(self):
+ """
+ True when a line of this order funds a paid-amount-off redemption on a
+ fulfilled order (hq#11846).
+
+ One query per order. Querysets that serialize many orders annotate
+ this name with ``funds_fulfilled_redemption_exists()`` instead: the
+ annotation lands in the instance dict, which is where a cached_property
+ reads from, so the per-order query never runs.
+ """
+ from ecommerce.discount_sources import ( # noqa: PLC0415
+ fulfilled_redemptions_funded_by,
+ )
+
+ return fulfilled_redemptions_funded_by(self).exists()
+
@cached_property
def latest_refund_request(self):
"""Return the learner's most recent refund request for this order, if any."""
@@ -1020,12 +1104,15 @@ def refund_status(self):
Return where this order sits in the self-service refund flow.
Precedence matters: a refund that already happened settles the question,
- then any request the learner has made, and only then whether they could
- make one right now.
-
- The final branch deliberately mirrors what `RefundRequestSerializer`
- accepts, so anything but `eligible` or `window_closed` means a request
- would be rejected.
+ then any request the learner has made, then whether the order is
+ refundable at all, then whether a person has to review it, and only then
+ the window.
+
+ `eligible` means the in-window request form with preset reasons.
+ `review_required` and `window_closed` are both still submittable through
+ the free-text path — `RefundRequestSerializer` gates on ownership,
+ fulfilled state, B2B and an existing pending request, never on the
+ window — and land in the manual queue instead.
"""
if self.state in (OrderStatus.REFUNDED, OrderStatus.PARTIALLY_REFUNDED):
return OrderRefundStatus.COMPLETED
@@ -1043,6 +1130,12 @@ def refund_status(self):
if self.state != OrderStatus.FULFILLED or self.is_b2b_order:
return OrderRefundStatus.INELIGIBLE
+ # Refunding this order would leave the credit it funded in place —
+ # clawback is out of scope — so the request is accepted but reviewed
+ # by a person regardless of the window.
+ if self.funds_fulfilled_redemption:
+ return OrderRefundStatus.REVIEW_REQUIRED
+
return (
OrderRefundStatus.ELIGIBLE
if self.is_within_refund_window
@@ -1150,6 +1243,8 @@ def _get_or_create(
# Apply any discounts to the PendingOrder
if discounts:
+ from ecommerce.discount_sources import source_line_for # noqa: PLC0415
+
now = now_in_utc()
for discount in discounts:
if discount:
@@ -1157,6 +1252,7 @@ def _get_or_create(
redemption_date=now,
redeemed_by=user,
redeemed_discount=discount,
+ source_line=source_line_for(discount, user, products),
)
# Create or get Line for each product. Calculate the Order total based on Lines and discount.
@@ -1405,16 +1501,22 @@ def total_price(self):
@staticmethod
def compute_discounted_unit_price_for(order, product_version):
"""Price of one unit of product_version under the discounts currently on order."""
+ from ecommerce.discount_sources import ( # noqa: PLC0415
+ resolved_amounts_from_redemptions,
+ )
from ecommerce.discounts import DiscountType # noqa: PLC0415
- discounts = [
- discount_redemption.redeemed_discount
- for discount_redemption in order.discounts.all()
- ]
+ # source_line is joined for the paid-amount-off rows rather than
+ # fetched lazily per row.
+ redemptions = list(
+ order.discounts.select_related("redeemed_discount", "source_line")
+ )
+ discounts = [redemption.redeemed_discount for redemption in redemptions]
return DiscountType.get_discounted_price(
discounts,
_product_from_version(product_version),
+ resolved_amounts=resolved_amounts_from_redemptions(redemptions),
).quantize(Decimal("0.01"))
def compute_discounted_unit_price(self):
diff --git a/ecommerce/models_test.py b/ecommerce/models_test.py
index 9a710e29bd..9b1d012832 100644
--- a/ecommerce/models_test.py
+++ b/ecommerce/models_test.py
@@ -6,15 +6,16 @@
import pytest
import reversion
from django.core.exceptions import ValidationError
-from django.db import IntegrityError, transaction
+from django.db import IntegrityError, connection, transaction
from django.db.models import ProtectedError
+from django.test.utils import CaptureQueriesContext
from django.urls import reverse
from freezegun import freeze_time
from mitol.common.utils import now_in_utc
from reversion.models import Version
from b2b.factories import ContractPageFactory
-from courses.factories import CourseRunFactory
+from courses.factories import CourseRunFactory, ProgramFactory
from ecommerce.constants import (
DISCOUNT_TYPE_DOLLARS_OFF,
DISCOUNT_TYPE_FIXED_PRICE,
@@ -25,6 +26,7 @@
REFUND_WINDOW_DAYS,
ZERO_PAYMENT_DATA,
)
+from ecommerce.discount_sources import resolve_program_child_purchase
from ecommerce.factories import (
BasketFactory,
BasketItemFactory,
@@ -39,6 +41,7 @@
ProgramProductFactory,
SetLimitDiscountFactory,
UnlimitedUseDiscountFactory,
+ make_purchase,
)
from ecommerce.fixtures import stripe_event
from ecommerce.models import (
@@ -1306,10 +1309,10 @@ def test_refund_window_extends_for_a_course_starting_after_purchase():
assert not order.is_within_refund_window
-def test_is_refund_eligible_only_when_a_request_would_be_accepted():
+def test_is_refund_eligible_means_the_in_window_self_service_case():
"""
- `is_refund_eligible` answers "could the learner request a refund now", not
- "is the window open" — being in window is necessary but not sufficient.
+ `is_refund_eligible` is `refund_status == eligible`, not "is the window
+ open" — being in window is necessary but not sufficient.
"""
fulfilled = OrderFactory.create(state=OrderStatus.FULFILLED)
assert fulfilled.is_refund_eligible
@@ -1446,6 +1449,56 @@ def test_refund_status_completed_outranks_everything(user, state):
assert order.refund_status == OrderRefundStatus.COMPLETED
+def test_refund_status_review_required_when_the_order_funded_a_credit():
+ """Refunding it would keep the credit alive — those requests need a human."""
+ line = _line_for(Decimal("100.00"))
+ order = line.order
+
+ assert order.refund_status == OrderRefundStatus.ELIGIBLE
+
+ redemption = DiscountRedemptionFactory.create(
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ source_line=line,
+ redeemed_order=OrderFactory.create(state=OrderStatus.PENDING),
+ )
+ # A pending consumer has not consumed anything yet.
+ assert order.refund_status == OrderRefundStatus.ELIGIBLE
+
+ redemption.redeemed_order.state = OrderStatus.FULFILLED
+ redemption.redeemed_order.save()
+
+ # funds_fulfilled_redemption is cached per instance, so ask a fresh one.
+ order = Order.objects.get(id=order.id)
+ assert order.refund_status == OrderRefundStatus.REVIEW_REQUIRED
+ assert order.is_refund_eligible is False
+
+ # Review outranks the window: window_closed would hide the credit from support.
+ with freeze_time(order.created_on + timedelta(days=REFUND_WINDOW_DAYS, seconds=1)):
+ assert order.refund_status == OrderRefundStatus.REVIEW_REQUIRED
+
+
+def test_refund_status_ignores_redemptions_that_spent_nothing_of_this_order():
+ """
+ Only a paid-amount-off redemption on another order counts: a standard
+ redemption may legally carry a source_line, and an order's own redemption is
+ not a credit it funded for someone else.
+ """
+ line = _line_for(Decimal("100.00"))
+ order = line.order
+
+ DiscountRedemptionFactory.create(
+ source_line=line,
+ redeemed_order=OrderFactory.create(state=OrderStatus.FULFILLED),
+ )
+ DiscountRedemptionFactory.create(
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ source_line=line,
+ redeemed_order=order,
+ )
+
+ assert order.refund_status == OrderRefundStatus.ELIGIBLE
+
+
def test_refund_reviewed_on_is_none_while_pending(user):
"""A request nobody has acted on has no review date."""
order = OrderFactory.create(purchaser=user, state=OrderStatus.FULFILLED)
@@ -1608,19 +1661,370 @@ def test_program_child_purchase_discount_tolerates_a_product_less_link_row():
discount.clean()
-def test_program_child_purchase_discount_is_not_redeemable_by_anyone(user):
+def test_program_child_purchase_discount_is_not_redeemable_without_products(user):
"""
- Eligibility is per qualifying purchase and has no resolver yet, so the
- generic redemption check has to refuse rather than treat the discount as
- unlimited-use.
+ Without product context there is nothing to resolve a source against, and a
+ program-child-purchase discount is automatic-only — so the code-redemption
+ path and the CMS finaid quote, which pass no products, can never attach one.
"""
discount = PaidAmountOffDiscountFactory.create()
assert discount.is_redeemable_by(user) is False
+def test_discount_product_quotes_the_resolved_credit(paid_amount_off_source):
+ """Discount.discount_product carries the per-user resolution."""
+ quoted = paid_amount_off_source.discount.discount_product(
+ paid_amount_off_source.program_product, paid_amount_off_source.user
+ )
+
+ assert quoted == Decimal("899.00")
+
+
+def test_discount_product_quotes_full_price_without_a_user(paid_amount_off_source):
+ """With no user there is nothing to resolve against, so the quote is the list price."""
+ quoted = paid_amount_off_source.discount.discount_product(
+ paid_amount_off_source.program_product
+ )
+
+ assert quoted == Decimal("999.00")
+
+
+def test_discount_product_declines_without_a_source(paid_amount_off_source):
+ """No source, no quote — the guard and the price agree."""
+ quoted = paid_amount_off_source.discount.discount_product(
+ paid_amount_off_source.program_product, UserFactory.create()
+ )
+
+ assert quoted is None
+
+
+def test_is_redeemable_by_requires_a_resolvable_source(paid_amount_off_source):
+ """A program-child-purchase discount is redeemable only with an available source."""
+ discount = paid_amount_off_source.discount
+ products = [paid_amount_off_source.program_product]
+
+ assert discount.is_redeemable_by(paid_amount_off_source.user, products) is True
+ assert discount.is_redeemable_by(UserFactory.create(), products) is False
+
+
+def test_is_redeemable_by_still_honors_max_redemptions(paid_amount_off_source):
+ """The program-child-purchase guard falls through to the limit checks, not past them."""
+ discount = paid_amount_off_source.discount
+ discount.max_redemptions = 1
+ discount.save()
+ DiscountRedemptionFactory.create(
+ redeemed_discount=discount,
+ redeemed_order=OrderFactory.create(state=OrderStatus.FULFILLED),
+ )
+
+ assert (
+ discount.is_redeemable_by(
+ paid_amount_off_source.user, [paid_amount_off_source.program_product]
+ )
+ is False
+ )
+
+
+def test_is_redeemable_by_fails_closed_without_linked_products(paid_amount_off_source):
+ """A program-child-purchase discount with no DiscountProduct rows resolves nothing."""
+ unlinked = PaidAmountOffDiscountFactory.create()
+
+ assert (
+ unlinked.is_redeemable_by(
+ paid_amount_off_source.user, [paid_amount_off_source.program_product]
+ )
+ is False
+ )
+
+
+def test_is_redeemable_by_checks_the_product_in_hand_not_every_link(
+ paid_amount_off_source,
+):
+ """A source for one linked program does not unlock a different linked program."""
+ with reversion.create_revision():
+ other_program_product = ProgramProductFactory.create()
+ DiscountProduct.objects.create(
+ discount=paid_amount_off_source.discount, product=other_program_product
+ )
+
+ assert (
+ paid_amount_off_source.discount.is_redeemable_by(
+ paid_amount_off_source.user, [other_program_product]
+ )
+ is False
+ )
+
+
+def test_is_redeemable_by_keys_eligibility_on_the_redemption_type(
+ paid_amount_off_source,
+):
+ """
+ A percent-off discount with the program-child-purchase redemption type is
+ gated by the same source check: the arm reads the redemption type, so a
+ stranger is refused instead of falling through to unlimited semantics.
+ """
+ percent_off = DiscountFactory.create(
+ discount_type=DISCOUNT_TYPE_PERCENT_OFF,
+ redemption_type=REDEMPTION_TYPE_PROGRAM_CHILD_PURCHASE,
+ automatic=True,
+ )
+ DiscountProduct.objects.create(
+ discount=percent_off, product=paid_amount_off_source.program_product
+ )
+
+ assert (
+ percent_off.is_redeemable_by(
+ paid_amount_off_source.user, [paid_amount_off_source.program_product]
+ )
+ is True
+ )
+
+
+def test_is_valid_for_basket_inherits_the_program_child_purchase_guard(
+ paid_amount_off_source,
+):
+ """
+ The auto-apply/attach path hands the basket's products to the guard: the
+ learner holding the source passes, a stranger does not.
+ """
+ own_basket = BasketFactory.create(user=paid_amount_off_source.user)
+ BasketItem.objects.create(
+ basket=own_basket, product=paid_amount_off_source.program_product, quantity=1
+ )
+ stranger_basket = BasketFactory.create(user=UserFactory.create())
+ BasketItem.objects.create(
+ basket=stranger_basket,
+ product=paid_amount_off_source.program_product,
+ quantity=1,
+ )
+
+ assert paid_amount_off_source.discount.is_valid_for_basket(own_basket) is True
+ assert paid_amount_off_source.discount.is_valid_for_basket(stranger_basket) is False
+
+
def test_friendly_format_for_paid_amount_off():
"""The label carries no amount — the true value is per-user."""
discount = PaidAmountOffDiscountFactory.create()
assert discount.friendly_format() == "the amount paid for a prior purchase"
+
+
+def test_basket_pricing_applies_the_resolved_paid_amount_off_credit(
+ paid_amount_off_source,
+):
+ """The cart shows the program at price minus what the child purchase cost."""
+ basket = BasketFactory.create(user=paid_amount_off_source.user)
+ item = BasketItem.objects.create(
+ basket=basket, product=paid_amount_off_source.program_product, quantity=1
+ )
+ BasketDiscount.objects.create(
+ redemption_date=now_in_utc(),
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=paid_amount_off_source.discount,
+ redeemed_basket=basket,
+ )
+
+ assert item.discounted_price == Decimal("899.00")
+
+
+def test_an_ordinary_basket_never_touches_the_resolver(user):
+ """Without a paid-amount-off discount, pricing does not even load the user."""
+ basket = BasketFactory.create(user=user)
+ item = BasketItem.objects.create(
+ basket=basket, product=ProductFactory.create(), quantity=1
+ )
+ BasketDiscount.objects.create(
+ redemption_date=now_in_utc(),
+ redeemed_by=user,
+ redeemed_discount=DiscountFactory.create(),
+ redeemed_basket=basket,
+ )
+ item = BasketItem.objects.get(id=item.id)
+
+ with CaptureQueriesContext(connection) as queries:
+ item.discounted_price # noqa: B018
+
+ assert not [q["sql"] for q in queries if "users_user" in q["sql"]]
+
+
+def test_line_pricing_reads_the_persisted_source_line(paid_amount_off_source):
+ """
+ Order pricing uses the frozen FK, not a fresh resolve: the source stays
+ credited even after another fulfilled order has consumed it.
+ """
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_order=OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.FULFILLED
+ ),
+ source_line=paid_amount_off_source.source_line,
+ )
+ order = OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.PENDING
+ )
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=paid_amount_off_source.discount,
+ redeemed_order=order,
+ source_line=paid_amount_off_source.source_line,
+ )
+
+ assert Line.compute_discounted_unit_price_for(
+ order, paid_amount_off_source.program_product_version
+ ) == Decimal("899.00")
+
+
+def test_a_source_line_on_a_standard_discount_is_ignored(paid_amount_off_source):
+ """A percent-off redemption may legally carry a source_line; pricing must
+ ignore it rather than hand a resolved amount to a type with no field for one.
+ """
+ order = OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.PENDING
+ )
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=DiscountFactory.create(
+ amount=Decimal("10"), discount_type=DISCOUNT_TYPE_PERCENT_OFF
+ ),
+ redeemed_order=order,
+ source_line=paid_amount_off_source.source_line,
+ )
+
+ assert Line.compute_discounted_unit_price_for(
+ order, paid_amount_off_source.program_product_version
+ ) == Decimal("899.10")
+
+
+def test_pending_order_persists_the_resolved_source_line(paid_amount_off_source):
+ """The redemption freezes which purchase funded it, and pricing uses it."""
+ order = PendingOrder.create_from_product(
+ paid_amount_off_source.program_product,
+ paid_amount_off_source.user,
+ paid_amount_off_source.discount,
+ )
+
+ redemption = order.discounts.get()
+ assert redemption.source_line == paid_amount_off_source.source_line
+ assert order.total_price_paid == Decimal("899.00")
+
+
+def test_pending_order_records_a_source_only_for_paid_amount_off(
+ paid_amount_off_source,
+):
+ """
+ A percent-off discount with the program-child-purchase redemption type is
+ eligibility-only: it prices as percent-off and freezes no source line.
+ """
+ percent_off = DiscountFactory.create(
+ amount=20,
+ discount_type=DISCOUNT_TYPE_PERCENT_OFF,
+ redemption_type=REDEMPTION_TYPE_PROGRAM_CHILD_PURCHASE,
+ automatic=True,
+ )
+ DiscountProduct.objects.create(
+ discount=percent_off, product=paid_amount_off_source.program_product
+ )
+
+ order = PendingOrder.create_from_product(
+ paid_amount_off_source.program_product,
+ paid_amount_off_source.user,
+ percent_off,
+ )
+
+ assert order.discounts.get().source_line is None
+ assert order.total_price_paid == Decimal("799.20")
+
+
+def test_reused_pending_order_re_resolves_the_source(paid_amount_off_source):
+ """Redemptions are deleted and recreated on reuse, so a source consumed in
+ the meantime drops away and the order re-prices to full.
+ """
+ first = PendingOrder.create_from_product(
+ paid_amount_off_source.program_product,
+ paid_amount_off_source.user,
+ paid_amount_off_source.discount,
+ )
+ assert first.discounts.get().source_line == paid_amount_off_source.source_line
+
+ # Another order consumes the source before this checkout completes.
+ DiscountRedemptionFactory.create(
+ redeemed_by=paid_amount_off_source.user,
+ redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_order=OrderFactory.create(
+ purchaser=paid_amount_off_source.user, state=OrderStatus.FULFILLED
+ ),
+ source_line=paid_amount_off_source.source_line,
+ )
+
+ second = PendingOrder.create_from_product(
+ paid_amount_off_source.program_product,
+ paid_amount_off_source.user,
+ paid_amount_off_source.discount,
+ )
+
+ assert second.pk == first.pk
+ assert second.discounts.get().source_line is None
+ assert second.total_price_paid == Decimal("999.00")
+
+
+def test_a_non_fulfilled_program_order_re_arms_the_source(paid_amount_off_source):
+ """Emergent but intended: a redemption on a refunded order stops counting,
+ so the source funds a fresh purchase.
+ """
+ order = PendingOrder.create_from_product(
+ paid_amount_off_source.program_product,
+ paid_amount_off_source.user,
+ paid_amount_off_source.discount,
+ )
+ Order.objects.filter(pk=order.pk).update(state=OrderStatus.FULFILLED)
+
+ assert (
+ resolve_program_child_purchase(
+ paid_amount_off_source.user, paid_amount_off_source.program_product
+ )
+ is None
+ )
+
+ Order.objects.filter(pk=order.pk).update(state=OrderStatus.REFUNDED)
+
+ assert (
+ resolve_program_child_purchase(
+ paid_amount_off_source.user, paid_amount_off_source.program_product
+ ).source_line
+ == paid_amount_off_source.source_line
+ )
+
+
+def test_chaining_credits_each_dollar_at_most_once(user):
+ """Course funds the vertical; the vertical's own paid amount funds the parent."""
+ parent = ProgramFactory.create()
+ vertical = ProgramFactory.create()
+ parent.add_requirement(vertical)
+ run = CourseRunFactory.create()
+ vertical.add_requirement(run.course)
+
+ # Buy the course at $100.
+ make_purchase(user, run, Decimal("100.00"))
+
+ # Buy the vertical ($300 list) with the course credited: pay $200.
+ with reversion.create_revision():
+ vertical_product = ProgramProductFactory.create(
+ purchasable_object=vertical, price=Decimal("300.00")
+ )
+ vertical_discount = PaidAmountOffDiscountFactory.create()
+ DiscountProduct.objects.create(discount=vertical_discount, product=vertical_product)
+ vertical_order = PendingOrder.create_from_product(
+ vertical_product, user, vertical_discount
+ )
+ assert vertical_order.total_price_paid == Decimal("200.00")
+ Order.objects.filter(pk=vertical_order.pk).update(state=OrderStatus.FULFILLED)
+
+ # The parent program credits what was actually paid for the vertical.
+ with reversion.create_revision():
+ parent_product = ProgramProductFactory.create(purchasable_object=parent)
+
+ assert resolve_program_child_purchase(user, parent_product).amount == Decimal(
+ "200.00"
+ )
diff --git a/ecommerce/serializers/__init__.py b/ecommerce/serializers/__init__.py
index 98c956fb8f..e547bbf74c 100644
--- a/ecommerce/serializers/__init__.py
+++ b/ecommerce/serializers/__init__.py
@@ -16,7 +16,6 @@
BULK_GENERATION_REDEMPTION_TYPES,
CYBERSOURCE_CARD_TYPES,
DISCOUNT_TYPE_DOLLARS_OFF,
- DISCOUNT_TYPE_PAID_AMOUNT_OFF,
DISCOUNT_TYPE_PERCENT_OFF,
PAYMENT_TYPES,
TRANSACTION_TYPE_REFUND,
@@ -187,18 +186,17 @@ def discount_is_price_neutral(discount) -> bool:
never move a price.
A percent-off or dollars-off discount of 0 subtracts nothing. A
- paid-amount-off discount stores 0 too, and its real per-user value is
- resolved elsewhere (hq#11846), so nothing renderable exists for it yet;
- the resolver revisits its visibility.
+ paid-amount-off discount stores 0 too but is not in this set: its value is
+ resolved per user, and it only attaches to a basket when a source resolves
+ (Discount.is_redeemable_by), so an attached one always moves the price.
A fixed-price discount is not in this set even at 0 — it sets the price
rather than reducing it, and a fixed price equal to the product price is a
real B2B-contract shape the shopper needs to see confirmed.
"""
- return discount.discount_type == DISCOUNT_TYPE_PAID_AMOUNT_OFF or (
- discount.amount == 0
- and discount.discount_type
- in (DISCOUNT_TYPE_PERCENT_OFF, DISCOUNT_TYPE_DOLLARS_OFF)
+ return discount.amount == 0 and discount.discount_type in (
+ DISCOUNT_TYPE_PERCENT_OFF,
+ DISCOUNT_TYPE_DOLLARS_OFF,
)
diff --git a/ecommerce/serializers/v0/serializers_test.py b/ecommerce/serializers/v0/serializers_test.py
index af272d0367..e68dc8f714 100644
--- a/ecommerce/serializers/v0/serializers_test.py
+++ b/ecommerce/serializers/v0/serializers_test.py
@@ -314,24 +314,24 @@ def test_v0_discount_serializer_rejects_converting_a_discount_with_a_courserun_p
assert "program products" in str(serializer.errors)
-def test_basket_serializer_hides_a_paid_amount_off_discount(user):
+def test_basket_serializer_lists_a_paid_amount_off_discount(user):
"""
- A paid-amount-off discount's stored amount is always 0 and its real
- per-user value is resolved elsewhere (hq#11846), so there is nothing to
- render for it yet.
+ A paid-amount-off discount stores 0 but is not price-neutral: it attaches
+ only when a source resolves, so the shopper must see it explain the price.
"""
basket = BasketFactory.create(user=user)
BasketItemFactory.create(basket=basket)
+ discount = PaidAmountOffDiscountFactory.create()
BasketDiscount.objects.create(
redemption_date=now_in_utc(),
redeemed_by=user,
- redeemed_discount=PaidAmountOffDiscountFactory.create(),
+ redeemed_discount=discount,
redeemed_basket=basket,
)
data = BasketWithProductSerializer(instance=basket).data
- assert data["discounts"] == []
+ assert [d["redeemed_discount"]["id"] for d in data["discounts"]] == [discount.id]
def test_basket_serializer_shows_a_fixed_price_discount_equal_to_the_product_price(
diff --git a/ecommerce/tasks.py b/ecommerce/tasks.py
index 9edc05b945..fcd599daca 100644
--- a/ecommerce/tasks.py
+++ b/ecommerce/tasks.py
@@ -57,9 +57,13 @@ def process_pending_order_resolutions():
@app.task(acks_late=True)
def perform_check_for_duplicate_discount_redemptions():
- from ecommerce.api import check_for_duplicate_discount_redemptions
+ from ecommerce.api import (
+ check_for_double_spent_sources,
+ check_for_duplicate_discount_redemptions,
+ )
check_for_duplicate_discount_redemptions()
+ check_for_double_spent_sources()
@app.task(acks_late=True)
diff --git a/ecommerce/templates/refund_order_confirm.html b/ecommerce/templates/refund_order_confirm.html
index e162857256..2d84b5d042 100644
--- a/ecommerce/templates/refund_order_confirm.html
+++ b/ecommerce/templates/refund_order_confirm.html
@@ -21,6 +21,18 @@ Refund Order
Reference: {{ order.reference_number }}
+ {% if used_source_redemptions %}
+
+ {% for redemption in used_source_redemptions %}
+ -
+ A purchase on this order funded a paid-amount-off discount on order
+ {{ redemption.redeemed_order.reference_number }}. Refunding this order
+ will not claw that discount back.
+
+ {% endfor %}
+
+ {% endif %}
+