Files

609 lines
23 KiB
Python

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 <username> 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 <code>{username}</code> 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: <code>{username}</code>\n"
f"Expire (UTC): <code>{expire_human}</code>\n"
f"Monthly traffic: <b>{bytes_to_gb(traffic)}</b>\n"
f"Price multiplier: <b>{multiplier}</b>"
)
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: <code>{username}</code>",
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): <code>{info['expire_human']}</code>\n"
f"Subscription URL:\n<code>{info['subscription_url']}</code>"
)
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 <code>{message.from_user.id}</code> ({message.from_user.full_name}) "
f"{info['action']} subscription for {months} month(s) and paid ⭐️{payment.total_amount}.\n"
f"Marzban username: <code>{info['username']}</code>\n"
f"Expires: <code>{info['expire_human']}</code>"
)
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(
"📊 <b>Stats</b>\n"
f"Payments: <b>{pay_count}</b>\n"
f"Stars earned: <b>⭐️{stars_total}</b>\n"
f"Sold months: <b>{months_total}</b>\n"
f"Marzban users: <b>{user_count}</b>\n"
f"Traffic used: <b>{bytes_to_gb(used)}</b>\n"
f"Traffic quota: <b>{bytes_to_gb(quota)}</b>"
)
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 <username|telegram_id>")
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 <YYYY-MM-DD or ISO datetime>")
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 <GB>")
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 <float>")
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())