fix: assign default tariff to promo bonuses

This commit is contained in:
3252a8
2026-06-07 22:23:42 +03:00
parent 2b6f25f0f8
commit 9de000e78c
4 changed files with 208 additions and 0 deletions
+80
View File
@@ -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",
)
@@ -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(