diff --git a/src/api/app.py b/src/api/app.py index 83763ef..3ec40d2 100644 --- a/src/api/app.py +++ b/src/api/app.py @@ -82,7 +82,9 @@ def refresh_payments(payload: dict = Body(...), db = Depends(get_session)): if provided != expected: raise HTTPException(status_code=403, detail="Invalid cron secret") - subs = db.exec(select(Subscription)).all() + subs = db.exec( + select(Subscription).where(Subscription.status == SubscriptionStatus.ACTIVE) + ).all() for s in subs: # skip subscriptions without a pricing plan diff --git a/src/api/roles/shared/account.py b/src/api/roles/shared/account.py index fce7466..a03ae4e 100644 --- a/src/api/roles/shared/account.py +++ b/src/api/roles/shared/account.py @@ -5,7 +5,7 @@ from src.database.account.models import Account, Availability, Notification from src.database.client.models import Client, FitnessGoals from src.database.coach.models import Coach, Experience, Certifications, CoachExperience, CoachCertifications -from src.database.payment.models import PricingPlan, PaymentInformation, Subscription, BillingCycle, Invoice +from src.database.payment.models import PricingPlan, PaymentInformation, Subscription, BillingCycle, Invoice, SubscriptionStatus from src.database.telemetry.models import ( HealthMetrics, ClientTelemetry, DailyProgressPicture, CompletedMealActivity, CompletedWorkout, @@ -32,7 +32,7 @@ from sqlalchemy import or_ from pydantic import BaseModel, EmailStr from typing import Optional, List -from datetime import datetime +from datetime import date, datetime router = APIRouter(prefix="/roles/shared/account", tags=["shared", "account"]) @@ -421,14 +421,8 @@ def notify_affected_accounts( """ Creates notification records for accounts affected by a user's deactivation. """ - role = "account" - if deactivated_account.client_id is not None: - role = "client" - elif deactivated_account.coach_id is not None: - role = "coach" - message = f"{deactivated_account.name} has deactivated their account." - details = "Shared plans or schedules involving this user may be affected." + details = "Subscription canceled" for affected_account in affected_accounts: if affected_account.id is None: @@ -445,6 +439,34 @@ def notify_affected_accounts( ) +def cancel_payments_for_request(db: Session, request: ClientCoachRequest): + subscriptions = db.exec( + select(Subscription) + .join(PricingPlan, Subscription.pricing_plan_id == PricingPlan.id) + .where( + Subscription.client_id == request.client_id, + PricingPlan.coach_id == request.coach_id, + Subscription.status == SubscriptionStatus.ACTIVE, + ) + ).all() + + for subscription in subscriptions: + subscription.status = SubscriptionStatus.CANCELED + subscription.canceled_at = date.today() + db.add(subscription) + + active_cycles = db.exec( + select(BillingCycle).where( + BillingCycle.subscription_id == subscription.id, + BillingCycle.active == True, + ) + ).all() + + for cycle in active_cycles: + cycle.active = False + db.add(cycle) + + def delete_client_coach_mappings(db: Session, account: Account): if account.client_id is not None: requests = db.exec( @@ -453,6 +475,8 @@ def delete_client_coach_mappings(db: Session, account: Account): ).all() for request in requests: + cancel_payments_for_request(db, request) + relationships = db.exec( select(ClientCoachRelationship) .where(ClientCoachRelationship.request_id == request.id) @@ -470,6 +494,8 @@ def delete_client_coach_mappings(db: Session, account: Account): ).all() for request in requests: + cancel_payments_for_request(db, request) + relationships = db.exec( select(ClientCoachRelationship) .where(ClientCoachRelationship.request_id == request.id) diff --git a/src/database/account/models.py b/src/database/account/models.py index 1b6c5c8..9546ee0 100644 --- a/src/database/account/models.py +++ b/src/database/account/models.py @@ -105,4 +105,4 @@ class Notification(SQLModelLU, table=True): message: str details: Optional[str] = None # if they do expandable dialogs we have it built in is_read: bool = False - created_at: date = Field(default_factory=date.today) \ No newline at end of file + created_at: date = Field(default_factory=date.today) diff --git a/tests/test_shared_account_notifications.py b/tests/test_shared_account_notifications.py index 5be1774..71c6776 100644 --- a/tests/test_shared_account_notifications.py +++ b/tests/test_shared_account_notifications.py @@ -1,15 +1,24 @@ from sqlmodel import select from datetime import datetime +import os from src.api.dependencies import create_jwt_token from src.database.account.models import Notification, Account +from src.database.client.models import Client from src.database.coach.models import Coach from src.database.coach_client_relationship.models import ( ClientCoachRequest, ClientCoachRelationship, ) +from src.database.payment.models import ( + BillingCycle, + PricingInterval, + PricingPlan, + Subscription, + SubscriptionStatus, +) -def create_client_coach_relationship(db_session): +def create_client_coach_relationship(db_session, with_subscription=False): client = db_session.exec( select(Account).where( Account.client_id.is_not(None), @@ -17,6 +26,23 @@ def create_client_coach_relationship(db_session): ) ).first() + if client is None: + client_profile = Client() + db_session.add(client_profile) + db_session.commit() + db_session.refresh(client_profile) + + client = Account( + name="Notification Test Client", + email=f"notification_client_{client_profile.id}@example.com", + hashed_password="test-hash", + client_id=client_profile.id, + is_active=True, + ) + db_session.add(client) + db_session.commit() + db_session.refresh(client) + assert client is not None coach = db_session.exec( @@ -64,6 +90,34 @@ def create_client_coach_relationship(db_session): db_session.add(relationship) db_session.commit() + if with_subscription: + pricing_plan = PricingPlan( + coach_id=coach.coach_id, + payment_interval=PricingInterval.MONTHLY, + price_cents=3000, + ) + db_session.add(pricing_plan) + db_session.commit() + db_session.refresh(pricing_plan) + + subscription = Subscription( + client_id=client.client_id, + pricing_plan_id=pricing_plan.id, + ) + db_session.add(subscription) + db_session.commit() + db_session.refresh(subscription) + + billing_cycle = BillingCycle( + active=True, + entry_date=datetime.utcnow().date(), + end_date=datetime.utcnow().date(), + subscription_id=subscription.id, + pricing_plan_id=pricing_plan.id, + ) + db_session.add(billing_cycle) + db_session.commit() + return client, coach, request, relationship @@ -71,7 +125,6 @@ def test_account_deactivate_sends_notification( test_client, db_session, client_auth_header, - coach_auth_header, ): client, coach, request, relationship = create_client_coach_relationship(db_session) @@ -116,11 +169,99 @@ def test_account_deactivate_sends_notification( assert db_session.get(ClientCoachRequest, request.id) is None +def test_account_deactivate_cancels_future_payments_for_relationship( + test_client, + db_session, + client_auth_header, +): + client, coach, request, relationship = create_client_coach_relationship( + db_session, + with_subscription=True, + ) + + client_auth_header = { + "Authorization": f"Bearer {create_jwt_token(client)}" + } + + resp = test_client.post( + "/roles/shared/account/deactivate", + headers=client_auth_header, + ) + + assert resp.status_code == 200, resp.text + + subscriptions = db_session.exec( + select(Subscription) + .join(PricingPlan, Subscription.pricing_plan_id == PricingPlan.id) + .where( + Subscription.client_id == client.client_id, + PricingPlan.coach_id == coach.coach_id, + ) + ).all() + active_cycles = db_session.exec( + select(BillingCycle) + .join(Subscription, BillingCycle.subscription_id == Subscription.id) + .join(PricingPlan, Subscription.pricing_plan_id == PricingPlan.id) + .where( + Subscription.client_id == client.client_id, + PricingPlan.coach_id == coach.coach_id, + BillingCycle.active == True, + ) + ).all() + + assert subscriptions + assert all(subscription.status == SubscriptionStatus.CANCELED for subscription in subscriptions) + assert all(subscription.canceled_at is not None for subscription in subscriptions) + assert active_cycles == [] + + +def test_refresh_payments_skips_canceled_subscription( + test_client, + db_session, + client_auth_header, +): + client, coach, request, relationship = create_client_coach_relationship( + db_session, + with_subscription=True, + ) + + subscriptions = db_session.exec( + select(Subscription) + .join(PricingPlan, Subscription.pricing_plan_id == PricingPlan.id) + .where( + Subscription.client_id == client.client_id, + PricingPlan.coach_id == coach.coach_id, + ) + ).all() + assert len(subscriptions) == 1 + + subscription = subscriptions[0] + subscription.status = SubscriptionStatus.CANCELED + db_session.add(subscription) + db_session.commit() + + cycles_before = db_session.exec( + select(BillingCycle).where(BillingCycle.subscription_id == subscription.id) + ).all() + + os.environ["CRON_SECRET"] = "test-cron-secret" + resp = test_client.post( + "/refresh_payments", + json={"cron_secret": "test-cron-secret"}, + ) + + cycles_after = db_session.exec( + select(BillingCycle).where(BillingCycle.subscription_id == subscription.id) + ).all() + + assert resp.status_code == 200, resp.text + assert len(cycles_after) == len(cycles_before) + + def test_account_deactivate_coach_notifies_client( test_client, db_session, client_auth_header, - coach_auth_header, ): client, coach, request, relationship = create_client_coach_relationship(db_session)