Merge pull request #499 from christianlouis/copilot/add-subscription-management-features
Merge main → subscription-management-features; fix migration chain collision
This commit is contained in:
@@ -491,3 +491,372 @@ def test_list_tiers_api(client):
|
||||
assert "starter" in ids
|
||||
assert "professional" in ids
|
||||
assert "business" in ids
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# New: subscription change management utilities
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
from app.utils.subscription import (
|
||||
SubscriptionChangeError,
|
||||
_tier_rank,
|
||||
apply_pending_subscription_changes,
|
||||
cancel_pending_subscription_change,
|
||||
request_subscription_change,
|
||||
)
|
||||
|
||||
# ---- _tier_rank -----------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_tier_rank_known_tiers():
|
||||
assert _tier_rank("free") == 0
|
||||
assert _tier_rank("starter") == 1
|
||||
assert _tier_rank("professional") == 2
|
||||
assert _tier_rank("business") == 3
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_tier_rank_unknown_defaults_to_zero():
|
||||
assert _tier_rank("unknown") == 0
|
||||
|
||||
|
||||
# ---- apply_pending_subscription_changes -----------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_apply_pending_no_profile_returns_false():
|
||||
db = MagicMock()
|
||||
db.query.return_value.filter.return_value.first.return_value = None
|
||||
assert apply_pending_subscription_changes(db, "nobody") is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_apply_pending_no_pending_returns_false():
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_change_pending_tier = None
|
||||
profile.subscription_change_pending_date = None
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
assert apply_pending_subscription_changes(db, "user1") is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_apply_pending_future_date_returns_false():
|
||||
from datetime import timedelta
|
||||
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_change_pending_tier = "free"
|
||||
profile.subscription_change_pending_date = datetime.now(timezone.utc) + timedelta(days=30)
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
assert apply_pending_subscription_changes(db, "user1") is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_apply_pending_due_date_applies_change():
|
||||
from datetime import timedelta
|
||||
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_tier = "starter"
|
||||
profile.subscription_change_pending_tier = "free"
|
||||
past = datetime.now(timezone.utc) - timedelta(days=1)
|
||||
profile.subscription_change_pending_date = past
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
|
||||
result = apply_pending_subscription_changes(db, "user1")
|
||||
|
||||
assert result is True
|
||||
assert profile.subscription_tier == "free"
|
||||
assert profile.subscription_period_start == past
|
||||
assert profile.subscription_change_pending_tier is None
|
||||
assert profile.subscription_change_pending_date is None
|
||||
db.commit.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_apply_pending_commit_failure_returns_false():
|
||||
from datetime import timedelta
|
||||
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_tier = "starter"
|
||||
profile.subscription_change_pending_tier = "free"
|
||||
profile.subscription_change_pending_date = datetime.now(timezone.utc) - timedelta(days=1)
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
db.commit.side_effect = Exception("DB error")
|
||||
|
||||
result = apply_pending_subscription_changes(db, "user1")
|
||||
|
||||
assert result is False
|
||||
db.rollback.assert_called_once()
|
||||
|
||||
|
||||
# ---- request_subscription_change -----------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_request_change_invalid_plan_raises():
|
||||
db = MagicMock()
|
||||
with patch("app.utils.subscription.get_all_tiers", return_value=[{"id": "free"}, {"id": "starter"}]):
|
||||
with pytest.raises(SubscriptionChangeError, match="Unknown subscription plan"):
|
||||
request_subscription_change(db, "user1", "galaxy_tier")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_request_change_same_plan_no_pending_raises():
|
||||
"""Requesting the active plan when no pending change exists should raise."""
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_tier = "starter"
|
||||
profile.subscription_change_pending_tier = None
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
|
||||
with patch("app.utils.subscription.get_all_tiers", return_value=[{"id": t} for t in TIER_ORDER]):
|
||||
with pytest.raises(SubscriptionChangeError, match="already on this plan"):
|
||||
request_subscription_change(db, "user1", "starter")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_request_change_same_plan_cancels_pending():
|
||||
"""Requesting the active plan when a downgrade is pending should cancel it."""
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_tier = "starter"
|
||||
profile.subscription_change_pending_tier = "free"
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
|
||||
with patch("app.utils.subscription.get_all_tiers", return_value=[{"id": t} for t in TIER_ORDER]):
|
||||
result = request_subscription_change(db, "user1", "starter")
|
||||
|
||||
assert result["immediate"] is True
|
||||
assert result["new_tier"] == "starter"
|
||||
assert "cancelled" in result["message"].lower()
|
||||
assert profile.subscription_change_pending_tier is None
|
||||
assert profile.subscription_change_pending_date is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_request_upgrade_is_immediate():
|
||||
"""Upgrading to a higher plan should take effect immediately."""
|
||||
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_tier = "free"
|
||||
profile.subscription_change_pending_tier = None
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
|
||||
with (
|
||||
patch("app.utils.subscription.get_all_tiers", return_value=[{"id": t} for t in TIER_ORDER]),
|
||||
patch("app.utils.subscription.get_tier", return_value={"id": "starter", "name": "Starter"}),
|
||||
):
|
||||
result = request_subscription_change(db, "user1", "starter")
|
||||
|
||||
assert result["immediate"] is True
|
||||
assert result["new_tier"] == "starter"
|
||||
assert profile.subscription_tier == "starter"
|
||||
# Period start should be set to approximately now
|
||||
assert profile.subscription_period_start is not None
|
||||
# Pending should be cleared
|
||||
assert profile.subscription_change_pending_tier is None
|
||||
db.commit.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_request_downgrade_schedules_for_future():
|
||||
"""Downgrade request within the billing period should be scheduled."""
|
||||
from datetime import timedelta
|
||||
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_tier = "professional"
|
||||
profile.subscription_change_pending_tier = None
|
||||
# Period started 10 days ago — we're mid-month
|
||||
profile.subscription_period_start = datetime.now(timezone.utc) - timedelta(days=10)
|
||||
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
|
||||
with (
|
||||
patch("app.utils.subscription.get_all_tiers", return_value=[{"id": t} for t in TIER_ORDER]),
|
||||
patch("app.utils.subscription.get_tier", return_value={"id": "starter", "name": "Starter"}),
|
||||
):
|
||||
result = request_subscription_change(db, "user1", "starter")
|
||||
|
||||
assert result["immediate"] is False
|
||||
assert result["effective_date"] is not None
|
||||
# The effective date should be in the future
|
||||
effective = datetime.fromisoformat(result["effective_date"])
|
||||
assert effective > datetime.now(timezone.utc)
|
||||
assert profile.subscription_change_pending_tier == "starter"
|
||||
db.commit.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_request_downgrade_immediate_when_period_elapsed():
|
||||
"""Downgrade request after billing period elapsed should apply immediately."""
|
||||
from datetime import timedelta
|
||||
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_tier = "professional"
|
||||
profile.subscription_change_pending_tier = None
|
||||
# Period started more than 1 month ago
|
||||
profile.subscription_period_start = datetime.now(timezone.utc) - timedelta(days=40)
|
||||
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
|
||||
with (
|
||||
patch("app.utils.subscription.get_all_tiers", return_value=[{"id": t} for t in TIER_ORDER]),
|
||||
patch("app.utils.subscription.get_tier", return_value={"id": "starter", "name": "Starter"}),
|
||||
):
|
||||
result = request_subscription_change(db, "user1", "starter")
|
||||
|
||||
assert result["immediate"] is True
|
||||
assert profile.subscription_tier == "starter"
|
||||
db.commit.assert_called_once()
|
||||
|
||||
|
||||
# ---- cancel_pending_subscription_change ----------------------------------
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cancel_pending_no_pending_returns_false():
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_change_pending_tier = None
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
assert cancel_pending_subscription_change(db, "user1") is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cancel_pending_clears_fields():
|
||||
db = MagicMock()
|
||||
profile = MagicMock()
|
||||
profile.subscription_change_pending_tier = "free"
|
||||
db.query.return_value.filter.return_value.first.return_value = profile
|
||||
|
||||
result = cancel_pending_subscription_change(db, "user1")
|
||||
|
||||
assert result is True
|
||||
assert profile.subscription_change_pending_tier is None
|
||||
assert profile.subscription_change_pending_date is None
|
||||
db.commit.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cancel_pending_no_profile_returns_false():
|
||||
db = MagicMock()
|
||||
db.query.return_value.filter.return_value.first.return_value = None
|
||||
assert cancel_pending_subscription_change(db, "ghost") is False
|
||||
|
||||
|
||||
# ---- API: POST /api/subscriptions/change ---------------------------------
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_api_change_subscription_upgrade(db_session):
|
||||
"""POST /api/subscriptions/change should immediately upgrade the plan."""
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.database import get_db
|
||||
from app.main import app
|
||||
from app.models import UserProfile
|
||||
from app.utils.subscription import seed_default_plans
|
||||
|
||||
# Create a user profile on the free tier
|
||||
profile = UserProfile(user_id="testuser", subscription_tier="free")
|
||||
db_session.add(profile)
|
||||
db_session.commit()
|
||||
seed_default_plans(db_session)
|
||||
|
||||
def override_get_db():
|
||||
yield db_session
|
||||
|
||||
app.dependency_overrides[get_db] = override_get_db
|
||||
with TestClient(app, base_url="http://localhost") as tc:
|
||||
with (
|
||||
patch("app.api.subscriptions._require_authenticated", return_value="testuser"),
|
||||
patch("app.config.settings.multi_user_enabled", True),
|
||||
):
|
||||
resp = tc.post(
|
||||
"/api/subscriptions/change",
|
||||
json={"plan_id": "starter", "billing_cycle": "monthly"},
|
||||
)
|
||||
app.dependency_overrides.pop(get_db, None)
|
||||
|
||||
assert resp.status_code == 200, resp.text
|
||||
data = resp.json()
|
||||
assert data["immediate"] is True
|
||||
assert data["new_tier"] == "starter"
|
||||
|
||||
db_session.refresh(profile)
|
||||
assert profile.subscription_tier == "starter"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_api_cancel_pending_change_no_pending_returns_404(db_session):
|
||||
"""DELETE /api/subscriptions/change should return 404 if nothing is pending."""
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.database import get_db
|
||||
from app.main import app
|
||||
from app.models import UserProfile
|
||||
|
||||
profile = UserProfile(user_id="testuser2", subscription_tier="starter")
|
||||
db_session.add(profile)
|
||||
db_session.commit()
|
||||
|
||||
def override_get_db():
|
||||
yield db_session
|
||||
|
||||
app.dependency_overrides[get_db] = override_get_db
|
||||
with TestClient(app, base_url="http://localhost") as tc:
|
||||
with (
|
||||
patch("app.api.subscriptions._require_authenticated", return_value="testuser2"),
|
||||
patch("app.config.settings.multi_user_enabled", True),
|
||||
):
|
||||
resp = tc.delete("/api/subscriptions/change")
|
||||
app.dependency_overrides.pop(get_db, None)
|
||||
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_api_cancel_pending_change_success(db_session):
|
||||
"""DELETE /api/subscriptions/change should cancel a pending downgrade."""
|
||||
from datetime import timedelta
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.database import get_db
|
||||
from app.main import app
|
||||
from app.models import UserProfile
|
||||
|
||||
profile = UserProfile(
|
||||
user_id="testuser3",
|
||||
subscription_tier="starter",
|
||||
subscription_change_pending_tier="free",
|
||||
subscription_change_pending_date=datetime.now(timezone.utc) + timedelta(days=20),
|
||||
)
|
||||
db_session.add(profile)
|
||||
db_session.commit()
|
||||
|
||||
def override_get_db():
|
||||
yield db_session
|
||||
|
||||
app.dependency_overrides[get_db] = override_get_db
|
||||
with TestClient(app, base_url="http://localhost") as tc:
|
||||
with (
|
||||
patch("app.api.subscriptions._require_authenticated", return_value="testuser3"),
|
||||
patch("app.config.settings.multi_user_enabled", True),
|
||||
):
|
||||
resp = tc.delete("/api/subscriptions/change")
|
||||
app.dependency_overrides.pop(get_db, None)
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["cancelled"] is True
|
||||
|
||||
db_session.refresh(profile)
|
||||
assert profile.subscription_change_pending_tier is None
|
||||
|
||||
Reference in New Issue
Block a user