formatting/3x-ui db import/cosmetic fixes

This commit is contained in:
Vladless
2025-07-19 01:14:42 +03:00
parent 8f471a1e43
commit c610fb118e
111 changed files with 2136 additions and 3816 deletions
+1 -5
View File
@@ -5,10 +5,6 @@ from database.models import BlockedUser
async def create_blocked_user(session: AsyncSession, tg_id: int):
stmt = (
insert(BlockedUser)
.values(tg_id=tg_id)
.on_conflict_do_nothing(index_elements=[BlockedUser.tg_id])
)
stmt = insert(BlockedUser).values(tg_id=tg_id).on_conflict_do_nothing(index_elements=[BlockedUser.tg_id])
await session.execute(stmt)
await session.commit()
+6 -18
View File
@@ -8,9 +8,7 @@ from database.models import Coupon, CouponUsage
from logger import logger
async def create_coupon(
session: AsyncSession, code: str, amount: int, usage_limit: int, days: int = None
) -> bool:
async def create_coupon(session: AsyncSession, code: str, amount: int, usage_limit: int, days: int = None) -> bool:
try:
exists = await session.scalar(select(Coupon.id).where(Coupon.code == code))
if exists:
@@ -42,9 +40,7 @@ async def get_coupon_by_code(session: AsyncSession, code: str) -> Coupon | None:
return result.scalar_one_or_none()
async def get_all_coupons(
session: AsyncSession, page: int = 1, per_page: int = 10
) -> dict:
async def get_all_coupons(session: AsyncSession, page: int = 1, per_page: int = 10) -> dict:
offset = (page - 1) * per_page
stmt = select(Coupon).order_by(Coupon.id.desc()).offset(offset).limit(per_page)
@@ -81,9 +77,7 @@ async def delete_coupon(session: AsyncSession, code: str) -> bool:
async def create_coupon_usage(session: AsyncSession, coupon_id: int, user_id: int):
try:
stmt = insert(CouponUsage).values(
coupon_id=coupon_id, user_id=user_id, used_at=datetime.utcnow()
)
stmt = insert(CouponUsage).values(coupon_id=coupon_id, user_id=user_id, used_at=datetime.utcnow())
await session.execute(stmt)
await session.commit()
logger.info(f"✅ Купон {coupon_id} использован пользователем {user_id}")
@@ -92,12 +86,8 @@ async def create_coupon_usage(session: AsyncSession, coupon_id: int, user_id: in
await session.rollback()
async def check_coupon_usage(
session: AsyncSession, coupon_id: int, user_id: int
) -> bool:
stmt = select(CouponUsage).where(
CouponUsage.coupon_id == coupon_id, CouponUsage.user_id == user_id
)
async def check_coupon_usage(session: AsyncSession, coupon_id: int, user_id: int) -> bool:
stmt = select(CouponUsage).where(CouponUsage.coupon_id == coupon_id, CouponUsage.user_id == user_id)
result = await session.execute(stmt)
return result.scalar_one_or_none() is not None
@@ -109,9 +99,7 @@ async def update_coupon_usage_count(session: AsyncSession, coupon_id: int):
.where(Coupon.id == coupon_id)
.values(
usage_count=Coupon.usage_count + 1,
is_used=case(
(Coupon.usage_count + 1 >= Coupon.usage_limit, True), else_=False
),
is_used=case((Coupon.usage_count + 1 >= Coupon.usage_limit, True), else_=False),
)
)
await session.commit()
+2 -3
View File
@@ -3,10 +3,9 @@ from sqlalchemy.orm import declarative_base
from config import DATABASE_URL
engine = create_async_engine(DATABASE_URL, echo=False, future=True)
async_session_maker = async_sessionmaker(
bind=engine, expire_on_commit=False, class_=AsyncSession
)
async_session_maker = async_sessionmaker(bind=engine, expire_on_commit=False, class_=AsyncSession)
Base = declarative_base()
+1 -5
View File
@@ -8,11 +8,7 @@ async def get_hot_leads(session: AsyncSession):
"""
Возвращает пользователей, у которых есть успешные оплаты, но нет активных ключей.
"""
subquery = (
select(Key.tg_id)
.where(Key.expiry_time > func.extract("epoch", func.now()) * 1000)
.distinct()
)
subquery = select(Key.tg_id).where(Key.expiry_time > func.extract("epoch", func.now()) * 1000).distinct()
stmt = (
select(Payment.tg_id)
+115
View File
@@ -0,0 +1,115 @@
import json
import sqlite3
import time
from datetime import datetime
from itertools import cycle
from sqlalchemy import select
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from database.models import Key, Server, User
async def import_keys_from_3xui_db(db_path: str, session: AsyncSession) -> tuple[int, int]:
imported = 0
skipped = 0
result = await session.execute(
select(Server.cluster_name)
.where(Server.enabled is True, Server.panel_type == "3x-ui", Server.cluster_name.isnot(None))
.distinct()
)
clusters = [row[0] for row in result.fetchall()]
if not clusters:
raise RuntimeError("❌ Не найдено доступных кластеров для 3x-ui")
cluster_cycle = cycle(clusters)
try:
conn = sqlite3.connect(db_path)
cursor = conn.cursor()
cursor.execute("SELECT id, remark, settings FROM inbounds")
inbounds = cursor.fetchall()
except Exception as e:
raise RuntimeError(f"Не удалось прочитать SQLite: {e}")
finally:
conn.close()
parsed_clients = []
for inbound_id, _remark, settings_raw in inbounds:
try:
settings = json.loads(settings_raw)
clients = settings.get("clients", [])
for c in clients:
expiry = c.get("expiryTime")
c["expiryTime"] = int(float(expiry)) if expiry else 0
c["limitIp"] = int(c.get("limitIp", 0) or 0)
c["inbound_id"] = inbound_id
parsed_clients.append(c)
except Exception:
continue
now_ts = int(time.time() * 1000)
for c in parsed_clients:
tg_id = c.get("tgId")
client_id = str(c.get("id"))
email = c.get("email")
expiry_time = int(c.get("expiryTime") or now_ts)
created_at = now_ts
server_id = next(cluster_cycle)
if not tg_id or not client_id:
continue
user_exists = await session.execute(select(User).where(User.tg_id == tg_id))
if not user_exists.scalar():
try:
session.add(
User(
tg_id=tg_id,
username=None,
first_name=None,
last_name=None,
language_code=None,
is_bot=False,
balance=0.0,
trial=1,
source_code=None,
created_at=datetime.utcnow(),
updated_at=datetime.utcnow(),
)
)
except SQLAlchemyError:
continue
key_exists = await session.execute(select(Key).where(Key.client_id == client_id))
if key_exists.scalar():
skipped += 1
continue
try:
session.add(
Key(
tg_id=tg_id,
client_id=client_id,
email=email,
created_at=created_at,
expiry_time=expiry_time,
key="",
server_id=server_id,
remnawave_link=None,
tariff_id=None,
is_frozen=False,
alias=None,
notified=False,
notified_24h=False,
)
)
imported += 1
except SQLAlchemyError:
continue
await session.commit()
return imported, skipped
+23 -19
View File
@@ -1,9 +1,12 @@
from database.db import engine, async_session_maker
from database.models import Base, Admin, User
from sqlalchemy import select
from config import ADMIN_ID
from datetime import datetime
from sqlalchemy import select
from config import ADMIN_ID
from database.db import async_session_maker, engine
from database.models import Admin, Base, User
async def init_db():
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
@@ -11,22 +14,23 @@ async def init_db():
async with async_session_maker() as session:
result = await session.execute(select(User).where(User.tg_id == 0))
if not result.scalar_one_or_none():
session.add(User(
tg_id=0,
username="system",
first_name="System",
is_bot=True,
created_at=datetime.utcnow(),
updated_at=datetime.utcnow()
))
session.add(
User(
tg_id=0,
username="system",
first_name="System",
is_bot=True,
created_at=datetime.utcnow(),
updated_at=datetime.utcnow(),
)
)
for tg_id in ADMIN_ID:
result = await session.execute(select(Admin).where(Admin.tg_id == tg_id))
if not result.scalar_one_or_none():
session.add(Admin(
tg_id=tg_id,
role="superadmin",
description="Imported from config",
added_at=datetime.utcnow()
))
session.add(
Admin(
tg_id=tg_id, role="superadmin", description="Imported from config", added_at=datetime.utcnow()
)
)
await session.commit()
+11 -39
View File
@@ -20,13 +20,9 @@ async def store_key(
tariff_id: int = None,
):
try:
exists = await session.execute(
select(Key).where(Key.tg_id == tg_id, Key.client_id == client_id)
)
exists = await session.execute(select(Key).where(Key.tg_id == tg_id, Key.client_id == client_id))
if exists.scalar_one_or_none():
logger.info(
f"[Store Key] Ключ уже существует — пропускаем: tg_id={tg_id}, client_id={client_id}"
)
logger.info(f"[Store Key] Ключ уже существует — пропускаем: tg_id={tg_id}, client_id={client_id}")
return
new_key = Key(
@@ -42,9 +38,7 @@ async def store_key(
)
session.add(new_key)
await session.commit()
logger.info(
f"✅ Ключ сохранён: tg_id={tg_id}, client_id={client_id}, server_id={server_id}"
)
logger.info(f"✅ Ключ сохранён: tg_id={tg_id}, client_id={client_id}, server_id={server_id}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при сохранении ключа: {e}")
await session.rollback()
@@ -67,9 +61,7 @@ async def get_key_by_server(session: AsyncSession, tg_id: int, client_id: str):
async def get_key_details(session: AsyncSession, email: str) -> dict | None:
stmt = (
select(Key, User).join(User, Key.tg_id == User.tg_id).where(Key.email == email)
)
stmt = select(Key, User).join(User, Key.tg_id == User.tg_id).where(Key.email == email)
result = await session.execute(stmt)
row = result.first()
if not row:
@@ -110,26 +102,18 @@ async def get_key_details(session: AsyncSession, email: str) -> dict | None:
async def get_key_count(session: AsyncSession, tg_id: int) -> int:
result = await session.execute(
select(func.count()).select_from(Key).where(Key.tg_id == tg_id)
)
result = await session.execute(select(func.count()).select_from(Key).where(Key.tg_id == tg_id))
return result.scalar() or 0
async def delete_key(session: AsyncSession, identifier: int | str):
stmt = delete(Key).where(
Key.tg_id == identifier
if str(identifier).isdigit()
else Key.client_id == identifier
)
stmt = delete(Key).where(Key.tg_id == identifier if str(identifier).isdigit() else Key.client_id == identifier)
await session.execute(stmt)
await session.commit()
logger.info(f"Ключ с идентификатором {identifier} удалён")
async def update_key_expiry(
session: AsyncSession, client_id: str, new_expiry_time: int
):
async def update_key_expiry(session: AsyncSession, client_id: str, new_expiry_time: int):
await session.execute(
update(Key)
.where(Key.client_id == client_id)
@@ -145,17 +129,11 @@ async def get_client_id_by_email(session: AsyncSession, email: str):
async def update_key_notified(session: AsyncSession, tg_id: int, client_id: str):
await session.execute(
update(Key)
.where(Key.tg_id == tg_id, Key.client_id == client_id)
.values(notified=True)
)
await session.execute(update(Key).where(Key.tg_id == tg_id, Key.client_id == client_id).values(notified=True))
await session.commit()
async def mark_key_as_frozen(
session: AsyncSession, tg_id: int, client_id: str, time_left: int
):
async def mark_key_as_frozen(session: AsyncSession, tg_id: int, client_id: str, time_left: int):
await session.execute(
text(
"""
@@ -170,9 +148,7 @@ async def mark_key_as_frozen(
)
async def mark_key_as_unfrozen(
session: AsyncSession, tg_id: int, client_id: str, new_expiry_time: int
):
async def mark_key_as_unfrozen(session: AsyncSession, tg_id: int, client_id: str, new_expiry_time: int):
await session.execute(
text(
"""
@@ -188,10 +164,6 @@ async def mark_key_as_unfrozen(
async def update_key_tariff(session: AsyncSession, client_id: str, tariff_id: int):
await session.execute(
update(Key)
.where(Key.client_id == client_id)
.values(tariff_id=tariff_id)
)
await session.execute(update(Key).where(Key.client_id == client_id).values(tariff_id=tariff_id))
await session.commit()
logger.info(f"Тариф ключа {client_id} обновлён на {tariff_id}")
+7 -11
View File
@@ -1,3 +1,6 @@
import secrets
import uuid
from datetime import datetime
from sqlalchemy import (
@@ -12,10 +15,9 @@ from sqlalchemy import (
String,
Text,
)
import secrets
import uuid
from sqlalchemy.orm import Mapped, declarative_base, mapped_column
Base = declarative_base()
@@ -27,9 +29,7 @@ class DictLikeMixin:
return getattr(self, key, default)
def to_dict(self):
return {
column.name: getattr(self, column.name) for column in self.__table__.columns
}
return {column.name: getattr(self, column.name) for column in self.__table__.columns}
class User(DictLikeMixin, Base):
@@ -139,11 +139,7 @@ class Referral(DictLikeMixin, Base):
class Notification(DictLikeMixin, Base):
__tablename__ = "notifications"
tg_id = Column(
BigInteger,
ForeignKey("users.tg_id", ondelete="CASCADE"),
primary_key=True
)
tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="CASCADE"), primary_key=True)
notification_type = Column(String, primary_key=True)
last_notification_time = Column(DateTime, default=datetime.utcnow)
@@ -219,4 +215,4 @@ class Admin(Base):
@staticmethod
def generate_token() -> str:
return secrets.token_urlsafe(32)
return secrets.token_urlsafe(32)
+10 -22
View File
@@ -25,17 +25,13 @@ async def add_notification(session: AsyncSession, tg_id: int, notification_type:
)
await session.execute(stmt)
await session.commit()
logger.info(
f"✅ Добавлено уведомление {notification_type} для пользователя {tg_id}"
)
logger.info(f"✅ Добавлено уведомление {notification_type} для пользователя {tg_id}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при добавлении уведомления: {e}")
await session.rollback()
async def delete_notification(
session: AsyncSession, tg_id: int, notification_type: str
):
async def delete_notification(session: AsyncSession, tg_id: int, notification_type: str):
await session.execute(
delete(Notification).where(
Notification.tg_id == tg_id,
@@ -46,9 +42,7 @@ async def delete_notification(
logger.info(f"🗑 Уведомление {notification_type} для пользователя {tg_id} удалено")
async def check_notification_time(
session: AsyncSession, tg_id: int, notification_type: str, hours: int = 12
) -> bool:
async def check_notification_time(session: AsyncSession, tg_id: int, notification_type: str, hours: int = 12) -> bool:
stmt = select(Notification.last_notification_time).where(
Notification.tg_id == tg_id, Notification.notification_type == notification_type
)
@@ -59,9 +53,7 @@ async def check_notification_time(
return datetime.utcnow() - last_time > timedelta(hours=hours)
async def get_last_notification_time(
session: AsyncSession, tg_id: int, notification_type: str
) -> int | None:
async def get_last_notification_time(session: AsyncSession, tg_id: int, notification_type: str) -> int | None:
stmt = select(Notification.last_notification_time).where(
Notification.tg_id == tg_id, Notification.notification_type == notification_type
)
@@ -79,16 +71,14 @@ async def check_notifications_bulk(
tg_ids: list[int] = None,
emails: list[str] = None,
) -> list[dict]:
from sqlalchemy import and_, func, select
from database.models import User, Key, Notification, BlockedUser
from sqlalchemy import select
from database.models import BlockedUser, Notification
try:
now = datetime.utcnow()
subq_last_notification = (
select(
Notification.tg_id,
func.max(Notification.last_notification_time).label("last_notification_time")
)
select(Notification.tg_id, func.max(Notification.last_notification_time).label("last_notification_time"))
.where(Notification.notification_type == notification_type)
.group_by(Notification.tg_id)
.subquery()
@@ -112,7 +102,7 @@ async def check_notifications_bulk(
and_(
User.trial.in_([0, -1]),
~User.tg_id.in_(select(BlockedUser.tg_id)),
~User.tg_id.in_(select(Key.tg_id.distinct()))
~User.tg_id.in_(select(Key.tg_id.distinct())),
)
)
@@ -135,9 +125,7 @@ async def check_notifications_bulk(
"username": row.username,
"first_name": row.first_name,
"last_name": row.last_name,
"last_notification_time": (
int(last_time.timestamp() * 1000) if last_time else None
)
"last_notification_time": (int(last_time.timestamp() * 1000) if last_time else None),
})
logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}")
+5 -14
View File
@@ -1,9 +1,9 @@
from datetime import datetime
from pytz import timezone
from sqlalchemy import insert, select
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from pytz import timezone
from database.models import Payment
from logger import logger
@@ -12,9 +12,7 @@ from logger import logger
MOSCOW_TZ = timezone("Europe/Moscow")
async def add_payment(
session: AsyncSession, tg_id: int, amount: float, payment_system: str
):
async def add_payment(session: AsyncSession, tg_id: int, amount: float, payment_system: str):
try:
now_moscow = datetime.now(MOSCOW_TZ).replace(tzinfo=None)
stmt = insert(Payment).values(
@@ -26,9 +24,7 @@ async def add_payment(
)
await session.execute(stmt)
await session.commit()
logger.info(
f"✅ Успешно добавлен платёж: {tg_id}, {amount}₽ через {payment_system}"
)
logger.info(f"✅ Успешно добавлен платёж: {tg_id}, {amount}₽ через {payment_system}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при добавлении платежа: {e}")
await session.rollback()
@@ -38,15 +34,10 @@ async def add_payment(
async def get_last_payments(session: AsyncSession, tg_id: int, limit: int = 3):
try:
result = await session.execute(
select(Payment)
.where(Payment.tg_id == tg_id)
.order_by(Payment.created_at.desc())
.limit(limit)
select(Payment).where(Payment.tg_id == tg_id).order_by(Payment.created_at.desc()).limit(limit)
)
payments = result.scalars().all()
logger.info(
f"✅ Получены последние платежи пользователя {tg_id}, всего: {len(payments)}"
)
logger.info(f"✅ Получены последние платежи пользователя {tg_id}, всего: {len(payments)}")
return [dict(p.__dict__) for p in payments]
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при получении платежей пользователя {tg_id}: {e}")
+15 -52
View File
@@ -13,23 +13,17 @@ async def add_referral(session: AsyncSession, referred_tg_id: int, referrer_tg_i
logger.warning(f"⚠️ Попытка самореферала: {referred_tg_id}")
return
stmt = insert(Referral).values(
referred_tg_id=referred_tg_id, referrer_tg_id=referrer_tg_id
)
stmt = insert(Referral).values(referred_tg_id=referred_tg_id, referrer_tg_id=referrer_tg_id)
await session.execute(stmt)
await session.commit()
logger.info(
f"✅ Добавлена реферальная связь: {referred_tg_id}{referrer_tg_id}"
)
logger.info(f"✅ Добавлена реферальная связь: {referred_tg_id}{referrer_tg_id}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при добавлении реферала: {e}")
await session.rollback()
raise
async def get_referral_by_referred_id(
session: AsyncSession, referred_tg_id: int
) -> dict | None:
async def get_referral_by_referred_id(session: AsyncSession, referred_tg_id: int) -> dict | None:
stmt = select(Referral).where(Referral.referred_tg_id == referred_tg_id)
result = await session.execute(stmt)
row = result.scalar_one_or_none()
@@ -37,11 +31,7 @@ async def get_referral_by_referred_id(
async def get_total_referrals(session: AsyncSession, referrer_tg_id: int) -> int:
stmt = (
select(func.count())
.select_from(Referral)
.where(Referral.referrer_tg_id == referrer_tg_id)
)
stmt = select(func.count()).select_from(Referral).where(Referral.referrer_tg_id == referrer_tg_id)
result = await session.execute(stmt)
return result.scalar()
@@ -62,17 +52,11 @@ async def get_active_referrals(session: AsyncSession, referrer_tg_id: int) -> in
async def mark_referral_reward_issued(session: AsyncSession, referred_tg_id: int):
await session.execute(
update(Referral)
.where(Referral.referred_tg_id == referred_tg_id)
.values(reward_issued=True)
)
await session.execute(update(Referral).where(Referral.referred_tg_id == referred_tg_id).values(reward_issued=True))
await session.commit()
async def get_total_referral_bonus(
session: AsyncSession, referrer_tg_id: int, max_levels: int
) -> float:
async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, max_levels: int) -> float:
if CHECK_REFERRAL_REWARD_ISSUED:
bonus_cte = """
WITH RECURSIVE
@@ -168,9 +152,7 @@ async def get_total_referral_bonus(
"""
)
result = await session.execute(
text(bonus_query), {"tg_id": referrer_tg_id, "max_levels": max_levels}
)
result = await session.execute(text(bonus_query), {"tg_id": referrer_tg_id, "max_levels": max_levels})
total_bonus_raw = result.scalar()
total_bonus = round(float(total_bonus_raw or 0), 2)
@@ -178,9 +160,7 @@ async def get_total_referral_bonus(
return total_bonus
async def get_referrals_by_level(
session: AsyncSession, referrer_tg_id: int, max_levels: int
) -> dict:
async def get_referrals_by_level(session: AsyncSession, referrer_tg_id: int, max_levels: int) -> dict:
query = """
WITH RECURSIVE referral_levels AS (
SELECT referred_tg_id, referrer_tg_id, 1 AS level
@@ -200,9 +180,7 @@ async def get_referrals_by_level(
GROUP BY level
ORDER BY level
"""
result = await session.execute(
text(query), {"referrer_tg_id": referrer_tg_id, "max_levels": max_levels}
)
result = await session.execute(text(query), {"referrer_tg_id": referrer_tg_id, "max_levels": max_levels})
return {
row["level"]: {
"total": row["level_count"],
@@ -214,19 +192,13 @@ async def get_referrals_by_level(
async def get_referral_stats(session: AsyncSession, referrer_tg_id: int):
try:
logger.info(
f"[ReferralStats] Получение статистики для пользователя {referrer_tg_id}"
)
logger.info(f"[ReferralStats] Получение статистики для пользователя {referrer_tg_id}")
total_referrals = await get_total_referrals(session, referrer_tg_id)
active_referrals = await get_active_referrals(session, referrer_tg_id)
max_levels = len(REFERRAL_BONUS_PERCENTAGES)
referrals_by_level = await get_referrals_by_level(
session, referrer_tg_id, max_levels
)
total_referral_bonus = await get_total_referral_bonus(
session, referrer_tg_id, max_levels
)
referrals_by_level = await get_referrals_by_level(session, referrer_tg_id, max_levels)
total_referral_bonus = await get_total_referral_bonus(session, referrer_tg_id, max_levels)
return {
"total_referrals": total_referrals,
@@ -236,18 +208,12 @@ async def get_referral_stats(session: AsyncSession, referrer_tg_id: int):
}
except Exception as e:
logger.error(
f"[ReferralStats] Ошибка при получении статистики для пользователя {referrer_tg_id}: {e}"
)
logger.error(f"[ReferralStats] Ошибка при получении статистики для пользователя {referrer_tg_id}: {e}")
raise
async def get_user_referral_count(session: AsyncSession, tg_id: int) -> int:
result = await session.execute(
select(func.count())
.select_from(Referral)
.where(Referral.referrer_tg_id == tg_id)
)
result = await session.execute(select(func.count()).select_from(Referral).where(Referral.referrer_tg_id == tg_id))
return result.scalar_one() or 0
@@ -272,7 +238,4 @@ async def get_top_referrals(session: AsyncSession, limit: int = 5):
.limit(limit)
)
result = await session.execute(query)
return [
{"referrer_tg_id": row.referrer_tg_id, "referral_count": row.referral_count}
for row in result.all()
]
return [{"referrer_tg_id": row.referrer_tg_id, "referral_count": row.referral_count} for row in result.all()]
+32 -56
View File
@@ -1,8 +1,8 @@
from sqlalchemy import delete, insert, select, update, func
from sqlalchemy import delete, func, insert, select, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from database.models import Server, Key
from database.models import Key, Server
from logger import logger
@@ -54,19 +54,17 @@ async def get_servers(session: AsyncSession, include_enabled: bool = False) -> d
if not include_enabled and not s.enabled:
continue
cluster = s.cluster_name
grouped.setdefault(cluster, []).append(
{
"server_name": s.server_name,
"api_url": s.api_url,
"subscription_url": s.subscription_url,
"inbound_id": s.inbound_id,
"panel_type": s.panel_type,
"enabled": s.enabled,
"max_keys": s.max_keys,
"tariff_group": s.tariff_group,
"cluster_name": cluster,
}
)
grouped.setdefault(cluster, []).append({
"server_name": s.server_name,
"api_url": s.api_url,
"subscription_url": s.subscription_url,
"inbound_id": s.inbound_id,
"panel_type": s.panel_type,
"enabled": s.enabled,
"max_keys": s.max_keys,
"tariff_group": s.tariff_group,
"cluster_name": cluster,
})
return grouped
except SQLAlchemyError as e:
@@ -80,9 +78,7 @@ async def get_clusters(session: AsyncSession) -> list[str]:
return [r[0] for r in result.all()]
async def check_unique_server_name(
session: AsyncSession, server_name: str, cluster_name: str | None = None
) -> bool:
async def check_unique_server_name(session: AsyncSession, server_name: str, cluster_name: str | None = None) -> bool:
stmt = select(Server).where(Server.server_name == server_name)
if cluster_name:
stmt = stmt.where(Server.cluster_name == cluster_name)
@@ -90,13 +86,9 @@ async def check_unique_server_name(
return result.scalar_one_or_none() is None
async def check_server_name_by_cluster(
session: AsyncSession, server_name: str
) -> dict | None:
async def check_server_name_by_cluster(session: AsyncSession, server_name: str) -> dict | None:
try:
result = await session.execute(
select(Server.cluster_name).where(Server.server_name == server_name)
)
result = await session.execute(select(Server.cluster_name).where(Server.server_name == server_name))
row = result.first()
return {"cluster_name": row[0]} if row else None
except SQLAlchemyError as e:
@@ -104,14 +96,10 @@ async def check_server_name_by_cluster(
return None
async def get_cluster_name_by_server(
session: AsyncSession, server_id_or_name: str
) -> str | None:
async def get_cluster_name_by_server(session: AsyncSession, server_id_or_name: str) -> str | None:
stmt = (
select(Server.cluster_name)
.where(
(Server.id == server_id_or_name) | (Server.server_name == server_id_or_name)
)
.where((Server.id == server_id_or_name) | (Server.server_name == server_id_or_name))
.limit(1)
)
@@ -125,7 +113,7 @@ async def get_server_by_name(session: AsyncSession, server_name: str) -> dict |
stmt = select(Server).where(Server.server_name == server_name)
result = await session.execute(stmt)
server = result.scalar_one_or_none()
if server:
return {
"id": server.id,
@@ -145,9 +133,7 @@ async def get_server_by_name(session: AsyncSession, server_name: str) -> dict |
return None
async def update_server_field(
session: AsyncSession, server_name: str, field: str, value: any
) -> bool:
async def update_server_field(session: AsyncSession, server_name: str, field: str, value: any) -> bool:
try:
stmt = update(Server).where(Server.server_name == server_name).values(**{field: value})
await session.execute(stmt)
@@ -160,11 +146,10 @@ async def update_server_field(
return False
async def update_server_name_with_keys(
session: AsyncSession, old_name: str, new_name: str
) -> bool:
async def update_server_name_with_keys(session: AsyncSession, old_name: str, new_name: str) -> bool:
try:
from sqlalchemy import update
from database.models import Key
if not await check_unique_server_name(session, new_name):
@@ -176,7 +161,7 @@ async def update_server_name_with_keys(
stmt_keys = update(Key).where(Key.server_id == old_name).values(server_id=new_name)
await session.execute(stmt_keys)
await session.commit()
logger.info(f"✅ Сервер переименован с {old_name} на {new_name}")
return True
@@ -196,16 +181,12 @@ async def get_available_clusters(session: AsyncSession) -> list[str]:
return []
async def update_server_cluster(
session: AsyncSession,
server_name: str,
new_cluster: str
) -> bool:
async def update_server_cluster(session: AsyncSession, server_name: str, new_cluster: str) -> bool:
try:
server_data = await get_server_by_name(session, server_name)
if not server_data:
return False
old_cluster = server_data["cluster_name"]
stmt_remaining = select(func.count()).where(
@@ -215,26 +196,21 @@ async def update_server_cluster(
remaining_servers = result.scalar_one()
if remaining_servers == 0:
stmt_update_keys = update(Key).where(
Key.server_id == old_cluster
).values(server_id=new_cluster)
stmt_update_keys = update(Key).where(Key.server_id == old_cluster).values(server_id=new_cluster)
await session.execute(stmt_update_keys)
stmt_new_cluster = select(Server.tariff_group).where(
Server.cluster_name == new_cluster
).limit(1)
stmt_new_cluster = select(Server.tariff_group).where(Server.cluster_name == new_cluster).limit(1)
result = await session.execute(stmt_new_cluster)
new_tariff_group = result.scalar_one_or_none()
stmt_update = update(Server).where(
Server.server_name == server_name
).values(
cluster_name=new_cluster,
tariff_group=new_tariff_group
stmt_update = (
update(Server)
.where(Server.server_name == server_name)
.values(cluster_name=new_cluster, tariff_group=new_tariff_group)
)
await session.execute(stmt_update)
await session.commit()
logger.info(f"✅ Сервер {server_name} перемещен в кластер {new_cluster} с обновлением тарифной группы")
return True
except SQLAlchemyError as e:
+20 -56
View File
@@ -1,6 +1,6 @@
from datetime import date, datetime
from sqlalchemy import and_, func, not_, select, exists
from sqlalchemy import and_, exists, func, not_, select
from sqlalchemy.ext.asyncio import AsyncSession
from database.models import Key, Payment, Referral, Tariff, User
@@ -11,24 +11,16 @@ async def count_total_users(session: AsyncSession) -> int:
async def count_users_updated_today(session: AsyncSession, today: date) -> int:
return await session.scalar(
select(func.count()).select_from(User).where(User.updated_at >= today)
)
return await session.scalar(select(func.count()).select_from(User).where(User.updated_at >= today))
async def count_users_registered_since(session: AsyncSession, since: date) -> int:
return await session.scalar(
select(func.count()).select_from(User).where(User.created_at >= since)
)
return await session.scalar(select(func.count()).select_from(User).where(User.created_at >= since))
async def count_users_registered_between(
session: AsyncSession, start: date, end: date
) -> int:
async def count_users_registered_between(session: AsyncSession, start: date, end: date) -> int:
return await session.scalar(
select(func.count())
.select_from(User)
.where(User.created_at >= start, User.created_at < end)
select(func.count()).select_from(User).where(User.created_at >= start, User.created_at < end)
)
@@ -38,78 +30,55 @@ async def count_total_keys(session: AsyncSession) -> int:
async def count_active_keys(session: AsyncSession) -> int:
current_time_ms = int(datetime.utcnow().timestamp() * 1000)
return await session.scalar(
select(func.count()).select_from(Key).where(Key.expiry_time > current_time_ms)
)
return await session.scalar(select(func.count()).select_from(Key).where(Key.expiry_time > current_time_ms))
async def count_trial_keys(session: AsyncSession) -> int:
subquery_success_payments = (
select(Payment.tg_id)
.where(and_(Payment.tg_id == Key.tg_id, Payment.status == "success"))
.exists()
select(Payment.tg_id).where(and_(Payment.tg_id == Key.tg_id, Payment.status == "success")).exists()
)
return await session.scalar(
select(func.count()).select_from(Key).where(not_(subquery_success_payments))
)
return await session.scalar(select(func.count()).select_from(Key).where(not_(subquery_success_payments)))
async def get_tariff_distribution(
session: AsyncSession, include_unbound: bool = False
) -> tuple[list[tuple[int, int]], list[dict]]:
result = await session.execute(
select(Key.tariff_id, func.count(Key.client_id))
.where(Key.tariff_id.isnot(None))
.group_by(Key.tariff_id)
select(Key.tariff_id, func.count(Key.client_id)).where(Key.tariff_id.isnot(None)).group_by(Key.tariff_id)
)
tariff_counts = result.all()
if not include_unbound:
return tariff_counts
result = await session.execute(
select(Key.expiry_time)
.where(Key.tariff_id.is_(None))
)
result = await session.execute(select(Key.expiry_time).where(Key.tariff_id.is_(None)))
no_tariff_keys = [{"expiry_time": row[0]} for row in result.all()]
return tariff_counts, no_tariff_keys
async def get_tariff_names(
session: AsyncSession, tariff_ids: list[int]
) -> dict[int, str]:
async def get_tariff_names(session: AsyncSession, tariff_ids: list[int]) -> dict[int, str]:
if not tariff_ids:
return {}
result = await session.execute(
select(Tariff.id, Tariff.name).where(Tariff.id.in_(tariff_ids))
)
result = await session.execute(select(Tariff.id, Tariff.name).where(Tariff.id.in_(tariff_ids)))
return dict(result.all())
async def get_tariff_groups(
session: AsyncSession, tariff_ids: list[int]
) -> dict[int, str]:
async def get_tariff_groups(session: AsyncSession, tariff_ids: list[int]) -> dict[int, str]:
if not tariff_ids:
return {}
result = await session.execute(
select(Tariff.id, Tariff.group_code).where(Tariff.id.in_(tariff_ids))
)
result = await session.execute(select(Tariff.id, Tariff.group_code).where(Tariff.id.in_(tariff_ids)))
return dict(result.all())
async def get_tariff_durations(
session: AsyncSession, tariff_ids: list[int]
) -> dict[int, int]:
async def get_tariff_durations(session: AsyncSession, tariff_ids: list[int]) -> dict[int, int]:
if not tariff_ids:
return {}
result = await session.execute(
select(Tariff.id, Tariff.duration_days).where(Tariff.id.in_(tariff_ids))
)
result = await session.execute(select(Tariff.id, Tariff.duration_days).where(Tariff.id.in_(tariff_ids)))
return dict(result.all())
@@ -120,10 +89,7 @@ async def count_total_referrals(session: AsyncSession) -> int:
async def sum_payments_since(session: AsyncSession, since: date) -> float:
result = await session.scalar(
select(func.coalesce(func.sum(Payment.amount), 0)).where(
and_(
Payment.created_at >= since,
Payment.payment_system.notin_(["referral", "coupon", "cashback"])
)
and_(Payment.created_at >= since, Payment.payment_system.notin_(["referral", "coupon", "cashback"]))
)
)
return round(float(result), 2)
@@ -135,7 +101,7 @@ async def sum_payments_between(session: AsyncSession, start: date, end: date) ->
and_(
Payment.created_at >= start,
Payment.created_at < end,
Payment.payment_system.notin_(["referral", "coupon", "cashback"])
Payment.payment_system.notin_(["referral", "coupon", "cashback"]),
)
)
)
@@ -153,9 +119,7 @@ async def sum_total_payments(session: AsyncSession) -> float:
async def count_hot_leads(session: AsyncSession) -> int:
subquery_active_keys = (
select(Key.tg_id)
.where(Key.expiry_time > int(datetime.utcnow().timestamp() * 1000))
.distinct()
select(Key.tg_id).where(Key.expiry_time > int(datetime.utcnow().timestamp() * 1000)).distinct()
)
stmt = (
@@ -168,4 +132,4 @@ async def count_hot_leads(session: AsyncSession) -> int:
)
result = await session.execute(select(func.count()).select_from(stmt.subquery()))
return result.scalar()
return result.scalar()
+15 -31
View File
@@ -1,6 +1,7 @@
from datetime import datetime
import hashlib
from datetime import datetime
from sqlalchemy import delete, insert, select, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
@@ -14,7 +15,7 @@ def create_subgroup_hash(subgroup_title: str, group_code: str) -> str:
return ""
unique_key = f"{subgroup_title}:{group_code}"
hash_object = hashlib.md5(unique_key.encode('utf-8'))
hash_object = hashlib.md5(unique_key.encode("utf-8"))
return hash_object.hexdigest()[:8]
@@ -25,24 +26,20 @@ async def find_subgroup_by_hash(session: AsyncSession, subgroup_hash: str, group
.distinct()
)
subgroups = [row[0] for row in result.fetchall()]
for subgroup_title in subgroups:
if create_subgroup_hash(subgroup_title, group_code) == subgroup_hash:
return subgroup_title
return None
async def get_tariffs(
session: AsyncSession, tariff_id: int = None, group_code: str = None
):
async def get_tariffs(session: AsyncSession, tariff_id: int = None, group_code: str = None):
try:
if tariff_id:
result = await session.execute(select(Tariff).where(Tariff.id == tariff_id))
elif group_code:
result = await session.execute(
select(Tariff).where(Tariff.group_code == group_code).order_by(Tariff.id)
)
result = await session.execute(select(Tariff).where(Tariff.group_code == group_code).order_by(Tariff.id))
else:
result = await session.execute(select(Tariff))
@@ -65,34 +62,26 @@ async def get_tariff_by_id(session: AsyncSession, tariff_id: int):
async def get_tariffs_for_cluster(session: AsyncSession, cluster_name: str):
try:
server_row = await session.execute(
select(Server.tariff_group)
.where(Server.cluster_name == cluster_name)
.limit(1)
select(Server.tariff_group).where(Server.cluster_name == cluster_name).limit(1)
)
row = server_row.first()
if not row:
server_row = await session.execute(
select(Server.tariff_group)
.where(Server.server_name == cluster_name)
.limit(1)
select(Server.tariff_group).where(Server.server_name == cluster_name).limit(1)
)
row = server_row.first()
if not row or not row[0]:
return []
group_code = row[0]
result = await session.execute(
select(Tariff)
.where(Tariff.group_code == group_code, Tariff.is_active.is_(True))
.order_by(Tariff.id)
select(Tariff).where(Tariff.group_code == group_code, Tariff.is_active.is_(True)).order_by(Tariff.id)
)
return [dict(r.__dict__) for r in result.scalars().all()]
except SQLAlchemyError as e:
logger.error(
f"[TARIFF] Ошибка при получении тарифов для кластера {cluster_name}: {e}"
)
logger.error(f"[TARIFF] Ошибка при получении тарифов для кластера {cluster_name}: {e}")
return []
@@ -116,9 +105,7 @@ async def update_tariff(session: AsyncSession, tariff_id: int, updates: dict):
return False
try:
updates["updated_at"] = datetime.utcnow()
await session.execute(
update(Tariff).where(Tariff.id == tariff_id).values(**updates)
)
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(**updates))
await session.commit()
return True
except SQLAlchemyError as e:
@@ -140,10 +127,7 @@ async def delete_tariff(session: AsyncSession, tariff_id: int):
async def check_tariff_exists(session: AsyncSession, tariff_id: int):
try:
result = await session.execute(
select(Tariff)
.where(Tariff.id == tariff_id, Tariff.is_active.is_(True))
)
result = await session.execute(select(Tariff).where(Tariff.id == tariff_id, Tariff.is_active.is_(True)))
tariff = result.scalar_one_or_none()
if tariff:
logger.info(f"[TARIFF] Тариф {tariff_id} найден в БД: {tariff.group_code}")
+1 -3
View File
@@ -9,9 +9,7 @@ from database.models import TemporaryData
from logger import logger
async def create_temporary_data(
session: AsyncSession, tg_id: int, state: str, data: dict
):
async def create_temporary_data(session: AsyncSession, tg_id: int, state: str, data: dict):
try:
stmt = (
insert(TemporaryData)
+14 -27
View File
@@ -1,4 +1,4 @@
from sqlalchemy import func, insert, select, not_
from sqlalchemy import func, insert, not_, select
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
@@ -6,9 +6,7 @@ from database.models import Payment, TrackingSource, User
from logger import logger
async def create_tracking_source(
session: AsyncSession, name: str, code: str, type_: str, created_by: int
):
async def create_tracking_source(session: AsyncSession, name: str, code: str, type_: str, created_by: int):
try:
stmt = insert(TrackingSource).values(
name=name,
@@ -42,9 +40,7 @@ async def get_all_tracking_sources(session: AsyncSession) -> list[dict]:
payments_subq = (
select(func.count(func.distinct(Payment.tg_id)))
.join(User, Payment.tg_id == User.tg_id)
.where(
(User.source_code == TrackingSource.code) & (Payment.status == "success")
)
.where((User.source_code == TrackingSource.code) & (Payment.status == "success"))
.correlate(TrackingSource)
.scalar_subquery()
)
@@ -74,9 +70,7 @@ async def get_all_tracking_sources(session: AsyncSession) -> list[dict]:
async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict | None:
source_result = await session.execute(
select(TrackingSource.created_at).where(TrackingSource.code == code)
)
source_result = await session.execute(select(TrackingSource.created_at).where(TrackingSource.code == code))
created_at_row = source_result.first()
if not created_at_row:
return None
@@ -85,20 +79,13 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict |
reg_subq = (
select(func.count(func.distinct(User.tg_id)))
.where(
(User.source_code == code) &
(User.created_at >= created_at)
)
.where((User.source_code == code) & (User.created_at >= created_at))
.scalar_subquery()
)
trial_subq = (
select(func.count(func.distinct(User.tg_id)))
.where(
(User.source_code == code) &
(User.trial == 1) &
(User.created_at >= created_at)
)
.where((User.source_code == code) & (User.trial == 1) & (User.created_at >= created_at))
.scalar_subquery()
)
@@ -106,10 +93,10 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict |
select(func.count(func.distinct(Payment.tg_id)))
.join(User, Payment.tg_id == User.tg_id)
.where(
(User.source_code == code) &
(Payment.status == "success") &
not_(Payment.payment_system.in_(["coupon", "referral", "cashback"])) &
(Payment.created_at >= created_at)
(User.source_code == code)
& (Payment.status == "success")
& not_(Payment.payment_system.in_(["coupon", "referral", "cashback"]))
& (Payment.created_at >= created_at)
)
.scalar_subquery()
)
@@ -118,10 +105,10 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict |
select(func.coalesce(func.sum(Payment.amount), 0))
.join(User, Payment.tg_id == User.tg_id)
.where(
(User.source_code == code) &
(Payment.status == "success") &
not_(Payment.payment_system.in_(["coupon", "referral", "cashback"])) &
(Payment.created_at >= created_at)
(User.source_code == code)
& (Payment.status == "success")
& not_(Payment.payment_system.in_(["coupon", "referral", "cashback"]))
& (Payment.created_at >= created_at)
)
.scalar_subquery()
)
+9 -30
View File
@@ -10,12 +10,12 @@ from database.models import (
BlockedUser,
CouponUsage,
Gift,
GiftUsage,
Notification,
Payment,
Referral,
TemporaryData,
User,
GiftUsage
)
from logger import logger
@@ -47,9 +47,7 @@ async def add_user(
await session.execute(stmt)
await session.commit()
logger.info(
f"[DB] Новый пользователь добавлен: {tg_id} (source: {source_code})"
)
logger.info(f"[DB] Новый пользователь добавлен: {tg_id} (source: {source_code})")
except SQLAlchemyError as e:
logger.error(f"[DB] Ошибка при добавлении пользователя {tg_id}: {e}")
await session.rollback()
@@ -61,13 +59,9 @@ async def update_balance(session: AsyncSession, tg_id: int, amount: float) -> No
result = await session.execute(select(User.balance).where(User.tg_id == tg_id))
current = result.scalar_one_or_none() or 0
new_balance = current + amount
await session.execute(
update(User).where(User.tg_id == tg_id).values(balance=new_balance)
)
await session.execute(update(User).where(User.tg_id == tg_id).values(balance=new_balance))
await session.commit()
logger.info(
f"[DB] Баланс пользователя {tg_id} обновлён: {current}{new_balance}"
)
logger.info(f"[DB] Баланс пользователя {tg_id} обновлён: {current}{new_balance}")
except SQLAlchemyError as e:
logger.error(f"[DB] Ошибка при обновлении баланса пользователя {tg_id}: {e}")
await session.rollback()
@@ -87,9 +81,7 @@ async def get_balance(session: AsyncSession, tg_id: int) -> float:
async def set_user_balance(session: AsyncSession, tg_id: int, balance: float) -> None:
try:
await session.execute(
update(User).where(User.tg_id == tg_id).values(balance=balance)
)
await session.execute(update(User).where(User.tg_id == tg_id).values(balance=balance))
await session.commit()
except SQLAlchemyError as e:
logger.error(f"Ошибка при установке баланса для пользователя {tg_id}: {e}")
@@ -98,9 +90,7 @@ async def set_user_balance(session: AsyncSession, tg_id: int, balance: float) ->
async def update_trial(session: AsyncSession, tg_id: int, status: int):
try:
await session.execute(
update(User).where(User.tg_id == tg_id).values(trial=status)
)
await session.execute(update(User).where(User.tg_id == tg_id).values(trial=status))
await session.commit()
logger.info(f"[DB] Триал статус обновлён для пользователя {tg_id}: {status}")
except SQLAlchemyError as e:
@@ -181,29 +171,18 @@ async def delete_user_data(session: AsyncSession, tg_id: int):
try:
await session.execute(delete(Notification).where(Notification.tg_id == tg_id))
result = await session.execute(
select(Gift.gift_id).where(Gift.sender_tg_id == tg_id)
)
result = await session.execute(select(Gift.gift_id).where(Gift.sender_tg_id == tg_id))
gift_ids = [row[0] for row in result.all()]
if gift_ids:
await session.execute(delete(GiftUsage).where(GiftUsage.gift_id.in_(gift_ids)))
await session.execute(delete(Gift).where(Gift.sender_tg_id == tg_id))
await session.execute(
update(Gift)
.where(Gift.recipient_tg_id == tg_id)
.values(recipient_tg_id=None)
)
await session.execute(update(Gift).where(Gift.recipient_tg_id == tg_id).values(recipient_tg_id=None))
await session.execute(delete(Payment).where(Payment.tg_id == tg_id))
await session.execute(
delete(Referral).where(
or_(
Referral.referrer_tg_id == tg_id,
Referral.referred_tg_id == tg_id
)
)
delete(Referral).where(or_(Referral.referrer_tg_id == tg_id, Referral.referred_tg_id == tg_id))
)
await session.execute(delete(CouponUsage).where(CouponUsage.user_id == tg_id))
await delete_key(session, tg_id)