Merge pull request #152 from JustYay/main

Optimize and fix refferals
This commit is contained in:
Vladislav Lisitsyn
2025-02-16 10:20:54 +03:00
committed by GitHub
3 changed files with 112 additions and 53 deletions
+24 -18
View File
@@ -537,18 +537,18 @@ async def get_balance(tg_id: int) -> float:
await conn.close()
async def update_balance(tg_id: int, amount: float, session: Any = None, is_admin: bool = False):
async def update_balance(
tg_id: int,
amount: float,
session: Any = None,
is_admin: bool = False,
skip_referral: bool = False, # <- флаг "пропустить реферальное начисление"
skip_cashback: bool = False # <- флаг "пропустить кэшбэк"
):
"""
Обновляет баланс пользователя в базе данных.
Кэшбек применяется только для положительных сумм, если пополнение НЕ через админку.
Args:
tg_id (int): Telegram ID пользователя.
amount (float): Сумма для обновления баланса.
session (Any, optional): Сессия базы данных. Если не передана, создается новая.
is_admin (bool, optional): Флаг, указывающий, что пополнение идёт через админку. По умолчанию False.
Raises:
Exception: В случае ошибки при подключении к базе данных или обновлении баланса.
- Кэшбек применяется только для положительных сумм, если пополнение НЕ через админку и не пропущен явно.
- Реферальный бонус тоже не срабатывает, если явно попросили пропустить (например, при начислении за купон).
"""
conn = None
try:
@@ -556,15 +556,20 @@ async def update_balance(tg_id: int, amount: float, session: Any = None, is_admi
conn = await asyncpg.connect(DATABASE_URL)
session = conn
extra = amount * (CASHBACK / 100.0) if (CASHBACK > 0 and amount > 0 and not is_admin) else 0
# Если пополнение не от админа и не сказали пропустить кэшбэк
if (CASHBACK > 0 and amount > 0 and not is_admin and not skip_cashback):
extra = amount * (CASHBACK / 100.0)
else:
extra = 0
total_amount = int(amount + extra)
current_balance = await session.fetchval("SELECT balance FROM connections WHERE tg_id = $1", tg_id)
current_balance = await session.fetchval(
"SELECT balance FROM connections WHERE tg_id = $1",
tg_id
) or 0
if current_balance is None:
current_balance = 0
new_balance = int(current_balance) + total_amount
new_balance = current_balance + total_amount
await session.execute(
"""
@@ -580,7 +585,8 @@ async def update_balance(tg_id: int, amount: float, session: Any = None, is_admi
f"({'+ кешбэк' if extra > 0 else 'без кешбэка'}), стало: {new_balance}"
)
if not is_admin:
# Если не админ и не пропустили реферальное начисление — обрабатываем реферальную цепочку
if not is_admin and not skip_referral:
await handle_referral_on_balance_update(tg_id, int(amount))
except Exception as e:
@@ -729,7 +735,7 @@ async def handle_referral_on_balance_update(tg_id: int, amount: float):
if bonus > 0:
logger.info(f"Начисление бонуса {bonus} рублей рефереру {referrer_tg_id} на уровне {level}.")
await update_balance(referrer_tg_id, bonus)
await update_balance(referrer_tg_id, bonus, skip_referral=True, skip_cashback=True)
if CHECK_REFERRAL_REWARD_ISSUED:
await conn.execute(
+24 -3
View File
@@ -14,14 +14,23 @@ async def create_key_on_cluster(
cluster_id: str, tg_id: int, client_id: str, email: str, expiry_timestamp: int, plan: int = None
):
"""
Создает ключ на всех серверах указанного кластера.
Создает ключ на всех серверах указанного кластера (или на конкретном сервере, если cluster_id — это имя сервера).
"""
try:
servers = await get_servers()
cluster = servers.get(cluster_id)
# Если не нашли кластер по ключу, ищем сервер по имени (аналогично renew_key_in_cluster, delete_key_from_cluster)
if not cluster:
raise ValueError(f"Кластер с ID {cluster_id} не найден.")
found_servers = []
for _key, server_list in servers.items():
for server_info in server_list:
if server_info.get("server_name", "").lower() == cluster_id.lower():
found_servers.append(server_info)
if found_servers:
cluster = found_servers
else:
raise ValueError(f"Кластер или сервер с ID/именем {cluster_id} не найден.")
semaphore = asyncio.Semaphore(2)
@@ -197,12 +206,24 @@ async def delete_key_from_cluster(cluster_id, email, client_id):
async def update_key_on_cluster(tg_id, client_id, email, expiry_time, cluster_id):
"""
Обновляет ключ на всех серверах указанного кластера (или сервера, если передано имя).
"""
try:
servers = await get_servers()
cluster = servers.get(cluster_id)
# Аналогичная логика поиска кластера или конкретного сервера
if not cluster:
raise ValueError(f"Кластер с ID {cluster_id} не найден.")
found_servers = []
for _key, server_list in servers.items():
for server_info in server_list:
if server_info.get("server_name", "").lower() == cluster_id.lower():
found_servers.append(server_info)
if found_servers:
cluster = found_servers
else:
raise ValueError(f"Кластер или сервер с ID/именем {cluster_id} не найден.")
tasks = []
for server_info in cluster:
+64 -32
View File
@@ -6,10 +6,14 @@ import aiohttp
import asyncpg
from aiohttp import web
from config import DATABASE_URL, PROJECT_NAME, SUB_MESSAGE, SUPERNODE, TRANSITION_DATE_STR, USE_COUNTRY_SELECTION
from config import (
DATABASE_URL, PROJECT_NAME, SUB_MESSAGE, SUPERNODE,
TRANSITION_DATE_STR, USE_COUNTRY_SELECTION
)
from database import get_key_details, get_servers
from logger import logger
db_pool = None
@@ -53,7 +57,10 @@ async def combine_unique_lines(urls, tg_id, query_string):
logger.info(f"Начинаем объединение подписок для tg_id: {tg_id}, запрос: {query_string}")
urls_with_query = [f"{url}?{query_string}" if query_string else url for url in urls]
urls_with_query = [
f"{url}?{query_string}" if query_string else url
for url in urls
]
logger.info(f"Составлены URL-адреса: {urls_with_query}")
tasks = [fetch_url_content(url, tg_id) for url in urls_with_query]
@@ -64,7 +71,6 @@ async def combine_unique_lines(urls, tg_id, query_string):
all_lines.update(filter(None, lines))
logger.info(f"Объединено {len(all_lines)} строк после фильтрации и удаления дубликатов для tg_id: {tg_id}")
return list(all_lines)
@@ -75,6 +81,40 @@ transition_timestamp_ms_adjusted = transition_timestamp_ms - (3 * 60 * 60 * 1000
logger.info(f"Время перехода (с поправкой на часовой пояс): {transition_timestamp_ms_adjusted}")
async def get_subscription_urls(server_id: str, email: str, conn) -> list:
"""
Универсальная функция, которая в зависимости от флага USE_COUNTRY_SELECTION
получает список URL-адресов для подписки. Возвращает пустой список, если
нужные данные не найдены.
"""
if USE_COUNTRY_SELECTION:
logger.info(f"Режим выбора страны активен. Ищем сервер {server_id} в БД.")
server_data = await conn.fetchrow(
"SELECT subscription_url FROM servers WHERE server_name = $1",
server_id
)
if not server_data:
logger.warning(f"Не найден сервер {server_id} в БД!")
return []
subscription_url = server_data["subscription_url"]
urls = [f"{subscription_url}/{email}"]
logger.info(f"Используем подписку {urls[0]}")
return urls
# Если режим выбора страны выключен, получаем все сервера кластера
servers = await get_servers()
logger.info(f"Режим выбора страны отключен. Используем кластер {server_id}.")
cluster_servers = servers.get(server_id, [])
if not cluster_servers:
logger.warning(f"Не найдены сервера для {server_id}")
return []
urls = [f"{server['subscription_url']}/{email}" for server in cluster_servers]
logger.info(f"Найдено {len(urls)} URL-адресов в кластере {server_id}")
return urls
async def handle_subscription(request, old_subscription=False):
"""Обрабатывает запрос на подписку (старую или новую)."""
email = request.match_info.get("email")
@@ -85,7 +125,8 @@ async def handle_subscription(request, old_subscription=False):
return web.Response(text="❌ Неверные параметры запроса.", status=400)
logger.info(
f"Обработка запроса для {'старого' if old_subscription else 'нового'} клиента: email={email}, tg_id={tg_id}"
f"Обработка запроса для {'старого' if old_subscription else 'нового'} клиента: "
f"email={email}, tg_id={tg_id}"
)
await init_db_pool()
@@ -100,6 +141,7 @@ async def handle_subscription(request, old_subscription=False):
stored_tg_id = client_data.get("tg_id")
server_id = client_data["server_id"]
# Проверяем, что tg_id из запроса совпадает с сохранённым в БД (для новой подписки)
if not old_subscription and str(tg_id) != str(stored_tg_id):
logger.warning(f"Неверный tg_id для клиента с email {email}.")
return web.Response(text="❌ Неверные данные. Получите свой ключ в боте.", status=403)
@@ -112,39 +154,29 @@ async def handle_subscription(request, old_subscription=False):
if created_at_ms >= transition_timestamp_ms_adjusted:
logger.info(f"Клиент с email {email} является новым.")
return web.Response(text="❌ Эта ссылка устарела. Пожалуйста, обновите ссылку.", status=400)
return web.Response(
text="❌ Эта ссылка устарела. Пожалуйста, обновите ссылку.",
status=400
)
urls = []
if USE_COUNTRY_SELECTION:
logger.info(f"Режим выбора страны активен. Ищем сервер {server_id} в БД.")
server_data = await conn.fetchrow("SELECT subscription_url FROM servers WHERE server_name = $1", server_id)
if not server_data:
logger.warning(f"Не найден сервер {server_id} в БД!")
return web.Response(text="❌ Сервер не найден.", status=404)
subscription_url = server_data["subscription_url"]
urls = [f"{subscription_url}/{email}"]
logger.info(f"Используем подписку {urls[0]}")
else:
servers = await get_servers()
logger.info(f"Режим выбора страны отключен. Используем кластер {server_id}.")
cluster_servers = servers.get(server_id, [])
if not cluster_servers:
logger.warning(f"Не найдены сервера для {server_id}")
return web.Response(text="❌ Сервер не найден.", status=404)
urls = [f"{server['subscription_url']}/{email}" for server in cluster_servers]
# Получаем список URL-адресов для подписки через универсальную функцию
urls = await get_subscription_urls(server_id, email, conn)
if not urls:
# Сообщаем, что сервер не найден или неверные данные
return web.Response(text="❌ Сервер не найден.", status=404)
query_string = request.query_string if not old_subscription else ""
combined_subscriptions = await combine_unique_lines(urls, tg_id or email, query_string)
combined_subscriptions = await combine_unique_lines(
urls,
tg_id or email, # Если tg_id нет, для лога используем email
query_string
)
base64_encoded = base64.b64encode(
"\n".join(combined_subscriptions).encode("utf-8")
).decode("utf-8")
base64_encoded = base64.b64encode("\n".join(combined_subscriptions).encode("utf-8")).decode("utf-8")
encoded_project_name = f"{PROJECT_NAME} - {SUB_MESSAGE}"
headers = {
"Content-Type": "text/plain; charset=utf-8",
"Content-Disposition": "inline",