722 lines
30 KiB
Python
722 lines
30 KiB
Python
import asyncio
|
|
import logging
|
|
import time
|
|
from datetime import datetime, timezone
|
|
from typing import Optional
|
|
|
|
from aiogram import Bot
|
|
from aiogram.types import InlineKeyboardButton, InlineKeyboardMarkup, WebAppInfo
|
|
from aiogram.utils.text_decorations import html_decoration as hd
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from sqlalchemy.orm import sessionmaker
|
|
|
|
from bot.infra.redis import redis_lock
|
|
from bot.middlewares.i18n import JsonI18n
|
|
from bot.services.panel_api_service import PanelApiService
|
|
from bot.services.subscription_service import SubscriptionService
|
|
from bot.utils.date_utils import month_start
|
|
from bot.utils.mini_app_url import subscription_mini_app_topup_url
|
|
from config.settings import Settings
|
|
from db.dal import subscription_dal, tariff_dal, user_dal
|
|
from db.models import Subscription
|
|
|
|
PREMIUM_WARNING_LEVEL_OFFSET = 1000
|
|
# Single warning per premium billing period when usage reached or exceeded the quota.
|
|
PREMIUM_WARNING_DEPLETED_LEVEL = PREMIUM_WARNING_LEVEL_OFFSET + 100
|
|
|
|
# Process active subscriptions in chunks and prefetch panel data concurrently
|
|
# to avoid an N+1 serial chain to the Remnawave panel each tick.
|
|
TARIFF_WORKER_BATCH_SIZE = 50
|
|
TARIFF_WORKER_PANEL_CONCURRENCY = 10
|
|
|
|
|
|
class TariffTrafficWorker:
|
|
def __init__(
|
|
self,
|
|
settings: Settings,
|
|
session_factory: sessionmaker,
|
|
panel_service: PanelApiService,
|
|
subscription_service: SubscriptionService,
|
|
bot: Optional[Bot] = None,
|
|
i18n: Optional[JsonI18n] = None,
|
|
):
|
|
self.settings = settings
|
|
self.session_factory = session_factory
|
|
self.panel_service = panel_service
|
|
self.subscription_service = subscription_service
|
|
self.bot = bot
|
|
self.i18n = i18n
|
|
self._stopped = asyncio.Event()
|
|
self._premium_nodes_cache = {}
|
|
self._premium_node_stats_tick_cache = {}
|
|
|
|
async def _user_lang(self, session: AsyncSession, user_id: int) -> str:
|
|
try:
|
|
row = await user_dal.get_user_by_id(session, user_id)
|
|
if row and getattr(row, "language_code", None):
|
|
code = str(row.language_code or "").strip()
|
|
if code:
|
|
return code
|
|
except Exception:
|
|
logging.exception("TariffTrafficWorker: failed to load user language for %s", user_id)
|
|
return self.settings.DEFAULT_LANGUAGE
|
|
|
|
def _usage_placeholders(self, used_bytes: int, limit_bytes: int) -> dict:
|
|
"""Formatted traffic stats for warning messages (HTML-safe quoted)."""
|
|
used_b = max(0, int(used_bytes or 0))
|
|
lim_b = max(0, int(limit_bytes or 0))
|
|
remaining_b = max(0, lim_b - used_b)
|
|
return {
|
|
"used": hd.quote(self._fmt_bytes(used_b)),
|
|
"remaining": hd.quote(self._fmt_bytes(remaining_b)),
|
|
"limit_total": hd.quote(self._fmt_bytes(lim_b)),
|
|
}
|
|
|
|
def _traffic_topup_markup(self, user_lang: str, kind: str) -> Optional[InlineKeyboardMarkup]:
|
|
if not self.bot:
|
|
return None
|
|
_ = lambda k, **kw: (
|
|
self.i18n.gettext(user_lang, k, **kw) if self.i18n else (lambda key, **_: key)
|
|
)
|
|
normalized = "premium" if str(kind or "").lower() == "premium" else "regular"
|
|
url = subscription_mini_app_topup_url(self.settings, normalized)
|
|
if normalized == "premium":
|
|
label_key = "traffic_warn_btn_topup_webapp_premium"
|
|
fallback_key = "traffic_warn_btn_topup_premium"
|
|
else:
|
|
label_key = "traffic_warn_btn_topup_webapp_regular"
|
|
fallback_key = "traffic_warn_btn_topup_regular"
|
|
# Mini App inside Telegram when SUBSCRIPTION_MINI_APP_URL is configured.
|
|
if url:
|
|
button = InlineKeyboardButton(text=_(label_key), web_app=WebAppInfo(url=url))
|
|
else:
|
|
button = InlineKeyboardButton(text=_(fallback_key), callback_data="tariff_topup:list")
|
|
return InlineKeyboardMarkup(inline_keyboard=[[button]])
|
|
|
|
async def run(self) -> None:
|
|
if not self.settings.tariffs_config:
|
|
return
|
|
while not self._stopped.is_set():
|
|
try:
|
|
async with redis_lock(
|
|
self.settings,
|
|
"tariff-traffic-worker",
|
|
ttl_seconds=self.settings.TARIFF_WORKER_LOCK_TTL_SECONDS,
|
|
) as acquired:
|
|
if not acquired:
|
|
logging.info("TariffTrafficWorker tick skipped: Redis lock is held")
|
|
else:
|
|
started = time.monotonic()
|
|
async with self.session_factory() as session:
|
|
await self.traffic_period_tick(session)
|
|
await session.commit()
|
|
async with self.session_factory() as session:
|
|
await self.legacy_throttle_recovery_tick(session)
|
|
await session.commit()
|
|
logging.info(
|
|
"metric worker_tick_duration_seconds=%.3f worker=tariff",
|
|
time.monotonic() - started,
|
|
)
|
|
except Exception:
|
|
logging.exception("TariffTrafficWorker tick failed")
|
|
try:
|
|
await asyncio.wait_for(
|
|
self._stopped.wait(),
|
|
timeout=self.settings.TARIFF_WORKER_TICK_SECONDS,
|
|
)
|
|
except asyncio.TimeoutError:
|
|
pass
|
|
|
|
def stop(self) -> None:
|
|
self._stopped.set()
|
|
|
|
async def traffic_period_tick(self, session: AsyncSession) -> None:
|
|
now = datetime.now(timezone.utc)
|
|
self._premium_node_stats_tick_cache = {}
|
|
warning_period_start = month_start(now)
|
|
result = await session.execute(
|
|
select(Subscription).where(
|
|
Subscription.is_active == True,
|
|
Subscription.end_date > now,
|
|
Subscription.tariff_key.is_not(None),
|
|
)
|
|
)
|
|
subs = list(result.scalars().all())
|
|
if not subs:
|
|
return
|
|
|
|
semaphore = asyncio.Semaphore(TARIFF_WORKER_PANEL_CONCURRENCY)
|
|
|
|
async def _fetch_panel(sub: Subscription) -> dict:
|
|
async with semaphore:
|
|
try:
|
|
data = await self.panel_service.get_user_by_uuid(
|
|
sub.panel_user_uuid, log_response=False
|
|
)
|
|
except Exception:
|
|
logging.exception(
|
|
"TariffTrafficWorker: failed to fetch panel user %s",
|
|
sub.panel_user_uuid,
|
|
)
|
|
return {}
|
|
return data or {}
|
|
|
|
for chunk_start in range(0, len(subs), TARIFF_WORKER_BATCH_SIZE):
|
|
chunk = subs[chunk_start : chunk_start + TARIFF_WORKER_BATCH_SIZE]
|
|
panel_payloads = await asyncio.gather(*(_fetch_panel(s) for s in chunk))
|
|
for sub, panel_data in zip(chunk, panel_payloads):
|
|
try:
|
|
tariff = self.settings.tariffs_config.require(sub.tariff_key)
|
|
except Exception:
|
|
continue
|
|
(
|
|
used,
|
|
limit,
|
|
panel_strategy,
|
|
) = self.subscription_service._extract_panel_traffic_details(panel_data)
|
|
panel_status = str(panel_data.get("status") or "").upper()
|
|
panel_username = (
|
|
panel_data.get("username") if isinstance(panel_data, dict) else None
|
|
)
|
|
if used is not None and used != sub.traffic_used_bytes:
|
|
sub.traffic_used_bytes = used
|
|
if limit is not None and limit != sub.traffic_limit_bytes:
|
|
sub.traffic_limit_bytes = limit
|
|
if panel_status and panel_status != (sub.status_from_panel or "").upper():
|
|
sub.status_from_panel = panel_status
|
|
|
|
if tariff.billing_model == "period":
|
|
await self._ensure_period_reset_strategy(sub, tariff, limit, panel_strategy)
|
|
await self._maybe_warn_or_throttle(
|
|
session,
|
|
sub,
|
|
tariff,
|
|
used,
|
|
limit,
|
|
warning_period_start=warning_period_start
|
|
if tariff.billing_model == "period"
|
|
else None,
|
|
)
|
|
|
|
await self._sync_premium_squad_limit(
|
|
session,
|
|
sub,
|
|
tariff,
|
|
now,
|
|
panel_username=panel_username,
|
|
panel_user_dict=panel_data,
|
|
)
|
|
|
|
async def _ensure_period_reset_strategy(
|
|
self,
|
|
sub: Subscription,
|
|
tariff,
|
|
limit: Optional[int],
|
|
panel_strategy: Optional[str],
|
|
) -> None:
|
|
if str(panel_strategy or "").upper() == "MONTH":
|
|
return
|
|
rb = int(getattr(sub, "regular_bonus_bytes", 0) or 0)
|
|
if bool(getattr(sub, "regular_unlimited_override", False)):
|
|
baseline = int(sub.tier_baseline_bytes or (tariff.monthly_bytes if tariff else 0) or 0)
|
|
traffic_limit_bytes = self.subscription_service._compute_main_traffic_limit_bytes(
|
|
tier_baseline_bytes=baseline,
|
|
topup_balance_bytes=int(sub.topup_balance_bytes or 0),
|
|
regular_bonus_bytes=rb,
|
|
regular_unlimited_override=True,
|
|
traffic_used_bytes=int(sub.traffic_used_bytes or 0),
|
|
)
|
|
else:
|
|
traffic_limit_bytes = int(
|
|
limit
|
|
or sub.traffic_limit_bytes
|
|
or (tariff.monthly_bytes + int(sub.topup_balance_bytes or 0) + rb)
|
|
)
|
|
payload = self.subscription_service._build_panel_update_payload(
|
|
panel_user_uuid=sub.panel_user_uuid,
|
|
expire_at=sub.end_date,
|
|
traffic_limit_bytes=traffic_limit_bytes,
|
|
traffic_limit_strategy="MONTH",
|
|
)
|
|
payload["activeInternalSquads"] = self.subscription_service._panel_squads_for_tariff(
|
|
tariff,
|
|
include_premium=not bool(getattr(sub, "premium_is_limited", False)),
|
|
)
|
|
await self.panel_service.update_user_details_on_panel(
|
|
sub.panel_user_uuid, payload, log_response=False
|
|
)
|
|
|
|
async def _maybe_warn_or_throttle(
|
|
self,
|
|
session: AsyncSession,
|
|
sub: Subscription,
|
|
tariff,
|
|
used: Optional[int],
|
|
limit: Optional[int],
|
|
*,
|
|
warning_period_start: Optional[datetime] = None,
|
|
) -> None:
|
|
if bool(getattr(sub, "regular_unlimited_override", False)):
|
|
return
|
|
used_val = int(used or sub.traffic_used_bytes or 0)
|
|
limit_val = int(limit or sub.traffic_limit_bytes or 0)
|
|
if limit_val <= 0:
|
|
return
|
|
ratio = used_val / limit_val
|
|
levels = list(getattr(self.settings, "tariff_traffic_warning_levels", [85, 90, 95]))
|
|
for level in levels:
|
|
threshold = level / 100
|
|
if ratio < threshold:
|
|
continue
|
|
warning = await tariff_dal.get_warning(
|
|
session,
|
|
subscription_id=sub.subscription_id,
|
|
period_start_at=warning_period_start if tariff.billing_model == "period" else None,
|
|
level=level,
|
|
traffic_limit_bytes=limit_val if tariff.billing_model == "traffic" else None,
|
|
)
|
|
if warning:
|
|
continue
|
|
await tariff_dal.create_warning(
|
|
session,
|
|
subscription_id=sub.subscription_id,
|
|
period_start_at=warning_period_start if tariff.billing_model == "period" else None,
|
|
level=level,
|
|
traffic_limit_bytes=limit_val if tariff.billing_model == "traffic" else None,
|
|
)
|
|
if self.bot:
|
|
try:
|
|
user_lang = await self._user_lang(session, sub.user_id)
|
|
_ = (
|
|
(lambda k, **kw: self.i18n.gettext(user_lang, k, **kw))
|
|
if self.i18n
|
|
else (lambda k, **kw: k)
|
|
)
|
|
left_pct = max(0, 100 - level)
|
|
tariff_name = hd.quote(str(tariff.name(user_lang)))
|
|
usage = self._usage_placeholders(used_val, limit_val)
|
|
if level < 100:
|
|
text = _(
|
|
"traffic_warning_regular_almost",
|
|
tariff_name=tariff_name,
|
|
left_pct=left_pct,
|
|
**usage,
|
|
)
|
|
else:
|
|
text = _(
|
|
"traffic_warning_regular_depleted",
|
|
tariff_name=tariff_name,
|
|
**usage,
|
|
)
|
|
markup = self._traffic_topup_markup(user_lang, "regular")
|
|
await self.bot.send_message(
|
|
sub.user_id,
|
|
text,
|
|
reply_markup=markup,
|
|
parse_mode="HTML",
|
|
)
|
|
except Exception:
|
|
logging.exception("Failed to send traffic warning to user %s", sub.user_id)
|
|
if ratio >= 1.0 and not sub.is_throttled:
|
|
logging.info(
|
|
"Tariff traffic limit reached for user %s subscription %s. "
|
|
"Leaving access control to Remnawave status handling.",
|
|
sub.user_id,
|
|
sub.subscription_id,
|
|
)
|
|
|
|
async def _sync_premium_squad_limit(
|
|
self,
|
|
session: AsyncSession,
|
|
sub: Subscription,
|
|
tariff,
|
|
now: datetime,
|
|
*,
|
|
panel_username: Optional[str] = None,
|
|
panel_user_dict: Optional[dict] = None,
|
|
) -> None:
|
|
if not getattr(tariff, "premium_squad_uuids", None):
|
|
if (
|
|
any(
|
|
int(value or 0) > 0
|
|
for value in (
|
|
sub.premium_baseline_bytes,
|
|
sub.premium_topup_balance_bytes,
|
|
sub.premium_used_bytes,
|
|
)
|
|
)
|
|
or sub.premium_is_limited
|
|
):
|
|
sub.premium_baseline_bytes = 0
|
|
sub.premium_topup_balance_bytes = 0
|
|
sub.premium_used_bytes = 0
|
|
sub.premium_is_limited = False
|
|
return
|
|
|
|
premium_period_start = month_start(now)
|
|
same_period = bool(getattr(sub, "premium_period_start_at", None) == premium_period_start)
|
|
premium_baseline = int(tariff.premium_monthly_bytes or 0)
|
|
premium_topup_balance = int(sub.premium_topup_balance_bytes or 0)
|
|
premium_topup_used = (
|
|
int(getattr(sub, "premium_topup_used_bytes", 0) or 0) if same_period else 0
|
|
)
|
|
# Admin-side overrides for free gifted premium traffic.
|
|
premium_unlimited_override = bool(getattr(sub, "premium_unlimited_override", False))
|
|
premium_bonus = max(0, int(getattr(sub, "premium_bonus_bytes", 0) or 0))
|
|
premium_limit = (
|
|
premium_baseline + premium_topup_balance + premium_topup_used + premium_bonus
|
|
)
|
|
if premium_limit <= 0 and not premium_unlimited_override:
|
|
return
|
|
|
|
node_uuids = await self._premium_node_uuids_for_tariff(tariff)
|
|
if not node_uuids:
|
|
logging.warning("Premium squads for tariff %s have no accessible nodes", tariff.key)
|
|
return
|
|
|
|
start_date = now.date().replace(day=1).isoformat()
|
|
end_date = now.date().isoformat()
|
|
premium_used = await self._premium_usage_for_user(
|
|
sub.panel_user_uuid,
|
|
node_uuids,
|
|
start_date,
|
|
end_date,
|
|
panel_username=panel_username,
|
|
)
|
|
if premium_used is None:
|
|
return
|
|
|
|
# Consume paid top-up balance only for overflow beyond baseline+bonus.
|
|
# Admin-granted bonus is "spent" against usage first along with baseline,
|
|
# so the user's paid top-up survives longer.
|
|
free_quota = premium_baseline + premium_bonus
|
|
overflow = max(0, int(premium_used) - free_quota)
|
|
delta_overflow = max(0, overflow - premium_topup_used)
|
|
consume_from_topup = min(premium_topup_balance, delta_overflow)
|
|
if consume_from_topup > 0:
|
|
premium_topup_balance -= consume_from_topup
|
|
premium_topup_used += consume_from_topup
|
|
premium_limit = (
|
|
premium_baseline + premium_topup_balance + premium_topup_used + premium_bonus
|
|
)
|
|
|
|
if premium_unlimited_override:
|
|
should_limit = False
|
|
else:
|
|
should_limit = premium_used >= premium_limit
|
|
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_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
|
|
sub.premium_used_bytes = int(premium_used)
|
|
sub.premium_is_limited = bool(should_limit)
|
|
sub.premium_period_start_at = premium_period_start
|
|
if not premium_unlimited_override:
|
|
await self._maybe_warn_premium_squad_limit(
|
|
session,
|
|
sub,
|
|
tariff,
|
|
premium_used,
|
|
premium_limit,
|
|
premium_period_start,
|
|
)
|
|
if not panel_needs_update:
|
|
return
|
|
|
|
squads = desired_squads
|
|
await self.panel_service.update_user_details_on_panel(
|
|
sub.panel_user_uuid,
|
|
{"uuid": sub.panel_user_uuid, "activeInternalSquads": squads},
|
|
log_response=False,
|
|
)
|
|
logging.info(
|
|
"Premium squad access %s for user %s tariff %s: %s/%s bytes",
|
|
"limited" if should_limit else "restored",
|
|
sub.user_id,
|
|
tariff.key,
|
|
premium_used,
|
|
premium_limit,
|
|
)
|
|
|
|
@staticmethod
|
|
def _internal_squad_uuid_set(raw) -> set[str]:
|
|
if not isinstance(raw, list):
|
|
return set()
|
|
out: set[str] = set()
|
|
for item in raw:
|
|
if isinstance(item, dict):
|
|
u = item.get("uuid") or item.get("internalSquadUuid") or item.get("squadUuid")
|
|
if u:
|
|
out.add(str(u))
|
|
elif item:
|
|
out.add(str(item))
|
|
return out
|
|
|
|
@staticmethod
|
|
def _fmt_bytes(value: int) -> str:
|
|
size = float(max(0, int(value or 0)))
|
|
for unit in ("B", "KB", "MB", "GB", "TB"):
|
|
if size < 1024 or unit == "TB":
|
|
return f"{size:.1f} {unit}" if unit != "B" else f"{int(size)} B"
|
|
size /= 1024
|
|
return f"{size:.1f} TB"
|
|
|
|
async def _maybe_warn_premium_squad_limit(
|
|
self,
|
|
session: AsyncSession,
|
|
sub: Subscription,
|
|
tariff,
|
|
used: int,
|
|
limit: int,
|
|
period_start_at: datetime,
|
|
) -> None:
|
|
if limit <= 0:
|
|
return
|
|
used_val = int(used or 0)
|
|
limit_val = int(limit)
|
|
ratio = used_val / limit_val
|
|
levels = list(getattr(self.settings, "tariff_traffic_warning_levels", [85, 90, 95]))
|
|
|
|
# Fully exhausted or over quota — one message per period (same idea as regular traffic at 100%). # noqa: E501
|
|
if ratio >= 1.0:
|
|
depleted_existing = await tariff_dal.get_warning(
|
|
session,
|
|
subscription_id=sub.subscription_id,
|
|
period_start_at=period_start_at,
|
|
level=PREMIUM_WARNING_DEPLETED_LEVEL,
|
|
)
|
|
if depleted_existing:
|
|
return
|
|
await tariff_dal.create_warning(
|
|
session,
|
|
subscription_id=sub.subscription_id,
|
|
period_start_at=period_start_at,
|
|
level=PREMIUM_WARNING_DEPLETED_LEVEL,
|
|
traffic_limit_bytes=None,
|
|
)
|
|
if self.bot:
|
|
try:
|
|
user_lang = await self._user_lang(session, sub.user_id)
|
|
_ = (
|
|
(lambda k, **kw: self.i18n.gettext(user_lang, k, **kw))
|
|
if self.i18n
|
|
else (lambda k, **kw: k)
|
|
)
|
|
access = await self.subscription_service.premium_access_for_tariff(tariff)
|
|
labels = access.get("node_labels") or access.get("squad_labels") or []
|
|
if labels:
|
|
visible = [hd.quote(str(x)) for x in labels[:8]]
|
|
servers = "\n".join(f"• {label}" for label in visible)
|
|
if len(labels) > len(visible):
|
|
more = len(labels) - len(visible)
|
|
servers += "\n" + _("traffic_warning_premium_servers_more", count=more)
|
|
else:
|
|
servers = _("traffic_warning_premium_generic_servers")
|
|
usage = self._usage_placeholders(used_val, limit_val)
|
|
text = _(
|
|
"traffic_warning_premium_depleted",
|
|
tariff_name=hd.quote(str(tariff.name(user_lang))),
|
|
servers=servers,
|
|
**usage,
|
|
)
|
|
markup = self._traffic_topup_markup(user_lang, "premium")
|
|
await self.bot.send_message(
|
|
sub.user_id,
|
|
text,
|
|
reply_markup=markup,
|
|
parse_mode="HTML",
|
|
)
|
|
except Exception:
|
|
logging.exception(
|
|
"Failed to send premium traffic depleted warning to user %s", sub.user_id
|
|
)
|
|
return
|
|
|
|
for level in levels:
|
|
if level >= 100:
|
|
continue
|
|
if ratio < level / 100:
|
|
continue
|
|
storage_level = PREMIUM_WARNING_LEVEL_OFFSET + int(level)
|
|
warning = await tariff_dal.get_warning(
|
|
session,
|
|
subscription_id=sub.subscription_id,
|
|
period_start_at=period_start_at,
|
|
level=storage_level,
|
|
)
|
|
if warning:
|
|
continue
|
|
await tariff_dal.create_warning(
|
|
session,
|
|
subscription_id=sub.subscription_id,
|
|
period_start_at=period_start_at,
|
|
level=storage_level,
|
|
traffic_limit_bytes=None,
|
|
)
|
|
if not self.bot:
|
|
continue
|
|
try:
|
|
user_lang = await self._user_lang(session, sub.user_id)
|
|
_ = (
|
|
(lambda k, **kw: self.i18n.gettext(user_lang, k, **kw))
|
|
if self.i18n
|
|
else (lambda k, **kw: k)
|
|
)
|
|
access = await self.subscription_service.premium_access_for_tariff(tariff)
|
|
labels = access.get("node_labels") or access.get("squad_labels") or []
|
|
if labels:
|
|
visible = [hd.quote(str(x)) for x in labels[:8]]
|
|
servers = "\n".join(f"• {label}" for label in visible)
|
|
if len(labels) > len(visible):
|
|
more = len(labels) - len(visible)
|
|
servers += "\n" + _("traffic_warning_premium_servers_more", count=more)
|
|
else:
|
|
servers = _("traffic_warning_premium_generic_servers")
|
|
left_pct = max(0, 100 - int(level))
|
|
usage = self._usage_placeholders(used_val, limit_val)
|
|
text = _(
|
|
"traffic_warning_premium_almost",
|
|
tariff_name=hd.quote(str(tariff.name(user_lang))),
|
|
left_pct=left_pct,
|
|
servers=servers,
|
|
**usage,
|
|
)
|
|
markup = self._traffic_topup_markup(user_lang, "premium")
|
|
await self.bot.send_message(
|
|
sub.user_id,
|
|
text,
|
|
reply_markup=markup,
|
|
parse_mode="HTML",
|
|
)
|
|
except Exception:
|
|
logging.exception("Failed to send premium traffic warning to user %s", sub.user_id)
|
|
|
|
async def _premium_node_uuids_for_tariff(self, tariff) -> list[str]:
|
|
cache_key = tuple(sorted(tariff.premium_squad_uuids or []))
|
|
cached = self._premium_nodes_cache.get(cache_key)
|
|
now_ts = datetime.now(timezone.utc).timestamp()
|
|
if cached and now_ts - cached["ts"] < 600:
|
|
return list(cached["nodes"])
|
|
|
|
nodes: list[str] = []
|
|
for squad_uuid in tariff.premium_squad_uuids or []:
|
|
accessible = (
|
|
await self.panel_service.get_internal_squad_accessible_nodes(squad_uuid) or []
|
|
)
|
|
for node in accessible:
|
|
if not isinstance(node, dict):
|
|
continue
|
|
node_uuid = node.get("uuid") or node.get("nodeUuid") or node.get("node_uuid")
|
|
if node_uuid:
|
|
nodes.append(str(node_uuid))
|
|
deduped = list(dict.fromkeys(nodes))
|
|
self._premium_nodes_cache[cache_key] = {"ts": now_ts, "nodes": deduped}
|
|
return deduped
|
|
|
|
async def _premium_usage_for_user(
|
|
self,
|
|
user_uuid: str,
|
|
node_uuids: list[str],
|
|
start_date: str,
|
|
end_date: str,
|
|
*,
|
|
panel_username: Optional[str] = None,
|
|
) -> Optional[int]:
|
|
total = 0
|
|
found = False
|
|
username = (panel_username or "").strip() or None
|
|
for node_uuid in node_uuids:
|
|
stats_cache_key = (node_uuid, start_date, end_date)
|
|
if stats_cache_key not in self._premium_node_stats_tick_cache:
|
|
self._premium_node_stats_tick_cache[stats_cache_key] = (
|
|
await self.panel_service.get_node_users_bandwidth_stats(
|
|
node_uuid,
|
|
start=start_date,
|
|
end=end_date,
|
|
)
|
|
)
|
|
stats = self._premium_node_stats_tick_cache.get(stats_cache_key)
|
|
if not stats:
|
|
continue
|
|
entries = stats.get("topUsers") or stats.get("usersStats") or stats.get("users") or []
|
|
if not isinstance(entries, list):
|
|
continue
|
|
for entry in entries:
|
|
if not isinstance(entry, dict):
|
|
continue
|
|
user_obj = entry.get("user") if isinstance(entry.get("user"), dict) else {}
|
|
entry_uuid = (
|
|
user_obj.get("uuid")
|
|
or entry.get("userUuid")
|
|
or entry.get("uuid")
|
|
or entry.get("user_uuid")
|
|
)
|
|
entry_username = (
|
|
user_obj.get("username") or entry.get("username") or entry.get("userUsername")
|
|
)
|
|
# Remnawave's /bandwidth-stats/nodes/{uuid}/users response
|
|
# currently exposes only {color, username, total}; match by
|
|
# username first, fall back to UUID if a future version
|
|
# adds it back.
|
|
matched = False
|
|
if entry_uuid and entry_uuid == user_uuid:
|
|
matched = True
|
|
elif username and entry_username and entry_username == username:
|
|
matched = True
|
|
if not matched:
|
|
continue
|
|
value = entry.get("total")
|
|
if value is None:
|
|
value = int(entry.get("download", 0) or 0) + int(entry.get("upload", 0) or 0)
|
|
total += int(value or 0)
|
|
found = True
|
|
if len(node_uuids) > 1:
|
|
await asyncio.sleep(0.1)
|
|
return total if found else 0
|
|
|
|
async def legacy_throttle_recovery_tick(self, session: AsyncSession) -> None:
|
|
"""Recover subscriptions throttled by older bot versions.
|
|
|
|
Current Remnawave versions enforce exhausted user traffic limits by
|
|
switching the user status to LIMITED, so new ticks must not remove users
|
|
from Internal Squads.
|
|
"""
|
|
result = await session.execute(
|
|
select(Subscription).where(
|
|
Subscription.is_active == True,
|
|
Subscription.is_throttled == True,
|
|
)
|
|
)
|
|
for sub in result.scalars().all():
|
|
try:
|
|
tariff = self.settings.tariffs_config.require(sub.tariff_key)
|
|
except Exception:
|
|
continue
|
|
if int(sub.traffic_limit_bytes or 0) <= int(sub.traffic_used_bytes or 0):
|
|
continue
|
|
for squad_uuid in tariff.squad_uuids:
|
|
await self.panel_service.add_users_to_internal_squad(
|
|
squad_uuid, [sub.panel_user_uuid]
|
|
)
|
|
await subscription_dal.update_subscription(
|
|
session,
|
|
sub.subscription_id,
|
|
{"is_throttled": False, "status_from_panel": "ACTIVE"},
|
|
)
|