import asyncio import calendar import logging import os from dataclasses import dataclass from datetime import datetime, timezone from decimal import Decimal, ROUND_HALF_UP, InvalidOperation from typing import Any import aiohttp import aiosqlite from aiogram import Bot, Dispatcher, F from aiogram.enums import ParseMode from aiogram.filters import Command from aiogram.types import ( CallbackQuery, LabeledPrice, Message, PreCheckoutQuery, ) from aiogram.utils.keyboard import InlineKeyboardBuilder from aiogram.client.default import DefaultBotProperties from aiogram.client.session.aiohttp import AiohttpSession logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) GB = 1024**3 PLANS = { 1: 100, 3: 270, 6: 500, } @dataclass class Settings: bot_token: str marzban_url: str marzban_username: str marzban_password: str admin_ids: set[int] proxy: str db_path: str = "stats.db" @classmethod def from_env(cls) -> "Settings": admin_raw = os.getenv("ADMIN_IDS", "") admin_ids = { int(part.strip()) for part in admin_raw.split(",") if part.strip().isdigit() } return cls( bot_token=os.environ["BOT_TOKEN"], marzban_url=os.environ["MARZBAN_URL"].rstrip("/"), marzban_username=os.environ["MARZBAN_USERNAME"], marzban_password=os.environ["MARZBAN_PASSWORD"], admin_ids=admin_ids, proxy=os.getenv("HTTP_PROXY", None), db_path=os.getenv("DB_PATH", "stats.db"), ) class MarzbanClient: def __init__(self, base_url: str, username: str, password: str): self.base_url = base_url self.username = username self.password = password self._session: aiohttp.ClientSession | None = None self._token: str | None = None async def __aenter__(self) -> "MarzbanClient": self._session = aiohttp.ClientSession(base_url=self.base_url) await self._authenticate() return self async def __aexit__(self, exc_type, exc, tb): if self._session: await self._session.close() async def _authenticate(self) -> None: assert self._session resp = await self._session.post( "/api/admin/token", data={"username": self.username, "password": self.password}, ) resp.raise_for_status() data = await resp.json() self._token = data["access_token"] async def _request(self, method: str, path: str, **kwargs) -> Any: assert self._session if self._token is None: await self._authenticate() headers = kwargs.pop("headers", {}) headers["Authorization"] = f"Bearer {self._token}" resp = await self._session.request(method, path, headers=headers, **kwargs) if resp.status == 401: await self._authenticate() headers["Authorization"] = f"Bearer {self._token}" resp = await self._session.request(method, path, headers=headers, **kwargs) if resp.status == 404: return None resp.raise_for_status() if resp.content_type.startswith("application/json"): return await resp.json() return await resp.text() async def get_user(self, username: str) -> dict[str, Any] | None: return await self._request("GET", f"/api/user/{username}") async def create_user(self, payload: dict[str, Any]) -> dict[str, Any]: return await self._request("POST", "/api/user", json=payload) async def update_user(self, username: str, payload: dict[str, Any]) -> dict[str, Any]: return await self._request("PUT", f"/api/user/{username}", json=payload) async def list_all_users(self) -> list[dict[str, Any]]: users: list[dict[str, Any]] = [] offset = 0 limit = 100 while True: chunk = await self._request("GET", f"/api/users?offset={offset}&limit={limit}") if not chunk: break batch = chunk.get("users", []) if isinstance(chunk, dict) else chunk if not batch: break users.extend(batch) if len(batch) < limit: break offset += len(batch) return users def add_months(dt: datetime, months: int) -> datetime: month = dt.month - 1 + months year = dt.year + month // 12 month = month % 12 + 1 day = min(dt.day, calendar.monthrange(year, month)[1]) return dt.replace(year=year, month=month, day=day) def bytes_to_gb(value: int | float) -> str: return f"{Decimal(value) / Decimal(GB):.2f} GB" class VpnBot: def __init__(self, settings: Settings): self.settings = settings self.session = AiohttpSession(proxy=settings.proxy) self.bot = Bot(settings.bot_token, default=DefaultBotProperties(parse_mode=ParseMode.HTML), session=self.session) self.dp = Dispatcher() self.marzban = MarzbanClient( settings.marzban_url, settings.marzban_username, settings.marzban_password, ) self.db: aiosqlite.Connection | None = None self.locks: dict[int, asyncio.Lock] = {} self.admin_selected_users: dict[int, str] = {} self._register_handlers() async def setup(self) -> None: self.db = await aiosqlite.connect(self.settings.db_path) await self.db.execute( """ CREATE TABLE IF NOT EXISTS payments ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, tg_username TEXT, marzban_username TEXT NOT NULL, months INTEGER NOT NULL, stars INTEGER NOT NULL, action TEXT NOT NULL, created_at TEXT NOT NULL ) """ ) await self.db.commit() async def run(self) -> None: async with self.marzban: await self.setup() await self.dp.start_polling(self.bot) def _register_handlers(self) -> None: self.dp.message.register(self.cmd_start, Command("start")) self.dp.message.register(self.cmd_stats, Command("stats")) self.dp.message.register(self.cmd_select_user, Command("select_user")) self.dp.message.register(self.cmd_selected_user, Command("selected_user")) self.dp.message.register(self.cmd_set_expire, Command("set_expire")) self.dp.message.register(self.cmd_set_traffic, Command("set_traffic")) self.dp.message.register(self.cmd_set_multiplier, Command("set_multiplier")) self.dp.callback_query.register(self.choose_plan, F.data.startswith("buy:")) self.dp.pre_checkout_query.register(self.pre_checkout) self.dp.message.register(self.successful_payment, F.successful_payment) @staticmethod def _extract_price_multiplier(note: str | None) -> float: if not note: return 1.0 for line in note.splitlines(): if line.startswith("price_multiplier="): try: value = float(line.split("=", 1)[1].strip()) return value if value > 0 else 1.0 except ValueError: return 1.0 return 1.0 @staticmethod def _set_price_multiplier_note(note: str | None, multiplier: float) -> str: lines = [line for line in (note or "").splitlines() if line and not line.startswith("price_multiplier=")] lines.append(f"price_multiplier={multiplier}") return "\n".join(lines) @staticmethod def _format_new_user_note(first_name: str | None, last_name: str | None, phone: str | None) -> str: parts: list[str] = [] if first_name: parts.append(f"First name: {first_name}") if last_name: parts.append(f"Last name: {last_name}") if phone: parts.append(f"Phone: {phone}") return "\n".join(parts) @staticmethod def _stars_with_multiplier(base_stars: int, multiplier: float) -> int: value = Decimal(base_stars) * Decimal(str(multiplier)) rounded = value.quantize(Decimal("1"), rounding=ROUND_HALF_UP) return max(1, int(rounded)) @staticmethod def _parse_command_arg(message: Message) -> str: text = (message.text or "").strip() parts = text.split(maxsplit=1) return parts[1].strip() if len(parts) > 1 else "" @staticmethod def _parse_utc_datetime(value: str) -> datetime: parsed = datetime.fromisoformat(value.strip()) if parsed.tzinfo is None: return parsed.replace(tzinfo=timezone.utc) return parsed.astimezone(timezone.utc) @staticmethod def _normalize_username(value: str) -> str: cleaned = value.strip() if cleaned.isdigit(): return f"tg{cleaned}" return cleaned async def _selected_user_or_reply(self, message: Message) -> str | None: username = self.admin_selected_users.get(message.from_user.id) if not username: await message.answer("No selected user. Use /select_user first.") return None return username async def _get_admin_user(self, message: Message, username: str) -> dict[str, Any] | None: user = await self.marzban.get_user(username) if not user: await message.answer(f"User {username} not found in Marzban.") return None return user async def _update_admin_user(self, username: str, current: dict[str, Any], **changes: Any) -> dict[str, Any]: payload = { "username": current.get("username", username), "status": current.get("status", "active"), "data_limit": int(current.get("data_limit") or 0), "data_limit_reset_strategy": current.get("data_limit_reset_strategy") or "month", "expire": int(current.get("expire") or 0), "proxies": current.get("proxies") or {}, "inbounds": current.get("inbounds") or {}, "note": current.get("note") or "", } payload.update(changes) return await self.marzban.update_user(username, payload) async def _send_selected_user_info(self, message: Message, username: str, user: dict[str, Any]) -> None: expire_ts = int(user.get("expire") or 0) expire_human = ( datetime.fromtimestamp(expire_ts, tz=timezone.utc).strftime("%Y-%m-%d %H:%M:%S") if expire_ts else "not set" ) traffic = int(user.get("data_limit") or 0) multiplier = self._extract_price_multiplier(user.get("note")) await message.answer( "Selected user:\n" f"Username: {username}\n" f"Expire (UTC): {expire_human}\n" f"Monthly traffic: {bytes_to_gb(traffic)}\n" f"Price multiplier: {multiplier}" ) async def cmd_start(self, message: Message) -> None: username = f"tg{message.from_user.id}" current = await self.marzban.get_user(username) multiplier = self._extract_price_multiplier((current or {}).get("note")) kb = InlineKeyboardBuilder() plan_lines: list[str] = [] for months, base_stars in PLANS.items(): month_label = "month" if months == 1 else "months" stars = self._stars_with_multiplier(base_stars, multiplier) kb.button(text=f"{months} {month_label} — ⭐️{stars}", callback_data=f"buy:{months}") plan_lines.append(f"• {months} {month_label} — {stars} stars") kb.adjust(1) await message.answer( "VPN subscription plans:\n" + "\n".join(plan_lines) + f"\n\nUsername in VPN panel: {username}", reply_markup=kb.as_markup(), ) async def choose_plan(self, query: CallbackQuery) -> None: if not query.data or not query.message: return months = int(query.data.split(":", 1)[1]) stars = PLANS[months] username = f"tg{query.from_user.id}" current = await self.marzban.get_user(username) multiplier = self._extract_price_multiplier((current or {}).get("note")) stars = self._stars_with_multiplier(stars, multiplier) await self.bot.send_invoice( chat_id=query.message.chat.id, title=f"VPN {months} month(s)", description="Buy or extend VPN subscription", payload=f"vpn:{months}", provider_token="", currency="XTR", prices=[LabeledPrice(label=f"{months} month(s)", amount=stars)], ) await query.answer() async def pre_checkout(self, pre_checkout_query: PreCheckoutQuery) -> None: payload = pre_checkout_query.invoice_payload ok = payload.startswith("vpn:") await self.bot.answer_pre_checkout_query( pre_checkout_query.id, ok=ok, error_message="Invalid plan" if not ok else None, ) async def successful_payment(self, message: Message) -> None: payment = message.successful_payment if not payment: return try: months = int(payment.invoice_payload.split(":", 1)[1]) except (IndexError, ValueError): await message.answer("Payment received, but payload is invalid. Contact admin.") return lock = self.locks.setdefault(message.from_user.id, asyncio.Lock()) order_info = getattr(payment, "order_info", None) phone = getattr(order_info, "phone_number", None) if order_info else None async with lock: info = await self.buy_or_extend(message.from_user, months, phone) await message.answer( "✅ Subscription updated!\n" f"Plan: {months} month(s), paid: ⭐️{payment.total_amount}\n" f"Expires at (UTC): {info['expire_human']}\n" f"Subscription URL:\n{info['subscription_url']}" ) await self.save_payment( user_id=message.from_user.id, tg_username=message.from_user.username, marzban_username=info["username"], months=months, stars=payment.total_amount, action=info["action"], ) await self.notify_admins( f"💸 User {message.from_user.id} ({message.from_user.full_name}) " f"{info['action']} subscription for {months} month(s) and paid ⭐️{payment.total_amount}.\n" f"Marzban username: {info['username']}\n" f"Expires: {info['expire_human']}" ) async def buy_or_extend(self, user: Any, months: int, phone: str | None = None) -> dict[str, str]: username = f"tg{user.id}" now = datetime.now(tz=timezone.utc) expire_from = now action = "created" current = await self.marzban.get_user(username) if current: current_expire = int(current.get("expire") or 0) if current_expire > int(now.timestamp()): expire_from = datetime.fromtimestamp(current_expire, tz=timezone.utc) action = "extended" expire_at = add_months(expire_from, months) data_limit = int(current.get("data_limit", 50 * GB)) if current else 50 * GB data_limit_reset_strategy = (current.get("data_limit_reset_strategy") if current else None) or "month" note = current.get("note") if current else self._format_new_user_note( user.first_name, user.last_name, phone ) payload = { "username": username, "status": "active", "data_limit": data_limit, "data_limit_reset_strategy": data_limit_reset_strategy, "expire": int(expire_at.timestamp()), "proxies": (current.get("proxies") if current else {"vless": {"flow": "xtls-rprx-vision"}}), "inbounds": (current.get("inbounds") if current else {"vless": ["VLESS TCP REALITY"]}), "note": note, } if current: updated = await self.marzban.update_user(username, payload) else: updated = await self.marzban.create_user(payload) subscription_url = updated.get("subscription_url") or "(not returned by API)" return { "username": username, "action": action, "subscription_url": subscription_url, "expire_human": expire_at.strftime("%Y-%m-%d %H:%M:%S"), } async def save_payment( self, user_id: int, tg_username: str | None, marzban_username: str, months: int, stars: int, action: str, ) -> None: assert self.db await self.db.execute( """ INSERT INTO payments (user_id, tg_username, marzban_username, months, stars, action, created_at) VALUES (?, ?, ?, ?, ?, ?, ?) """, ( user_id, tg_username, marzban_username, months, stars, action, datetime.now(tz=timezone.utc).isoformat(), ), ) await self.db.commit() async def notify_admins(self, text: str) -> None: for admin_id in self.settings.admin_ids: try: await self.bot.send_message(admin_id, text) except Exception as exc: logger.warning("Failed to notify admin %s: %s", admin_id, exc) async def cmd_stats(self, message: Message) -> None: if message.from_user.id not in self.settings.admin_ids: await message.answer("Admins only") return assert self.db cursor = await self.db.execute( "SELECT COUNT(*), COALESCE(SUM(stars), 0), COALESCE(SUM(months), 0) FROM payments" ) pay_count, stars_total, months_total = await cursor.fetchone() users = await self.marzban.list_all_users() user_count = len(users) used = sum(int(u.get("used_traffic") or 0) for u in users) quota = sum(int(u.get("data_limit") or 0) for u in users) await message.answer( "📊 Stats\n" f"Payments: {pay_count}\n" f"Stars earned: ⭐️{stars_total}\n" f"Sold months: {months_total}\n" f"Marzban users: {user_count}\n" f"Traffic used: {bytes_to_gb(used)}\n" f"Traffic quota: {bytes_to_gb(quota)}" ) async def cmd_select_user(self, message: Message) -> None: if message.from_user.id not in self.settings.admin_ids: await message.answer("Admins only") return arg = self._parse_command_arg(message) if not arg: await message.answer("Usage: /select_user ") return username = self._normalize_username(arg) user = await self._get_admin_user(message, username) if not user: return self.admin_selected_users[message.from_user.id] = username await self._send_selected_user_info(message, username, user) async def cmd_selected_user(self, message: Message) -> None: if message.from_user.id not in self.settings.admin_ids: await message.answer("Admins only") return username = await self._selected_user_or_reply(message) if not username: return user = await self._get_admin_user(message, username) if not user: return await self._send_selected_user_info(message, username, user) async def cmd_set_expire(self, message: Message) -> None: if message.from_user.id not in self.settings.admin_ids: await message.answer("Admins only") return arg = self._parse_command_arg(message) if not arg: await message.answer("Usage: /set_expire ") return username = await self._selected_user_or_reply(message) if not username: return try: expire_dt = self._parse_utc_datetime(arg) except ValueError: await message.answer("Invalid datetime format. Example: 2026-12-31 or 2026-12-31T15:30:00") return current = await self._get_admin_user(message, username) if not current: return updated = await self._update_admin_user(username, current, expire=int(expire_dt.timestamp())) await self._send_selected_user_info(message, username, updated) async def cmd_set_traffic(self, message: Message) -> None: if message.from_user.id not in self.settings.admin_ids: await message.answer("Admins only") return arg = self._parse_command_arg(message) if not arg: await message.answer("Usage: /set_traffic ") return username = await self._selected_user_or_reply(message) if not username: return try: gb_value = Decimal(arg) except (ValueError, InvalidOperation): await message.answer("Invalid value. Example: /set_traffic 75") return if gb_value <= 0: await message.answer("Traffic must be greater than 0.") return data_limit = int((gb_value * Decimal(GB)).quantize(Decimal("1"), rounding=ROUND_HALF_UP)) current = await self._get_admin_user(message, username) if not current: return updated = await self._update_admin_user(username, current, data_limit=data_limit, data_limit_reset_strategy="month") await self._send_selected_user_info(message, username, updated) async def cmd_set_multiplier(self, message: Message) -> None: if message.from_user.id not in self.settings.admin_ids: await message.answer("Admins only") return arg = self._parse_command_arg(message) if not arg: await message.answer("Usage: /set_multiplier ") return username = await self._selected_user_or_reply(message) if not username: return try: multiplier = float(arg) except ValueError: await message.answer("Invalid multiplier. Example: /set_multiplier 1.25") return if multiplier <= 0: await message.answer("Multiplier must be greater than 0.") return current = await self._get_admin_user(message, username) if not current: return updated = await self._update_admin_user( username, current, note=self._set_price_multiplier_note(current.get("note"), multiplier), ) await self._send_selected_user_info(message, username, updated) async def main() -> None: settings = Settings.from_env() bot = VpnBot(settings) await bot.run() if __name__ == "__main__": asyncio.run(main())