fix: stop premium squad sync churn

This commit is contained in:
3252a8
2026-05-25 11:54:03 +03:00
parent 11b7823188
commit 340afe1b80
2 changed files with 370 additions and 15 deletions
+174 -14
View File
@@ -2,7 +2,7 @@ import asyncio
import logging import logging
import time import time
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Optional from typing import Any, Optional
from aiogram import Bot from aiogram import Bot
from aiogram.types import InlineKeyboardButton, InlineKeyboardMarkup, WebAppInfo from aiogram.types import InlineKeyboardButton, InlineKeyboardMarkup, WebAppInfo
@@ -30,6 +30,7 @@ PREMIUM_WARNING_DEPLETED_LEVEL = PREMIUM_WARNING_LEVEL_OFFSET + 100
TARIFF_WORKER_BATCH_SIZE = 50 TARIFF_WORKER_BATCH_SIZE = 50
TARIFF_WORKER_PANEL_CONCURRENCY = 10 TARIFF_WORKER_PANEL_CONCURRENCY = 10
TARIFF_WORKER_BULK_PANEL_FETCH_THRESHOLD = 50 TARIFF_WORKER_BULK_PANEL_FETCH_THRESHOLD = 50
TARIFF_WORKER_SQUAD_CONFIRMATION_CACHE_TTL_SECONDS = 900
class TariffTrafficWorker: class TariffTrafficWorker:
@@ -51,6 +52,7 @@ class TariffTrafficWorker:
self._stopped = asyncio.Event() self._stopped = asyncio.Event()
self._premium_nodes_cache = {} self._premium_nodes_cache = {}
self._premium_node_usage_tick_cache = {} self._premium_node_usage_tick_cache = {}
self._premium_squad_match_cache = {}
async def _user_lang(self, session: AsyncSession, user_id: int) -> str: async def _user_lang(self, session: AsyncSession, user_id: int) -> str:
try: try:
@@ -148,6 +150,7 @@ class TariffTrafficWorker:
return return
panel_users_by_uuid = await self._prefetch_panel_users_by_uuid(subs) panel_users_by_uuid = await self._prefetch_panel_users_by_uuid(subs)
panel_view = "list" if panel_users_by_uuid is not None else "full_fetch"
semaphore = asyncio.Semaphore(TARIFF_WORKER_PANEL_CONCURRENCY) semaphore = asyncio.Semaphore(TARIFF_WORKER_PANEL_CONCURRENCY)
async def _fetch_panel(sub: Subscription) -> dict: async def _fetch_panel(sub: Subscription) -> dict:
@@ -230,6 +233,7 @@ class TariffTrafficWorker:
now, now,
panel_username=panel_username, panel_username=panel_username,
panel_user_dict=panel_data, panel_user_dict=panel_data,
panel_view=panel_view,
) )
async def _prefetch_panel_users_by_uuid( async def _prefetch_panel_users_by_uuid(
@@ -461,6 +465,7 @@ class TariffTrafficWorker:
*, *,
panel_username: Optional[str] = None, panel_username: Optional[str] = None,
panel_user_dict: Optional[dict] = None, panel_user_dict: Optional[dict] = None,
panel_view: str = "unknown",
) -> None: ) -> None:
if not getattr(tariff, "premium_squad_uuids", None): if not getattr(tariff, "premium_squad_uuids", None):
if ( if (
@@ -531,22 +536,59 @@ class TariffTrafficWorker:
should_limit = False should_limit = False
else: else:
should_limit = premium_used >= premium_limit should_limit = premium_used >= premium_limit
panel_needs_update = bool(sub.premium_is_limited) != should_limit access_state_changed = bool(sub.premium_is_limited) != should_limit
desired_squads = self.subscription_service._panel_squads_for_tariff( desired_squads = self.subscription_service._panel_squads_for_tariff(
tariff, tariff,
include_premium=not should_limit, include_premium=not should_limit,
) )
desired_set = self._internal_squad_uuid_set(desired_squads) desired_set = self._internal_squad_uuid_set(desired_squads)
if isinstance(panel_user_dict, dict): squad_match_cache_key = self._premium_squad_match_cache_key(
current_known = False sub.panel_user_uuid,
current_raw = None desired_set,
for key in ("activeInternalSquads", "active_internal_squads"): )
if key in panel_user_dict: panel_needs_update = access_state_changed
current_raw = panel_user_dict.get(key) panel_user_for_report = panel_user_dict
current_known = True panel_view_for_report = panel_view
break panel_update_reasons: list[str] = []
if current_known and desired_set != self._internal_squad_uuid_set(current_raw): if access_state_changed:
panel_update_reasons.append(
"premium_access_limited" if should_limit else "premium_access_restored"
)
current_known, current_set = self._panel_active_squad_uuid_set(panel_user_dict)
if current_known:
current_mismatch = desired_set != current_set
if not current_mismatch:
panel_needs_update = False
elif panel_view == "list":
if self._premium_squad_match_cache_is_fresh(squad_match_cache_key):
panel_needs_update = False
else:
full_panel_user = await self._get_full_panel_user_for_squad_confirmation(
sub.panel_user_uuid,
)
full_known, full_set = self._panel_active_squad_uuid_set(full_panel_user)
if full_known:
panel_user_for_report = full_panel_user
panel_view_for_report = "full_fetch"
if desired_set != full_set:
panel_needs_update = True
panel_update_reasons.append("activeInternalSquads_mismatch")
else:
self._remember_premium_squad_match(squad_match_cache_key)
panel_needs_update = False
elif not access_state_changed:
panel_needs_update = False
else:
panel_needs_update = True panel_needs_update = True
panel_update_reasons.append("activeInternalSquads_mismatch")
if (
not panel_needs_update
and current_known
and desired_set == current_set
and panel_view != "list"
):
self._remember_premium_squad_match(squad_match_cache_key)
sub.premium_baseline_bytes = premium_baseline sub.premium_baseline_bytes = premium_baseline
sub.premium_topup_balance_bytes = premium_topup_balance sub.premium_topup_balance_bytes = premium_topup_balance
sub.premium_topup_used_bytes = premium_topup_used sub.premium_topup_used_bytes = premium_topup_used
@@ -566,11 +608,21 @@ class TariffTrafficWorker:
return return
squads = desired_squads squads = desired_squads
await self.panel_service.update_user_details_on_panel( self._log_premium_squad_panel_patch(
sub=sub,
panel_uuid=sub.panel_user_uuid,
update_payload={"uuid": sub.panel_user_uuid, "activeInternalSquads": squads},
current_panel_user=panel_user_for_report,
reasons=panel_update_reasons or ["premium_squad_sync"],
panel_view=panel_view_for_report,
)
updated_panel_user = await self.panel_service.update_user_details_on_panel(
sub.panel_user_uuid, sub.panel_user_uuid,
{"uuid": sub.panel_user_uuid, "activeInternalSquads": squads}, {"uuid": sub.panel_user_uuid, "activeInternalSquads": squads},
log_response=False, log_response=False,
) )
if updated_panel_user:
self._remember_premium_squad_match(squad_match_cache_key)
logging.info( logging.info(
"Premium squad access %s for user %s tariff %s: %s/%s bytes", "Premium squad access %s for user %s tariff %s: %s/%s bytes",
"limited" if should_limit else "restored", "limited" if should_limit else "restored",
@@ -580,20 +632,128 @@ class TariffTrafficWorker:
premium_limit, premium_limit,
) )
async def _get_full_panel_user_for_squad_confirmation(
self,
panel_user_uuid: str,
) -> Optional[dict]:
try:
return await self.panel_service.get_user_by_uuid(
panel_user_uuid,
log_response=False,
)
except Exception:
logging.exception(
"TariffTrafficWorker: failed to confirm panel squads for user %s",
panel_user_uuid,
)
return None
@staticmethod
def _premium_squad_match_cache_key(
panel_user_uuid: str,
desired_set: set[str],
) -> tuple[str, tuple[str, ...]]:
return str(panel_user_uuid), tuple(sorted(desired_set))
def _premium_squad_match_cache_is_fresh(self, cache_key: tuple[str, tuple[str, ...]]) -> bool:
cached_at = self._premium_squad_match_cache.get(cache_key)
if not cached_at:
return False
return (
time.monotonic() - float(cached_at)
< TARIFF_WORKER_SQUAD_CONFIRMATION_CACHE_TTL_SECONDS
)
def _remember_premium_squad_match(self, cache_key: tuple[str, tuple[str, ...]]) -> None:
self._premium_squad_match_cache[cache_key] = time.monotonic()
@classmethod
def _panel_active_squad_uuid_set(
cls,
panel_user_dict: Optional[dict],
) -> tuple[bool, set[str]]:
current_known, current_raw = cls._panel_active_squads_raw(panel_user_dict)
return current_known, cls._internal_squad_uuid_set(current_raw)
@staticmethod
def _panel_active_squads_raw(panel_user_dict: Optional[dict]) -> tuple[bool, Any]:
if not isinstance(panel_user_dict, dict):
return False, None
for key in (
"activeInternalSquads",
"active_internal_squads",
"activeInternalSquadUuids",
"active_internal_squad_uuids",
):
if key in panel_user_dict:
return True, panel_user_dict.get(key)
return False, None
def _log_premium_squad_panel_patch(
self,
*,
sub: Subscription,
panel_uuid: str,
update_payload: dict[str, Any],
current_panel_user: Optional[dict],
reasons: list[str],
panel_view: str,
) -> None:
current_known, current_set = self._panel_active_squad_uuid_set(current_panel_user)
desired_set = self._internal_squad_uuid_set(update_payload.get("activeInternalSquads"))
fields = "none" if current_known and current_set == desired_set else "activeInternalSquads"
logging.info(
"Sync panel PATCH: source=%s user_id=%s telegram_id=%s panel_uuid=%s "
"panel_view=%s reasons=%s fields=%s payload_fields=%s changes=%s",
"premium_squad_limit",
getattr(sub, "user_id", None),
getattr(sub, "user_id", None),
panel_uuid,
panel_view,
",".join(reasons),
fields,
"activeInternalSquads",
"activeInternalSquads:%s->%s"
% (
self._format_squad_uuid_set(current_set if current_known else None),
self._format_squad_uuid_set(desired_set),
),
)
@staticmethod @staticmethod
def _internal_squad_uuid_set(raw) -> set[str]: def _internal_squad_uuid_set(raw) -> set[str]:
if not isinstance(raw, list): if not isinstance(raw, (list, tuple, set)):
return set() return set()
out: set[str] = set() out: set[str] = set()
for item in raw: for item in raw:
if isinstance(item, dict): if isinstance(item, dict):
u = item.get("uuid") or item.get("internalSquadUuid") or item.get("squadUuid") nested_squad = item.get("internalSquad") or item.get("squad")
if not isinstance(nested_squad, dict):
nested_squad = {}
u = (
item.get("uuid")
or item.get("internalSquadUuid")
or item.get("squadUuid")
or nested_squad.get("uuid")
)
if u: if u:
out.add(str(u)) out.add(str(u))
elif item: elif item:
out.add(str(item)) out.add(str(item))
return out return out
@staticmethod
def _format_squad_uuid_set(value: Optional[set[str]]) -> str:
if value is None:
return "missing"
values = sorted(str(item) for item in value)
preview = ",".join(values[:4])
suffix = ",..." if len(values) > 4 else ""
text = f"[{len(values)}:{preview}{suffix}]"
if len(text) > 96:
return f"{text[:93]}..."
return text
@staticmethod @staticmethod
def _fmt_bytes(value: int) -> str: def _fmt_bytes(value: int) -> str:
size = float(max(0, int(value or 0))) size = float(max(0, int(value or 0)))
+196 -1
View File
@@ -304,6 +304,14 @@ class TariffWorkerTests(unittest.IsolatedAsyncioTestCase):
] ]
} }
) )
panel_service.get_user_by_uuid = AsyncMock(
return_value={
"activeInternalSquads": [
{"uuid": "squad-1"},
{"uuid": "premium-squad"},
]
}
)
panel_service.update_user_details_on_panel = AsyncMock(return_value={"response": {}}) panel_service.update_user_details_on_panel = AsyncMock(return_value={"response": {}})
subscription_service = SubscriptionService(settings, panel_service) subscription_service = SubscriptionService(settings, panel_service)
worker = TariffTrafficWorker( worker = TariffTrafficWorker(
@@ -345,8 +353,9 @@ class TariffWorkerTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(sub.premium_used_bytes, 5 * (1024**3)) self.assertEqual(sub.premium_used_bytes, 5 * (1024**3))
self.assertFalse(sub.premium_is_limited) self.assertFalse(sub.premium_is_limited)
panel_service.update_user_details_on_panel.assert_not_awaited() panel_service.update_user_details_on_panel.assert_not_awaited()
panel_service.get_user_by_uuid.assert_not_awaited()
async def test_premium_sync_patches_panel_when_current_squads_are_known_and_wrong(self): async def test_premium_sync_trusts_full_fetch_over_bulk_list_squad_mismatch(self):
payload = _tariffs_config_payload() payload = _tariffs_config_payload()
payload["tariffs"][0]["premium_squad_uuids"] = ["premium-squad"] payload["tariffs"][0]["premium_squad_uuids"] = ["premium-squad"]
payload["tariffs"][0]["premium_monthly_gb"] = 25 payload["tariffs"][0]["premium_monthly_gb"] = 25
@@ -373,6 +382,14 @@ class TariffWorkerTests(unittest.IsolatedAsyncioTestCase):
] ]
} }
) )
panel_service.get_user_by_uuid = AsyncMock(
return_value={
"activeInternalSquads": [
{"uuid": "squad-1"},
{"uuid": "premium-squad"},
]
}
)
panel_service.update_user_details_on_panel = AsyncMock(return_value={"response": {}}) panel_service.update_user_details_on_panel = AsyncMock(return_value={"response": {}})
subscription_service = SubscriptionService(settings, panel_service) subscription_service = SubscriptionService(settings, panel_service)
worker = TariffTrafficWorker( worker = TariffTrafficWorker(
@@ -404,11 +421,185 @@ class TariffWorkerTests(unittest.IsolatedAsyncioTestCase):
now, now,
panel_username="tg_123", panel_username="tg_123",
panel_user_dict={"activeInternalSquads": [{"uuid": "squad-1"}]}, panel_user_dict={"activeInternalSquads": [{"uuid": "squad-1"}]},
panel_view="list",
) )
panel_service.get_user_by_uuid.assert_awaited_once_with(
"panel-uuid",
log_response=False,
)
panel_service.update_user_details_on_panel.assert_not_awaited()
panel_service.get_user_by_uuid.reset_mock()
await worker._sync_premium_squad_limit(
AsyncMock(),
sub,
tariff,
now,
panel_username="tg_123",
panel_user_dict={"activeInternalSquads": [{"uuid": "squad-1"}]},
panel_view="list",
)
panel_service.get_user_by_uuid.assert_not_awaited()
panel_service.update_user_details_on_panel.assert_not_awaited()
async def test_premium_state_change_skips_panel_patch_when_full_user_already_matches(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.get_user_by_uuid = AsyncMock(
return_value={
"activeInternalSquads": [
{"uuid": "squad-1"},
{"uuid": "premium-squad"},
]
}
)
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=True,
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_view="list",
)
self.assertFalse(sub.premium_is_limited)
panel_service.get_user_by_uuid.assert_awaited_once_with(
"panel-uuid",
log_response=False,
)
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.get_user_by_uuid = AsyncMock(
return_value={"activeInternalSquads": [{"uuid": "squad-1"}]}
)
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")
with self.assertLogs(level="INFO") as logs:
await worker._sync_premium_squad_limit(
AsyncMock(),
sub,
tariff,
now,
panel_username="tg_123",
panel_user_dict={"activeInternalSquads": [{"uuid": "squad-1"}]},
panel_view="list",
)
panel_service.update_user_details_on_panel.assert_awaited_once() panel_service.update_user_details_on_panel.assert_awaited_once()
panel_service.get_user_by_uuid.assert_awaited_once_with(
"panel-uuid",
log_response=False,
)
payload_sent = panel_service.update_user_details_on_panel.await_args.args[1] payload_sent = panel_service.update_user_details_on_panel.await_args.args[1]
self.assertEqual(payload_sent["activeInternalSquads"], ["squad-1", "premium-squad"]) self.assertEqual(payload_sent["activeInternalSquads"], ["squad-1", "premium-squad"])
self.assertTrue(
any(
"Sync panel PATCH: source=premium_squad_limit" in line
and "reasons=activeInternalSquads_mismatch" in line
and "fields=activeInternalSquads" in line
for line in logs.output
)
)
async def test_premium_unlimited_override_never_throttles(self): async def test_premium_unlimited_override_never_throttles(self):
payload = _tariffs_config_payload() payload = _tariffs_config_payload()
@@ -437,6 +628,9 @@ class TariffWorkerTests(unittest.IsolatedAsyncioTestCase):
] ]
} }
) )
panel_service.get_user_by_uuid = AsyncMock(
return_value={"activeInternalSquads": [{"uuid": "squad-1"}]}
)
panel_service.update_user_details_on_panel = AsyncMock(return_value={"response": {}}) panel_service.update_user_details_on_panel = AsyncMock(return_value={"response": {}})
subscription_service = SubscriptionService(settings, panel_service) subscription_service = SubscriptionService(settings, panel_service)
worker = TariffTrafficWorker( worker = TariffTrafficWorker(
@@ -467,6 +661,7 @@ class TariffWorkerTests(unittest.IsolatedAsyncioTestCase):
datetime.now(timezone.utc), datetime.now(timezone.utc),
panel_username="tg_42", panel_username="tg_42",
panel_user_dict={"activeInternalSquads": [{"uuid": "squad-1"}]}, panel_user_dict={"activeInternalSquads": [{"uuid": "squad-1"}]},
panel_view="list",
) )
self.assertFalse(sub.premium_is_limited) self.assertFalse(sub.premium_is_limited)