diff --git a/squarelet/conftest.py b/squarelet/conftest.py index 020106a14..dffabffcc 100644 --- a/squarelet/conftest.py +++ b/squarelet/conftest.py @@ -22,6 +22,7 @@ ProfessionalPlanFactory, ProfileChangeRequestFactory, SubscriptionFactory, + SubscriptionItemFactory, ) from squarelet.users.tests.factories import UserFactory @@ -38,6 +39,7 @@ register(ProfessionalPlanFactory) register(ProfileChangeRequestFactory) register(SubscriptionFactory) +register(SubscriptionItemFactory) register(CustomerFactory) register(PaymentMethodFactory) diff --git a/squarelet/core/management/commands/export_orgs.py b/squarelet/core/management/commands/export_orgs.py index b4f0d60b4..3b3015948 100644 --- a/squarelet/core/management/commands/export_orgs.py +++ b/squarelet/core/management/commands/export_orgs.py @@ -44,7 +44,7 @@ def handle(self, *args, **kwargs): ) for org in orgs: subtypes = ", ".join(str(s) for s in org.subtypes.all()) - plans = ", ".join(str(p) for p in org.plans.all()) + plans = ", ".join(str(p) for p in org.get_plans()) email_domains = [ e.email.split("@")[1] for e in EmailAddress.objects.filter(user__organizations=org) diff --git a/squarelet/core/management/commands/import_documentcloud.py b/squarelet/core/management/commands/import_documentcloud.py index da5d74a30..e89304d69 100644 --- a/squarelet/core/management/commands/import_documentcloud.py +++ b/squarelet/core/management/commands/import_documentcloud.py @@ -58,7 +58,9 @@ def handle(self, *args, **kwargs): if first_admin: organization.set_billing_email(first_admin.email) if organization.user_count() > organization.max_users: - active_subs = list(organization.subscriptions.select_related("plan")) + active_subs = list( + organization.subscription_items.select_related("plan") + ) paid_subs = [s for s in active_subs if not s.plan.free] if not paid_subs: organization.max_users = organization.user_count() diff --git a/squarelet/core/management/commands/sync_odoo.py b/squarelet/core/management/commands/sync_odoo.py index 2772311f6..a26f4c12c 100644 --- a/squarelet/core/management/commands/sync_odoo.py +++ b/squarelet/core/management/commands/sync_odoo.py @@ -170,7 +170,7 @@ def _build_org_vals(org, odoo_plan_ids, sunlight_status, member_tag_ids): def _compute_org_plans_and_status(org, inherited_plan_ids): """Return (odoo_plan_ids, sunlight_status) for an org.""" - plans = list(org.plans.values_list("name", "wix")) + plans = list(org.get_plans().values_list("name", "wix")) has_sunlight = any(wix for _, wix in plans) own_plan_ids = [ pid for pid in (_resolve_plan_id(name) for name, _ in plans) if pid is not None @@ -292,7 +292,9 @@ def get_or_create_org(org, dry_run=False, member_tag_ids=None, inherited_plan_id def _member_desired_plans(user, org_plan_ids): """Union of the org's inherited plans and the user's own personal plans.""" - personal = list(user.individual_organization.plans.values_list("name", flat=True)) + personal = list( + user.individual_organization.get_plans().values_list("name", flat=True) + ) personal_plan_ids = [ pid for pid in (_resolve_plan_id(name) for name in personal) if pid is not None ] @@ -603,7 +605,7 @@ def _load_collaborative_data(): collab_orgs = Organization.objects.filter( collective_enabled=True, individual=False, - ).prefetch_related("plans", "members") + ).prefetch_related("subscriptions__plans", "members") for collab_org in collab_orgs: tag_id = settings.COLLABORATIVE_TAGS.get(collab_org.slug) if tag_id is None: @@ -619,9 +621,9 @@ def _load_collaborative_data(): pid for pid in ( _resolve_plan_id(name) - for name in collab_org.plans.filter(wix=True).values_list( - "name", flat=True - ) + for name in collab_org.get_plans() + .filter(wix=True) + .values_list("name", flat=True) ) if pid is not None ] @@ -636,7 +638,7 @@ def _build_org_queryset(collaborative_data): """Return the queryset of all orgs to sync.""" sunlight_slugs = set( Organization.objects.filter( - plans__wix=True, + subscriptions__plans__wix=True, individual=False, ) .values_list("slug", flat=True) diff --git a/squarelet/core/tests/test_odoo_sync.py b/squarelet/core/tests/test_odoo_sync.py index c1f2fe72b..961a1efc8 100644 --- a/squarelet/core/tests/test_odoo_sync.py +++ b/squarelet/core/tests/test_odoo_sync.py @@ -213,7 +213,10 @@ class TestComputeOrgPlansAndStatus: def test_confirmed_when_any_wix_plan(self): """An own wix plan sets status Confirmed and resolves all plan ids.""" org = Mock() - org.plans.values_list.return_value = [("Pro", True), ("Free", False)] + org.get_plans.return_value.values_list.return_value = [ + ("Pro", True), + ("Free", False), + ] with patch.object(sync_odoo, "_resolve_plan_id", side_effect=[10, 20]): ids, status = sync_odoo._compute_org_plans_and_status(org, None) assert ids == [10, 20] @@ -222,7 +225,7 @@ def test_confirmed_when_any_wix_plan(self): def test_no_status_without_wix_plan(self): """No own wix plan leaves the sunlight status unset.""" org = Mock() - org.plans.values_list.return_value = [("Free", False)] + org.get_plans.return_value.values_list.return_value = [("Free", False)] with patch.object(sync_odoo, "_resolve_plan_id", return_value=20): _, status = sync_odoo._compute_org_plans_and_status(org, None) assert status is None @@ -230,7 +233,7 @@ def test_no_status_without_wix_plan(self): def test_inherited_plans_merged_and_deduped(self): """Inherited ids are merged with own ids and duplicates removed.""" org = Mock() - org.plans.values_list.return_value = [("Pro", True)] + org.get_plans.return_value.values_list.return_value = [("Pro", True)] with patch.object(sync_odoo, "_resolve_plan_id", return_value=10): ids, _ = sync_odoo._compute_org_plans_and_status(org, [10, 30]) assert ids == [10, 30] @@ -240,7 +243,7 @@ def test_inherited_plans_do_not_set_sunlight_status(self): inherited (collaborative/enterprise) plans must not confirm it.""" org = Mock() # own plans: none of them wix - org.plans.values_list.return_value = [("Free", False)] + org.get_plans.return_value.values_list.return_value = [("Free", False)] with patch.object(sync_odoo, "_resolve_plan_id", return_value=20): ids, status = sync_odoo._compute_org_plans_and_status( org, inherited_plan_ids=[101, 102] @@ -254,7 +257,7 @@ def test_own_wix_plan_confirms_even_with_inherited(self): """An own wix plan sets Confirmed; inherited plans are additive, not the trigger.""" org = Mock() - org.plans.values_list.return_value = [("Sunlight Basic", True)] + org.get_plans.return_value.values_list.return_value = [("Sunlight Basic", True)] with patch.object(sync_odoo, "_resolve_plan_id", return_value=10): ids, status = sync_odoo._compute_org_plans_and_status( org, inherited_plan_ids=[101] @@ -269,7 +272,9 @@ class TestMemberDesiredPlans: def test_unions_org_and_personal_plans(self): """Org plans and the user's personal plans are unioned.""" user = Mock() - user.individual_organization.plans.values_list.return_value = ["Personal"] + user.individual_organization.get_plans.return_value.values_list.return_value = [ + "Personal" + ] with patch.object(sync_odoo, "_resolve_plan_id", return_value=30): assert sync_odoo._member_desired_plans(user, [10, 20]) == [10, 20, 30] @@ -278,7 +283,9 @@ def test_drops_unresolved_personal_plan(self): This shouldn't ever happen as we ensure all plans at the beginning, but it is important we still have a test case.""" user = Mock() - user.individual_organization.plans.values_list.return_value = ["Broken"] + user.individual_organization.get_plans.return_value.values_list.return_value = [ + "Broken" + ] with patch.object(sync_odoo, "_resolve_plan_id", return_value=None): assert sync_odoo._member_desired_plans(user, [10]) == [10] diff --git a/squarelet/core/views.py b/squarelet/core/views.py index fbf3d2698..e6a9fc8cf 100644 --- a/squarelet/core/views.py +++ b/squarelet/core/views.py @@ -30,9 +30,9 @@ def get_context_data(self, **kwargs): pro_plan = None org_plans = None if not user.is_anonymous: - pro_plan = user.individual_organization.subscriptions.first() + pro_plan = user.individual_organization.subscription_items.first() org_plans = user.organizations.filter( - subscriptions__isnull=False, + subscriptions__items__isnull=False, individual=False, ).distinct() context["user"] = user diff --git a/squarelet/oidc/utils.py b/squarelet/oidc/utils.py index 4017317aa..bcdc8f86c 100644 --- a/squarelet/oidc/utils.py +++ b/squarelet/oidc/utils.py @@ -66,7 +66,9 @@ def send_cache_invalidations(model, uuids): def oidc_login_hook(request, user, client): """Log which client users login to""" # take an arbitrary non-individual organization, since most users will have one org - organizations = list(user.organizations.values("id", "name", plan=F("plans__name"))) + organizations = list( + user.organizations.values("id", "name", plan=F("subscriptions__plans__name")) + ) user.logins.create( client=client, metadata={ diff --git a/squarelet/organizations/admin.py b/squarelet/organizations/admin.py index 614c4c3fd..a0791dd73 100644 --- a/squarelet/organizations/admin.py +++ b/squarelet/organizations/admin.py @@ -43,6 +43,7 @@ ProfileChangeRequest, ReceiptEmail, Subscription, + SubscriptionItem, ) from squarelet.organizations.payments.factory import get_payment_provider from squarelet.users.models import User @@ -78,7 +79,7 @@ def format_value(self, value): class SubscriptionInline(admin.TabularInline): model = Subscription - readonly_fields = ("plan", "subscription_id", "cancelled", "quantity") + readonly_fields = ("subscription_id", "interval", "collection_method", "cancelled") extra = 0 can_delete = False @@ -240,8 +241,8 @@ def queryset(self, request, queryset): if value is None: return queryset if value == "none": - return queryset.filter(subscriptions__isnull=True) - return queryset.filter(subscriptions__plan_id=value) + return queryset.filter(subscriptions__items__isnull=True) + return queryset.filter(subscriptions__items__plan_id=value) class OverdueInvoiceFilter(admin.SimpleListFilter): @@ -459,12 +460,22 @@ def get_queryset(self, request): ) plan_value = request.GET.get("plan") if plan_value and plan_value != "none": + # `to_attr` attaches to the last relation in the path, so the + # subscriptions are collected on the organization and their + # matching lines prefetched underneath. qs = qs.prefetch_related( Prefetch( "subscriptions", - queryset=Subscription.objects.filter( - plan_id=plan_value - ).select_related("plan"), + queryset=Subscription.objects.filter(items__plan_id=plan_value) + .distinct() + .prefetch_related( + Prefetch( + "items", + queryset=SubscriptionItem.objects.filter( + plan_id=plan_value + ).select_related("plan"), + ) + ), to_attr="plan_subscriptions", ) ) @@ -527,7 +538,11 @@ def get_subscription_renews(self, obj): # A subscription renews only if it hasn't been cancelled and its plan # is set to auto-renew (plans with auto_renew=False are created to # cancel at period end). - return not any(s.cancelled or not s.plan.auto_renew for s in subs) + return not any( + sub.cancelled or not item.plan.auto_renew + for sub in subs + for item in sub.items.all() + ) get_subscription_renews.short_description = "Will Renew" get_subscription_renews.boolean = True diff --git a/squarelet/organizations/forms.py b/squarelet/organizations/forms.py index a3f6a296d..b490f3cc3 100644 --- a/squarelet/organizations/forms.py +++ b/squarelet/organizations/forms.py @@ -225,7 +225,7 @@ class MergeForm(forms.Form): ) bad_organization = forms.ModelChoiceField( queryset=Organization.objects.filter( - subscriptions__isnull=True, + subscriptions__items__isnull=True, individual=False, merged=None, ), diff --git a/squarelet/organizations/management/commands/audit_subscriptions.py b/squarelet/organizations/management/commands/audit_subscriptions.py index 7d6db6108..d2c539ecb 100644 --- a/squarelet/organizations/management/commands/audit_subscriptions.py +++ b/squarelet/organizations/management/commands/audit_subscriptions.py @@ -12,7 +12,7 @@ import stripe # Squarelet -from squarelet.organizations.models.payment import Customer, Subscription +from squarelet.organizations.models.payment import Customer, SubscriptionItem logger = logging.getLogger(__name__) @@ -40,7 +40,7 @@ def _datetimes_match(a, b): class Command(BaseCommand): - """Compare local Subscription records against Stripe and report mismatches. + """Compare local SubscriptionItem records against Stripe and report mismatches. Checks for: - Local subscriptions with no subscription_id @@ -114,7 +114,7 @@ def handle(self, *args, **options): def _load_local_subs(self, org_filter): """Return (local_subs, id→sub map) for subscriptions with a Stripe ID.""" - qs = Subscription.objects.select_related("plan", "organization").exclude( + qs = SubscriptionItem.objects.select_related("plan", "organization").exclude( subscription_id=None ) if org_filter: @@ -127,7 +127,7 @@ def _load_local_subs(self, org_filter): def _report_no_stripe_id(self, org_filter): """Print paid subscriptions with no subscription_id; return count.""" - qs = Subscription.objects.select_related("plan", "organization").filter( + qs = SubscriptionItem.objects.select_related("plan", "organization").filter( subscription_id=None, plan__base_price__gt=0 ) if org_filter: diff --git a/squarelet/organizations/management/commands/sync_subscriptions.py b/squarelet/organizations/management/commands/sync_subscriptions.py index f95861f43..69b68bf97 100644 --- a/squarelet/organizations/management/commands/sync_subscriptions.py +++ b/squarelet/organizations/management/commands/sync_subscriptions.py @@ -8,12 +8,12 @@ import stripe # Squarelet -from squarelet.organizations.models.payment import Subscription +from squarelet.organizations.models.payment import SubscriptionItem from squarelet.organizations.payments.factory import get_payment_provider class Command(BaseCommand): - """Sync local Subscription fields from Stripe. + """Sync local SubscriptionItem fields from Stripe. Fetches the live Stripe subscription for each local record and updates stripe_status and current_period_end. Safe to re-run — skips records @@ -39,7 +39,7 @@ def handle(self, *args, **options): org_filter = options["org"] dry_run = options["dry_run"] - qs = Subscription.objects.select_related("plan", "organization").exclude( + qs = SubscriptionItem.objects.select_related("plan", "organization").exclude( subscription_id=None ) if org_filter: diff --git a/squarelet/organizations/migrations/0083_rename_subscription_to_item.py b/squarelet/organizations/migrations/0083_rename_subscription_to_item.py new file mode 100644 index 000000000..9277cde1e --- /dev/null +++ b/squarelet/organizations/migrations/0083_rename_subscription_to_item.py @@ -0,0 +1,27 @@ +from django.db import migrations + + +class Migration(migrations.Migration): + """Rename Subscription to SubscriptionItem. + + Hand-written: makemigrations cannot infer a rename without being asked + interactively, and non-interactively it emits CreateModel + DeleteModel, + which would drop every subscription. + + First half of splitting the model in two. A SubscriptionItem is one line + on a Stripe subscription; the Subscription that owns those lines arrives + next. Renaming on its own first means every existing reference fails + loudly rather than silently binding to a `Subscription` that now means + something different. + """ + + dependencies = [ + ("organizations", "0083_merge_20260909_1055"), + ] + + operations = [ + migrations.RenameModel( + old_name="Subscription", + new_name="SubscriptionItem", + ), + ] diff --git a/squarelet/organizations/migrations/0084_subscriptionitem_related_names.py b/squarelet/organizations/migrations/0084_subscriptionitem_related_names.py new file mode 100644 index 000000000..6a9316b87 --- /dev/null +++ b/squarelet/organizations/migrations/0084_subscriptionitem_related_names.py @@ -0,0 +1,47 @@ +# Generated by Django 5.2.12 on 2026-08-26 17:20 + +import django.db.models.deletion +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ("organizations", "0083_rename_subscription_to_item"), + ] + + operations = [ + migrations.AlterField( + model_name="subscriptionitem", + name="organization", + field=models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + related_name="subscription_items", + to="organizations.organization", + verbose_name="organization", + ), + ), + migrations.AlterField( + model_name="subscriptionitem", + name="plan", + field=models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + related_name="subscription_items", + to="organizations.plan", + verbose_name="plan", + ), + ), + migrations.AlterField( + model_name="subscriptionitem", + name="plan_price", + field=models.ForeignKey( + blank=True, + help_text="The price this subscription is billed at. Nullable until every subscription has been migrated off the legacy plan foreign key.", + null=True, + on_delete=django.db.models.deletion.PROTECT, + related_name="subscription_items", + to="organizations.planprice", + verbose_name="plan price", + ), + ), + ] diff --git a/squarelet/organizations/migrations/0085_subscription_parent.py b/squarelet/organizations/migrations/0085_subscription_parent.py new file mode 100644 index 000000000..9584bd766 --- /dev/null +++ b/squarelet/organizations/migrations/0085_subscription_parent.py @@ -0,0 +1,336 @@ +# Generated by Django 5.2.12 on 2026-08-26 18:10 + +from collections import defaultdict + +import django.db.models.deletion +from django.db import migrations, models + + +def adopt_items_into_subscriptions(apps, schema_editor): + """Give every item a parent, one per billing shape. + + Stripe requires every item on a subscription to share a billing interval + and a collection method, which is what the new uniqueness constraint + encodes: one Subscription per (organization, interval, collection + method). Items agreeing on all three belong on the same parent. + + In practice this is a one-to-one lift - every organization holds exactly + one subscription today, multi-plan support having existed in the schema + without ever being used. It is written as a grouping anyway: a parent + per item would meet that constraint as an opaque IntegrityError, halfway + through a deploy, if a second line of the same shape appeared between + now and then. + + Interval and collection method are inferred from the plan: annual plans + bill annually and are the only ones invoiced rather than charged. + """ + SubscriptionItem = apps.get_model("organizations", "SubscriptionItem") + Subscription = apps.get_model("organizations", "Subscription") + Invoice = apps.get_model("organizations", "Invoice") + + groups = defaultdict(list) + for item in SubscriptionItem.objects.select_related("plan").order_by("pk"): + annual = bool(item.plan and item.plan.annual) + interval = "annual" if annual else "monthly" + collection_method = "send_invoice" if annual else "charge_automatically" + groups[(item.organization_id, interval, collection_method)].append(item) + + for (organization_id, interval, collection_method), items in groups.items(): + stripe_ids = { + item.legacy_subscription_id for item in items + if item.legacy_subscription_id + } + if len(stripe_ids) > 1: + # Two Stripe subscriptions of the same shape for one + # organization. They cannot share a parent - the row holds one + # Stripe id, so the other would still be billing with nothing + # naming it - and the constraint leaves no room for two parents. + # Refuse, naming both, rather than pick one and lose the other. + raise RuntimeError( + "Organization %s holds more than one %s/%s Stripe " + "subscription (%s). Consolidate them before migrating." + % ( + organization_id, + interval, + collection_method, + ", ".join(sorted(stripe_ids)), + ) + ) + + # A parent stops only once every line on it has stopped; one + # cancelled line among several is that line's own business. + cancelled = all(item.cancelled for item in items) + cancel_ats = [item.cancel_at for item in items if item.cancel_at] + # The Stripe-facing fields describe a single subscription, so they + # come from the line naming it - identical across the group, by the + # check above. + source = next( + (item for item in items if item.legacy_subscription_id), items[0] + ) + subscription = Subscription.objects.create( + organization_id=organization_id, + subscription_id=source.legacy_subscription_id or "", + interval=interval, + collection_method=collection_method, + cancelled=cancelled, + cancel_at=max(cancel_ats) if cancelled and cancel_ats else None, + stripe_status=source.stripe_status, + current_period_end=source.current_period_end, + ) + SubscriptionItem.objects.filter(pk__in=[item.pk for item in items]).update( + subscription=subscription + ) + + # Invoices billed these lines; they now belong to the parent. + Invoice.objects.filter(subscription__in=items).update( + new_subscription=subscription + ) + + +def split_back_out(apps, schema_editor): + """Copy subscription-level state back onto each item, then drop parents. + + This does not make the migration reversible on a database with rows in + it, and is not meant to: reversing the operations below re-adds + `organization_id` as NOT NULL before this function can fill it in, so + Postgres rejects the column before the data ever gets here. Rolling this + deploy back means restoring from a backup. What follows documents the + shape of the reverse, and works on an empty database. + + `legacy_subscription_id` is unique, so a Stripe id can go back onto only + one line. A parent carrying several lines is precisely the shape the old + schema could not hold - it is what the split exists to make possible - so + the id returns to the oldest line and the rest name nothing. Rolling + back a deploy that had already built such a subscription needs those + extra lines dealt with by hand, which beats refusing to roll back. + """ + SubscriptionItem = apps.get_model("organizations", "SubscriptionItem") + Subscription = apps.get_model("organizations", "Subscription") + + for parent in Subscription.objects.all(): + for index, item in enumerate(parent.items.order_by("pk")): + item.organization_id = parent.organization_id + item.legacy_subscription_id = ( + (parent.subscription_id or None) if index == 0 else None + ) + item.cancelled = parent.cancelled + item.cancel_at = parent.cancel_at + item.stripe_status = parent.stripe_status + item.current_period_end = parent.current_period_end + item.save() + Subscription.objects.all().delete() + + +class Migration(migrations.Migration): + + dependencies = [ + ("organizations", "0084_subscriptionitem_related_names"), + ] + + operations = [ + migrations.RemoveField( + model_name="organization", + name="plans", + ), + migrations.AddField( + model_name="subscriptionitem", + name="stripe_item_id", + field=models.CharField( + blank=True, + default="", + help_text="The subscription item ID on stripe. Blank for items that never reach Stripe, which is every comped one.", + max_length=255, + verbose_name="stripe item id", + ), + ), + migrations.CreateModel( + name="Subscription", + fields=[ + ( + "id", + models.AutoField( + auto_created=True, + primary_key=True, + serialize=False, + verbose_name="ID", + ), + ), + ( + "subscription_id", + models.CharField( + blank=True, + default="", + help_text="The subscription ID on stripe. Blank for subscriptions that never reach Stripe, which is every comped one.", + max_length=255, + verbose_name="subscription id", + ), + ), + ( + "interval", + models.CharField( + choices=[("monthly", "Monthly"), ("annual", "Annual")], + default="monthly", + help_text="Billing interval shared by every item", + max_length=20, + verbose_name="interval", + ), + ), + ( + "collection_method", + models.CharField( + choices=[ + ("charge_automatically", "Charge automatically"), + ("send_invoice", "Send invoice"), + ], + default="charge_automatically", + help_text="How Stripe collects payment, shared by every item", + max_length=30, + verbose_name="collection method", + ), + ), + ("cancelled", models.BooleanField(default=False)), + ( + "cancel_at", + models.DateField( + blank=True, + help_text="Date when Stripe will terminate this subscription. Set when cancel() is called. Null for free subscriptions.", + null=True, + verbose_name="cancel at", + ), + ), + ( + "stripe_status", + models.CharField(blank=True, default="", max_length=30), + ), + ("current_period_end", models.DateTimeField(blank=True, null=True)), + ( + "organization", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + related_name="subscriptions", + to="organizations.organization", + verbose_name="organization", + ), + ), + ( + "plans", + models.ManyToManyField( + blank=True, + help_text="Plans billed on this subscription", + related_name="subscriptions", + through="organizations.SubscriptionItem", + to="organizations.plan", + verbose_name="plans", + ), + ), + ], + options={ + "ordering": ("organization", "interval"), + }, + ), + migrations.AlterUniqueTogether( + name="subscriptionitem", + unique_together=set(), + ), + # The FK below wants the column name `subscription_id`, which the old + # CharField still occupies until it is dropped further down. Rename + # it out of the way rather than dropping it early - the data + # migration still needs to read it. + migrations.RenameField( + model_name="subscriptionitem", + old_name="subscription_id", + new_name="legacy_subscription_id", + ), + migrations.AddField( + model_name="subscriptionitem", + name="subscription", + field=models.ForeignKey( + blank=True, + help_text="The Stripe subscription this is a line on", + null=True, + on_delete=django.db.models.deletion.CASCADE, + related_name="items", + to="organizations.subscription", + verbose_name="subscription", + ), + ), + migrations.AlterUniqueTogether( + name="subscriptionitem", + unique_together={("subscription", "plan")}, + ), + migrations.AddConstraint( + model_name="subscription", + constraint=models.UniqueConstraint( + condition=models.Q(("subscription_id", ""), _negated=True), + fields=("subscription_id",), + name="unique_stripe_subscription_id_when_set", + ), + ), + migrations.AddConstraint( + model_name="subscription", + constraint=models.UniqueConstraint( + fields=("organization", "interval", "collection_method"), + name="unique_subscription_per_billing_shape", + ), + ), + # Give every existing item a parent, carrying its Stripe linkage + # across before the columns holding it are dropped below. + migrations.AddField( + model_name="invoice", + name="new_subscription", + field=models.ForeignKey( + blank=True, + null=True, + on_delete=django.db.models.deletion.SET_NULL, + related_name="+", + to="organizations.subscription", + ), + ), + migrations.RunPython(adopt_items_into_subscriptions, split_back_out), + migrations.RemoveField( + model_name="invoice", + name="subscription", + ), + migrations.RenameField( + model_name="invoice", + old_name="new_subscription", + new_name="subscription", + ), + migrations.AlterField( + model_name="invoice", + name="subscription", + field=models.ForeignKey( + blank=True, + help_text="The subscription this invoice is for, if applicable", + null=True, + on_delete=django.db.models.deletion.SET_NULL, + related_name="invoices", + to="organizations.subscription", + verbose_name="subscription", + ), + ), + migrations.RemoveField( + model_name="subscriptionitem", + name="cancel_at", + ), + migrations.RemoveField( + model_name="subscriptionitem", + name="cancelled", + ), + migrations.RemoveField( + model_name="subscriptionitem", + name="current_period_end", + ), + migrations.RemoveField( + model_name="subscriptionitem", + name="organization", + ), + migrations.RemoveField( + model_name="subscriptionitem", + name="stripe_status", + ), + migrations.RemoveField( + model_name="subscriptionitem", + name="legacy_subscription_id", + ), + ] diff --git a/squarelet/organizations/models/membership.py b/squarelet/organizations/models/membership.py index 6a962ab21..9f2f90342 100644 --- a/squarelet/organizations/models/membership.py +++ b/squarelet/organizations/models/membership.py @@ -54,7 +54,7 @@ def save(self, *args, **kwargs): is_new = self.pk is None direct_wix_plan_pks = ( list( - self.organization.subscriptions.filter(plan__wix=True).values_list( + self.organization.subscription_items.filter(plan__wix=True).values_list( "plan_id", flat=True ) ) @@ -115,7 +115,9 @@ def delete(self, *args, **kwargs): user_pk = self.user.pk user_uuid = self.user.uuid direct_plan_pks = list( - org.subscriptions.filter(plan__wix=True).values_list("plan_id", flat=True) + org.subscription_items.filter(plan__wix=True).values_list( + "plan_id", flat=True + ) ) group_wix_plans = [(g.pk, p.pk) for g, p in org.get_wix_plans_from_groups()] diff --git a/squarelet/organizations/models/organization.py b/squarelet/organizations/models/organization.py index 105d8fdc0..b144a6220 100644 --- a/squarelet/organizations/models/organization.py +++ b/squarelet/organizations/models/organization.py @@ -165,14 +165,41 @@ class Organization(AvatarMixin, models.Model): ), ) - plans = models.ManyToManyField( - verbose_name=_("plans"), - to="organizations.Plan", - through="organizations.Subscription", - related_name="organizations", - help_text=_("Plans this organization is subscribed to"), - blank=True, - ) + @property + def subscription_items(self): + """Every subscription line belonging to this organization. + + A queryset rather than a related manager: items hang off Subscription + now, so there is no direct relation from here. Reads behave the same + (`filter`, `get`, `exists`, ...), but creating an item needs a parent + subscription, so use `add_subscription()` for that. + + This does not participate in `prefetch_related`. To avoid a query per + organization, prefetch `subscriptions__items` and walk that instead. + """ + # pylint: disable=import-outside-toplevel + # Squarelet + from squarelet.organizations.models.payment import SubscriptionItem + + return SubscriptionItem.objects.filter(subscription__organization=self) + + def get_plans(self): + """Plans this organization is subscribed to. + + Replaces the old `plans` many-to-many. SubscriptionItem no longer + carries an organization of its own - it reaches one through its + subscription - so a through-model relation is no longer possible from + here. The equivalent relation lives on Subscription as `plans`. + + Callers that need this for many organizations at once should prefetch + `subscriptions__plans` and walk that instead, so this stays a single + query per organization rather than one per call. + """ + # pylint: disable=import-outside-toplevel + # Squarelet + from squarelet.organizations.models.payment import Plan + + return Plan.objects.filter(subscriptions__organization=self).distinct() # Every user has an individual organization # created when creating an account. Its UUID @@ -338,7 +365,7 @@ def save(self, *args, **kwargs): # If share_resources was toggled ON, sync all wix-enabled plans to members if share_resources_toggled_on and self.collective_enabled: wix_plan_pks = list( - self.plans.filter(wix=True).values_list("pk", flat=True) + self.get_plans().filter(wix=True).values_list("pk", flat=True) ) if wix_plan_pks: org_pk = self.pk @@ -511,7 +538,7 @@ def add_subscription(self, plan, max_users, user, token=None, payment_method=Non # check and the INSERT. Organization.objects.select_for_update().filter(pk=self.pk).get() - if self.subscriptions.filter(plan=plan).exists(): + if self.subscription_items.filter(plan=plan).exists(): raise SubscriptionError( f"Organization already has an active subscription to {plan}" ) @@ -520,7 +547,7 @@ def add_subscription(self, plan, max_users, user, token=None, payment_method=Non if max_users is None: max_users = plan.minimum_users - is_first = not self.subscriptions.exists() + is_first = not self.subscription_items.exists() payment_method = self._resolve_payment_method(payment_method, token) @@ -536,7 +563,7 @@ def add_subscription(self, plan, max_users, user, token=None, payment_method=Non # receives no billing_cycle_anchor (Stripe sets its own anchor). Only after # the subscription exists do we record the anchor for subsequent subscriptions # to align to. - _, stripe_subscription = self.subscriptions.start( + _, stripe_subscription = self.subscription_items.start( organization=self, plan=plan, payment_method=payment_method, @@ -598,15 +625,15 @@ def _dispatch_wix_sync(self, plan): sync_wix_for_group_member.delay(child_org.pk, self.pk, plan.pk) def remove_subscription(self, plan_or_subscription, user=None): - """Cancel the subscription for the given plan or Subscription instance.""" + """Cancel the subscription for the given plan or SubscriptionItem instance.""" # pylint: disable=import-outside-toplevel # Squarelet - from squarelet.organizations.models.payment import Subscription as Sub + from squarelet.organizations.models.payment import SubscriptionItem as Sub if isinstance(plan_or_subscription, Sub): sub = plan_or_subscription else: - sub = self.subscriptions.get(plan=plan_or_subscription) + sub = self.subscription_items.get(plan=plan_or_subscription) wix_unsync_plan = sub.plan if sub.plan and sub.plan.wix else None @@ -635,8 +662,8 @@ def modify_subscription(self, old_plan, new_plan, max_users, user): add/remove actions — but verify that use case is gone before removing. """ try: - sub = self.subscriptions.get(plan=old_plan) - except self.subscriptions.model.DoesNotExist: + sub = self.subscription_items.get(plan=old_plan) + except self.subscription_items.model.DoesNotExist: raise ValueError( f"Organization does not have an active subscription to {old_plan}" ) @@ -693,7 +720,9 @@ def subscription_cancelled(self, subscription): """The subscription was cancelled due to payment failure Args: - subscription: The specific Subscription instance to cancel. + subscription: The Subscription to cancel. Every line it carries + goes with it, because they all bill on the one invoice that + failed. """ if subscription is None: logger.error( @@ -703,18 +732,20 @@ def subscription_cancelled(self, subscription): ) return - # Create change log entry - self.change_logs.create( - reason=ChangeLogReason.failed, - from_plan=subscription.plan, - from_max_users=self.max_users, - to_max_users=self.max_users, - ) + # One log entry per line, since each names its own plan + cancelled_plans = [ + item.plan for item in subscription.items.select_related("plan") if item.plan + ] + for plan in cancelled_plans: + self.change_logs.create( + reason=ChangeLogReason.failed, + from_plan=plan, + from_max_users=self.max_users, + to_max_users=self.max_users, + ) - # Capture plan before subscription delete clears it - cancelled_plan = ( - subscription.plan if subscription.plan and subscription.plan.wix else None - ) + # Capture the plans before the delete cascades the lines away + wix_plans = [plan for plan in cancelled_plans if plan.wix] # Cancel subscription in Stripe if it exists if subscription.subscription_id: @@ -746,12 +777,12 @@ def subscription_cancelled(self, subscription): subscription.delete() # Remove Wix labels now that subscription is gone - if cancelled_plan: - self._dispatch_wix_unsync(cancelled_plan) + for plan in wix_plans: + self._dispatch_wix_unsync(plan) def has_active_subscription(self, plan=None): """Check if the organization has an active subscription""" - qs = self.subscriptions.all() + qs = self.subscription_items.all() if plan is not None: qs = qs.filter(plan=plan) return qs.exists() @@ -864,12 +895,12 @@ def get_wix_plans_from_groups(self): # Check membership groups for group in self.groups.filter(share_resources=True): - for plan in group.plans.filter(wix=True): + for plan in group.get_plans().filter(wix=True): wix_plans.append((group, plan)) # Check parent hierarchy (recursive) if self.parent and self.parent.share_resources: - for plan in self.parent.plans.filter(wix=True): + for plan in self.parent.get_plans().filter(wix=True): wix_plans.append((self.parent, plan)) # Also get parent's groups recursively wix_plans.extend(self.parent.get_wix_plans_from_groups()) @@ -893,12 +924,14 @@ def _add(source): if source.pk in _seen: return _seen.add(source.pk) - for plan in source.plans.all(): + for plan in source.get_plans(): if not plan.free: inherited.append((source, plan)) # Membership groups that share resources - for group in self.groups.filter(share_resources=True).prefetch_related("plans"): + for group in self.groups.filter(share_resources=True).prefetch_related( + "subscriptions__plans" + ): _add(group) # Parent hierarchy (recursive) @@ -933,7 +966,7 @@ def has_member_org(self, org): def merge(self, org, user): """Merge another organization into this one""" - if org.subscriptions.exists(): + if org.subscription_items.exists(): raise ValueError(f"{org} has active subscriptions and may not be merged") if org.merged is not None: raise ValueError( diff --git a/squarelet/organizations/models/payment.py b/squarelet/organizations/models/payment.py index d1166d4a5..690949596 100644 --- a/squarelet/organizations/models/payment.py +++ b/squarelet/organizations/models/payment.py @@ -27,7 +27,7 @@ EntitlementGrantQuerySet, EntitlementQuerySet, PlanQuerySet, - SubscriptionQuerySet, + SubscriptionItemQuerySet, ) logger = logging.getLogger(__name__) @@ -370,9 +370,29 @@ def add_source(self, token): class Subscription(models.Model): - """Through table for organization plans""" + """A subscription on Stripe. + + One row per Stripe subscription; its lines are SubscriptionItems. Stripe + requires every item on a subscription to share a billing interval and a + collection method, so an organization needs a separate subscription for + each combination it holds - a monthly MuckRock plan and an annual Sunlight + plan cannot sit on the same one. That is what the uniqueness constraint + below encodes. + + Fields here are subscription-level: status, period end and cancellation + apply to every item at once. Keeping them in one place means a renewal + webhook updates a single row rather than fanning out across items that + could then disagree. + """ - objects = SubscriptionQuerySet.as_manager() + INTERVAL_CHOICES = [ + ("monthly", _("Monthly")), + ("annual", _("Annual")), + ] + COLLECTION_CHOICES = [ + ("charge_automatically", _("Charge automatically")), + ("send_invoice", _("Send invoice")), + ] organization = models.ForeignKey( verbose_name=_("organization"), @@ -380,91 +400,56 @@ class Subscription(models.Model): on_delete=models.CASCADE, related_name="subscriptions", ) - plan = models.ForeignKey( - verbose_name=_("plan"), - to="organizations.Plan", - on_delete=models.CASCADE, - related_name="subscriptions", - ) - subscription_id = models.CharField( _("subscription id"), max_length=255, - unique=True, blank=True, - null=True, - help_text=_("The subscription ID on stripe"), + default="", + help_text=_( + "The subscription ID on stripe. Blank for subscriptions that " + "never reach Stripe, which is every comped one." + ), + ) + interval = models.CharField( + _("interval"), + max_length=20, + choices=INTERVAL_CHOICES, + default="monthly", + help_text=_("Billing interval shared by every item"), + ) + collection_method = models.CharField( + _("collection method"), + max_length=30, + choices=COLLECTION_CHOICES, + default="charge_automatically", + help_text=_("How Stripe collects payment, shared by every item"), ) - # The cancelled flag is used to mark subscriptions that are ready for cancellation. - # Cancellation happens at the end of the billing period; at that point, - # the subscription is deleted from the database. + # The cancelled flag marks a subscription as ready for cancellation. + # Cancellation happens at the end of the billing period; at that point the + # record is deleted. cancelled = models.BooleanField(default=False) - cancel_at = models.DateField( _("cancel at"), null=True, blank=True, help_text=_( - "Date when Stripe will terminate this subscription. " - "Set when cancel() is called. Null for free plans or legacy records." - ), - ) - - quantity = models.PositiveIntegerField( - _("quantity"), - default=1, - help_text=_( - "Number of units of this plan's resources granted to the organization" + "Date when Stripe will terminate this subscription. Set when " + "cancel() is called. Null for free subscriptions." ), ) + stripe_status = models.CharField(max_length=30, blank=True, default="") + current_period_end = models.DateTimeField(null=True, blank=True) - plan_price = models.ForeignKey( - verbose_name=_("plan price"), - to="organizations.PlanPrice", - on_delete=models.PROTECT, + plans = models.ManyToManyField( + verbose_name=_("plans"), + to="organizations.Plan", + through="organizations.SubscriptionItem", related_name="subscriptions", + help_text=_("Plans billed on this subscription"), blank=True, - null=True, - help_text=_( - "The price this subscription is billed at. Nullable until every " - "subscription has been migrated off the legacy plan foreign key." - ), ) - granted_reason = models.TextField( - _("granted reason"), - blank=True, - default="", - help_text=_( - "Why this subscription received non-standard pricing (comped, or a " - "partner coupon). Blank for ordinary self-serve subscriptions." - ), - ) - granted_by = models.ForeignKey( - verbose_name=_("granted by"), - to="users.User", - on_delete=models.PROTECT, - related_name="granted_subscriptions", - blank=True, - null=True, - help_text=_( - "Staff user who authorized the non-standard pricing. Blank for " - "ordinary self-serve subscriptions." - ), - ) - - stripe_status = models.CharField(max_length=30, blank=True, default="") - current_period_end = models.DateTimeField(null=True, blank=True) - - class Meta: - unique_together = ("organization", "plan") - ordering = ("plan",) - - def __str__(self): - plan_name = self.plan.name if self.plan else "Free" - return f"Subscription: {self.organization} to {plan_name}" - @cached_property def stripe_subscription(self): if self.subscription_id: @@ -475,6 +460,30 @@ def stripe_subscription(self): ) return None + @property + def free(self): + """A subscription costs nothing when every line does.""" + return all(item.plan is None or item.plan.free for item in self.items.all()) + + @property + def auto_renew(self): + """Renew unless some line says otherwise.""" + return all(item.plan.auto_renew for item in self.items.all()) + + def stripe_items(self, include_ids=False): + """Stripe line specs for every item on this subscription. + + Pass include_ids when modifying an existing subscription: Stripe needs + each line's own id to update it in place rather than replace it. + """ + specs = [] + for item in self.items.select_related("plan"): + spec = {"plan": item.plan.stripe_id, "quantity": item.quantity} + if include_ids and item.stripe_item_id: + spec["id"] = item.stripe_item_id + specs.append(spec) + return specs + def cache_stripe_subscription_fields(self, stripe_sub): """Cache subscription status and period end from a Stripe subscription.""" self.stripe_status = stripe_sub.status or "" @@ -488,8 +497,11 @@ def cache_stripe_subscription_fields(self, stripe_sub): ) def start(self, payment_method="card", anchor_day=None): - """Start the Stripe subscription. Returns the Stripe subscription object - for paid plans, or None for free plans.""" + """Create this subscription on Stripe, with all of its items. + + Returns the Stripe subscription for paid subscriptions, or None when + every line is free - those never reach Stripe at all. + """ if self.stripe_subscription: logger.error( "Trying to start an existing subscription: %s %s", @@ -497,66 +509,43 @@ def start(self, payment_method="card", anchor_day=None): self.subscription_id, ) return None - stripe_subscription = None - if self.plan and not self.plan.free: - # Annual plans support payment by invoice - if self.plan.annual and payment_method == "invoice": - billing = "send_invoice" - days_until_due = 30 - else: - billing = "charge_automatically" - days_until_due = None - - stripe_subscription = ( - get_payment_provider() - .get_subscription_service() - .create( - stripe_customer=self.organization.customer().stripe_customer, - plan_id=self.plan.stripe_id, - quantity=self.quantity, - billing=billing, - metadata={"action": f"Subscription ({self.plan})"}, - days_until_due=days_until_due, - anchor_day=anchor_day, - cancel_at_period_end=not self.plan.auto_renew, - ) - ) - self.subscription_id = stripe_subscription.id - self.cache_stripe_subscription_fields(stripe_subscription) - if not self.plan.auto_renew and self.current_period_end: - self.cancel_at = self.current_period_end.date() - # Save subscription before creating invoice - self.save() - - # Check for 3DS/SCA on the first invoice payment. - if stripe_subscription.status == "incomplete": - self._check_3ds_action_required(stripe_subscription) + if self.free: + return None - # Create Invoice record synchronously; webhook is the fallback. - self._sync_latest_invoice(stripe_subscription) + # Annual subscriptions support payment by invoice + if self.interval == "annual" and payment_method == "invoice": + billing = "send_invoice" + days_until_due = 30 + else: + billing = "charge_automatically" + days_until_due = None + self.collection_method = billing - # Trigger respective mailchimp journeys if this is the organization plan, - # but only if the org doesn't already have another active subscription - # granting the same entitlement (avoid duplicate journey triggers). - if self.plan_id and self.plan.entitlements.filter(slug="organization").exists(): - already_has_org_entitlement = ( - self.organization.subscriptions.exclude(pk=self.pk) - .filter( - plan__entitlements__slug="organization", - ) - .exists() + stripe_subscription = ( + get_payment_provider() + .get_subscription_service() + .create( + stripe_customer=self.organization.customer().stripe_customer, + items=self.stripe_items(), + billing=billing, + metadata={"action": f"Subscription ({self.organization})"}, + days_until_due=days_until_due, + anchor_day=anchor_day, + cancel_at_period_end=not self.auto_renew, ) - if not already_has_org_entitlement: - journey_key = ( - "verified_premium_org" - if self.organization.verified_journalist - else "unverified_premium_org" - ) - for user in self.organization.users.all(): - mailchimp_journey(user.email, journey_key) + ) + self.subscription_id = stripe_subscription.id + self.cache_stripe_subscription_fields(stripe_subscription) + if not self.auto_renew and self.current_period_end: + self.cancel_at = self.current_period_end.date() + # Save before creating the invoice + self.save() - # Slack notification for new subscription - self.send_slack_notification("started") + # Check for 3DS/SCA on the first invoice payment. + if stripe_subscription.status == "incomplete": + self._check_3ds_action_required(stripe_subscription) + + self._sync_latest_invoice(stripe_subscription) return stripe_subscription def _check_3ds_action_required(self, stripe_subscription): @@ -638,8 +627,10 @@ def cancel(self): self.cancel_at = self.current_period_end.date() self.save() - # Slack notification for cancellation - self.send_slack_notification("cancelled") + # The notification names a plan, so it belongs to the lines, not to + # the subscription that carries them. + for item in self.items.select_related("plan"): + item.send_slack_notification("cancelled") def uncancel(self): """Re-enable renewal for a subscription that was pending cancellation. @@ -667,62 +658,221 @@ def uncancel(self): self.cancel_at = None self.save() - def modify(self, plan): - """Modify an existing plan - Note - this should never be used to switch from a MR to a PP plan or vice versa - """ - old_plan = self.plan - self.plan = plan - self.save() - - if old_plan.free and not plan.free: - # start subscription on stripe - self.start() - elif not old_plan.free and plan.free: - # cancel subscription on stripe - get_payment_provider().get_subscription_service().delete( - self.stripe_subscription - ) - self.subscription_id = None - self.cancel_at = None - elif not old_plan.free and not plan.free: - # modify plan - self.stripe_modify() - - self.save() - def stripe_modify(self): - """Update stripe subscription to match local subscription""" + """Push local state to Stripe for every item on this subscription.""" if self.stripe_subscription: updated = ( get_payment_provider() .get_subscription_service() .modify( self.subscription_id, - cancel_at_period_end=not self.plan.auto_renew, - items=[ - { - "id": self.stripe_subscription["items"]["data"][0].id, - "plan": self.plan.stripe_id, - "quantity": self.quantity, - } - ], + cancel_at_period_end=not self.auto_renew, + items=self.stripe_items(include_ids=True), billing=( - "send_invoice" if self.plan.annual else "charge_automatically" + "send_invoice" + if self.interval == "annual" + else "charge_automatically" ), - metadata={"action": f"Subscription ({self.plan})"}, - days_until_due=(30 if self.plan.annual else None), + metadata={"action": f"Subscription ({self.organization})"}, + days_until_due=(30 if self.interval == "annual" else None), ) ) self.cancelled = False if updated: self.cache_stripe_subscription_fields(updated) - if not self.plan.auto_renew and self.current_period_end: + if not self.auto_renew and self.current_period_end: self.cancel_at = self.current_period_end.date() else: self.cancel_at = None self.save() + class Meta: + ordering = ("organization", "interval") + constraints = [ + # Every real Stripe subscription id is unique; any number of + # comped subscriptions may leave it blank. + models.UniqueConstraint( + fields=["subscription_id"], + condition=~models.Q(subscription_id=""), + name="unique_stripe_subscription_id_when_set", + ), + # One subscription per organization per billing shape. Anything + # that would need a second one for the same shape should be an + # item on the existing subscription instead. + models.UniqueConstraint( + fields=["organization", "interval", "collection_method"], + name="unique_subscription_per_billing_shape", + ), + ] + + def __str__(self): + return ( + f"{self.organization.name}: {self.get_interval_display()}, " + f"{self.get_collection_method_display()}" + ) + + +class SubscriptionItem(models.Model): + """One line on a Stripe subscription. + + The organization is reached through `subscription`, deliberately not + duplicated here: a denormalized copy that has to agree with its parent is + exactly the kind of drift this migration exists to remove. + """ + + objects = SubscriptionItemQuerySet.as_manager() + + plan = models.ForeignKey( + verbose_name=_("plan"), + to="organizations.Plan", + on_delete=models.CASCADE, + related_name="subscription_items", + ) + + subscription = models.ForeignKey( + verbose_name=_("subscription"), + to="organizations.Subscription", + on_delete=models.CASCADE, + related_name="items", + blank=True, + null=True, + help_text=_("The Stripe subscription this is a line on"), + ) + stripe_item_id = models.CharField( + _("stripe item id"), + max_length=255, + blank=True, + default="", + help_text=_( + "The subscription item ID on stripe. Blank for items that never " + "reach Stripe, which is every comped one." + ), + ) + + quantity = models.PositiveIntegerField( + _("quantity"), + default=1, + help_text=_( + "Number of units of this plan's resources granted to the organization" + ), + ) + + plan_price = models.ForeignKey( + verbose_name=_("plan price"), + to="organizations.PlanPrice", + on_delete=models.PROTECT, + related_name="subscription_items", + blank=True, + null=True, + help_text=_( + "The price this subscription is billed at. Nullable until every " + "subscription has been migrated off the legacy plan foreign key." + ), + ) + + granted_reason = models.TextField( + _("granted reason"), + blank=True, + default="", + help_text=_( + "Why this subscription received non-standard pricing (comped, or a " + "partner coupon). Blank for ordinary self-serve subscriptions." + ), + ) + granted_by = models.ForeignKey( + verbose_name=_("granted by"), + to="users.User", + on_delete=models.PROTECT, + related_name="granted_subscriptions", + blank=True, + null=True, + help_text=_( + "Staff user who authorized the non-standard pricing. Blank for " + "ordinary self-serve subscriptions." + ), + ) + + class Meta: + # One line per plan per subscription. "An organization may not hold + # the same plan twice" is now wider than a single subscription can + # see, so add_subscription() enforces that part. + unique_together = ("subscription", "plan") + ordering = ("plan",) + + def __str__(self): + plan_name = self.plan.name if self.plan else "Free" + return f"SubscriptionItem: {self.subscription.organization} to {plan_name}" + + @property + def organization(self): + """The owning organization, reached through the parent subscription. + + Read-only on purpose: the column lives on `Subscription` so a line can + never disagree with the subscription it bills on. Select or prefetch + `subscription__organization` before touching this in a loop. + """ + return self.subscription.organization + + def modify(self, plan): + """Change which plan this line bills. + + Never use this to move between products - that is an add plus a + remove, since the two subscriptions bill separately. + """ + self.plan = plan + self.save() + self.subscription.stripe_modify() + + def cancel(self): + """Stop billing this line. + + The last line on a subscription cancels the whole subscription at + period end, so the customer keeps what they already paid for. Any + other line is dropped from the Stripe subscription right away with + proration suppressed - the next invoice simply omits it, and no + mid-period credit or charge is generated. + """ + if self.subscription.items.count() <= 1: + self.subscription.cancel() + return + + if self.subscription.stripe_subscription and self.stripe_item_id: + get_payment_provider().get_subscription_service().modify( + self.subscription.subscription_id, + items=[{"id": self.stripe_item_id, "deleted": True}], + proration_behavior="none", + ) + self.send_slack_notification("cancelled") + self.delete() + + def notify_started(self): + """Announce a newly added line. + + The Mailchimp journey fires only for the line that first grants the + organization entitlement, so an org that already has it through + another line is not enrolled twice. + """ + organization = self.subscription.organization + if self.plan_id and self.plan.entitlements.filter(slug="organization").exists(): + already_has_org_entitlement = ( + SubscriptionItem.objects.filter( + subscription__organization=organization, + plan__entitlements__slug="organization", + ) + .exclude(pk=self.pk) + .exists() + ) + if not already_has_org_entitlement: + journey_key = ( + "verified_premium_org" + if organization.verified_journalist + else "unverified_premium_org" + ) + for user in organization.users.all(): + mailchimp_journey(user.email, journey_key) + + self.send_slack_notification("started") + def send_slack_notification(self, event, **kwargs): """Queue a Slack notification asynchronously for subscription events.""" if not is_production_env(): @@ -1005,7 +1155,7 @@ def has_available_slots(self): """Check if new subscriptions are allowed for this plan""" # Only Sunlight plans have subscription limits if self.slug.startswith("sunlight-") and self.wix: - current_count = Subscription.objects.sunlight_active_count() + current_count = SubscriptionItem.objects.sunlight_active_count() return current_count < settings.MAX_SUNLIGHT_SUBSCRIPTIONS return True @@ -1655,7 +1805,7 @@ def matching_organizations(self): rule_clauses.append(Q(verified_journalist=True)) if self.require_active_subscription: # Mirrors org.has_active_subscription() = bool(subscriptions.first()) - rule_clauses.append(Q(subscriptions__isnull=False)) + rule_clauses.append(Q(subscriptions__items__isnull=False)) if rule_clauses: rule_q = rule_clauses[0] diff --git a/squarelet/organizations/payments/base.py b/squarelet/organizations/payments/base.py index aac25779c..61906b384 100644 --- a/squarelet/organizations/payments/base.py +++ b/squarelet/organizations/payments/base.py @@ -74,18 +74,22 @@ class SubscriptionService(ABC): """Manages Stripe Subscription objects.""" @abstractmethod - def create( # pylint: disable=too-many-arguments,too-many-positional-arguments + def create( # pylint: disable=too-many-positional-arguments self, stripe_customer, - plan_id, - quantity, + items, billing, metadata, days_until_due, anchor_day=None, cancel_at_period_end=False, ): - """Create a new subscription for a customer.""" + """Create a new subscription for a customer. + + `items` is a list of Stripe line specs, e.g. + `[{"plan": plan_id, "quantity": 1}]` - one per SubscriptionItem, since + a subscription may bill several plans at once. + """ @abstractmethod def retrieve(self, subscription_id): diff --git a/squarelet/organizations/payments/providers/stripe_modern.py b/squarelet/organizations/payments/providers/stripe_modern.py index 279c25b9b..d7742353f 100644 --- a/squarelet/organizations/payments/providers/stripe_modern.py +++ b/squarelet/organizations/payments/providers/stripe_modern.py @@ -137,11 +137,10 @@ def retrieve_payment_method(self, pm_id): class StripeModernSubscriptionService(SubscriptionService): """Subscription operations using current Stripe API.""" - def create( # pylint: disable=too-many-arguments,too-many-positional-arguments + def create( # pylint: disable=too-many-positional-arguments self, stripe_customer, - plan_id, - quantity, + items, billing, metadata, days_until_due, @@ -154,7 +153,7 @@ def create( # pylint: disable=too-many-arguments,too-many-positional-arguments # for SCA/3DS detection as of API version 2025-03-31.basil params = { "customer": stripe_customer.id, - "items": [{"plan": plan_id, "quantity": quantity}], + "items": items, "collection_method": billing, "metadata": metadata, "days_until_due": days_until_due, diff --git a/squarelet/organizations/querysets.py b/squarelet/organizations/querysets.py index 6f8456d71..20b0ffbeb 100644 --- a/squarelet/organizations/querysets.py +++ b/squarelet/organizations/querysets.py @@ -85,7 +85,7 @@ def create_individual(self, user, uuid=None): user.individual_organization.change_logs.create( reason=ChangeLogReason.created, user=user, - to_plan=user.individual_organization.plans.first(), + to_plan=user.individual_organization.get_plans().first(), to_max_users=user.individual_organization.max_users, ) return user.individual_organization @@ -107,7 +107,7 @@ def get_viewable(self, user): elif user.is_authenticated: return self.filter( Q(public=True) - | Q(organizations__in=user.organizations.all()) + | Q(subscriptions__organization__in=user.organizations.all()) | Q(private_organizations__in=user.organizations.all()) ).distinct() else: @@ -127,7 +127,7 @@ def choices(self, organization): # to which they have been granted explicit access return queryset.filter( Q(public=True) - | Q(organizations=organization) + | Q(subscriptions__organization=organization) | Q(private_organizations=organization) ).distinct() @@ -151,7 +151,7 @@ def get_public(self): def get_subscribed(self, user): if user.is_authenticated: return self.filter( - plans__organizations__in=user.organizations.all() + plans__subscriptions__organization__in=user.organizations.all() ).distinct() else: return self.none() @@ -171,7 +171,9 @@ def for_organization(self, org, client=None): from squarelet.organizations.models.payment import EntitlementGrant matching_grants = EntitlementGrant.objects.for_org(org) - qs = self.filter(Q(plans__organizations=org) | Q(grants__in=matching_grants)) + qs = self.filter( + Q(plans__subscriptions__organization=org) | Q(grants__in=matching_grants) + ) if client is not None: qs = qs.filter(client=client) return qs.distinct() @@ -413,20 +415,54 @@ def confirm_payment_intent( return charge -class SubscriptionQuerySet(models.QuerySet): +class SubscriptionItemQuerySet(models.QuerySet): def start(self, organization, plan, payment_method="card", quantity=1): - subscription = self.model( + """Add a line for `plan` and make sure Stripe knows about it. + + Stripe requires every item on a subscription to share a billing + interval and a collection method, so those two fields decide which + subscription the line joins. A line that matches an existing + subscription is added to it and bills on the same invoice; one that + does not starts a new subscription. + + Returns the new line and the Stripe subscription carrying it, which + is None for a subscription that costs nothing. + """ + # Lazy import to avoid a circular import (payment.py imports this module) + # pylint: disable=import-outside-toplevel + # Squarelet + from squarelet.organizations.models.payment import Subscription + + interval = "annual" if plan.annual else "monthly" + collection_method = ( + "send_invoice" + if interval == "annual" and payment_method == "invoice" + else "charge_automatically" + ) + subscription, created = Subscription.objects.get_or_create( organization=organization, - plan=plan, - quantity=quantity, + interval=interval, + collection_method=collection_method, ) - anchor = organization.billing_anchor - stripe_subscription = subscription.start( - payment_method=payment_method, - anchor_day=anchor.day if anchor else None, + item = self.model.objects.create( + subscription=subscription, plan=plan, quantity=quantity ) - subscription.save() - return subscription, stripe_subscription + + if created or not subscription.subscription_id: + anchor = organization.billing_anchor + stripe_subscription = subscription.start( + payment_method=payment_method, + anchor_day=anchor.day if anchor else None, + ) + else: + # The subscription is already live on Stripe, so the new line is + # pushed onto it rather than opening a second subscription. + subscription.stripe_modify() + stripe_subscription = subscription.stripe_subscription + + if not plan.free: + item.notify_started() + return item, stripe_subscription def sunlight_active_count(self): """Count active Sunlight subscriptions across all variants""" diff --git a/squarelet/organizations/rules/__init__.py b/squarelet/organizations/rules/__init__.py index 24fc382c2..43dff13bf 100644 --- a/squarelet/organizations/rules/__init__.py +++ b/squarelet/organizations/rules/__init__.py @@ -3,4 +3,4 @@ import squarelet.organizations.rules.invitations import squarelet.organizations.rules.memberships import squarelet.organizations.rules.organizations -import squarelet.organizations.rules.subscriptions +import squarelet.organizations.rules.subscription_items diff --git a/squarelet/organizations/rules/subscription_items.py b/squarelet/organizations/rules/subscription_items.py new file mode 100644 index 000000000..c0c695b1d --- /dev/null +++ b/squarelet/organizations/rules/subscription_items.py @@ -0,0 +1,23 @@ +# Third Party +from rules import add_perm, is_authenticated, predicate + +# Squarelet +from squarelet.core.rules import skip_if_not_obj + + +@predicate +@skip_if_not_obj +def is_member(user, item): + return item.organization.has_member(user) + + +@predicate +@skip_if_not_obj +def is_admin(user, item): + return item.organization.has_admin(user) + + +add_perm("organizations.view_subscriptionitem", is_authenticated & is_member) +add_perm("organizations.add_subscriptionitem", is_authenticated) +add_perm("organizations.change_subscriptionitem", is_authenticated & is_admin) +add_perm("organizations.delete_subscriptionitem", is_authenticated & is_admin) diff --git a/squarelet/organizations/rules/subscriptions.py b/squarelet/organizations/rules/subscriptions.py deleted file mode 100644 index 2c87c909e..000000000 --- a/squarelet/organizations/rules/subscriptions.py +++ /dev/null @@ -1,23 +0,0 @@ -# Third Party -from rules import add_perm, is_authenticated, predicate - -# Squarelet -from squarelet.core.rules import skip_if_not_obj - - -@predicate -@skip_if_not_obj -def is_member(user, subscription): - return subscription.organization.has_member(user) - - -@predicate -@skip_if_not_obj -def is_admin(user, subscription): - return subscription.organization.has_admin(user) - - -add_perm("organizations.view_subscription", is_authenticated & is_member) -add_perm("organizations.add_subscription", is_authenticated) -add_perm("organizations.change_subscription", is_authenticated & is_admin) -add_perm("organizations.delete_subscription", is_authenticated & is_admin) diff --git a/squarelet/organizations/serializers.py b/squarelet/organizations/serializers.py index 0bc41c42b..80e012b6f 100644 --- a/squarelet/organizations/serializers.py +++ b/squarelet/organizations/serializers.py @@ -127,7 +127,7 @@ def get_entitlements(self, obj): result = [] # Plan-based: one entry per subscription's entitlements (no dedup) - for sub in obj.subscriptions.prefetch_related("plan__entitlements").all(): + for sub in obj.subscription_items.prefetch_related("plan__entitlements").all(): for ent in sub.plan.entitlements.filter(client=client): result.append( { diff --git a/squarelet/organizations/signals.py b/squarelet/organizations/signals.py index 31058b3f1..e908366a3 100644 --- a/squarelet/organizations/signals.py +++ b/squarelet/organizations/signals.py @@ -32,7 +32,7 @@ def should_sync_wix(org): return ( org and org.share_resources - and org.subscriptions.filter(plan__wix=True).exists() + and org.subscription_items.filter(plan__wix=True).exists() ) @@ -102,7 +102,7 @@ def sync_wix_on_parent_change(sender, instance, **kwargs): child_pk = instance.pk parent_pk = parent.pk - for sub in parent.subscriptions.filter(plan__wix=True).select_related("plan"): + for sub in parent.subscription_items.filter(plan__wix=True).select_related("plan"): plan_pk = sub.plan.pk transaction.on_commit( lambda c=child_pk, par=parent_pk, p=plan_pk: ( @@ -136,7 +136,9 @@ def sync_wix_on_member_add(sender, instance, action, pk_set, reverse, **kwargs): if not should_sync_wix(group): return group_pk = group.pk - for sub in group.subscriptions.filter(plan__wix=True).select_related("plan"): + for sub in group.subscription_items.filter(plan__wix=True).select_related( + "plan" + ): plan_pk = sub.plan.pk for member_pk in pk_set: transaction.on_commit( @@ -150,9 +152,9 @@ def sync_wix_on_member_add(sender, instance, action, pk_set, reverse, **kwargs): for group_pk in pk_set: group = Organization.objects.filter(pk=group_pk).first() if should_sync_wix(group): - for sub in group.subscriptions.filter(plan__wix=True).select_related( - "plan" - ): + for sub in group.subscription_items.filter( + plan__wix=True + ).select_related("plan"): plan_pk = sub.plan.pk transaction.on_commit( lambda g=group_pk, p=plan_pk: sync_wix_for_group_member.delay( diff --git a/squarelet/organizations/tasks.py b/squarelet/organizations/tasks.py index 114af5210..94b52c8aa 100644 --- a/squarelet/organizations/tasks.py +++ b/squarelet/organizations/tasks.py @@ -431,7 +431,9 @@ def handle_invoice_created(invoice_data): if subscription: metadata = { "organization": str(organization.uuid), - "plan": str(subscription.plan), + "plan": ", ".join( + str(item.plan) for item in subscription.items.select_related("plan") + ), "subscription_id": subscription.subscription_id, } try: diff --git a/squarelet/organizations/tests/factories.py b/squarelet/organizations/tests/factories.py index d595ad230..247e0f8a8 100644 --- a/squarelet/organizations/tests/factories.py +++ b/squarelet/organizations/tests/factories.py @@ -42,8 +42,9 @@ def admins(self, create, extracted, **kwargs): @factory.post_generation def plans(self, create, extracted, **kwargs): if create and extracted: + subscription = SubscriptionFactory(organization=self) for plan in extracted: - SubscriptionFactory(plan=plan, organization=self) + SubscriptionItemFactory(subscription=subscription, plan=plan) # Clear the memoized cache for the plan/subscription properties # since they may have been accessed before subscriptions were created for attr_name in list(vars(self).keys()): @@ -106,13 +107,40 @@ def _create(cls, model_class, *args, **kwargs): class SubscriptionFactory(factory.django.DjangoModelFactory): + """A Stripe subscription for one organization. + + Keyed on the billing shape the way production is: asking for a second + subscription with the same interval and collection method returns the + one that already exists, rather than tripping the unique constraint. + """ + organization = factory.SubFactory( "squarelet.organizations.tests.factories.OrganizationFactory" ) - plan = factory.SubFactory("squarelet.organizations.tests.factories.PlanFactory") + # Declared so django_get_or_create can key on them; both match the + # model defaults. + interval = "monthly" + collection_method = "charge_automatically" class Meta: model = "organizations.Subscription" + django_get_or_create = ("organization", "interval", "collection_method") + + +class SubscriptionItemFactory(factory.django.DjangoModelFactory): + """A line on a subscription. + + Pass `subscription__organization=` or `subscription__cancelled=` to steer + the parent; a bare call builds one for you. + """ + + subscription = factory.SubFactory( + "squarelet.organizations.tests.factories.SubscriptionFactory" + ) + plan = factory.SubFactory("squarelet.organizations.tests.factories.PlanFactory") + + class Meta: + model = "organizations.SubscriptionItem" @factory.django.mute_signals(signals.pre_save, signals.post_save) diff --git a/squarelet/organizations/tests/models/test_entitlement_grant.py b/squarelet/organizations/tests/models/test_entitlement_grant.py index 0b2b7cc04..b03f8ce98 100644 --- a/squarelet/organizations/tests/models/test_entitlement_grant.py +++ b/squarelet/organizations/tests/models/test_entitlement_grant.py @@ -12,7 +12,7 @@ EntitlementGrantFactory, OrganizationFactory, PlanFactory, - SubscriptionFactory, + SubscriptionItemFactory, ) @@ -57,7 +57,7 @@ def test_grant_matches_active_subscription_org(self): We can grant entitlements to our active subscribers. """ subscribed = OrganizationFactory() - SubscriptionFactory(organization=subscribed) + SubscriptionItemFactory(subscription__organization=subscribed) unsubscribed = OrganizationFactory() grant = EntitlementGrantFactory(require_active_subscription=True) assert grant.matches(subscribed) is True @@ -69,10 +69,10 @@ def test_grant_requires_all_checked_criteria(self): Grant rules are combined with "AND" logic. """ verified_and_sub = OrganizationFactory(verified_journalist=True) - SubscriptionFactory(organization=verified_and_sub) + SubscriptionItemFactory(subscription__organization=verified_and_sub) verified_only = OrganizationFactory(verified_journalist=True) sub_only = OrganizationFactory(verified_journalist=False) - SubscriptionFactory(organization=sub_only) + SubscriptionItemFactory(subscription__organization=sub_only) grant = EntitlementGrantFactory( require_verified=True, require_active_subscription=True @@ -174,7 +174,7 @@ def test_verified_rule(self): @pytest.mark.django_db() def test_active_subscription_rule(self): subscribed = OrganizationFactory() - SubscriptionFactory(organization=subscribed) + SubscriptionItemFactory(subscription__organization=subscribed) unsubscribed = OrganizationFactory() grant = EntitlementGrantFactory(require_active_subscription=True) matched = list(grant.matching_organizations()) @@ -212,7 +212,7 @@ def test_explicit_membership_respects_org_type_filter(self): @pytest.mark.django_db() def test_both_rules_and_logic(self): verified_and_sub = OrganizationFactory(verified_journalist=True) - SubscriptionFactory(organization=verified_and_sub) + SubscriptionItemFactory(subscription__organization=verified_and_sub) verified_only = OrganizationFactory(verified_journalist=True) grant = EntitlementGrantFactory( require_verified=True, require_active_subscription=True @@ -275,7 +275,7 @@ def test_verified_rule_matches(self): @pytest.mark.django_db() def test_active_subscription_rule_matches(self): subscribed = OrganizationFactory() - SubscriptionFactory(organization=subscribed) + SubscriptionItemFactory(subscription__organization=subscribed) unsubscribed = OrganizationFactory() entitlement = EntitlementFactory() EntitlementGrantFactory( @@ -287,7 +287,7 @@ def test_active_subscription_rule_matches(self): @pytest.mark.django_db() def test_both_rules_require_both_at_db_level(self): verified_and_sub = OrganizationFactory(verified_journalist=True) - SubscriptionFactory(organization=verified_and_sub) + SubscriptionItemFactory(subscription__organization=verified_and_sub) verified_only = OrganizationFactory(verified_journalist=True) entitlement = EntitlementFactory() EntitlementGrantFactory( diff --git a/squarelet/organizations/tests/models/test_organization.py b/squarelet/organizations/tests/models/test_organization.py index c51f364da..dca6013ef 100644 --- a/squarelet/organizations/tests/models/test_organization.py +++ b/squarelet/organizations/tests/models/test_organization.py @@ -9,7 +9,7 @@ # Squarelet from squarelet.organizations.models import ( Organization, - Subscription, + SubscriptionItem, consolidate_inherited_benefits, ) from squarelet.organizations.payments.exceptions import SubscriptionError @@ -186,7 +186,7 @@ def test_customer_new(self, organization_factory, mocker): @pytest.mark.django_db() def test_subscription_blank(self, organization_factory): organization = organization_factory() - assert organization.subscriptions.first() is None + assert organization.subscription_items.first() is None @pytest.mark.django_db() def test_save_card(self, organization_factory, mocker, user_factory): @@ -219,7 +219,7 @@ def _setup_stripe_mock(self, mocker, period_end=1_800_000_000, status="active"): status=status, latest_invoice=None, ) - # Patch the subscription service used by Subscription.start() in payment.py + # Patch the subscription service used by SubscriptionItem.start() in payment.py mock_sub_service = mocker.patch( "squarelet.organizations.models.payment.get_payment_provider" ).return_value.get_subscription_service.return_value @@ -236,7 +236,7 @@ def _setup_stripe_mock(self, mocker, period_end=1_800_000_000, status="active"): def test_add_subscription( self, organization_factory, mocker, user_factory, professional_plan_factory ): - """Adding a subscription creates a Subscription record, passes quantity to + """Adding a subscription creates a SubscriptionItem record, passes quantity to Stripe, and sets org.update_on from Stripe's current_period_end.""" user = user_factory() @@ -252,7 +252,7 @@ def test_add_subscription( organization.add_subscription(plan, max_users, user, token="tok_visa") - sub = organization.subscriptions.get(plan=plan) + sub = organization.subscription_items.get(plan=plan) assert sub.quantity == max_users organization.refresh_from_db() @@ -262,10 +262,9 @@ def test_add_subscription( mock_sub_service.create.assert_called_with( stripe_customer=mock_customer.stripe_customer, - plan_id=plan.stripe_id, - quantity=max_users, + items=[{"plan": plan.stripe_id, "quantity": max_users}], billing="charge_automatically", - metadata={"action": f"Subscription ({plan.name})"}, + metadata={"action": f"Subscription ({organization})"}, days_until_due=None, anchor_day=None, cancel_at_period_end=False, @@ -299,10 +298,9 @@ def test_add_second_subscription_uses_billing_anchor( mock_sub_service.create.assert_called_with( stripe_customer=mock_customer.stripe_customer, - plan_id=plan_b.stripe_id, - quantity=1, + items=[{"plan": plan_b.stripe_id, "quantity": 1}], billing="charge_automatically", - metadata={"action": f"Subscription ({plan_b.name})"}, + metadata={"action": f"Subscription ({organization})"}, days_until_due=None, anchor_day=15, cancel_at_period_end=False, @@ -322,13 +320,12 @@ def test_add_subscription_with_invoice_payment_method( organization.add_subscription(plan, 1, user, payment_method="invoice") - assert organization.subscriptions.filter(plan=plan).exists() + assert organization.subscription_items.filter(plan=plan).exists() mock_sub_service.create.assert_called_with( stripe_customer=mock_customer.stripe_customer, - plan_id=plan.stripe_id, - quantity=1, + items=[{"plan": plan.stripe_id, "quantity": 1}], billing="send_invoice", - metadata={"action": f"Subscription ({plan.name})"}, + metadata={"action": f"Subscription ({organization})"}, days_until_due=30, anchor_day=None, cancel_at_period_end=False, @@ -347,13 +344,12 @@ def test_add_subscription_with_existing_card_payment_method( organization.add_subscription(plan, 3, user, payment_method="existing-card") - assert organization.subscriptions.filter(plan=plan).exists() + assert organization.subscription_items.filter(plan=plan).exists() mock_sub_service.create.assert_called_with( stripe_customer=mock_customer.stripe_customer, - plan_id=plan.stripe_id, - quantity=3, + items=[{"plan": plan.stripe_id, "quantity": 3}], billing="charge_automatically", - metadata={"action": f"Subscription ({plan.name})"}, + metadata={"action": f"Subscription ({organization})"}, days_until_due=None, anchor_day=None, cancel_at_period_end=False, @@ -379,13 +375,12 @@ def test_add_subscription_with_new_card_payment_method( ) mocked_save_card.assert_called_with(token, user) - assert organization.subscriptions.filter(plan=plan).exists() + assert organization.subscription_items.filter(plan=plan).exists() mock_sub_service.create.assert_called_with( stripe_customer=mock_customer.stripe_customer, - plan_id=plan.stripe_id, - quantity=2, + items=[{"plan": plan.stripe_id, "quantity": 2}], billing="charge_automatically", - metadata={"action": f"Subscription ({plan.name})"}, + metadata={"action": f"Subscription ({organization})"}, days_until_due=None, anchor_day=None, cancel_at_period_end=False, @@ -406,13 +401,12 @@ def test_add_subscription_auto_detects_card( organization.add_subscription(plan, 4, user) - assert organization.subscriptions.filter(plan=plan).exists() + assert organization.subscription_items.filter(plan=plan).exists() mock_sub_service.create.assert_called_with( stripe_customer=mock_customer.stripe_customer, - plan_id=plan.stripe_id, - quantity=4, + items=[{"plan": plan.stripe_id, "quantity": 4}], billing="charge_automatically", - metadata={"action": f"Subscription ({plan.name})"}, + metadata={"action": f"Subscription ({organization})"}, days_until_due=None, anchor_day=None, cancel_at_period_end=False, @@ -424,24 +418,26 @@ def test_subscription_cancelled( organization_factory, mocker, professional_plan_factory, - subscription_factory, + subscription_item_factory, ): mocker.patch("stripe.Plan.create") plan = professional_plan_factory() organization = organization_factory() - sub = subscription_factory( - organization=organization, plan=plan, subscription_id="sub_test123" + sub = subscription_item_factory( + subscription__organization=organization, + plan=plan, + subscription__subscription_id="sub_test123", ) # Inject mock Stripe subscription via cached_property's __dict__ slot mock_stripe_sub = mocker.MagicMock() - sub.__dict__["stripe_subscription"] = mock_stripe_sub + sub.subscription.__dict__["stripe_subscription"] = mock_stripe_sub - organization.subscription_cancelled(subscription=sub) + organization.subscription_cancelled(subscription=sub.subscription) # Should cancel in Stripe first by calling delete on the stripe_subscription mock_stripe_sub.delete.assert_called_once() # Local subscription should be deleted - assert not Subscription.objects.filter(pk=sub.pk).exists() + assert not SubscriptionItem.objects.filter(pk=sub.pk).exists() @pytest.mark.django_db def test_subscription_cancelled_without_subscription_id( @@ -449,20 +445,22 @@ def test_subscription_cancelled_without_subscription_id( organization_factory, mocker, professional_plan_factory, - subscription_factory, + subscription_item_factory, ): """Should still delete local subscription even if no subscription_id""" mocker.patch("stripe.Plan.create") plan = professional_plan_factory() organization = organization_factory() - sub = subscription_factory( - organization=organization, plan=plan, subscription_id=None + sub = subscription_item_factory( + subscription__organization=organization, + plan=plan, + subscription__subscription_id="", ) - organization.subscription_cancelled(subscription=sub) + organization.subscription_cancelled(subscription=sub.subscription) # Should still delete local subscription - assert not Subscription.objects.filter(pk=sub.pk).exists() + assert not SubscriptionItem.objects.filter(pk=sub.pk).exists() @pytest.mark.django_db def test_subscription_cancelled_stripe_error( @@ -470,14 +468,16 @@ def test_subscription_cancelled_stripe_error( organization_factory, mocker, professional_plan_factory, - subscription_factory, + subscription_item_factory, ): """Should handle Stripe errors gracefully and still delete local subscription""" mocker.patch("stripe.Plan.create") plan = professional_plan_factory() organization = organization_factory() - sub = subscription_factory( - organization=organization, plan=plan, subscription_id="sub_test123" + sub = subscription_item_factory( + subscription__organization=organization, + plan=plan, + subscription__subscription_id="sub_test123", ) mock_stripe_sub = mocker.MagicMock() @@ -493,12 +493,12 @@ def test_subscription_cancelled_stripe_error( return_value=mock_provider, ) - organization.subscription_cancelled(subscription=sub) + organization.subscription_cancelled(subscription=sub.subscription) # Should attempt to delete the Stripe subscription mock_stripe_sub.delete.assert_called_once() # Should still delete local subscription despite error - assert not Subscription.objects.filter(pk=sub.pk).exists() + assert not SubscriptionItem.objects.filter(pk=sub.pk).exists() @pytest.mark.django_db def test_subscription_cancelled_correct_stripe_pattern( @@ -506,15 +506,17 @@ def test_subscription_cancelled_correct_stripe_pattern( organization_factory, mocker, professional_plan_factory, - subscription_factory, + subscription_item_factory, ): """Test subscription_cancelled uses correct Stripe API pattern (retrieve then delete)""" mocker.patch("stripe.Plan.create") plan = professional_plan_factory() organization = organization_factory() - sub = subscription_factory( - organization=organization, plan=plan, subscription_id="sub_test123" + sub = subscription_item_factory( + subscription__organization=organization, + plan=plan, + subscription__subscription_id="sub_test123", ) mock_stripe_sub = mocker.MagicMock() mock_provider = mocker.MagicMock() @@ -526,12 +528,12 @@ def test_subscription_cancelled_correct_stripe_pattern( return_value=mock_provider, ) - organization.subscription_cancelled(subscription=sub) + organization.subscription_cancelled(subscription=sub.subscription) # Verify delete was called on the stripe_subscription instance mock_stripe_sub.delete.assert_called_once() # Verify local subscription was deleted - assert not Subscription.objects.filter(pk=sub.pk).exists() + assert not SubscriptionItem.objects.filter(pk=sub.pk).exists() @pytest.mark.django_db def test_subscription_cancelled_nonexistent_stripe_subscription( @@ -539,14 +541,16 @@ def test_subscription_cancelled_nonexistent_stripe_subscription( organization_factory, mocker, professional_plan_factory, - subscription_factory, + subscription_item_factory, ): """Test graceful handling when Stripe subscription doesn't exist""" mocker.patch("stripe.Plan.create") plan = professional_plan_factory() organization = organization_factory() - sub = subscription_factory( - organization=organization, plan=plan, subscription_id="sub_nonexistent" + sub = subscription_item_factory( + subscription__organization=organization, + plan=plan, + subscription__subscription_id="sub_nonexistent", ) mock_provider = mocker.MagicMock() # stripe_subscription returns None (subscription doesn't exist on Stripe) @@ -557,10 +561,10 @@ def test_subscription_cancelled_nonexistent_stripe_subscription( ) # Should not raise an error - organization.subscription_cancelled(subscription=sub) + organization.subscription_cancelled(subscription=sub.subscription) # Verify local subscription was still deleted - assert not Subscription.objects.filter(pk=sub.pk).exists() + assert not SubscriptionItem.objects.filter(pk=sub.pk).exists() @pytest.mark.django_db(transaction=True) def test_subscription_cancelled_removes_wix_labels( @@ -572,12 +576,12 @@ def test_subscription_cancelled_removes_wix_labels( user2 = user_factory() organization = organization_factory(plans=[wix_plan], users=[user1, user2]) - sub = organization.subscriptions.first() + sub = organization.subscription_items.first() sub.__dict__["stripe_subscription"] = None # no Stripe sub to cancel mock_unsync = mocker.patch("squarelet.organizations.tasks.unsync_wix.delay") - organization.subscription_cancelled(subscription=sub) + organization.subscription_cancelled(subscription=sub.subscription) # Should call unsync for each user assert mock_unsync.call_count == 2 @@ -593,12 +597,12 @@ def test_subscription_cancelled_no_wix_no_unsync( user = user_factory() organization = organization_factory(plans=[non_wix_plan], users=[user]) - sub = organization.subscriptions.first() + sub = organization.subscription_items.first() sub.__dict__["stripe_subscription"] = None mock_unsync = mocker.patch("squarelet.organizations.tasks.unsync_wix.delay") - organization.subscription_cancelled(subscription=sub) + organization.subscription_cancelled(subscription=sub.subscription) mock_unsync.assert_not_called() @@ -614,7 +618,7 @@ def test_modify_subscription_changes_plan( organization.modify_subscription(old_plan, new_plan, 5, user) - assert organization.subscriptions.filter(plan=new_plan).exists() + assert organization.subscription_items.filter(plan=new_plan).exists() @pytest.mark.django_db(transaction=True) def test_remove_subscription_wix_removes_labels( @@ -712,7 +716,7 @@ def test_merge_fks(self): if f.many_to_many and not f.auto_created ] ) - == 4 + == 3 ) @pytest.mark.django_db() @@ -1327,7 +1331,7 @@ class TestMultipleSubscriptions: def test_add_subscription_creates_new( self, organization_factory, plan_factory, user_factory, mocker ): - """add_subscription creates a new Subscription for an org with none.""" + """add_subscription creates a new SubscriptionItem for an org with none.""" org = organization_factory() plan = plan_factory() user = user_factory() @@ -1340,17 +1344,23 @@ def test_add_subscription_creates_new( # Pass payment_method explicitly to skip card detection (avoids Stripe call) org.add_subscription(plan, org.max_users, user, payment_method="invoice") - assert Subscription.objects.filter(organization=org, plan=plan).exists() + assert SubscriptionItem.objects.filter( + subscription__organization=org, plan=plan + ).exists() @pytest.mark.django_db def test_add_subscription_same_plan_raises( - self, organization_factory, plan_factory, subscription_factory, user_factory + self, + organization_factory, + plan_factory, + subscription_item_factory, + user_factory, ): """add_subscription raises ValueError if org already has active sub for plan.""" org = organization_factory() plan = plan_factory() user = user_factory() - subscription_factory(organization=org, plan=plan) + subscription_item_factory(subscription__organization=org, plan=plan) with pytest.raises( SubscriptionError, match="already has an active subscription" @@ -1376,7 +1386,9 @@ def test_add_subscription_two_different_plans( org.add_subscription(plan_a, org.max_users, user, payment_method="invoice") org.add_subscription(plan_b, org.max_users, user, payment_method="invoice") - assert Subscription.objects.filter(organization=org).count() == 2 + assert ( + SubscriptionItem.objects.filter(subscription__organization=org).count() == 2 + ) @pytest.mark.django_db def test_add_subscription_none_max_users_uses_plan_minimum( @@ -1399,7 +1411,7 @@ def test_add_subscription_none_max_users_uses_plan_minimum( org.add_subscription(plan, None, user, payment_method="invoice") - sub = Subscription.objects.get(organization=org, plan=plan) + sub = SubscriptionItem.objects.get(subscription__organization=org, plan=plan) assert sub.quantity == plan.minimum_users @pytest.mark.django_db @@ -1407,7 +1419,7 @@ def test_remove_subscription_by_plan( self, organization_factory, plan_factory, - subscription_factory, + subscription_item_factory, user_factory, mocker, ): @@ -1416,8 +1428,8 @@ def test_remove_subscription_by_plan( plan_a = plan_factory() plan_b = plan_factory() user = user_factory() - sub_a = subscription_factory(organization=org, plan=plan_a) - sub_b = subscription_factory(organization=org, plan=plan_b) + sub_a = subscription_item_factory(subscription__organization=org, plan=plan_a) + sub_b = subscription_item_factory(subscription__organization=org, plan=plan_b) mocker.patch( "squarelet.organizations.models.Subscription.stripe_subscription", @@ -1426,16 +1438,17 @@ def test_remove_subscription_by_plan( org.remove_subscription(plan_a, user) - sub_a.refresh_from_db() - assert sub_a.cancelled - assert Subscription.objects.filter(pk=sub_b.pk).exists() + # Both lines share one subscription, so removing one drops just that + # line and leaves the other billing. + assert not SubscriptionItem.objects.filter(pk=sub_a.pk).exists() + assert SubscriptionItem.objects.filter(pk=sub_b.pk).exists() @pytest.mark.django_db def test_modify_subscription( self, organization_factory, plan_factory, - subscription_factory, + subscription_item_factory, user_factory, ): """modify_subscription updates the plan on the matching subscription.""" @@ -1443,11 +1456,11 @@ def test_modify_subscription( plan_a = plan_factory() plan_b = plan_factory() user = user_factory() - subscription_factory(organization=org, plan=plan_a) + subscription_item_factory(subscription__organization=org, plan=plan_a) org.modify_subscription(plan_a, plan_b, org.max_users, user) - assert org.subscriptions.filter(plan=plan_b).exists() + assert org.subscription_items.filter(plan=plan_b).exists() @pytest.mark.django_db def test_modify_subscription_missing_plan_raises( @@ -1464,25 +1477,27 @@ def test_modify_subscription_missing_plan_raises( @pytest.mark.django_db def test_has_active_subscription_with_plan_arg( - self, organization_factory, plan_factory, subscription_factory + self, organization_factory, plan_factory, subscription_item_factory ): """has_active_subscription(plan=X) is True for X, False for others.""" org = organization_factory() plan_a = plan_factory() plan_b = plan_factory() - subscription_factory(organization=org, plan=plan_a) + subscription_item_factory(subscription__organization=org, plan=plan_a) assert org.has_active_subscription(plan=plan_a) assert not org.has_active_subscription(plan=plan_b) @pytest.mark.django_db def test_has_active_subscription_includes_cancelled( - self, organization_factory, plan_factory, subscription_factory + self, organization_factory, plan_factory, subscription_item_factory ): """cancelled=True means pending cancellation at period end — still active.""" org = organization_factory() plan = plan_factory() - subscription_factory(organization=org, plan=plan, cancelled=True) + subscription_item_factory( + subscription__organization=org, plan=plan, subscription__cancelled=True + ) assert org.has_active_subscription(plan=plan) assert org.has_active_subscription() diff --git a/squarelet/organizations/tests/models/test_plan.py b/squarelet/organizations/tests/models/test_plan.py index 19b95aa9e..8efe48d25 100644 --- a/squarelet/organizations/tests/models/test_plan.py +++ b/squarelet/organizations/tests/models/test_plan.py @@ -86,46 +86,52 @@ def test_has_available_slots_sunlight_no_wix(self, plan_factory): @override_settings(MAX_SUNLIGHT_SUBSCRIPTIONS=15) @pytest.mark.django_db def test_has_available_slots_sunlight_under_limit( - self, plan_factory, subscription_factory + self, plan_factory, subscription_item_factory ): """Sunlight wix plan under limit has available slots""" sunlight_plan = plan_factory(slug="sunlight-essential-monthly", wix=True) # Create 10 active subscriptions (under limit of 15) - subscription_factory.create_batch(10, plan=sunlight_plan, cancelled=False) + subscription_item_factory.create_batch( + 10, plan=sunlight_plan, subscription__cancelled=False + ) assert sunlight_plan.has_available_slots() is True @override_settings(MAX_SUNLIGHT_SUBSCRIPTIONS=15) @pytest.mark.django_db def test_has_available_slots_sunlight_at_limit( - self, plan_factory, subscription_factory + self, plan_factory, subscription_item_factory ): """Sunlight wix plan at limit has no available slots""" sunlight_plan = plan_factory(slug="sunlight-essential-monthly", wix=True) # Create 15 active subscriptions (at limit) - subscription_factory.create_batch(15, plan=sunlight_plan, cancelled=False) + subscription_item_factory.create_batch( + 15, plan=sunlight_plan, subscription__cancelled=False + ) assert sunlight_plan.has_available_slots() is False @override_settings(MAX_SUNLIGHT_SUBSCRIPTIONS=15) @pytest.mark.django_db def test_has_available_slots_sunlight_over_limit( - self, plan_factory, subscription_factory + self, plan_factory, subscription_item_factory ): """Sunlight wix plan over limit has no available slots""" sunlight_plan = plan_factory(slug="sunlight-essential-monthly", wix=True) # Create 20 active subscriptions (over limit) - subscription_factory.create_batch(20, plan=sunlight_plan, cancelled=False) + subscription_item_factory.create_batch( + 20, plan=sunlight_plan, subscription__cancelled=False + ) assert sunlight_plan.has_available_slots() is False @override_settings(MAX_SUNLIGHT_SUBSCRIPTIONS=15) @pytest.mark.django_db def test_has_available_slots_counts_all_sunlight_variants( - self, plan_factory, subscription_factory + self, plan_factory, subscription_item_factory ): """Limit is shared across all Sunlight plan variants""" sunlight_basic = plan_factory(slug="sunlight-essential-monthly", wix=True) @@ -133,9 +139,13 @@ def test_has_available_slots_counts_all_sunlight_variants( # Create 10 subscriptions for basic, 5 for premium (total 15) for _ in range(10): - subscription_factory(plan=sunlight_basic, cancelled=False) + subscription_item_factory( + plan=sunlight_basic, subscription__cancelled=False + ) for _ in range(5): - subscription_factory(plan=sunlight_premium, cancelled=False) + subscription_item_factory( + plan=sunlight_premium, subscription__cancelled=False + ) # Both plans should show no slots available assert sunlight_basic.has_available_slots() is False @@ -144,16 +154,16 @@ def test_has_available_slots_counts_all_sunlight_variants( @override_settings(MAX_SUNLIGHT_SUBSCRIPTIONS=15) @pytest.mark.django_db def test_has_available_slots_includes_cancelled( - self, plan_factory, subscription_factory + self, plan_factory, subscription_item_factory ): """cancelled=True means pending cancellation — counts toward limit.""" sunlight_plan = plan_factory(slug="sunlight-essential-monthly", wix=True) # Create 10 active and 5 pending-cancellation subscriptions (total 15 = limit) for _ in range(10): - subscription_factory(plan=sunlight_plan, cancelled=False) + subscription_item_factory(plan=sunlight_plan, subscription__cancelled=False) for _ in range(5): - subscription_factory(plan=sunlight_plan, cancelled=True) + subscription_item_factory(plan=sunlight_plan, subscription__cancelled=True) # 15 total subscriptions = at the limit, no slots available assert sunlight_plan.has_available_slots() is False diff --git a/squarelet/organizations/tests/models/test_subscription.py b/squarelet/organizations/tests/models/test_subscription.py index ba0bfec1b..591485c9c 100644 --- a/squarelet/organizations/tests/models/test_subscription.py +++ b/squarelet/organizations/tests/models/test_subscription.py @@ -9,26 +9,28 @@ import pytest import stripe +# Squarelet +from squarelet.organizations.models import SubscriptionItem + # Local from .test_invoice import Invoice, create_mock_stripe_invoice class TestSubscription: - """Unit tests for the Subscription model""" + """Unit tests for the Subscription and SubscriptionItem models""" - def test_str(self, subscription_factory): - subscription = subscription_factory.build() + def test_str(self, subscription_item_factory): + subscription = subscription_item_factory.build() assert ( - str(subscription) - == f"Subscription: {subscription.organization} to {subscription.plan.name}" + str(subscription) == f"SubscriptionItem: {subscription.organization} to " + f"{subscription.plan.name}" ) def test_stripe_subscription(self, subscription_factory, mocker): mocked = mocker.patch("stripe.Subscription.retrieve") stripe_subscription = "stripe_subscription" mocked.return_value = stripe_subscription - subscription_id = "subscription_id" - subscription = subscription_factory.build(subscription_id=subscription_id) + subscription = subscription_factory.build(subscription_id="subscription_id") assert subscription.stripe_subscription == stripe_subscription def test_stripe_subscription_empty(self, subscription_factory): @@ -36,9 +38,9 @@ def test_stripe_subscription_empty(self, subscription_factory): assert subscription.stripe_subscription is None @pytest.mark.django_db() - def test_start(self, subscription_factory, professional_plan_factory, mocker): + def test_start(self, subscription_item_factory, professional_plan_factory, mocker): plan = professional_plan_factory() - subscription = subscription_factory(plan=plan) + subscription = subscription_item_factory(plan=plan).subscription # Mock stripe subscription creation stripe_subscription_id = "sub_test123" @@ -62,10 +64,9 @@ def test_start(self, subscription_factory, professional_plan_factory, mocker): mock_sub_service.create.assert_called_with( stripe_customer=mocked_customer.stripe_customer, - plan_id=subscription.plan.stripe_id, - quantity=subscription.quantity, + items=subscription.stripe_items(), billing="charge_automatically", - metadata={"action": f"Subscription ({plan.name})"}, + metadata={"action": f"Subscription ({subscription.organization})"}, days_until_due=None, anchor_day=None, cancel_at_period_end=False, @@ -74,7 +75,7 @@ def test_start(self, subscription_factory, professional_plan_factory, mocker): @pytest.mark.django_db() def test_start_no_auto_renew( - self, subscription_factory, professional_plan_factory, mocker + self, subscription_item_factory, professional_plan_factory, mocker ): """A plan with auto_renew disabled starts the Stripe subscription with cancel_at_period_end=True so it does not automatically renew.""" @@ -85,7 +86,7 @@ def test_start_no_auto_renew( plan = professional_plan_factory() plan.auto_renew = False plan.save() - subscription = subscription_factory(plan=plan) + subscription = subscription_item_factory(plan=plan).subscription period_end_ts = 1_800_000_000 mock_stripe_subscription = Mock( @@ -109,12 +110,9 @@ def test_start_no_auto_renew( mock_sub_svc.create.assert_called_with( stripe_customer=mocked_customer.stripe_customer, - plan_id=subscription.plan.stripe_id, - quantity=subscription.quantity, + items=subscription.stripe_items(), billing="charge_automatically", - metadata={ - "action": f"Subscription ({plan.name})", - }, + metadata={"action": f"Subscription ({subscription.organization})"}, days_until_due=None, anchor_day=None, cancel_at_period_end=True, @@ -124,22 +122,26 @@ def test_start_no_auto_renew( ).date() assert subscription.cancel_at == expected_date - def test_start_existing(self, subscription_factory, mocker): + @pytest.mark.django_db() + def test_start_existing(self, subscription_item_factory, mocker): """If there is an existing subscription, do not start another one""" - subscription = subscription_factory.build() + subscription = subscription_item_factory().subscription mocked = mocker.patch("squarelet.organizations.models.Organization.customer") mocker.patch("squarelet.organizations.models.Subscription.stripe_subscription") subscription.start() - mocked.subscriptions.create.assert_not_called() + mocked.subscription_items.create.assert_not_called() - def test_start_free(self, subscription_factory, mocker): + @pytest.mark.django_db() + def test_start_free(self, subscription_item_factory, mocker): """If there is an existing subscription, do not start another one""" - subscription = subscription_factory.build() + subscription = subscription_item_factory().subscription mocked = mocker.patch("squarelet.organizations.models.Organization.customer") subscription.start() - mocked.subscriptions.create.assert_not_called() + mocked.subscription_items.create.assert_not_called() - def test_cancel(self, subscription_factory, mocker): + @pytest.mark.django_db() + def test_cancel(self, subscription_item_factory, mocker): + subscription = subscription_item_factory().subscription mocked_save = mocker.patch("squarelet.organizations.models.Subscription.save") mocked_stripe_subscription = mocker.patch( "squarelet.organizations.models.Subscription.stripe_subscription" @@ -153,7 +155,6 @@ def test_cancel(self, subscription_factory, mocker): mock_sub_svc = mock_provider.get_subscription_service.return_value mock_sub_svc.cancel_at_period_end.return_value = mock_updated mock_sub_svc.get_current_period_end.return_value = period_end_ts - subscription = subscription_factory.build() subscription.cancel() mock_sub_svc.cancel_at_period_end.assert_called_once_with( mocked_stripe_subscription, @@ -166,105 +167,143 @@ def test_cancel(self, subscription_factory, mocker): assert subscription.cancel_at == expected_date mocked_save.assert_called() - def test_cancel_no_stripe_subscription(self, subscription_factory, mocker): + @pytest.mark.django_db() + def test_cancel_no_stripe_subscription(self, subscription_item_factory, mocker): """cancel_at stays None when there is no Stripe subscription (free plan).""" + subscription = subscription_item_factory().subscription mocked_save = mocker.patch("squarelet.organizations.models.Subscription.save") mocker.patch( "squarelet.organizations.models.Subscription.stripe_subscription", new=None, ) - subscription = subscription_factory.build() subscription.cancel() assert subscription.cancelled assert subscription.cancel_at is None mocked_save.assert_called() - def test_modify_start( - self, subscription_factory, professional_plan_factory, mocker + @pytest.mark.django_db() + def test_modify_pushes_the_new_plan_to_stripe( + self, subscription_item_factory, professional_plan_factory, mocker ): - mocked_save = mocker.patch("squarelet.organizations.models.Subscription.save") - mocked_start = mocker.patch("squarelet.organizations.models.Subscription.start") - plan = professional_plan_factory.build() - subscription = subscription_factory.build() - subscription.modify(plan) - mocked_save.assert_called() - mocked_start.assert_called() + """Changing a line's plan saves it and re-syncs the whole subscription.""" + item = subscription_item_factory() + plan = professional_plan_factory() + mocked_modify = mocker.patch( + "squarelet.organizations.models.Subscription.stripe_modify" + ) + item.modify(plan) + item.refresh_from_db() + assert item.plan == plan + mocked_modify.assert_called_once() - def test_modify_cancel( - self, subscription_factory, professional_plan_factory, plan_factory, mocker + @pytest.mark.django_db() + def test_cancel_last_item_cancels_the_subscription( + self, subscription_item_factory, mocker ): - mocked_save = mocker.patch("squarelet.organizations.models.Subscription.save") - mocked_stripe_subscription = mocker.patch( - "squarelet.organizations.models.Subscription.stripe_subscription" + """The only line left cancels the whole subscription at period end.""" + item = subscription_item_factory() + mocked_cancel = mocker.patch( + "squarelet.organizations.models.Subscription.cancel" ) - plan = professional_plan_factory.build() - free_plan = plan_factory.build() - subscription = subscription_factory.build(plan=plan, subscription_id="id") - subscription.modify(free_plan) - mocked_save.assert_called() - mocked_stripe_subscription.delete.assert_called() - assert subscription.subscription_id is None + item.cancel() + mocked_cancel.assert_called_once() + assert SubscriptionItem.objects.filter(pk=item.pk).exists() - def test_modify_modify( - self, subscription_factory, professional_plan_factory, mocker + @pytest.mark.django_db() + def test_cancel_one_of_several_items_drops_only_that_line( + self, subscription_item_factory, plan_factory, mocker ): - mocked_save = mocker.patch("squarelet.organizations.models.Subscription.save") + """Other lines keep billing; the removed line leaves no proration.""" + item = subscription_item_factory( + subscription__subscription_id="sub_multi", stripe_item_id="si_one" + ) + subscription_item_factory( + subscription=item.subscription, plan=plan_factory(name="Second Plan") + ) + mocker.patch("squarelet.organizations.models.Subscription.stripe_subscription") + mock_sub_svc = mocker.patch( + "squarelet.organizations.models.payment.get_payment_provider" + ).return_value.get_subscription_service.return_value + + item.cancel() + + mock_sub_svc.modify.assert_called_once_with( + "sub_multi", + items=[{"id": "si_one", "deleted": True}], + proration_behavior="none", + ) + assert not SubscriptionItem.objects.filter(pk=item.pk).exists() + + @pytest.mark.django_db() + def test_stripe_modify_sends_every_line( + self, subscription_item_factory, professional_plan_factory, mocker + ): + """stripe_modify pushes all of the subscription's lines, with their ids.""" + item = subscription_item_factory( + plan=professional_plan_factory(), + subscription__subscription_id="sub_mod", + stripe_item_id="si_mod", + ) + subscription = item.subscription mock_sub_svc = mocker.patch( "squarelet.organizations.models.payment.get_payment_provider" ).return_value.get_subscription_service.return_value mock_sub_svc.modify.return_value = Mock(status="active") mock_sub_svc.get_current_period_end.return_value = None mocker.patch("squarelet.organizations.models.Subscription.stripe_subscription") - plan = professional_plan_factory.build() - subscription = subscription_factory.build(plan=plan) - subscription.modify(plan) - mocked_save.assert_called() + + subscription.stripe_modify() + assert subscription.cancel_at is None mock_sub_svc.modify.assert_called_with( - subscription.subscription_id, + "sub_mod", cancel_at_period_end=False, items=[ { - "id": subscription.stripe_subscription["items"]["data"][0].id, - "plan": subscription.plan.stripe_id, - "quantity": subscription.quantity, + "id": "si_mod", + "plan": item.plan.stripe_id, + "quantity": item.quantity, } ], billing="charge_automatically", - metadata={"action": f"Subscription ({plan.name})"}, + metadata={"action": f"Subscription ({subscription.organization})"}, days_until_due=None, ) - def test_modify_modify_no_auto_renew( - self, subscription_factory, professional_plan_factory, mocker + @pytest.mark.django_db() + def test_stripe_modify_no_auto_renew( + self, subscription_item_factory, professional_plan_factory, mocker ): - """Modifying to a plan with auto_renew disabled flags the Stripe - subscription to cancel at period end and sets cancel_at.""" - mocker.patch("squarelet.organizations.models.Subscription.save") + """A plan with auto_renew off flags the Stripe subscription to end.""" + plan = professional_plan_factory() + plan.auto_renew = False + plan.save() + item = subscription_item_factory( + plan=plan, subscription__subscription_id="sub_norenew" + ) period_end_ts = 1_800_000_000 - mock_updated = Mock(status="active") mock_sub_svc = mocker.patch( "squarelet.organizations.models.payment.get_payment_provider" ).return_value.get_subscription_service.return_value - mock_sub_svc.modify.return_value = mock_updated + mock_sub_svc.modify.return_value = Mock(status="active") mock_sub_svc.get_current_period_end.return_value = period_end_ts mocker.patch("squarelet.organizations.models.Subscription.stripe_subscription") - plan = professional_plan_factory.build(auto_renew=False) - subscription = subscription_factory.build(plan=plan) - subscription.modify(plan) + + item.subscription.stripe_modify() + assert mock_sub_svc.modify.call_args.kwargs["cancel_at_period_end"] is True expected_date = datetime.fromtimestamp( period_end_ts, tz=get_current_timezone() ).date() - assert subscription.cancel_at == expected_date + assert item.subscription.cancel_at == expected_date @pytest.mark.django_db() def test_start_creates_invoice_with_card( - self, subscription_factory, professional_plan_factory, mocker + self, subscription_item_factory, professional_plan_factory, mocker ): """Test that subscription.start() creates an Invoice record for card payment""" plan = professional_plan_factory() - subscription = subscription_factory(plan=plan) + subscription = subscription_item_factory(plan=plan).subscription # Mock Stripe subscription creation stripe_subscription_id = "sub_test123" @@ -311,7 +350,7 @@ def test_start_creates_invoice_with_card( @pytest.mark.django_db() def test_start_creates_invoice_with_invoice_payment( - self, subscription_factory, plan_factory, mocker + self, subscription_item_factory, plan_factory, mocker ): """Test that subscription.start() creates Invoice for invoice payment method""" # Mock Stripe Plan creation @@ -324,7 +363,9 @@ def test_start_creates_invoice_with_invoice_payment( base_price=240, minimum_users=1, ) - subscription = subscription_factory(plan=plan) + subscription = subscription_item_factory( + plan=plan, subscription__interval="annual" + ).subscription # Mock Stripe subscription creation stripe_subscription_id = "sub_annual123" @@ -362,10 +403,9 @@ def test_start_creates_invoice_with_invoice_payment( # Verify subscription was created with send_invoice billing mock_provider.get_subscription_service.return_value.create.assert_called_with( stripe_customer=mocked_customer.stripe_customer, - plan_id=subscription.plan.stripe_id, - quantity=subscription.quantity, + items=subscription.stripe_items(), billing="send_invoice", - metadata={"action": f"Subscription ({plan.name})"}, + metadata={"action": f"Subscription ({subscription.organization})"}, days_until_due=30, anchor_day=None, cancel_at_period_end=False, @@ -380,12 +420,12 @@ def test_start_creates_invoice_with_invoice_payment( @pytest.mark.django_db() def test_start_free_plan_no_invoice( - self, subscription_factory, plan_factory, mocker + self, subscription_item_factory, plan_factory, mocker ): """Test that free plans don't create invoices""" mocker.patch("stripe.Plan.create") plan = plan_factory() # Free plan (no base_price = free) - subscription = subscription_factory(plan=plan) + subscription = subscription_item_factory(plan=plan).subscription mocked_customer = mocker.patch( "squarelet.organizations.models.Organization.customer" @@ -402,11 +442,11 @@ def test_start_free_plan_no_invoice( @pytest.mark.django_db() def test_start_handles_stripe_invoice_retrieval_error( - self, subscription_factory, professional_plan_factory, mocker + self, subscription_item_factory, professional_plan_factory, mocker ): """Test that subscription still succeeds if invoice retrieval fails""" plan = professional_plan_factory() - subscription = subscription_factory(plan=plan) + subscription = subscription_item_factory(plan=plan).subscription # Mock Stripe subscription creation stripe_subscription_id = "sub_test123" @@ -440,12 +480,12 @@ def test_start_handles_stripe_invoice_retrieval_error( @pytest.mark.django_db() def test_start_caches_stripe_status( - self, subscription_factory, professional_plan_factory, mocker + self, subscription_item_factory, professional_plan_factory, mocker ): """start() caches stripe_status and current_period_end from Stripe response""" plan = professional_plan_factory() - subscription = subscription_factory(plan=plan) + subscription = subscription_item_factory(plan=plan).subscription period_end_ts = 1800000000 mock_stripe_sub = Mock( diff --git a/squarelet/organizations/tests/payments/test_stripe_modern.py b/squarelet/organizations/tests/payments/test_stripe_modern.py index ea697f29b..5098992f3 100644 --- a/squarelet/organizations/tests/payments/test_stripe_modern.py +++ b/squarelet/organizations/tests/payments/test_stripe_modern.py @@ -276,8 +276,7 @@ def test_create_uses_collection_method(self, subscription_service, mocker): mock_create = mocker.patch("stripe.Subscription.create") subscription_service.create( stripe_customer=mock_customer, - plan_id="plan_123", - quantity=5, + items=[{"plan": "plan_123", "quantity": 5}], billing="charge_automatically", metadata={"action": "test"}, days_until_due=None, @@ -297,8 +296,7 @@ def test_create_translates_send_invoice(self, subscription_service, mocker): mock_create = mocker.patch("stripe.Subscription.create") subscription_service.create( stripe_customer=mock_customer, - plan_id="plan_123", - quantity=1, + items=[{"plan": "plan_123", "quantity": 1}], billing="send_invoice", metadata={}, days_until_due=30, diff --git a/squarelet/organizations/tests/test_admin.py b/squarelet/organizations/tests/test_admin.py index a1e0a8e14..ecb55bc05 100644 --- a/squarelet/organizations/tests/test_admin.py +++ b/squarelet/organizations/tests/test_admin.py @@ -376,14 +376,14 @@ def test_filter_by_plan_returns_orgs_with_that_subscription( request_factory, organization_factory, plan_factory, - subscription_factory, + subscription_item_factory, ): pro_plan = plan_factory(name="Pro") free_plan = plan_factory(name="Free") pro_org = organization_factory() free_org = organization_factory() - subscription_factory(organization=pro_org, plan=pro_plan) - subscription_factory(organization=free_org, plan=free_plan) + subscription_item_factory(subscription__organization=pro_org, plan=pro_plan) + subscription_item_factory(subscription__organization=free_org, plan=free_plan) request = request_factory.get("/") result = self._filter({"plan": [str(pro_plan.pk)]}).queryset( @@ -398,12 +398,14 @@ def test_filter_by_plan_includes_cancelled_subscriptions( request_factory, organization_factory, plan_factory, - subscription_factory, + subscription_item_factory, ): """Cancelled subs still count as subscribed per design decision.""" pro_plan = plan_factory(name="Pro") org = organization_factory() - subscription_factory(organization=org, plan=pro_plan, cancelled=True) + subscription_item_factory( + subscription__organization=org, plan=pro_plan, subscription__cancelled=True + ) request = request_factory.get("/") result = self._filter({"plan": [str(pro_plan.pk)]}).queryset( @@ -417,12 +419,12 @@ def test_filter_by_none_returns_orgs_without_subscriptions( request_factory, organization_factory, plan_factory, - subscription_factory, + subscription_item_factory, ): plan = plan_factory(name="Pro") subscribed = organization_factory() unsubscribed = organization_factory() - subscription_factory(organization=subscribed, plan=plan) + subscription_item_factory(subscription__organization=subscribed, plan=plan) request = request_factory.get("/") result = self._filter({"plan": ["none"]}).queryset( @@ -437,12 +439,12 @@ def test_unset_filter_is_a_no_op( request_factory, organization_factory, plan_factory, - subscription_factory, + subscription_item_factory, ): plan = plan_factory(name="Pro") subscribed = organization_factory() unsubscribed = organization_factory() - subscription_factory(organization=subscribed, plan=plan) + subscription_item_factory(subscription__organization=subscribed, plan=plan) request = request_factory.get("/") result = self._filter({}).queryset(request, Organization.objects.all()) @@ -455,13 +457,13 @@ def test_filter_result_has_no_duplicate_rows( request_factory, organization_factory, plan_factory, - subscription_factory, + subscription_item_factory, ): plan = plan_factory(name="Pro") other_plan = plan_factory(name="Other") org = organization_factory() - subscription_factory(organization=org, plan=plan) - subscription_factory(organization=org, plan=other_plan) + subscription_item_factory(subscription__organization=org, plan=plan) + subscription_item_factory(subscription__organization=org, plan=other_plan) request = request_factory.get("/") result = self._filter({"plan": [str(plan.pk)]}).queryset( diff --git a/squarelet/organizations/tests/test_migrations.py b/squarelet/organizations/tests/test_migrations.py new file mode 100644 index 000000000..22f865017 --- /dev/null +++ b/squarelet/organizations/tests/test_migrations.py @@ -0,0 +1,177 @@ +"""Tests for data migrations. + +Migrations run once, on production, with no way to try again - so the ones +that move data are worth exercising against real rows rather than reasoning +about. These roll the schema back to just before the migration under test, +build the rows it will find, and roll forward. +""" + +# Historical models come back from `apps.get_model` as classes, and are named +# like classes. +# pylint: disable=invalid-name + +# Django +from django.db import connection +from django.db.migrations.executor import MigrationExecutor +from django.db.migrations.loader import MigrationLoader + +# Third Party +import pytest + +APP = "organizations" +# The migration under test, found by name rather than by number: the squash +# further up the stack renumbers it and removes the one it currently depends +# on, so anything pinned here would break on rebase. +UNDER_TEST = "subscription_parent" + + +def _bracket(): + """Return the migration under test and the one immediately before it.""" + loader = MigrationLoader(connection) + target = next( + name + for app, name in loader.graph.nodes + if app == APP and name.endswith(UNDER_TEST) + ) + before = next( + name + for app, name in loader.graph.node_map[(APP, target)].parents + if app == APP + ) + return before, target + + +def migrate_to(target): + """Run the organizations app to `target` and return its model registry.""" + executor = MigrationExecutor(connection) + executor.loader.build_graph() + executor.migrate([(APP, target)]) + executor.loader.build_graph() + return executor.loader.project_state([(APP, target)]).apps + + +def migrate_to_latest(): + """Put the app back at its newest migration, whatever that is. + + Not the migration under test: branches above this one add migrations + after it, and leaving the schema short of them would break every test + that ran afterwards. + """ + executor = MigrationExecutor(connection) + executor.loader.build_graph() + executor.migrate(executor.loader.graph.leaf_nodes(APP)) + + +@pytest.mark.django_db(transaction=True) +class TestAdoptItemsIntoSubscriptions: + """0085 gives every subscription line a parent, one per billing shape.""" + + @pytest.fixture(autouse=True) + def _leave_the_database_migrated(self): + """Put the schema back for whatever runs next. + + The rows go first: a test that deliberately leaves data 0085 refuses + would otherwise be refused again on the way back up, and every test + after it would run against the wrong schema. + """ + yield + with connection.cursor() as cursor: + cursor.execute("DELETE FROM organizations_subscriptionitem") + migrate_to_latest() + + @staticmethod + def _organization(apps, name): + Organization = apps.get_model(APP, "Organization") + return Organization.objects.create(name=name, slug=name.lower()) + + @staticmethod + def _plan(apps, name, annual=False): + Plan = apps.get_model(APP, "Plan") + return Plan.objects.create(name=name, slug=name.lower(), annual=annual) + + def test_lines_of_one_shape_share_a_parent(self): + """The point of the split: two monthly lines, one Stripe subscription.""" + old = migrate_to(_bracket()[0]) + SubscriptionItem = old.get_model(APP, "SubscriptionItem") + organization = self._organization(old, "Shared") + SubscriptionItem.objects.create( + organization=organization, + plan=self._plan(old, "Paid"), + subscription_id="sub_shared", + ) + # Comped lines never reached Stripe, so they carry no id. The old + # schema made `subscription_id` unique, which is why a group can hold + # at most one line that names a Stripe subscription. + SubscriptionItem.objects.create( + organization=organization, plan=self._plan(old, "Comped") + ) + + new = migrate_to(_bracket()[1]) + + Subscription = new.get_model(APP, "Subscription") + subscriptions = Subscription.objects.filter(organization=organization.pk) + assert subscriptions.count() == 1 + assert subscriptions.first().items.count() == 2 + assert subscriptions.first().subscription_id == "sub_shared" + + def test_a_different_shape_gets_its_own_parent(self): + """Stripe cannot bill monthly and annual on one subscription.""" + old = migrate_to(_bracket()[0]) + SubscriptionItem = old.get_model(APP, "SubscriptionItem") + organization = self._organization(old, "Mixed") + SubscriptionItem.objects.create( + organization=organization, plan=self._plan(old, "Monthly") + ) + SubscriptionItem.objects.create( + organization=organization, + plan=self._plan(old, "Yearly", annual=True), + ) + + new = migrate_to(_bracket()[1]) + + Subscription = new.get_model(APP, "Subscription") + shapes = set( + Subscription.objects.filter(organization=organization.pk).values_list( + "interval", "collection_method" + ) + ) + assert shapes == { + ("monthly", "charge_automatically"), + ("annual", "send_invoice"), + } + + def test_two_stripe_subscriptions_of_one_shape_are_refused(self): + """One row holds one Stripe id, so merging would orphan the other.""" + old = migrate_to(_bracket()[0]) + SubscriptionItem = old.get_model(APP, "SubscriptionItem") + organization = self._organization(old, "Doubled") + for plan_name, stripe_id in (("A", "sub_a"), ("B", "sub_b")): + SubscriptionItem.objects.create( + organization=organization, + plan=self._plan(old, plan_name), + subscription_id=stripe_id, + ) + + with pytest.raises(Exception, match="more than one"): + migrate_to(_bracket()[1]) + + def test_a_parent_stops_only_when_every_line_has(self): + """One cancelled line among several is that line's own business.""" + old = migrate_to(_bracket()[0]) + SubscriptionItem = old.get_model(APP, "SubscriptionItem") + organization = self._organization(old, "Partly") + SubscriptionItem.objects.create( + organization=organization, + plan=self._plan(old, "Staying"), + subscription_id="sub_partly", + cancelled=False, + ) + SubscriptionItem.objects.create( + organization=organization, plan=self._plan(old, "Going"), cancelled=True + ) + + new = migrate_to(_bracket()[1]) + + Subscription = new.get_model(APP, "Subscription") + subscription = Subscription.objects.get(organization=organization.pk) + assert not subscription.cancelled diff --git a/squarelet/organizations/tests/test_querysets.py b/squarelet/organizations/tests/test_querysets.py index 68b5ca8bb..9ad252cfc 100644 --- a/squarelet/organizations/tests/test_querysets.py +++ b/squarelet/organizations/tests/test_querysets.py @@ -19,7 +19,7 @@ Membership, Organization, Plan, - Subscription, + SubscriptionItem, ) from squarelet.organizations.tests.factories import ( ChargeFactory, @@ -27,7 +27,7 @@ MembershipFactory, OrganizationFactory, PlanFactory, - SubscriptionFactory, + SubscriptionItemFactory, ) from squarelet.users.tests.factories import UserFactory @@ -207,7 +207,7 @@ def test_create_individual_basic(self): assert org.change_logs.filter( reason=ChangeLogReason.created, user=user, - to_plan=org.plans.first(), + to_plan=org.get_plans().first(), to_max_users=1, ).exists() @@ -303,7 +303,7 @@ def test_membership_create_with_wix_plan_triggers_sync(mocker): user = UserFactory() wix_plan = PlanFactory(wix=True) org = OrganizationFactory() - SubscriptionFactory(organization=org, plan=wix_plan) + SubscriptionItemFactory(subscription__organization=org, plan=wix_plan) # Create membership membership = Membership.objects.create(user=user, organization=org) @@ -324,7 +324,7 @@ def test_membership_create_without_wix_plan_no_sync(mocker): user = UserFactory() non_wix_plan = PlanFactory(wix=False) org = OrganizationFactory() - SubscriptionFactory(organization=org, plan=non_wix_plan) + SubscriptionItemFactory(subscription__organization=org, plan=non_wix_plan) # Create membership membership = Membership.objects.create(user=user, organization=org) @@ -356,17 +356,17 @@ def test_membership_create_without_plan_no_sync(mocker): mock_sync.assert_not_called() -class TestSubscriptionQuerySet(TestCase): - """Unit tests for Subscription queryset""" +class TestSubscriptionItemQuerySet(TestCase): + """Unit tests for SubscriptionItem queryset""" @pytest.mark.django_db def test_sunlight_active_count_zero(self): """Test count returns zero when no Sunlight subscriptions exist""" # Create some non-Sunlight subscriptions regular_plan = PlanFactory(slug="professional", wix=False) - SubscriptionFactory(plan=regular_plan, cancelled=False) + SubscriptionItemFactory(plan=regular_plan, subscription__cancelled=False) - count = Subscription.objects.sunlight_active_count() + count = SubscriptionItem.objects.sunlight_active_count() assert count == 0 @pytest.mark.django_db @@ -377,11 +377,11 @@ def test_sunlight_active_count_basic(self): sunlight_plan2 = PlanFactory(slug="sunlight-enhanced-annual", wix=True) # Create active subscriptions - SubscriptionFactory(plan=sunlight_plan1, cancelled=False) - SubscriptionFactory(plan=sunlight_plan1, cancelled=False) - SubscriptionFactory(plan=sunlight_plan2, cancelled=False) + SubscriptionItemFactory(plan=sunlight_plan1, subscription__cancelled=False) + SubscriptionItemFactory(plan=sunlight_plan1, subscription__cancelled=False) + SubscriptionItemFactory(plan=sunlight_plan2, subscription__cancelled=False) - count = Subscription.objects.sunlight_active_count() + count = SubscriptionItem.objects.sunlight_active_count() assert count == 3 @pytest.mark.django_db @@ -390,12 +390,12 @@ def test_sunlight_active_count_includes_cancelled(self): sunlight_plan = PlanFactory(slug="sunlight-essential-monthly", wix=True) # Create active and pending-cancellation subscriptions - SubscriptionFactory(plan=sunlight_plan, cancelled=False) - SubscriptionFactory(plan=sunlight_plan, cancelled=False) - SubscriptionFactory(plan=sunlight_plan, cancelled=True) - SubscriptionFactory(plan=sunlight_plan, cancelled=True) + SubscriptionItemFactory(plan=sunlight_plan, subscription__cancelled=False) + SubscriptionItemFactory(plan=sunlight_plan, subscription__cancelled=False) + SubscriptionItemFactory(plan=sunlight_plan, subscription__cancelled=True) + SubscriptionItemFactory(plan=sunlight_plan, subscription__cancelled=True) - count = Subscription.objects.sunlight_active_count() + count = SubscriptionItem.objects.sunlight_active_count() assert count == 4 @pytest.mark.django_db @@ -404,10 +404,10 @@ def test_sunlight_active_count_excludes_non_wix(self): sunlight_wix = PlanFactory(slug="sunlight-essential-monthly", wix=True) sunlight_no_wix = PlanFactory(slug="sunlight-enhanced-annual", wix=False) - SubscriptionFactory(plan=sunlight_wix, cancelled=False) - SubscriptionFactory(plan=sunlight_no_wix, cancelled=False) + SubscriptionItemFactory(plan=sunlight_wix, subscription__cancelled=False) + SubscriptionItemFactory(plan=sunlight_no_wix, subscription__cancelled=False) - count = Subscription.objects.sunlight_active_count() + count = SubscriptionItem.objects.sunlight_active_count() assert count == 1 @pytest.mark.django_db @@ -417,12 +417,12 @@ def test_sunlight_active_count_mixed_subscriptions(self): regular_plan = PlanFactory(slug="professional", wix=False) # Create mix of subscriptions (cancelled Sunlight still counts) - SubscriptionFactory(plan=sunlight_plan, cancelled=False) - SubscriptionFactory(plan=sunlight_plan, cancelled=False) - SubscriptionFactory(plan=sunlight_plan, cancelled=True) - SubscriptionFactory(plan=regular_plan, cancelled=False) + SubscriptionItemFactory(plan=sunlight_plan, subscription__cancelled=False) + SubscriptionItemFactory(plan=sunlight_plan, subscription__cancelled=False) + SubscriptionItemFactory(plan=sunlight_plan, subscription__cancelled=True) + SubscriptionItemFactory(plan=regular_plan, subscription__cancelled=False) - count = Subscription.objects.sunlight_active_count() + count = SubscriptionItem.objects.sunlight_active_count() assert count == 3 @@ -452,7 +452,7 @@ def test_get_viewable_authenticated(self): public_plan = PlanFactory(public=True) org_plan = PlanFactory(public=False) # Create subscription to associate plan with org - SubscriptionFactory(organization=org, plan=org_plan) + SubscriptionItemFactory(subscription__organization=org, plan=org_plan) unrelated_private_plan = PlanFactory(public=False) viewable = Plan.objects.get_viewable(user) @@ -636,7 +636,7 @@ def test_get_subscribed_authenticated(self): ) entitlement.plans.add(plan) # Create subscription to associate plan with org - SubscriptionFactory(organization=org, plan=plan) + SubscriptionItemFactory(subscription__organization=org, plan=plan) unrelated_entitlement = Entitlement.objects.create( name="Unrelated Feature", slug="unrelated-feature-auth", client=client diff --git a/squarelet/organizations/tests/test_serializers.py b/squarelet/organizations/tests/test_serializers.py index 70ccffaa7..e638dbb27 100644 --- a/squarelet/organizations/tests/test_serializers.py +++ b/squarelet/organizations/tests/test_serializers.py @@ -16,7 +16,7 @@ EntitlementGrantFactory, OrganizationFactory, PlanFactory, - SubscriptionFactory, + SubscriptionItemFactory, ) @@ -85,7 +85,7 @@ def test_serializer_update_on_uses_org_anchor(self): entitlement.plans.set([plan]) org_update_on = date.today() + timedelta(days=7) org = OrganizationFactory(update_on=org_update_on) - SubscriptionFactory(organization=org, plan=plan) + SubscriptionItemFactory(subscription__organization=org, plan=plan) serializer = OrganizationDetailSerializer(org, context={"client": client}) rows = [ @@ -204,7 +204,7 @@ def test_quantity_is_emitted_for_a_pack(self): client = ClientFactory() org = OrganizationFactory() plan, entitlement = self._pack(client) - SubscriptionFactory(organization=org, plan=plan, quantity=3) + SubscriptionItemFactory(subscription__organization=org, plan=plan, quantity=3) rows = OrganizationDetailSerializer( org, context={"client": client} @@ -220,7 +220,7 @@ def test_pack_resources_are_emitted_untouched(self): client = ClientFactory() org = OrganizationFactory() plan, entitlement = self._pack(client) - SubscriptionFactory(organization=org, plan=plan, quantity=5) + SubscriptionItemFactory(subscription__organization=org, plan=plan, quantity=5) rows = OrganizationDetailSerializer( org, context={"client": client} @@ -255,8 +255,12 @@ def test_a_pack_alongside_a_base_plan_is_a_separate_row(self): base_plan.entitlements.add(base_entitlement) pack_plan, pack_entitlement = self._pack(client) - SubscriptionFactory(organization=org, plan=base_plan, quantity=5) - SubscriptionFactory(organization=org, plan=pack_plan, quantity=2) + SubscriptionItemFactory( + subscription__organization=org, plan=base_plan, quantity=5 + ) + SubscriptionItemFactory( + subscription__organization=org, plan=pack_plan, quantity=2 + ) rows = OrganizationDetailSerializer( org, context={"client": client} diff --git a/squarelet/organizations/tests/test_tasks.py b/squarelet/organizations/tests/test_tasks.py index 50c49cbd0..df1f5db0b 100644 --- a/squarelet/organizations/tests/test_tasks.py +++ b/squarelet/organizations/tests/test_tasks.py @@ -20,12 +20,17 @@ # Squarelet from squarelet.organizations import tasks -from squarelet.organizations.models import Charge, Invoice, PaymentMethod, Subscription +from squarelet.organizations.models import ( + Charge, + Invoice, + PaymentMethod, + SubscriptionItem, +) from squarelet.organizations.tests.factories import ( EntitlementGrantFactory, InvoiceFactory, OrganizationFactory, - SubscriptionFactory, + SubscriptionItemFactory, ) # pylint:disable=too-many-lines @@ -41,24 +46,24 @@ def test_restore_organization(organization_plan_factory, mocker): organization_plan = organization_plan_factory() # Org whose anchor is due today and whose only sub is cancelled -> anchor clears - subsc_update_cancel = SubscriptionFactory( + subsc_update_cancel = SubscriptionItemFactory( plan=organization_plan, - cancelled=True, - organization__update_on=today - timedelta(1), + subscription__cancelled=True, + subscription__organization__update_on=today - timedelta(1), ) org_cancel = subsc_update_cancel.organization # Org whose anchor is not due yet -> untouched - subsc_update_later = SubscriptionFactory( + subsc_update_later = SubscriptionItemFactory( plan=organization_plan, - organization__update_on=today + timedelta(1), + subscription__organization__update_on=today + timedelta(1), ) org_later = subsc_update_later.organization # Org whose anchor is due today and still has an active sub -> anchor advances - subsc_update = SubscriptionFactory( + subsc_update = SubscriptionItemFactory( plan=organization_plan, - organization__update_on=today - timedelta(1), + subscription__organization__update_on=today - timedelta(1), ) org_due = subsc_update.organization @@ -69,7 +74,7 @@ def test_restore_organization(organization_plan_factory, mocker): org_due.refresh_from_db() # cancelled sub in due org should have been deleted - assert not Subscription.objects.filter(pk=subsc_update_cancel.pk).exists() + assert not SubscriptionItem.objects.filter(pk=subsc_update_cancel.pk).exists() # org with no remaining subs loses its anchor assert org_cancel.update_on is None @@ -99,18 +104,18 @@ def test_restore_organization_annual_sub_not_deleted_early( # Annual sub cancelled but cancel_at is ~12 months in the future future_cancel_at = today + timedelta(days=365) - sub = SubscriptionFactory( + sub = SubscriptionItemFactory( plan=organization_plan, - cancelled=True, - cancel_at=future_cancel_at, - organization__update_on=today - timedelta(1), + subscription__cancelled=True, + subscription__cancel_at=future_cancel_at, + subscription__organization__update_on=today - timedelta(1), ) org = sub.organization tasks.restore_organization() # Sub should still exist — cancel_at has not passed - assert Subscription.objects.filter(pk=sub.pk).exists() + assert SubscriptionItem.objects.filter(pk=sub.pk).exists() # update_on should advance (org still has a sub) org.refresh_from_db() assert org.update_on == (today - timedelta(1)) + relativedelta(months=1) @@ -128,18 +133,18 @@ def test_restore_organization_annual_sub_deleted_after_cancel_at( # Annual sub cancelled and cancel_at is in the past past_cancel_at = today - timedelta(days=1) - sub = SubscriptionFactory( + sub = SubscriptionItemFactory( plan=organization_plan, - cancelled=True, - cancel_at=past_cancel_at, - organization__update_on=today - timedelta(1), + subscription__cancelled=True, + subscription__cancel_at=past_cancel_at, + subscription__organization__update_on=today - timedelta(1), ) org = sub.organization tasks.restore_organization() # Sub should be deleted — cancel_at has passed - assert not Subscription.objects.filter(pk=sub.pk).exists() + assert not SubscriptionItem.objects.filter(pk=sub.pk).exists() # org lost its last sub, so update_on clears org.refresh_from_db() assert org.update_on is None @@ -206,7 +211,7 @@ def test_org_with_subscription_not_double_counted(self, mocker): verified_journalist=True, update_on=date(2026, 7, 1), ) - SubscriptionFactory(organization=org) + SubscriptionItemFactory(subscription__organization=org) EntitlementGrantFactory(require_verified=True) tasks.restore_organization() @@ -730,7 +735,7 @@ def test_handle_invoice_failed(organization_factory, user_factory, mailoutbox): def test_handle_invoice_failed_attempt_4_subscription_cancelled( organization_factory, user_factory, - subscription_factory, + subscription_item_factory, plan_factory, mailoutbox, mocker, @@ -744,8 +749,10 @@ def test_handle_invoice_failed_attempt_4_subscription_cancelled( ) plan = plan_factory() stripe_sub_id = "sub_cancel123" - subscription_factory( - organization=organization, plan=plan, subscription_id=stripe_sub_id + subscription_item_factory( + subscription__organization=organization, + plan=plan, + subscription__subscription_id=stripe_sub_id, ) # Mock subscription_cancelled at the class level @@ -775,7 +782,7 @@ def test_handle_invoice_failed_attempt_4_subscription_cancelled( @pytest.mark.django_db() def test_handle_invoice_failed_subscription_cancel_error( - organization_factory, user_factory, subscription_factory, plan_factory, mocker + organization_factory, user_factory, subscription_item_factory, plan_factory, mocker ): """Test that error is raised if subscription_cancelled raises error""" user = user_factory() @@ -786,8 +793,10 @@ def test_handle_invoice_failed_subscription_cancel_error( ) plan = plan_factory() stripe_sub_id = "sub_cancel_err123" - subscription_factory( - organization=organization, plan=plan, subscription_id=stripe_sub_id + subscription_item_factory( + subscription__organization=organization, + plan=plan, + subscription__subscription_id=stripe_sub_id, ) # Mock subscription_cancelled to raise an error at class level @@ -877,12 +886,13 @@ class TestHandleInvoiceCreated: """Unit tests for the handle_invoice_created task""" @pytest.mark.django_db(transaction=True) - def test_creates_invoice(self, organization_factory, subscription_factory): + def test_creates_invoice(self, organization_factory, subscription_item_factory): timestamp = timezone.now().replace(microsecond=0) due_timestamp = (timestamp + timedelta(days=30)).replace(microsecond=0) organization = organization_factory(customer__customer_id="cus_123") - subscription = subscription_factory( - organization=organization, subscription_id="sub_123" + subscription = subscription_item_factory( + subscription__organization=organization, + subscription__subscription_id="sub_123", ) invoice_data = { @@ -902,7 +912,7 @@ def test_creates_invoice(self, organization_factory, subscription_factory): invoice = Invoice.objects.get(invoice_id="in_123") assert invoice.organization == organization - assert invoice.subscription == subscription + assert invoice.subscription == subscription.subscription assert invoice.amount == 10000 assert invoice.status == "draft" @@ -973,16 +983,19 @@ def test_skips_invoice_with_non_subscription_parent(self, organization_factory): assert not Invoice.objects.filter(invoice_id="in_quote").exists() @pytest.mark.django_db(transaction=True) - def test_updates_existing_invoice(self, organization_factory, subscription_factory): + def test_updates_existing_invoice( + self, organization_factory, subscription_item_factory + ): timestamp = timezone.now().replace(microsecond=0) organization = organization_factory(customer__customer_id="cus_123") - subscription = subscription_factory( - organization=organization, subscription_id="sub_123" + subscription = subscription_item_factory( + subscription__organization=organization, + subscription__subscription_id="sub_123", ) existing_invoice = InvoiceFactory( invoice_id="in_123", organization=organization, - subscription=subscription, + subscription=subscription.subscription, amount=5000, ) @@ -1367,14 +1380,14 @@ def test_sets_payment_failed_for_overdue_within_grace_period( @pytest.mark.django_db(transaction=True) @override_settings(OVERDUE_INVOICE_GRACE_PERIOD_DAYS=30) def test_cancels_at_grace_period_threshold( - self, invoice_factory, organization_factory, subscription_factory, mocker + self, invoice_factory, organization_factory, subscription_item_factory, mocker ): """Invoice exactly at grace period threshold should cancel subscription""" org = organization_factory() - subscription = subscription_factory(organization=org) + subscription = subscription_item_factory(subscription__organization=org) invoice = invoice_factory( organization=org, - subscription=subscription, + subscription=subscription.subscription, status="open", due_date=date.today() - timedelta(days=30), ) @@ -1400,15 +1413,15 @@ def test_cancels_at_grace_period_threshold( @pytest.mark.django_db(transaction=True) @override_settings(OVERDUE_INVOICE_GRACE_PERIOD_DAYS=30) def test_does_not_resend_email_if_payment_failed_already_set( - self, invoice_factory, organization_factory, subscription_factory, mocker + self, invoice_factory, organization_factory, subscription_item_factory, mocker ): """Should not resend overdue email if payment_failed flag is already set (still cancels at threshold)""" org = organization_factory(payment_failed=True) - subscription = subscription_factory(organization=org) + subscription = subscription_item_factory(subscription__organization=org) invoice = invoice_factory( organization=org, - subscription=subscription, + subscription=subscription.subscription, status="open", due_date=date.today() - timedelta(days=30), ) @@ -1426,14 +1439,14 @@ def test_does_not_resend_email_if_payment_failed_already_set( @pytest.mark.django_db(transaction=True) @override_settings(OVERDUE_INVOICE_GRACE_PERIOD_DAYS=30) def test_cancels_subscription_past_grace_period( - self, invoice_factory, organization_factory, subscription_factory, mocker + self, invoice_factory, organization_factory, subscription_item_factory, mocker ): """Invoice past grace period should cancel subscription""" org = organization_factory() - subscription = subscription_factory(organization=org) + subscription = subscription_item_factory(subscription__organization=org) invoice = invoice_factory( organization=org, - subscription=subscription, + subscription=subscription.subscription, status="open", due_date=date.today() - timedelta(days=35), ) @@ -1465,14 +1478,14 @@ def test_cancels_subscription_past_grace_period( @pytest.mark.django_db(transaction=True) @override_settings(OVERDUE_INVOICE_GRACE_PERIOD_DAYS=30) def test_handles_stripe_error_when_marking_uncollectible( - self, invoice_factory, organization_factory, subscription_factory, mocker + self, invoice_factory, organization_factory, subscription_item_factory, mocker ): """Should handle Stripe errors gracefully when marking uncollectible""" org = organization_factory() - subscription = subscription_factory(organization=org) + subscription = subscription_item_factory(subscription__organization=org) invoice = invoice_factory( organization=org, - subscription=subscription, + subscription=subscription.subscription, status="open", due_date=date.today() - timedelta(days=35), ) @@ -1527,11 +1540,11 @@ def test_handles_organization_without_subscription( @pytest.mark.django_db(transaction=True) @override_settings(OVERDUE_INVOICE_GRACE_PERIOD_DAYS=30) def test_does_not_cancel_org_subscription_when_invoice_has_none( - self, invoice_factory, organization_factory, subscription_factory, mocker + self, invoice_factory, organization_factory, subscription_item_factory, mocker ): """Invoice with no subscription should not cancel the org's subscription""" org = organization_factory() - org_subscription = subscription_factory(organization=org) + org_subscription = subscription_item_factory(subscription__organization=org) invoice = invoice_factory( organization=org, subscription=None, @@ -1554,20 +1567,22 @@ def test_does_not_cancel_org_subscription_when_invoice_has_none( mock_subscription_cancelled.assert_not_called() # Org's subscription should still exist - assert Subscription.objects.filter(pk=org_subscription.pk).exists() + assert SubscriptionItem.objects.filter(pk=org_subscription.pk).exists() @pytest.mark.django_db(transaction=True) @override_settings(OVERDUE_INVOICE_GRACE_PERIOD_DAYS=30) def test_cancels_invoice_subscription_not_org_subscription( - self, invoice_factory, organization_factory, subscription_factory, mocker + self, invoice_factory, organization_factory, subscription_item_factory, mocker ): """Should cancel the invoice's subscription, not the org's current one""" org = organization_factory() - org_subscription = subscription_factory(organization=org) - invoice_subscription = subscription_factory(organization=org) + org_subscription = subscription_item_factory(subscription__organization=org) + invoice_subscription = subscription_item_factory( + subscription__organization=org, subscription__interval="annual" + ) invoice = invoice_factory( organization=org, - subscription=invoice_subscription, + subscription=invoice_subscription.subscription, status="open", due_date=date.today() - timedelta(days=35), ) @@ -1587,10 +1602,10 @@ def test_cancels_invoice_subscription_not_org_subscription( tasks.process_overdue_invoice(invoice.id) # The invoice's subscription should be deleted - assert not Subscription.objects.filter(pk=invoice_subscription.pk).exists() + assert not SubscriptionItem.objects.filter(pk=invoice_subscription.pk).exists() # The org's other subscription should still exist - assert Subscription.objects.filter(pk=org_subscription.pk).exists() + assert SubscriptionItem.objects.filter(pk=org_subscription.pk).exists() @pytest.mark.django_db @override_settings(OVERDUE_INVOICE_GRACE_PERIOD_DAYS=45) @@ -2238,23 +2253,28 @@ class TestSubscriptionCancelledWithExplicitSubscription: @pytest.mark.django_db def test_with_explicit_sub_cancels_only_that_sub( - self, organization_factory, plan_factory, subscription_factory + self, organization_factory, plan_factory, subscription_item_factory ): """Org with two subscriptions: cancelling one leaves the other.""" org = organization_factory() plan_a = plan_factory() plan_b = plan_factory() - sub_a = subscription_factory(organization=org, plan=plan_a) - sub_b = subscription_factory(organization=org, plan=plan_b) + # Different billing shapes, so these are two Stripe subscriptions + sub_a = subscription_item_factory(subscription__organization=org, plan=plan_a) + sub_b = subscription_item_factory( + subscription__organization=org, + plan=plan_b, + subscription__interval="annual", + ) - org.subscription_cancelled(subscription=sub_a) + org.subscription_cancelled(subscription=sub_a.subscription) - assert not Subscription.objects.filter(pk=sub_a.pk).exists() - assert Subscription.objects.filter(pk=sub_b.pk).exists() + assert not SubscriptionItem.objects.filter(pk=sub_a.pk).exists() + assert SubscriptionItem.objects.filter(pk=sub_b.pk).exists() @pytest.mark.django_db def test_handle_invoice_failed_passes_subscription( - self, organization_factory, plan_factory, subscription_factory, mocker + self, organization_factory, plan_factory, subscription_item_factory, mocker ): """handle_invoice_failed passes the matching subscription to subscription_cancelled on 4th failure.""" @@ -2262,7 +2282,11 @@ def test_handle_invoice_failed_passes_subscription( org = organization_factory(customer__customer_id=customer_id) plan = plan_factory() stripe_sub_id = "sub_test123" - subscription_factory(organization=org, plan=plan, subscription_id=stripe_sub_id) + subscription_item_factory( + subscription__organization=org, + plan=plan, + subscription__subscription_id=stripe_sub_id, + ) mocked_cancel = mocker.patch( "squarelet.organizations.models.organization." @@ -2670,7 +2694,7 @@ def test_canceled_status_reconciles_subscription( {"id": "sub_upd_cancel", "status": "canceled"} ) - assert not Subscription.objects.filter(pk=subscription.pk).exists() + assert not SubscriptionItem.objects.filter(pk=subscription.pk).exists() patched.assert_called_once_with("organization", [org_uuid]) @pytest.mark.django_db @@ -2695,7 +2719,7 @@ def test_deletes_subscription_and_invalidates_cache( tasks.handle_subscription_deleted({"id": "sub_del", "status": "canceled"}) - assert not Subscription.objects.filter(pk=subscription.pk).exists() + assert not SubscriptionItem.objects.filter(pk=subscription.pk).exists() patched.assert_called_once_with("organization", [org_uuid]) @pytest.mark.django_db diff --git a/squarelet/organizations/tests/test_templates.py b/squarelet/organizations/tests/test_templates.py index 86a134a78..7145bdc7a 100644 --- a/squarelet/organizations/tests/test_templates.py +++ b/squarelet/organizations/tests/test_templates.py @@ -93,7 +93,7 @@ def test_own_benefits_listed_with_own_plans( org = organization_factory(plans=[plan]) html = self._render( - subscriptions=list(org.subscriptions.all()), + subscriptions=list(org.subscription_items.all()), subscription_benefits=["100 requests each month"], ) @@ -110,7 +110,7 @@ def test_own_and_inherited_benefits_render_in_separate_cards( parent = organization_factory(name="Parent Org") html = self._render( - subscriptions=list(org.subscriptions.all()), + subscriptions=list(org.subscription_items.all()), subscription_benefits=["Own benefit"], inherited_orgs=[parent], inherited_benefits=["Inherited benefit"], @@ -142,7 +142,7 @@ def test_active_subscription_shows_price_and_renewal_date( """An active subscription shows its price and next renewal date""" plan = plan_factory(name="Org Plan", annual=False, base_price=30) org = organization_factory(plans=[plan]) - subscription = org.subscriptions.get() + subscription = org.subscription_items.get() subscription.current_period_end = datetime.datetime( 2026, 2, 15, 12, 0, tzinfo=datetime.timezone.utc ) @@ -161,7 +161,7 @@ def test_annual_subscription_shows_per_year_price( plan = plan_factory(name="Org Plan", annual=True, base_price=300) org = organization_factory(plans=[plan]) - html = self._render(subscriptions=list(org.subscriptions.all())) + html = self._render(subscriptions=list(org.subscription_items.all())) assert "$300 per year" in html @@ -172,7 +172,7 @@ def test_free_subscription_shows_free_instead_of_price( plan = plan_factory(name="Org Plan", base_price=0, price_per_user=0) org = organization_factory(plans=[plan]) - html = self._render(subscriptions=list(org.subscriptions.all())) + html = self._render(subscriptions=list(org.subscription_items.all())) assert "Free" in html @@ -183,7 +183,7 @@ def test_cancelled_subscription_shows_price_and_expiration_date( instead of a renewal date""" plan = plan_factory(name="Org Plan", annual=False, base_price=30) org = organization_factory(plans=[plan]) - subscription = org.subscriptions.get() + subscription = org.subscription_items.get() subscription.cancelled = True subscription.cancel_at = datetime.date(2026, 1, 15) subscription.current_period_end = datetime.datetime( @@ -204,7 +204,7 @@ def test_no_card_on_file_shows_fallback_text( plan = plan_factory(name="Org Plan") org = organization_factory(plans=[plan]) - html = self._render(subscriptions=list(org.subscriptions.all())) + html = self._render(subscriptions=list(org.subscription_items.all())) assert "No card on file" in html @@ -216,7 +216,7 @@ def test_card_on_file_shows_brand_and_last4( org = organization_factory(plans=[plan]) html = self._render( - subscriptions=list(org.subscriptions.all()), + subscriptions=list(org.subscription_items.all()), card_brand="visa", card_last4="4242", ) @@ -230,7 +230,7 @@ def test_plan_name_links_to_plan_detail(self, organization_factory, plan_factory plan = plan_factory(name="Org Plan") org = organization_factory(plans=[plan]) - html = self._render(subscriptions=list(org.subscriptions.all())) + html = self._render(subscriptions=list(org.subscription_items.all())) expected_url = reverse("plan_detail", kwargs={"pk": plan.pk, "slug": plan.slug}) assert f'href="{expected_url}"' in html diff --git a/squarelet/organizations/tests/test_views.py b/squarelet/organizations/tests/test_views.py index 2c95aaea6..9677aa829 100644 --- a/squarelet/organizations/tests/test_views.py +++ b/squarelet/organizations/tests/test_views.py @@ -1428,7 +1428,7 @@ def test_blocked_by_active_subscription( organization_factory, user_factory, plan_factory, - subscription_factory, + subscription_item_factory, mocker, ): """A non-cancelled subscription blocks removal""" @@ -1438,8 +1438,10 @@ def test_blocked_by_active_subscription( ) user = user_factory() organization = organization_factory(admins=[user]) - subscription_factory( - organization=organization, plan=plan_factory(), cancelled=False + subscription_item_factory( + subscription__organization=organization, + plan=plan_factory(), + subscription__cancelled=False, ) response = self.call_view(rf, user, {}, slug=organization.slug) assert response.status_code == 302 @@ -1456,7 +1458,7 @@ def test_allowed_when_all_subscriptions_cancelled( organization_factory, user_factory, plan_factory, - subscription_factory, + subscription_item_factory, mocker, ): """Removal is allowed when every subscription is cancelled""" @@ -1466,8 +1468,10 @@ def test_allowed_when_all_subscriptions_cancelled( ) user = user_factory() organization = organization_factory(admins=[user]) - subscription_factory( - organization=organization, plan=plan_factory(), cancelled=True + subscription_item_factory( + subscription__organization=organization, + plan=plan_factory(), + subscription__cancelled=True, ) response = self.call_view(rf, user, {}, slug=organization.slug) assert response.status_code == 302 @@ -1506,7 +1510,7 @@ def test_ajax_blocked_by_active_subscription( organization_factory, user_factory, plan_factory, - subscription_factory, + subscription_item_factory, mocker, ): """AJAX removal returns a 400 with the error when blocked""" @@ -1516,8 +1520,10 @@ def test_ajax_blocked_by_active_subscription( ) user = user_factory() organization = organization_factory(admins=[user]) - subscription_factory( - organization=organization, plan=plan_factory(), cancelled=False + subscription_item_factory( + subscription__organization=organization, + plan=plan_factory(), + subscription__cancelled=False, ) response = self._ajax_post(rf, user, organization.slug) assert response.status_code == 400 @@ -1618,7 +1624,7 @@ def test_non_staff_update_card_no_action( @pytest.mark.django_db() class TestCancelSubscription(ViewTestMixin): - """Test staff-action logging on the Organization Cancel Subscription view""" + """Test staff-action logging on the Organization Cancel SubscriptionItem view""" view = views.CancelSubscription url = "/organizations/{slug}/subscriptions/{pk}/cancel" @@ -1629,7 +1635,7 @@ def test_staff_cancel_subscription_creates_action( organization_factory, user_factory, plan_factory, - subscription_factory, + subscription_item_factory, mocker, ): """Staff cancelling a subscription on someone's behalf is logged""" @@ -1638,7 +1644,9 @@ def test_staff_cancel_subscription_creates_action( staff_member = _assign_org_perm(staff_member, "can_edit_subscription") organization = organization_factory() plan = plan_factory(name="Professional") - subscription = subscription_factory(organization=organization, plan=plan) + subscription = subscription_item_factory( + subscription__organization=organization, plan=plan + ) response = self.call_view( rf, staff_member, {}, slug=organization.slug, pk=subscription.pk @@ -1661,15 +1669,15 @@ def test_non_staff_cancel_subscription_no_action( organization_factory, user_factory, plan_factory, - subscription_factory, + subscription_item_factory, mocker, ): """A regular admin cancelling their own subscription is not logged""" mocker.patch("squarelet.organizations.models.Organization.remove_subscription") admin = user_factory(is_staff=False) organization = organization_factory(admins=[admin]) - subscription = subscription_factory( - organization=organization, plan=plan_factory() + subscription = subscription_item_factory( + subscription__organization=organization, plan=plan_factory() ) response = self.call_view( @@ -1752,7 +1760,7 @@ def test_post_creates_org(self, rf, user_factory): user = user_factory(email_verified=True) self.call_view(rf, user, {"name": "test"}) organization = user.organizations.get(individual=False) - assert not organization.subscriptions.exists() + assert not organization.subscription_items.exists() assert organization.has_admin(user) assert organization.receipt_email.email == user.email diff --git a/squarelet/organizations/views/__init__.py b/squarelet/organizations/views/__init__.py index 2df57b47b..a50c89d38 100644 --- a/squarelet/organizations/views/__init__.py +++ b/squarelet/organizations/views/__init__.py @@ -31,7 +31,7 @@ "Detail", "List", "autocomplete", - # Subscription views + # SubscriptionItem views "ManageSubscriptions", "Resubscribe", "UpdateCard", diff --git a/squarelet/organizations/views/create.py b/squarelet/organizations/views/create.py index 627fedff2..cb75ecfbf 100644 --- a/squarelet/organizations/views/create.py +++ b/squarelet/organizations/views/create.py @@ -23,7 +23,7 @@ def form_valid(self, form): organization.change_logs.create( reason=ChangeLogReason.created, user=self.request.user, - to_plan=organization.plans.first(), + to_plan=organization.get_plans().first(), to_max_users=organization.max_users, ) return redirect(organization) diff --git a/squarelet/organizations/views/detail.py b/squarelet/organizations/views/detail.py index edadb1199..b5a30eb57 100644 --- a/squarelet/organizations/views/detail.py +++ b/squarelet/organizations/views/detail.py @@ -62,7 +62,7 @@ def get_context_data(self, **kwargs): # Get subscriptions, if any, along with the benefits they add up to upgrade_plan = Plan.objects.get(slug="organization") subscriptions = list( - org.subscriptions.select_related("plan").prefetch_related( + org.subscription_items.select_related("plan").prefetch_related( "plan__entitlements" ) ) @@ -94,7 +94,7 @@ def get_context_data(self, **kwargs): } context["show_wix_sync"] = bool( - org.subscriptions.filter(plan__wix=True).exists() + org.subscription_items.filter(plan__wix=True).exists() or org.get_wix_plans_from_groups() ) inherited_orgs, inherited_benefits = consolidate_inherited_benefits( @@ -320,7 +320,7 @@ def handle_sync_wix(self, request): triggered = False # Direct Wix plans on this org - for sub in org.subscriptions.filter(plan__wix=True).select_related("plan"): + for sub in org.subscription_items.filter(plan__wix=True).select_related("plan"): self._sync_wix_for_org(org, sub.plan) triggered = True diff --git a/squarelet/organizations/wix.py b/squarelet/organizations/wix.py index 6aa7cd8f2..8ed862af9 100644 --- a/squarelet/organizations/wix.py +++ b/squarelet/organizations/wix.py @@ -213,9 +213,11 @@ def unsync_wix(organization, plan, user): def get_wix_labels_for_user(user): """Get all Wix labels a user qualifies for across all their memberships.""" labels = set() - for membership in user.memberships.prefetch_related("organization__plans").all(): + for membership in user.memberships.prefetch_related( + "organization__subscriptions__plans" + ).all(): org = membership.organization - for plan in org.plans.all(): + for plan in org.get_plans(): if plan and plan.wix: tier = get_tier_from_plan(plan) labels.add(f"custom.{tier}-member") diff --git a/squarelet/payments/forms.py b/squarelet/payments/forms.py index d600d8c73..b5342af1e 100644 --- a/squarelet/payments/forms.py +++ b/squarelet/payments/forms.py @@ -165,7 +165,7 @@ def _configure_organization_field(self): # Exclude organizations already subscribed to this plan if self.plan: subscribed_orgs = Organization.objects.filter( - subscriptions__plan=self.plan, + subscriptions__items__plan=self.plan, subscriptions__cancelled=False, ) base_queryset = base_queryset.exclude(pk__in=subscribed_orgs) diff --git a/squarelet/payments/tests/test_forms.py b/squarelet/payments/tests/test_forms.py index 317c57298..adc2fe739 100644 --- a/squarelet/payments/tests/test_forms.py +++ b/squarelet/payments/tests/test_forms.py @@ -46,7 +46,11 @@ def test_init_with_user_shows_admin_orgs( assert org in form.fields["organization"].queryset def test_init_excludes_already_subscribed_orgs( - self, user_factory, organization_factory, plan_factory, subscription_factory + self, + user_factory, + organization_factory, + plan_factory, + subscription_item_factory, ): """Form excludes organizations already subscribed to the plan""" user = user_factory() @@ -55,7 +59,9 @@ def test_init_excludes_already_subscribed_orgs( plan = plan_factory(public=True, for_groups=True) # Create an active subscription for the org - subscription_factory(organization=org, plan=plan, cancelled=False) + subscription_item_factory( + subscription__organization=org, plan=plan, subscription__cancelled=False + ) form = PlanPurchaseForm(plan=plan, user=user) diff --git a/squarelet/payments/tests/test_views.py b/squarelet/payments/tests/test_views.py index d789b9b82..11fdeaf8a 100644 --- a/squarelet/payments/tests/test_views.py +++ b/squarelet/payments/tests/test_views.py @@ -194,7 +194,12 @@ def test_unauthenticated_user_redirected_to_login(self, rf, plan_factory): assert "/accounts/login/" in response.url def test_already_subscribed_organization( - self, rf, user_factory, organization_factory, plan_factory, subscription_factory + self, + rf, + user_factory, + organization_factory, + plan_factory, + subscription_item_factory, ): """Test that already subscribed organizations are excluded from form""" user = user_factory() @@ -203,7 +208,9 @@ def test_already_subscribed_organization( org.add_creator(user) # Create existing subscription - subscription_factory(organization=org, plan=plan, cancelled=False) + subscription_item_factory( + subscription__organization=org, plan=plan, subscription__cancelled=False + ) data = { "organization": str(org.pk), diff --git a/squarelet/payments/views.py b/squarelet/payments/views.py index e5c82a869..1d0492c85 100644 --- a/squarelet/payments/views.py +++ b/squarelet/payments/views.py @@ -32,7 +32,7 @@ from squarelet.organizations.models import Charge, Organization, Plan from squarelet.organizations.models.payment import ( ReceiptEmail, - Subscription, + SubscriptionItem, get_payment_brand, ) from squarelet.organizations.payments.base import PaymentActionRequired @@ -138,7 +138,7 @@ def get_context_data(self, **kwargs): # Check user's individual organization individual_org = user.individual_organization - individual_subscription = individual_org.subscriptions.filter( + individual_subscription = individual_org.subscription_items.filter( plan=plan ).first() if individual_subscription: @@ -158,7 +158,7 @@ def get_context_data(self, **kwargs): admin_orgs = admin_orgs_base for org in admin_orgs: - org_subscription = org.subscriptions.filter(plan=plan).first() + org_subscription = org.subscription_items.filter(plan=plan).first() if org_subscription: existing_subscriptions.append((org_subscription, org)) @@ -231,7 +231,7 @@ def post( result = form.save(request.user) organization = result["organization"] - if organization.subscriptions.filter(plan=plan).exists(): + if organization.subscription_items.filter(plan=plan).exists(): messages.warning(request, _("Already subscribed")) return redirect(plan) @@ -281,7 +281,9 @@ def _handle_sunlight_subscription(self, request, plan, result): stripe_token = result["stripe_token"] payment_method = result["payment_method"] - locked_count = Subscription.objects.select_for_update().sunlight_active_count() + locked_count = ( + SubscriptionItem.objects.select_for_update().sunlight_active_count() + ) if locked_count >= settings.MAX_SUNLIGHT_SUBSCRIPTIONS: transaction.on_commit( lambda: add_to_waitlist.delay(organization.pk, plan.pk, request.user.pk) @@ -317,7 +319,7 @@ def _handle_regular_subscription(self, request, plan, result): ) return None except PaymentActionRequired as exc: - # Subscription saved but first invoice needs 3DS. + # SubscriptionItem saved but first invoice needs 3DS. redirect_url = organization.get_absolute_url() if self._is_ajax(): return JsonResponse( @@ -371,7 +373,7 @@ def get_context_data(self, **kwargs): if self.request.user.is_authenticated: # Check user's individual organization individual_org = self.request.user.individual_organization - individual_subscriptions = individual_org.subscriptions.filter( + individual_subscriptions = individual_org.subscription_items.filter( plan__slug__startswith="sunlight-", plan__wix=True ).select_related("plan") @@ -384,7 +386,7 @@ def get_context_data(self, **kwargs): ).distinct() for org in admin_orgs: - org_subscriptions = org.subscriptions.filter( + org_subscriptions = org.subscription_items.filter( plan__slug__startswith="sunlight-", plan__wix=True ).select_related("plan") @@ -554,7 +556,7 @@ def get_context_data(self, **kwargs): context = super().get_context_data(**kwargs) # Get subscriptions and add renewal/cancellation date and cost data - subscriptions = self.object.subscriptions.all() + subscriptions = self.object.subscription_items.all() for subscription in subscriptions: subscription.next_date = get_subscription_next_date(subscription) subscription.cost = subscription.plan.cost(self.object.max_users) @@ -690,7 +692,9 @@ class BaseCancelSubscription(SubscriptionObjectMixin, UpdateView): def get_context_data(self, **kwargs): context = super().get_context_data(**kwargs) - subscription = self.object.subscriptions.filter(id=self.kwargs["pk"]).first() + subscription = self.object.subscription_items.filter( + id=self.kwargs["pk"] + ).first() if subscription: context["subscription"] = subscription context["next_date"] = get_subscription_next_date(subscription) @@ -698,7 +702,9 @@ def get_context_data(self, **kwargs): def form_valid(self, form): organization = self.object - subscription = self.object.subscriptions.filter(id=self.kwargs["pk"]).first() + subscription = self.object.subscription_items.filter( + id=self.kwargs["pk"] + ).first() if subscription: organization.remove_subscription(subscription) self.log_staff_action( @@ -772,7 +778,9 @@ def _error(self, message, redirect_url="subscriptions"): def post(self, request, *args, **kwargs): organization = self.get_object() redirect_url = self.reverse_subject("subscriptions") - subscription = organization.subscriptions.filter(id=self.kwargs["pk"]).first() + subscription = organization.subscription_items.filter( + id=self.kwargs["pk"] + ).first() print(subscription) try: subscription.uncancel() diff --git a/squarelet/statistics/tasks.py b/squarelet/statistics/tasks.py index 8a201ff2e..934630804 100644 --- a/squarelet/statistics/tasks.py +++ b/squarelet/statistics/tasks.py @@ -30,14 +30,14 @@ def store_statistics(): is_agency=True ).count() kwargs["total_users_pro"] = User.objects.filter( - organizations__plans__slug="professional" + organizations__subscriptions__plans__slug="professional" ).count() kwargs["total_users_org"] = User.objects.filter( - organizations__plans__slug="organization" + organizations__subscriptions__plans__slug="organization" ).count() kwargs["total_users_mfa"] = Authenticator.objects.distinct("user").count() kwargs["total_orgs"] = Organization.objects.exclude( - individual=True, plans=None + individual=True, subscriptions__plans__isnull=True ).count() kwargs["verified_orgs"] = Organization.objects.filter( verified_journalist=True @@ -48,7 +48,9 @@ def store_statistics(): stats.users_today.set( User.objects.filter(last_login__range=(yesterday_midnight, today_midnight)) ) - stats.pro_users.set(User.objects.filter(organizations__plans__slug="professional")) + stats.pro_users.set( + User.objects.filter(organizations__subscriptions__plans__slug="professional") + ) stats.save() diff --git a/squarelet/users/tests/test_adapters.py b/squarelet/users/tests/test_adapters.py index 9d217141c..4de8842b7 100644 --- a/squarelet/users/tests/test_adapters.py +++ b/squarelet/users/tests/test_adapters.py @@ -102,7 +102,7 @@ def test_post_login_direct_redirect(self): "email_check_completed": True, # Email is already verified "mfa_step": "completed", # MFA completed "join_org": True, # Organization joining completed - "subscription": "completed", # Subscription completed + "subscription": "completed", # SubscriptionItem completed } # Call post_login with a destination URL diff --git a/squarelet/users/tests/test_views.py b/squarelet/users/tests/test_views.py index 5cc9c41bf..bde9bdd57 100644 --- a/squarelet/users/tests/test_views.py +++ b/squarelet/users/tests/test_views.py @@ -159,14 +159,18 @@ def test_get_bad(self, rf, user_factory): self.call_view(rf, other_user, username=user.username) def test_own_subscription_benefits_are_consolidated( - self, rf, user_factory, plan_factory, subscription_factory + self, rf, user_factory, plan_factory, subscription_item_factory ): """Every plan the user pays for is summarized into one benefits list""" user = user_factory() plan_a = plan_factory(name="Plan A", benefits=["Shared", "A only"]) plan_b = plan_factory(name="Plan B", benefits=["Shared", "B only"]) - subscription_factory(organization=user.individual_organization, plan=plan_a) - subscription_factory(organization=user.individual_organization, plan=plan_b) + subscription_item_factory( + subscription__organization=user.individual_organization, plan=plan_a + ) + subscription_item_factory( + subscription__organization=user.individual_organization, plan=plan_b + ) response = self.call_view(rf, user, username=user.username) @@ -182,12 +186,19 @@ def test_own_subscription_benefits_are_consolidated( ] def test_own_subscription_benefits_separate_from_inherited( - self, rf, user_factory, organization_factory, plan_factory, subscription_factory + self, + rf, + user_factory, + organization_factory, + plan_factory, + subscription_item_factory, ): """Benefits the user pays for are not mixed in with their orgs'""" user = user_factory() own_plan = plan_factory(name="Own Plan", benefits=["Own benefit"]) - subscription_factory(organization=user.individual_organization, plan=own_plan) + subscription_item_factory( + subscription__organization=user.individual_organization, plan=own_plan + ) org_plan = plan_factory( name="Org Plan", benefits=["Org benefit"], base_price=100 ) @@ -598,7 +609,7 @@ class TestIndividualSubscriptionStaffActions(ViewTestMixin): url = "/users/{username}/cancel/{pk}/" def test_staff_cancel_subscription_creates_action( - self, rf, user_factory, plan_factory, subscription_factory, mocker + self, rf, user_factory, plan_factory, subscription_item_factory, mocker ): """Staff cancelling a user's subscription targets the individual org""" mocker.patch("squarelet.organizations.models.Organization.remove_subscription") @@ -606,7 +617,9 @@ def test_staff_cancel_subscription_creates_action( staff = user_factory(is_staff=True) organization = user.individual_organization plan = plan_factory(name="Professional") - subscription = subscription_factory(organization=organization, plan=plan) + subscription = subscription_item_factory( + subscription__organization=organization, plan=plan + ) response = self.call_view( rf, staff, {}, username=user.username, pk=subscription.pk @@ -623,13 +636,13 @@ def test_staff_cancel_subscription_creates_action( assert action.public is False def test_owner_cancel_subscription_no_action( - self, rf, user_factory, plan_factory, subscription_factory, mocker + self, rf, user_factory, plan_factory, subscription_item_factory, mocker ): """A user cancelling their own subscription is not logged""" mocker.patch("squarelet.organizations.models.Organization.remove_subscription") user = user_factory(username="dotted.name") - subscription = subscription_factory( - organization=user.individual_organization, plan=plan_factory() + subscription = subscription_item_factory( + subscription__organization=user.individual_organization, plan=plan_factory() ) response = self.call_view( @@ -640,14 +653,15 @@ def test_owner_cancel_subscription_no_action( assert not Action.objects.filter(verb="cancelled a subscription").exists() def test_staff_managing_own_account_no_action( - self, rf, user_factory, plan_factory, subscription_factory, mocker + self, rf, user_factory, plan_factory, subscription_item_factory, mocker ): """A staff member managing their own individual account is not logged — they are the owner (admin) of their own individual organization""" mocker.patch("squarelet.organizations.models.Organization.remove_subscription") staff = user_factory(is_staff=True, username="staffer") - subscription = subscription_factory( - organization=staff.individual_organization, plan=plan_factory() + subscription = subscription_item_factory( + subscription__organization=staff.individual_organization, + plan=plan_factory(), ) response = self.call_view( diff --git a/squarelet/users/views.py b/squarelet/users/views.py index 4e24dd416..133a8325c 100644 --- a/squarelet/users/views.py +++ b/squarelet/users/views.py @@ -176,7 +176,7 @@ def get_context_data(self, **kwargs): individual_org = user.individual_organization upgrade_plan = Plan.objects.filter(slug="professional").first() subscriptions = list( - individual_org.subscriptions.select_related("plan").prefetch_related( + individual_org.subscription_items.select_related("plan").prefetch_related( "plan__entitlements" ) ) @@ -219,7 +219,7 @@ def _get_premium_org_plans(user): for org in user.organizations.filter(individual=False): plans.extend( (org, sub.plan) - for sub in org.subscriptions.select_related("plan") + for sub in org.subscription_items.select_related("plan") if not sub.plan.free ) plans.extend((org, plan) for _source, plan in org.get_inherited_plans())