From a96007d48ad212c8bdf04d7a8350950e1720808e Mon Sep 17 00:00:00 2001 From: 3252a8 <3252a8@proton.me> Date: Wed, 13 May 2026 15:11:56 +0300 Subject: [PATCH] fix: avoid redundant premium squad restore updates --- bot/services/tariff_worker.py | 25 +++--- tests/test_tariff_worker.py | 140 +++++++++++++++++++++++++++++++++- 2 files changed, 150 insertions(+), 15 deletions(-) diff --git a/bot/services/tariff_worker.py b/bot/services/tariff_worker.py index 3fc2c4d..3ddf2e6 100644 --- a/bot/services/tariff_worker.py +++ b/bot/services/tariff_worker.py @@ -357,25 +357,22 @@ class TariffTrafficWorker: should_limit = False else: should_limit = premium_used >= premium_limit - changed = ( - int(sub.premium_baseline_bytes or 0) != premium_baseline - or int(sub.premium_topup_balance_bytes or 0) != premium_topup_balance - or int(getattr(sub, "premium_topup_used_bytes", 0) or 0) != premium_topup_used - or int(sub.premium_used_bytes or 0) != premium_used - or bool(sub.premium_is_limited) != should_limit - or getattr(sub, "premium_period_start_at", None) != premium_period_start - ) + panel_needs_update = bool(sub.premium_is_limited) != should_limit desired_squads = self.subscription_service._panel_squads_for_tariff( tariff, include_premium=not should_limit, ) desired_set = self._internal_squad_uuid_set(desired_squads) if isinstance(panel_user_dict, dict): - current_raw = panel_user_dict.get("activeInternalSquads") or panel_user_dict.get( - "active_internal_squads" - ) - if desired_set != self._internal_squad_uuid_set(current_raw): - changed = True + current_known = False + current_raw = None + for key in ("activeInternalSquads", "active_internal_squads"): + if key in panel_user_dict: + current_raw = panel_user_dict.get(key) + current_known = True + break + if current_known and desired_set != self._internal_squad_uuid_set(current_raw): + panel_needs_update = True sub.premium_baseline_bytes = premium_baseline sub.premium_topup_balance_bytes = premium_topup_balance sub.premium_topup_used_bytes = premium_topup_used @@ -391,7 +388,7 @@ class TariffTrafficWorker: premium_limit, premium_period_start, ) - if not changed: + if not panel_needs_update: return squads = desired_squads diff --git a/tests/test_tariff_worker.py b/tests/test_tariff_worker.py index 1b623ce..070e2f6 100644 --- a/tests/test_tariff_worker.py +++ b/tests/test_tariff_worker.py @@ -276,6 +276,139 @@ class TariffWorkerTests(unittest.IsolatedAsyncioTestCase): self.assertEqual(sub.premium_topup_used_bytes, 0) self.assertEqual(sub.premium_period_start_at, datetime(2026, 6, 1, tzinfo=timezone.utc)) + async def test_premium_usage_update_does_not_patch_panel_when_access_state_unchanged(self): + payload = _tariffs_config_payload() + payload["tariffs"][0]["premium_squad_uuids"] = ["premium-squad"] + payload["tariffs"][0]["premium_monthly_gb"] = 25 + with tempfile.TemporaryDirectory() as tmpdir: + config_path = Path(tmpdir) / "tariffs.json" + config_path.write_text(json.dumps(payload), encoding="utf-8") + + settings = Settings( + _env_file=None, + BOT_TOKEN="token", + POSTGRES_USER="app_user", + POSTGRES_PASSWORD="app_password", + TARIFFS_CONFIG_PATH=str(config_path), + TARIFF_TRAFFIC_WARNING_LEVELS="101", + ) + panel_service = AsyncMock(spec=PanelApiService) + panel_service.get_internal_squad_accessible_nodes = AsyncMock( + return_value=[{"uuid": "node-1"}] + ) + panel_service.get_node_users_bandwidth_stats = AsyncMock( + return_value={ + "topUsers": [ + {"username": "tg_123", "total": 5 * (1024**3)}, + ] + } + ) + panel_service.update_user_details_on_panel = AsyncMock(return_value={"response": {}}) + subscription_service = SubscriptionService(settings, panel_service) + worker = TariffTrafficWorker( + settings=settings, + session_factory=SimpleNamespace(), + panel_service=panel_service, + subscription_service=subscription_service, + ) + now = datetime(2026, 5, 9, tzinfo=timezone.utc) + sub = SimpleNamespace( + subscription_id=1, + user_id=123, + panel_user_uuid="panel-uuid", + premium_baseline_bytes=25 * (1024**3), + premium_topup_balance_bytes=0, + premium_topup_used_bytes=0, + premium_used_bytes=1 * (1024**3), + premium_is_limited=False, + premium_period_start_at=datetime(2026, 5, 1, tzinfo=timezone.utc), + premium_unlimited_override=False, + premium_bonus_bytes=0, + ) + tariff = settings.tariffs_config.require("standard") + + await worker._sync_premium_squad_limit( + AsyncMock(), + sub, + tariff, + now, + panel_username="tg_123", + panel_user_dict={ + "activeInternalSquads": [ + {"uuid": "squad-1"}, + {"uuid": "premium-squad"}, + ] + }, + ) + + self.assertEqual(sub.premium_used_bytes, 5 * (1024**3)) + self.assertFalse(sub.premium_is_limited) + panel_service.update_user_details_on_panel.assert_not_awaited() + + async def test_premium_sync_patches_panel_when_current_squads_are_known_and_wrong(self): + payload = _tariffs_config_payload() + payload["tariffs"][0]["premium_squad_uuids"] = ["premium-squad"] + payload["tariffs"][0]["premium_monthly_gb"] = 25 + with tempfile.TemporaryDirectory() as tmpdir: + config_path = Path(tmpdir) / "tariffs.json" + config_path.write_text(json.dumps(payload), encoding="utf-8") + + settings = Settings( + _env_file=None, + BOT_TOKEN="token", + POSTGRES_USER="app_user", + POSTGRES_PASSWORD="app_password", + TARIFFS_CONFIG_PATH=str(config_path), + TARIFF_TRAFFIC_WARNING_LEVELS="101", + ) + panel_service = AsyncMock(spec=PanelApiService) + panel_service.get_internal_squad_accessible_nodes = AsyncMock( + return_value=[{"uuid": "node-1"}] + ) + panel_service.get_node_users_bandwidth_stats = AsyncMock( + return_value={ + "topUsers": [ + {"username": "tg_123", "total": 5 * (1024**3)}, + ] + } + ) + panel_service.update_user_details_on_panel = AsyncMock(return_value={"response": {}}) + subscription_service = SubscriptionService(settings, panel_service) + worker = TariffTrafficWorker( + settings=settings, + session_factory=SimpleNamespace(), + panel_service=panel_service, + subscription_service=subscription_service, + ) + now = datetime(2026, 5, 9, tzinfo=timezone.utc) + sub = SimpleNamespace( + subscription_id=1, + user_id=123, + panel_user_uuid="panel-uuid", + premium_baseline_bytes=25 * (1024**3), + premium_topup_balance_bytes=0, + premium_topup_used_bytes=0, + premium_used_bytes=5 * (1024**3), + premium_is_limited=False, + premium_period_start_at=datetime(2026, 5, 1, tzinfo=timezone.utc), + premium_unlimited_override=False, + premium_bonus_bytes=0, + ) + tariff = settings.tariffs_config.require("standard") + + await worker._sync_premium_squad_limit( + AsyncMock(), + sub, + tariff, + now, + panel_username="tg_123", + panel_user_dict={"activeInternalSquads": [{"uuid": "squad-1"}]}, + ) + + panel_service.update_user_details_on_panel.assert_awaited_once() + payload_sent = panel_service.update_user_details_on_panel.await_args.args[1] + self.assertEqual(payload_sent["activeInternalSquads"], ["squad-1", "premium-squad"]) + async def test_premium_unlimited_override_never_throttles(self): payload = _tariffs_config_payload() payload["tariffs"][0]["premium_squad_uuids"] = ["premium-squad"] @@ -327,7 +460,12 @@ class TariffWorkerTests(unittest.IsolatedAsyncioTestCase): tariff = settings.tariffs_config.require("standard") await worker._sync_premium_squad_limit( - AsyncMock(), sub, tariff, datetime.now(timezone.utc), panel_username="tg_42" + AsyncMock(), + sub, + tariff, + datetime.now(timezone.utc), + panel_username="tg_42", + panel_user_dict={"activeInternalSquads": [{"uuid": "squad-1"}]}, ) self.assertFalse(sub.premium_is_limited)