From 9de000e78c2bc8cd93bc26dc697b111fabf7077c Mon Sep 17 00:00:00 2001 From: 3252a8 <3252a8@proton.me> Date: Sun, 7 Jun 2026 22:23:42 +0300 Subject: [PATCH] fix: assign default tariff to promo bonuses --- backend/bot/services/promo_code_service.py | 5 ++ .../subscription_service_impl/lifecycle.py | 57 +++++++++++++ tests/test_promo_code_service.py | 80 +++++++++++++++++++ tests/test_subscription_service_behavior.py | 66 +++++++++++++++ 4 files changed, 208 insertions(+) create mode 100644 tests/test_promo_code_service.py diff --git a/backend/bot/services/promo_code_service.py b/backend/bot/services/promo_code_service.py index 947550d..be96b3b 100644 --- a/backend/bot/services/promo_code_service.py +++ b/backend/bot/services/promo_code_service.py @@ -87,12 +87,17 @@ class PromoCodeService: return False, _("promo_code_already_used_by_user", code=code_display) bonus_days = promo_data.bonus_days + default_tariff_key = None + tariffs_config = getattr(self.settings, "tariffs_config", None) + if tariffs_config: + default_tariff_key = getattr(tariffs_config, "default_tariff", None) new_end_date = await self.subscription_service.extend_active_subscription_days( session=session, user_id=user_id, bonus_days=bonus_days, reason=f"promo code {applied_code}", + tariff_key=default_tariff_key, ) if new_end_date: diff --git a/backend/bot/services/subscription_service_impl/lifecycle.py b/backend/bot/services/subscription_service_impl/lifecycle.py index 40c62ab..fbd2b79 100644 --- a/backend/bot/services/subscription_service_impl/lifecycle.py +++ b/backend/bot/services/subscription_service_impl/lifecycle.py @@ -783,6 +783,7 @@ class SubscriptionLifecycleMixin: bonus_days: int, reason: str = "bonus", extend_hwid_devices: bool = True, + tariff_key: Optional[str] = None, ) -> Optional[datetime]: reason_lower = (reason or "").lower() apply_main_traffic_limit = any( @@ -809,6 +810,17 @@ class SubscriptionLifecycleMixin: preserve_tariff_limits = bool( active_sub and active_sub.tariff_key and self._tariffs_config() ) + bonus_tariff = None + if not active_sub and tariff_key and self._tariffs_config(): + try: + bonus_tariff = self._resolve_tariff(tariff_key) + except Exception: + logging.warning( + "Unable to resolve bonus tariff %s for user %s.", + tariff_key, + user_id, + exc_info=True, + ) if not active_sub or not active_sub.end_date: logging.info( f"No active subscription found for user {user_id}. Creating new one for {bonus_days} days." # noqa: E501 @@ -818,10 +830,17 @@ class SubscriptionLifecycleMixin: # Apply main traffic limit for admin/referral/promo bonuses, fallback to trial limit otherwise # noqa: E501 traffic_limit = ( + self._traffic_limit_for_period_tariff(bonus_tariff) + if bonus_tariff + else self.settings.user_traffic_limit_bytes if apply_main_traffic_limit else self.settings.trial_traffic_limit_bytes ) + premium_baseline_bytes = bonus_tariff.premium_monthly_bytes if bonus_tariff else 0 + base_hwid_limit = ( + self._base_hwid_limit_for_tariff(bonus_tariff) if bonus_tariff else None + ) bonus_sub_payload = { "user_id": user_id, @@ -834,6 +853,21 @@ class SubscriptionLifecycleMixin: "status_from_panel": "ACTIVE_BONUS", "traffic_limit_bytes": traffic_limit, "auto_renew_enabled": False, + "tariff_key": bonus_tariff.key if bonus_tariff else None, + "tier_baseline_bytes": bonus_tariff.monthly_bytes if bonus_tariff else None, + "topup_balance_bytes": 0, + "regular_bonus_bytes": 0, + "regular_unlimited_override": False, + "premium_baseline_bytes": premium_baseline_bytes, + "premium_topup_balance_bytes": 0, + "premium_topup_used_bytes": 0, + "premium_used_bytes": 0, + "premium_is_limited": False, + "premium_period_start_at": None, + "period_start_at": None, + "is_throttled": False, + "hwid_device_limit": base_hwid_limit, + "extra_hwid_devices": 0, # Registration/referral bonus grants are short-lived, like a # trial: only warn a few hours before they end, not days ahead. "suppress_early_expiry_notifications": True, @@ -886,13 +920,36 @@ class SubscriptionLifecycleMixin: panel_update_payload = self._build_panel_update_payload( expire_at=new_end_date_obj, traffic_limit_bytes=( + updated_sub_model.traffic_limit_bytes + if bonus_tariff + else self.settings.user_traffic_limit_bytes if apply_main_traffic_limit and not preserve_tariff_limits else None ), + traffic_limit_strategy=( + "MONTH" + if bonus_tariff and bonus_tariff.billing_model == "period" + else self.settings.USER_TRAFFIC_STRATEGY + if bonus_tariff + else None + ), + hwid_device_limit=( + self._effective_hwid_limit(updated_sub_model.hwid_device_limit, 0) + if bonus_tariff + else None + ), include_uuid=False, include_default_squads=False, ) + if bonus_tariff: + panel_update_payload["activeInternalSquads"] = self._panel_squads_for_tariff( + bonus_tariff + ) + if self.settings.parsed_user_external_squad_uuid: + panel_update_payload["externalSquadUuid"] = ( + self.settings.parsed_user_external_squad_uuid + ) panel_update_success = await self.panel_service.update_user_details_on_panel( panel_uuid, diff --git a/tests/test_promo_code_service.py b/tests/test_promo_code_service.py new file mode 100644 index 0000000..2297978 --- /dev/null +++ b/tests/test_promo_code_service.py @@ -0,0 +1,80 @@ +from datetime import datetime, timezone +from types import SimpleNamespace +from unittest import IsolatedAsyncioTestCase +from unittest.mock import AsyncMock, patch + +from bot.services.promo_code_service import PromoCodeService + + +class PromoCodeServiceTests(IsolatedAsyncioTestCase): + async def test_apply_promo_passes_default_tariff_for_new_bonus_subscription(self): + end_date = datetime(2026, 1, 8, tzinfo=timezone.utc) + settings = SimpleNamespace( + MIGRATION_REMNASHOP_PROMO_CODE_COMPAT_ENABLED=False, + BRUTE_FORCE_LOCK_SECONDS=60, + BRUTE_FORCE_MAX_FAILURES=5, + BRUTE_FORCE_WINDOW_SECONDS=300, + tariffs_config=SimpleNamespace(default_tariff="standard"), + ) + subscription_service = SimpleNamespace( + extend_active_subscription_days=AsyncMock(return_value=end_date) + ) + i18n = SimpleNamespace(gettext=lambda lang, key, **kw: key) + service = PromoCodeService(settings, subscription_service, AsyncMock(), i18n) + session = AsyncMock() + promo = SimpleNamespace( + promo_code_id=5, + code="HELLO", + bonus_days=7, + ) + + with ( + patch( + "bot.services.promo_code_service.security_dal.check_throttle", + AsyncMock(return_value=SimpleNamespace(locked=False, retry_after=None)), + ), + patch( + "bot.services.promo_code_service.promo_code_dal.get_active_promo_code_by_code_str", + AsyncMock(return_value=promo), + ), + patch( + "bot.services.promo_code_service.promo_code_dal.get_user_activation_for_promo", + AsyncMock(return_value=None), + ), + patch( + "bot.services.promo_code_service.promo_code_dal.record_promo_activation", + AsyncMock(return_value=True), + ), + patch( + "bot.services.promo_code_service.promo_code_dal.increment_promo_code_usage", + AsyncMock(return_value=True), + ), + patch( + "bot.services.promo_code_service.security_dal.clear_throttle_state", + AsyncMock(), + ), + patch( + "bot.services.promo_code_service.NotificationService", + return_value=SimpleNamespace(notify_promo_activation=AsyncMock()), + ), + patch( + "bot.services.promo_code_service.user_dal.get_user_by_id", + AsyncMock(return_value=None), + ), + ): + success, result = await service.apply_promo_code( + session=session, + user_id=42, + code_input="hello", + user_lang="en", + ) + + self.assertTrue(success) + self.assertEqual(result, end_date) + subscription_service.extend_active_subscription_days.assert_awaited_once_with( + session=session, + user_id=42, + bonus_days=7, + reason="promo code HELLO", + tariff_key="standard", + ) diff --git a/tests/test_subscription_service_behavior.py b/tests/test_subscription_service_behavior.py index 5caf06a..042a468 100644 --- a/tests/test_subscription_service_behavior.py +++ b/tests/test_subscription_service_behavior.py @@ -532,6 +532,72 @@ class SubscriptionServiceActivationDispatchTests(unittest.IsolatedAsyncioTestCas class SubscriptionServiceBonusExtensionTests(unittest.IsolatedAsyncioTestCase): + async def test_promo_bonus_without_active_subscription_uses_default_tariff_squads(self): + with tempfile.TemporaryDirectory() as tmpdir: + settings = _make_settings( + _tariffs_config_payload(), + tmpdir, + USER_TRAFFIC_LIMIT_GB=999, + USER_EXTERNAL_SQUAD_UUID="external-squad", + ) + service = _make_service(settings) + service._get_or_create_panel_user_link_details = AsyncMock( + return_value=("panel-user", "short-uuid", "short", False) + ) + service.panel_service.update_user_details_on_panel = AsyncMock( + return_value={"ok": True} + ) + updated_sub = SimpleNamespace( + subscription_id=10, + end_date=datetime.now(timezone.utc) + timedelta(days=7), + traffic_limit_bytes=100 * GIB, + tariff_key="standard", + hwid_device_limit=3, + ) + + with ( + patch( + "bot.services.subscription_service_impl.lifecycle.user_dal.get_user_by_id", + AsyncMock(return_value=SimpleNamespace(user_id=42)), + ), + patch( + "bot.services.subscription_service_impl.lifecycle.subscription_dal.get_active_subscription_by_user_id", + AsyncMock(return_value=None), + ), + patch( + "bot.services.subscription_service_impl.lifecycle.subscription_dal.deactivate_other_active_subscriptions", + AsyncMock(), + ), + patch( + "bot.services.subscription_service_impl.lifecycle.subscription_dal.upsert_subscription", + AsyncMock(return_value=updated_sub), + ) as upsert_subscription, + ): + await service.extend_active_subscription_days( + session=AsyncMock(), + user_id=42, + bonus_days=7, + reason="promo code HELLO", + tariff_key="standard", + ) + + sub_payload = upsert_subscription.await_args.args[1] + self.assertEqual(sub_payload["tariff_key"], "standard") + self.assertEqual(sub_payload["traffic_limit_bytes"], 100 * GIB) + self.assertEqual(sub_payload["tier_baseline_bytes"], 100 * GIB) + self.assertEqual(sub_payload["premium_baseline_bytes"], 25 * GIB) + self.assertEqual(sub_payload["hwid_device_limit"], 3) + + panel_payload = service.panel_service.update_user_details_on_panel.await_args.args[1] + self.assertEqual(panel_payload["trafficLimitBytes"], 100 * GIB) + self.assertEqual(panel_payload["trafficLimitStrategy"], "MONTH") + self.assertEqual(panel_payload["hwidDeviceLimit"], 3) + self.assertEqual( + panel_payload["activeInternalSquads"], + ["main-squad", "shared-squad", "premium-squad"], + ) + self.assertEqual(panel_payload["externalSquadUuid"], "external-squad") + async def test_referral_extension_preserves_existing_tariff_limit(self): with tempfile.TemporaryDirectory() as tmpdir: settings = _make_settings(