WEB-APP/ Optimization/ Build fix/ Hotkey edit mode/ Log rotation/ Form a11y/ E2E non-blocking

This commit is contained in:
Vladless
2026-04-13 00:38:29 +00:00
parent c1b9b20ea3
commit c693c28ee7
270 changed files with 20933 additions and 6365 deletions
+1 -1
View File
@@ -5,7 +5,7 @@ from .db import Base, async_session_maker, engine, reset_async_db_engine
from .gifts import *
from . import identities
from .hot_leads import *
from .init_db import *
from .setup.init_db import *
from .keys import *
from .notifications import *
from .payments import *
+2
View File
@@ -0,0 +1,2 @@
from .resolution import *
from .tg_mirror import *
+87
View File
@@ -0,0 +1,87 @@
from dataclasses import dataclass
from enum import Enum
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from database.models import Identity, User
class ActorSurface(str, Enum):
TELEGRAM = "telegram"
WEB = "web"
UNKNOWN = "unknown"
@dataclass(frozen=True)
class ResolvedActor:
surface: ActorSurface
billing_user_id: int | None
telegram_chat_id: int | None
identity_id: str | None
def telegram_chat_id(user: User | None) -> int | None:
if user is None:
return None
return user.tg_id
async def resolve_user_optional(session: AsyncSession, legacy_id: int) -> User | None:
r = await session.execute(select(User).where(User.tg_id == legacy_id))
u = r.scalar_one_or_none()
if u is not None:
return u
r2 = await session.execute(select(User).where(User.id == legacy_id))
return r2.scalar_one_or_none()
async def notify_telegram_chat_id(session: AsyncSession, legacy_ref: int) -> int | None:
payer = await resolve_user_optional(session, legacy_ref)
tg = telegram_chat_id(payer)
if tg is not None:
return tg
if payer is None:
return legacy_ref
return None
async def resolve_actor_from_legacy_ref(session: AsyncSession, legacy_ref: int) -> ResolvedActor:
user = await resolve_user_optional(session, legacy_ref)
if user is None:
return ResolvedActor(
surface=ActorSurface.UNKNOWN,
billing_user_id=None,
telegram_chat_id=legacy_ref,
identity_id=None,
)
user_tg = telegram_chat_id(user)
if user_tg is not None and int(user_tg) == int(legacy_ref):
surface = ActorSurface.TELEGRAM
elif int(user.id) == int(legacy_ref):
surface = ActorSurface.WEB
elif user_tg is None:
surface = ActorSurface.WEB
else:
surface = ActorSurface.UNKNOWN
return ResolvedActor(
surface=surface,
billing_user_id=int(user.id),
telegram_chat_id=user_tg,
identity_id=user.identity_id,
)
async def resolve_actor_from_identity(session: AsyncSession, identity: Identity) -> ResolvedActor:
from database.identities import ensure_billing_user_for_identity
billing_uid = await ensure_billing_user_for_identity(session, identity)
user = await resolve_user_optional(session, billing_uid)
return ResolvedActor(
surface=ActorSurface.WEB,
billing_user_id=billing_uid,
telegram_chat_id=telegram_chat_id(user),
identity_id=identity.id,
)
+50
View File
@@ -0,0 +1,50 @@
from __future__ import annotations
from sqlalchemy import select, update
from sqlalchemy.ext.asyncio import AsyncSession
from database.models import (
BlockedUser,
CouponUsage,
Gift,
GiftUsage,
Key,
ManualBan,
Notification,
Payment,
Referral,
TemporaryData,
User,
)
def mirror_telegram_id(user: User | None) -> int | None:
if user is None:
return None
return user.tg_id
async def refresh_tg_mirrors_for_user(session: AsyncSession, user_id: int) -> None:
r = await session.execute(select(User.tg_id).where(User.id == user_id))
tg = r.scalar_one_or_none()
await session.execute(update(Key).where(Key.user_id == user_id).values(tg_id=tg))
await session.execute(update(Payment).where(Payment.user_id == user_id).values(tg_id=tg))
await session.execute(update(Notification).where(Notification.user_id == user_id).values(tg_id=tg))
await session.execute(update(GiftUsage).where(GiftUsage.user_id == user_id).values(tg_id=tg))
await session.execute(update(CouponUsage).where(CouponUsage.user_id == user_id).values(tg_id=tg))
await session.execute(update(TemporaryData).where(TemporaryData.user_id == user_id).values(tg_id=tg))
await session.execute(update(BlockedUser).where(BlockedUser.user_id == user_id).values(tg_id=tg))
await session.execute(update(ManualBan).where(ManualBan.user_id == user_id).values(tg_id=tg))
await session.execute(
update(Referral).where(Referral.referred_user_id == user_id).values(referred_tg_id=tg)
)
await session.execute(
update(Referral).where(Referral.referrer_user_id == user_id).values(referrer_tg_id=tg)
)
await session.execute(update(Gift).where(Gift.sender_user_id == user_id).values(sender_tg_id=tg))
await session.execute(
update(Gift).where(Gift.recipient_user_id == user_id).values(recipient_tg_id=tg)
)
+1 -1
View File
@@ -77,7 +77,7 @@ async def fetch_successful_payment_rows_db(
Payment.created_at,
)
stmt = (
select(Payment.payment_system, Payment.payment_id, Payment.tg_id)
select(Payment.payment_system, Payment.payment_id, Payment.user_id)
.where(
Payment.status == "success",
Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED),
+26 -10
View File
@@ -1,27 +1,43 @@
from sqlalchemy import select
from sqlalchemy.dialects.postgresql import insert
from sqlalchemy.ext.asyncio import AsyncSession
from database.models import BlockedUser
from database.models import BlockedUser, User
from database.access.resolution import resolve_user_optional
from logger import logger
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])
async def create_blocked_user(session: AsyncSession, legacy_user_ref: int):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
stmt = (
insert(BlockedUser)
.values(user_id=u.id, tg_id=u.tg_id)
.on_conflict_do_nothing(index_elements=[BlockedUser.user_id])
)
await session.execute(stmt)
await session.commit()
async def save_blocked_user_ids(session: AsyncSession, tg_ids: list[int]) -> None:
"""Вставка списка tg_id в таблицу BlockedUser батчами по 500. Вызывать только из основного event loop."""
"""Вставка списка telegram id в таблицу blocked_users батчами по 500."""
if not tg_ids:
return
batch_size = 500
total = 0
for i in range(0, len(tg_ids), batch_size):
batch = tg_ids[i : i + batch_size]
values = [{"tg_id": tg_id} for tg_id in batch]
stmt = insert(BlockedUser).values(values).on_conflict_do_nothing(index_elements=[BlockedUser.tg_id])
res = await session.execute(select(User.id, User.tg_id).where(User.tg_id.in_(batch)))
rows = res.all()
uid_by_tg = {int(tgid): int(uid) for uid, tgid in rows if tgid is not None}
values = [
{"user_id": uid_by_tg[int(tg)], "tg_id": int(tg)}
for tg in batch
if int(tg) in uid_by_tg
]
if not values:
continue
stmt = insert(BlockedUser).values(values).on_conflict_do_nothing(index_elements=[BlockedUser.user_id])
await session.execute(stmt)
await session.commit()
total += len(batch)
logger.info(f"📝 Добавлено {total} пользователей в blocked_users")
total += len(values)
logger.info(f"📝 Добавлено до {total} пользователей в blocked_users")
+112 -65
View File
@@ -1,9 +1,9 @@
from datetime import datetime
from sqlalchemy import case, delete, func, insert, select, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy import case, delete, func, insert, or_, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from database.access.resolution import resolve_user_optional
from database.models import Coupon, CouponUsage
from logger import logger
@@ -19,48 +19,42 @@ async def create_coupon(
max_discount_amount: int | None = None,
min_order_amount: int | None = None,
) -> bool:
try:
exists = await session.scalar(select(Coupon.id).where(Coupon.code == code))
if exists:
logger.warning(f"[Coupon] ⚠️ Купон с кодом {code} уже существует.")
exists = await session.scalar(select(Coupon.id).where(Coupon.code == code))
if exists:
logger.warning(f"[Coupon] ⚠️ Купон с кодом {code} уже существует.")
return False
if percent is not None:
try:
percent_value = int(percent)
except (TypeError, ValueError):
logger.warning(f"[Coupon] ⚠️ Некорректный процент для купона {code}.")
return False
if percent is not None:
try:
percent_value = int(percent)
except (TypeError, ValueError):
logger.warning(f"[Coupon] ⚠️ Некорректный процент для купона {code}.")
return False
if percent_value <= 0 or percent_value > 100:
logger.warning(f"[Coupon] ⚠️ процент должен быть в диапазоне 1..100 для купона {code}.")
return False
if percent_value <= 0 or percent_value > 100:
logger.warning(f"[Coupon] ⚠️ процент должен быть в диапазоне 1..100 для купона {code}.")
return False
if (amount or 0) > 0 or (days or 0) > 0:
logger.warning(f"[Coupon] ⚠️ Купон {code} не может одновременно иметь percent и amount/days.")
return False
if (amount or 0) > 0 or (days or 0) > 0:
logger.warning(f"[Coupon] ⚠️ Купон {code} не может одновременно иметь percent и amount/days.")
return False
await session.execute(
insert(Coupon).values(
code=code,
amount=int(amount) if amount is not None else 0,
usage_limit=usage_limit,
usage_count=0,
is_used=False,
days=days,
new_users_only=new_users_only,
percent=percent,
max_discount_amount=max_discount_amount,
min_order_amount=min_order_amount,
)
await session.execute(
insert(Coupon).values(
code=code,
amount=int(amount) if amount is not None else 0,
usage_limit=usage_limit,
usage_count=0,
is_used=False,
days=days,
new_users_only=new_users_only,
percent=percent,
max_discount_amount=max_discount_amount,
min_order_amount=min_order_amount,
)
await session.commit()
logger.info(f"[Coupon] ✅ Купон {code} успешно создан.")
return True
except SQLAlchemyError as e:
await session.rollback()
logger.error(f"[Coupon] ❌ Ошибка при создании купона {code}: {e}")
return False
)
logger.info(f"[Coupon] ✅ Купон {code} успешно создан.")
return True
async def get_coupon_by_code(session: AsyncSession, code: str) -> Coupon | None:
@@ -69,6 +63,15 @@ async def get_coupon_by_code(session: AsyncSession, code: str) -> Coupon | None:
return result.scalar_one_or_none()
async def get_coupon_by_code_ci(session: AsyncSession, code: str) -> Coupon | None:
normalized = str(code or "").strip()
if not normalized:
return None
stmt = select(Coupon).where(func.lower(Coupon.code) == normalized.lower())
result = await session.execute(stmt)
return result.scalar_one_or_none()
async def get_all_coupons(session: AsyncSession, page: int = 1, per_page: int = 10) -> dict:
offset = (page - 1) * per_page
@@ -99,45 +102,89 @@ async def delete_coupon(session: AsyncSession, code: str) -> bool:
await session.execute(delete(CouponUsage).where(CouponUsage.coupon_id == coupon.id))
await session.delete(coupon)
await session.commit()
logger.info(f"🗑 Купон {code} удалён вместе с его использованиями")
return True
async def _coupon_usage_billing_match(session: AsyncSession, legacy_user_ref: int):
u = await resolve_user_optional(session, legacy_user_ref)
if u is not None:
opts = [CouponUsage.user_id == u.id]
if u.tg_id is not None:
opts.append(CouponUsage.tg_id == u.tg_id)
return or_(*opts)
return or_(CouponUsage.user_id == legacy_user_ref, CouponUsage.tg_id == legacy_user_ref)
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())
await session.execute(stmt)
await session.commit()
logger.info(f"✅ Купон {coupon_id} использован пользователем {user_id}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при сохранении использования купона: {e}")
await session.rollback()
raise
u = await resolve_user_optional(session, user_id)
uid = u.id if u is not None else user_id
stmt = insert(CouponUsage).values(
coupon_id=coupon_id,
user_id=uid,
tg_id=u.tg_id if u is not None else None,
used_at=datetime.utcnow(),
)
await session.execute(stmt)
logger.info(f"✅ Купон {coupon_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)
async def check_coupon_usage(session: AsyncSession, coupon_id: int, legacy_user_ref: int) -> bool:
m = await _coupon_usage_billing_match(session, legacy_user_ref)
stmt = select(CouponUsage).where(CouponUsage.coupon_id == coupon_id).where(m)
result = await session.execute(stmt)
return result.scalar_one_or_none() is not None
async def has_any_coupon_usage(session: AsyncSession, legacy_user_ref: int) -> bool:
m = await _coupon_usage_billing_match(session, legacy_user_ref)
stmt = select(CouponUsage.coupon_id).where(m).limit(1)
result = await session.execute(stmt)
return result.first() is not None
async def update_coupon_usage_count(session: AsyncSession, coupon_id: int):
try:
await session.execute(
update(Coupon)
.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),
)
await session.execute(
update(Coupon)
.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),
)
await session.commit()
logger.info(f"🔁 Обновлён счётчик купона {coupon_id}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при обновлении купона {coupon_id}: {e}")
await session.rollback()
raise
)
logger.info(f"🔁 Обновлён счётчик купона {coupon_id}")
async def mark_coupon_used(session: AsyncSession, coupon_id: int, legacy_user_ref: int):
u = await resolve_user_optional(session, legacy_user_ref)
uid = u.id if u is not None else legacy_user_ref
match = [CouponUsage.user_id == int(uid)]
if u is not None and u.tg_id is not None:
match.append(CouponUsage.tg_id == int(u.tg_id))
existing = await session.execute(
select(CouponUsage).where(
CouponUsage.coupon_id == int(coupon_id),
or_(*match),
)
)
if existing.scalar_one_or_none() is not None:
return
await session.execute(
insert(CouponUsage).values(
coupon_id=coupon_id,
user_id=uid,
tg_id=u.tg_id if u is not None else None,
used_at=datetime.utcnow(),
)
)
await session.execute(
update(Coupon)
.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),
)
)
def apply_percent_coupon(price_rub: int, coupon: Coupon) -> tuple[int, int]:
+6
View File
@@ -20,6 +20,12 @@ if USE_PGBOUNCER and "+asyncpg" in DATABASE_URL:
_pool_recycle = 60 if USE_PGBOUNCER else 300
_QUERY_TIMEOUT_SEC = 30
if "+asyncpg" in _db_url:
_connect_args.setdefault("command_timeout", _QUERY_TIMEOUT_SEC)
_connect_args.setdefault("timeout", _QUERY_TIMEOUT_SEC)
def _create_engine():
return create_async_engine(
+95 -31
View File
@@ -1,17 +1,17 @@
from datetime import datetime
from sqlalchemy import insert
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy import func, insert, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from database.models import Gift
from database.access.resolution import resolve_user_optional
from database.models import Gift, GiftUsage
from logger import logger
async def store_gift_link(
session: AsyncSession,
gift_id: str,
sender_tg_id: int,
sender_legacy_ref: int,
selected_months: int,
expiry_time: datetime,
gift_link: str,
@@ -21,33 +21,97 @@ async def store_gift_link(
selected_device_limit: int | None = None,
selected_traffic_gb: int | None = None,
selected_price_rub: int | None = None,
):
try:
stmt = insert(Gift).values(
) -> bool:
u = await resolve_user_optional(session, sender_legacy_ref)
if u is None:
raise ValueError(f"sender not found for gift: {sender_legacy_ref}")
stmt = insert(Gift).values(
gift_id=gift_id,
sender_user_id=u.id,
sender_tg_id=u.tg_id,
recipient_user_id=None,
selected_months=selected_months,
expiry_time=expiry_time,
gift_link=gift_link,
created_at=datetime.utcnow(),
is_used=False,
tariff_id=tariff_id,
is_unlimited=is_unlimited,
max_usages=max_usages,
selected_device_limit=selected_device_limit,
selected_traffic_gb=selected_traffic_gb,
selected_price_rub=selected_price_rub,
)
await session.execute(stmt)
logger.info(
f"🎁 Подарок {gift_id} сохранён "
f"(tariff_id={tariff_id}, max_usages={max_usages}, "
f"device={selected_device_limit}, traffic={selected_traffic_gb}, price={selected_price_rub})"
)
return True
async def get_gift_locked(session: AsyncSession, gift_id: str) -> Gift | None:
"""SELECT FOR UPDATE по gift_id — берёт row-lock для atomic redemption.
Используется в `services.gifts.redeem_gift` чтобы два параллельных запроса
на активацию одного и того же подарка не смогли обойти проверку `is_used`.
"""
result = await session.execute(select(Gift).where(Gift.gift_id == gift_id).with_for_update())
return result.scalar_one_or_none()
async def get_gift_usage(session: AsyncSession, gift_id: str, user_id: int) -> GiftUsage | None:
"""Возвращает запись об использовании подарка конкретным пользователем, если есть."""
result = await session.execute(
select(GiftUsage).where(
GiftUsage.gift_id == gift_id,
GiftUsage.user_id == user_id,
)
)
return result.scalar_one_or_none()
async def count_gift_usages(session: AsyncSession, gift_id: str) -> int:
"""Сколько раз подарок был активирован (для `is_unlimited=False` с лимитом)."""
result = await session.execute(
select(func.count()).select_from(GiftUsage).where(GiftUsage.gift_id == gift_id)
)
return int(result.scalar_one() or 0)
async def record_gift_usage(
session: AsyncSession,
gift_id: str,
user_id: int,
tg_id: int | None,
) -> None:
"""Вставляет запись о применении подарка. Композитный ключ (gift_id, user_id)."""
await session.execute(
insert(GiftUsage).values(
gift_id=gift_id,
sender_tg_id=sender_tg_id,
recipient_tg_id=None,
selected_months=selected_months,
expiry_time=expiry_time,
gift_link=gift_link,
created_at=datetime.utcnow(),
is_used=False,
tariff_id=tariff_id,
is_unlimited=is_unlimited,
max_usages=max_usages,
selected_device_limit=selected_device_limit,
selected_traffic_gb=selected_traffic_gb,
selected_price_rub=selected_price_rub,
user_id=user_id,
tg_id=tg_id,
)
await session.execute(stmt)
await session.commit()
logger.info(
f"🎁 Подарок {gift_id} сохранён "
f"(tariff_id={tariff_id}, max_usages={max_usages}, "
f"device={selected_device_limit}, traffic={selected_traffic_gb}, price={selected_price_rub})"
)
async def mark_gift_fully_redeemed(
session: AsyncSession,
gift_id: str,
recipient_user_id: int,
recipient_tg_id: int | None,
) -> None:
"""Помечает подарок как полностью использованный (is_used=True) и фиксирует получателя.
Вызывается для non-unlimited подарков, когда набрали max_usages.
"""
await session.execute(
update(Gift)
.where(Gift.gift_id == gift_id)
.values(
is_used=True,
recipient_user_id=recipient_user_id,
recipient_tg_id=recipient_tg_id,
)
return True
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при сохранении подарка {gift_id}: {e}")
await session.rollback()
raise
)
+4 -4
View File
@@ -8,17 +8,17 @@ from database.models import Key, Payment, User
async def get_hot_leads(session: AsyncSession):
now_ms = func.extract("epoch", func.now()) * 1000
sub_active = select(Key.tg_id).where(Key.expiry_time > now_ms).distinct()
sub_active = select(Key.user_id).where(Key.expiry_time > now_ms).distinct()
stmt = (
select(Payment.tg_id)
.join(User, User.tg_id == Payment.tg_id)
select(Payment.user_id)
.join(User, User.id == Payment.user_id)
.distinct()
.where(User.trial == 1)
.where(Payment.amount > 0)
.where(Payment.status == "success")
.where(Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED))
.where(~Payment.tg_id.in_(sub_active))
.where(~Payment.user_id.in_(sub_active))
)
result = await session.execute(stmt)
+208 -10
View File
@@ -3,11 +3,12 @@ import secrets
from datetime import datetime, timedelta
import bcrypt
from sqlalchemy import select
from sqlalchemy import delete, func, select, text, update
from sqlalchemy.ext.asyncio import AsyncSession
from config import API_TOKEN_TTL_DAYS
from core.executor import run_cpu, run_io
from database.access.tg_mirror import refresh_tg_mirrors_for_user
from database.models import Admin, Identity, User
@@ -57,7 +58,6 @@ async def create_identity(
await session.flush()
if tg_id:
await session.execute(User.__table__.update().where(User.tg_id == tg_id).values(identity_id=identity.id))
await session.commit()
await session.refresh(identity)
return identity
@@ -93,7 +93,6 @@ async def issue_token_for_identity(session: AsyncSession, identity: Identity) ->
token = generate_token()
identity.api_token_hash = await run_io(hash_token, token)
identity.token_issued_at = datetime.utcnow()
await session.commit()
await session.refresh(identity)
return token
@@ -116,7 +115,6 @@ async def create_identity_with_token(
identity = await create_identity(session, email=email, tg_id=tg_id)
if password:
identity.password_hash = await run_cpu(hash_password, password)
await session.commit()
await session.refresh(identity)
token = await issue_token_for_identity(session, identity)
return identity, token
@@ -146,10 +144,209 @@ async def login_by_email(session: AsyncSession, email: str, password: str) -> tu
return identity, token
async def resolve_tg_id(session: AsyncSession, identity_id: str) -> int | None:
"""По identity_id возвращает tg_id, если привязан."""
async def set_initial_password(
session: AsyncSession,
identity_id: str,
password: str,
) -> Identity | None:
identity = await get_identity_by_id(session, identity_id)
return identity.tg_id if identity else None
if not identity or identity.password_hash:
return None
identity.password_hash = await run_cpu(hash_password, password)
await session.refresh(identity)
return identity
async def set_password_for_identity(
session: AsyncSession,
identity_id: str,
new_password: str,
) -> Identity | None:
identity = await get_identity_by_id(session, identity_id)
if not identity:
return None
identity.password_hash = await run_cpu(hash_password, new_password)
await session.refresh(identity)
return identity
async def change_identity_password(
session: AsyncSession,
identity_id: str,
current_password: str,
new_password: str,
) -> str | None:
"""Возвращает None при успехе, иначе код: no_password | wrong_password."""
identity = await get_identity_by_id(session, identity_id)
if not identity:
return "wrong_password"
if not identity.password_hash:
return "no_password"
if not await run_cpu(check_password, current_password, identity.password_hash):
return "wrong_password"
identity.password_hash = await run_cpu(hash_password, new_password)
await session.refresh(identity)
return None
async def ensure_billing_user_for_identity(session: AsyncSession, identity: Identity) -> int:
from database.users import add_user, check_user_exists
if identity.tg_id is not None:
tid = int(identity.tg_id)
if not await check_user_exists(session, tid):
await add_user(session, tid)
ur = await session.execute(select(User).where(User.tg_id == tid).limit(1))
u = ur.scalar_one()
await session.execute(update(User).where(User.id == u.id).values(identity_id=identity.id))
return int(u.id)
res = await session.execute(select(User).where(User.identity_id == identity.id))
row = res.scalars().first()
if row is not None:
return int(row.id)
new_u = User(identity_id=identity.id, tg_id=None)
session.add(new_u)
await session.flush()
return int(new_u.id)
async def merge_billing_user_into_telegram(session: AsyncSession, identity_id: str, telegram_tg_id: int) -> None:
from database.models import (
CouponUsage,
Gift,
GiftUsage,
Key,
Notification,
Payment,
Referral,
ScheduledBroadcast,
TemporaryData,
)
from database.access.resolution import resolve_user_optional
from database.users import invalidate_balance_cache, invalidate_profile_cache, update_balance
res = await session.execute(select(User).where(User.identity_id == identity_id))
rows = res.scalars().all()
if not rows:
return
billing = rows[0]
src_uid = int(billing.id)
dst_tg = int(telegram_tg_id)
if billing.tg_id is not None and int(billing.tg_id) > 0:
return
dst_u = await resolve_user_optional(session, dst_tg)
if dst_u is None:
new_u = User(
tg_id=dst_tg,
identity_id=identity_id,
username=billing.username,
first_name=billing.first_name,
last_name=billing.last_name,
language_code=billing.language_code,
is_bot=billing.is_bot or False,
balance=float(billing.balance or 0.0),
trial=int(billing.trial or 0),
preferred_currency=billing.preferred_currency or "RUB",
source_code=billing.source_code,
)
session.add(new_u)
await session.flush()
dst_uid = int(new_u.id)
else:
dst_uid = int(dst_u.id)
bal = float(billing.balance or 0.0)
if bal:
await update_balance(session, dst_uid, bal)
st = int(billing.trial or 0)
dt_r = await session.execute(select(User.trial).where(User.id == dst_uid))
dt_val = dt_r.scalar_one_or_none()
if dt_val is not None and st > int(dt_val or 0):
await session.execute(update(User).where(User.id == dst_uid).values(trial=st))
await session.execute(update(Key).where(Key.user_id == src_uid).values(user_id=dst_uid))
await session.execute(update(Payment).where(Payment.user_id == src_uid).values(user_id=dst_uid))
await session.execute(
text(
"DELETE FROM notifications AS n1 USING notifications AS n2 "
"WHERE n1.user_id = :src AND n2.user_id = :dst AND n1.notification_type = n2.notification_type"
),
{"src": src_uid, "dst": dst_uid},
)
await session.execute(update(Notification).where(Notification.user_id == src_uid).values(user_id=dst_uid))
await session.execute(update(Gift).where(Gift.sender_user_id == src_uid).values(sender_user_id=dst_uid))
await session.execute(
update(Gift).where(Gift.recipient_user_id == src_uid).values(recipient_user_id=dst_uid)
)
await session.execute(
text(
"DELETE FROM gift_usages AS g1 USING gift_usages AS g2 "
"WHERE g1.user_id = :src AND g2.user_id = :dst AND g1.gift_id = g2.gift_id"
),
{"src": src_uid, "dst": dst_uid},
)
await session.execute(update(GiftUsage).where(GiftUsage.user_id == src_uid).values(user_id=dst_uid))
await session.execute(
text(
"DELETE FROM coupon_usages AS c1 USING coupon_usages AS c2 "
"WHERE c1.user_id = :src AND c2.user_id = :dst AND c1.coupon_id = c2.coupon_id"
),
{"src": src_uid, "dst": dst_uid},
)
await session.execute(update(CouponUsage).where(CouponUsage.user_id == src_uid).values(user_id=dst_uid))
await session.execute(update(TemporaryData).where(TemporaryData.user_id == src_uid).values(user_id=dst_uid))
await session.execute(
update(ScheduledBroadcast)
.where(ScheduledBroadcast.created_by_user_id == src_uid)
.values(created_by_user_id=dst_uid)
)
await session.execute(
text(
"DELETE FROM referrals AS r1 USING referrals AS r2 "
"WHERE r1.referred_user_id = :src AND r2.referred_user_id = :dst "
"AND r1.referrer_user_id = r2.referrer_user_id"
),
{"src": src_uid, "dst": dst_uid},
)
await session.execute(
text(
"DELETE FROM referrals AS r1 USING referrals AS r2 "
"WHERE r1.referrer_user_id = :src AND r2.referrer_user_id = :dst "
"AND r1.referred_user_id = r2.referred_user_id"
),
{"src": src_uid, "dst": dst_uid},
)
await session.execute(
update(Referral).where(Referral.referred_user_id == src_uid).values(referred_user_id=dst_uid)
)
await session.execute(
update(Referral).where(Referral.referrer_user_id == src_uid).values(referrer_user_id=dst_uid)
)
await refresh_tg_mirrors_for_user(session, dst_uid)
await session.execute(delete(User).where(User.id == src_uid))
await session.execute(update(User).where(User.id == dst_uid).values(identity_id=identity_id))
await invalidate_balance_cache(src_uid)
await invalidate_profile_cache(src_uid)
await invalidate_balance_cache(dst_uid)
await invalidate_profile_cache(dst_uid)
async def resolve_tg_id(session: AsyncSession, identity_id: str) -> int | None:
"""По identity_id возвращает внутренний user id (users.id) для биллинга и ключей."""
identity = await get_identity_by_id(session, identity_id)
if not identity:
return None
return await ensure_billing_user_for_identity(session, identity)
async def attach_email(session: AsyncSession, identity_id: str, email: str) -> Identity | None:
@@ -164,7 +361,6 @@ async def attach_email(session: AsyncSession, identity_id: str, email: str) -> I
if existing and existing.id != identity_id:
return None
identity.email = email_clean
await session.commit()
await session.refresh(identity)
return identity
@@ -177,12 +373,15 @@ async def attach_telegram(session: AsyncSession, identity_id: str, tg_id: int) -
existing = await get_identity_by_tg_id(session, tg_id)
if existing and existing.id != identity_id:
return None
await merge_billing_user_into_telegram(session, identity_id, tg_id)
identity = await get_identity_by_id(session, identity_id)
if not identity:
return None
identity.tg_id = tg_id
admin_row = await session.execute(select(Admin).where(Admin.tg_id == tg_id))
if admin_row.scalar_one_or_none():
identity.is_admin = True
await session.execute(User.__table__.update().where(User.tg_id == tg_id).values(identity_id=identity_id))
await session.commit()
await session.refresh(identity)
return identity
@@ -196,6 +395,5 @@ async def get_or_create_identity_for_tg(session: AsyncSession, tg_id: int) -> Id
session.add(identity)
await session.flush()
await session.execute(User.__table__.update().where(User.tg_id == tg_id).values(identity_id=identity.id))
await session.commit()
await session.refresh(identity)
return identity
+36 -41
View File
@@ -6,7 +6,6 @@ from datetime import datetime
from itertools import cycle
from sqlalchemy import select
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from config import USE_COUNTRY_SELECTION
@@ -76,53 +75,49 @@ async def import_keys_from_3xui_db(db_path: str, session: AsyncSession) -> tuple
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(),
)
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 as e:
await session.rollback()
raise RuntimeError(f"Ошибка при импорте пользователя tg_id={tg_id}") from e
)
await session.flush()
user_row = await session.execute(select(User.id).where(User.tg_id == tg_id))
bill_uid = user_row.scalar_one()
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,
)
session.add(
Key(
user_id=bill_uid,
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 as e:
await session.rollback()
raise RuntimeError(f"Ошибка при импорте ключа client_id={client_id}") from e
)
imported += 1
await session.commit()
return imported, skipped
+325 -124
View File
@@ -1,9 +1,8 @@
import asyncio
from datetime import datetime
from datetime import UTC, datetime
from types import SimpleNamespace
from sqlalchemy import delete, func, select, text, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from core.cache_config import (
@@ -12,6 +11,7 @@ from core.cache_config import (
KEYS_LIST_CACHE_TTL_SEC,
)
from core.redis_cache import cache_delete, cache_get, cache_key, cache_set
from database.access.resolution import resolve_user_optional
from database.models import Key, Tariff, User
from database.users import invalidate_profile_cache, invalidate_user_snapshot
from logger import logger
@@ -25,10 +25,22 @@ async def invalidate_key_email(client_id: str) -> None:
await cache_delete(cache_key("key_email", client_id))
async def invalidate_keys_list(tg_id: int) -> None:
await cache_delete(cache_key("keys_list", tg_id))
await cache_delete(cache_key("key_count", tg_id))
await invalidate_profile_cache(tg_id)
async def _purge_keys_cache_ids(*ids: int) -> None:
for i in ids:
await cache_delete(cache_key("keys_list", i))
await cache_delete(cache_key("key_count", i))
await invalidate_profile_cache(i)
async def invalidate_keys_list(session: AsyncSession, legacy_user_ref: int) -> None:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
await _purge_keys_cache_ids(legacy_user_ref)
return
if u.tg_id is not None:
await _purge_keys_cache_ids(u.id, u.tg_id)
else:
await _purge_keys_cache_ids(u.id)
async def invalidate_key_details_by_client_id(session: AsyncSession, client_id: str) -> None:
@@ -45,7 +57,7 @@ async def invalidate_key_details_by_client_id(session: AsyncSession, client_id:
async def store_key(
session: AsyncSession,
tg_id: int,
legacy_user_ref: int,
client_id: str,
email: str,
expiry_time: int,
@@ -61,71 +73,72 @@ async def store_key(
current_traffic_limit: int | None = None,
):
"""Сохраняет или обновляет ключ подписки."""
try:
exists = await session.execute(select(Key).where(Key.tg_id == tg_id, Key.client_id == client_id))
existing_key = exists.scalar_one_or_none()
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
raise ValueError(f"Пользователь не найден для ключа: {legacy_user_ref}")
uid = u.id
exists = await session.execute(select(Key).where(Key.user_id == uid, Key.client_id == client_id))
existing_key = exists.scalar_one_or_none()
if existing_key:
values: dict = {
"email": email,
"expiry_time": expiry_time,
"key": key,
"server_id": server_id,
"remnawave_link": remnawave_link,
"tariff_id": tariff_id,
"alias": alias,
}
if existing_key:
values: dict = {
"email": email,
"expiry_time": expiry_time,
"key": key,
"server_id": server_id,
"remnawave_link": remnawave_link,
"tariff_id": tariff_id,
"alias": alias,
"tg_id": u.tg_id,
}
if selected_device_limit is not None:
values["selected_device_limit"] = selected_device_limit
if selected_traffic_limit is not None:
values["selected_traffic_limit"] = selected_traffic_limit
if selected_price_rub is not None:
values["selected_price_rub"] = selected_price_rub
if current_device_limit is not None:
values["current_device_limit"] = current_device_limit
if current_traffic_limit is not None:
values["current_traffic_limit"] = current_traffic_limit
if selected_device_limit is not None:
values["selected_device_limit"] = selected_device_limit
if selected_traffic_limit is not None:
values["selected_traffic_limit"] = selected_traffic_limit
if selected_price_rub is not None:
values["selected_price_rub"] = selected_price_rub
if current_device_limit is not None:
values["current_device_limit"] = current_device_limit
if current_traffic_limit is not None:
values["current_traffic_limit"] = current_traffic_limit
await session.execute(update(Key).where(Key.tg_id == tg_id, Key.client_id == client_id).values(**values))
logger.info(f"[Store Key] Ключ обновлён: tg_id={tg_id}, client_id={client_id}, server_id={server_id}")
else:
if current_device_limit is None:
current_device_limit = selected_device_limit
if current_traffic_limit is None:
current_traffic_limit = selected_traffic_limit
await session.execute(update(Key).where(Key.user_id == uid, Key.client_id == client_id).values(**values))
logger.info(f"[Store Key] Ключ обновлён: user_id={uid}, client_id={client_id}, server_id={server_id}")
else:
if current_device_limit is None:
current_device_limit = selected_device_limit
if current_traffic_limit is None:
current_traffic_limit = selected_traffic_limit
new_key = Key(
tg_id=tg_id,
client_id=client_id,
email=email,
created_at=int(datetime.utcnow().timestamp() * 1000),
expiry_time=expiry_time,
key=key,
server_id=server_id,
remnawave_link=remnawave_link,
tariff_id=tariff_id,
alias=alias,
selected_device_limit=selected_device_limit,
selected_traffic_limit=selected_traffic_limit,
selected_price_rub=selected_price_rub,
current_device_limit=current_device_limit,
current_traffic_limit=current_traffic_limit,
)
add_result = session.add(new_key)
if asyncio.iscoroutine(add_result):
await add_result
logger.info(f"[Store Key] Ключ создан: tg_id={tg_id}, client_id={client_id}, server_id={server_id}")
new_key = Key(
user_id=uid,
tg_id=u.tg_id,
client_id=client_id,
email=email,
created_at=int(datetime.now(UTC).timestamp() * 1000),
expiry_time=expiry_time,
key=key,
server_id=server_id,
remnawave_link=remnawave_link,
tariff_id=tariff_id,
alias=alias,
selected_device_limit=selected_device_limit,
selected_traffic_limit=selected_traffic_limit,
selected_price_rub=selected_price_rub,
current_device_limit=current_device_limit,
current_traffic_limit=current_traffic_limit,
)
add_result = session.add(new_key)
if asyncio.iscoroutine(add_result):
await add_result
logger.info(f"[Store Key] Ключ создан: user_id={uid}, client_id={client_id}, server_id={server_id}")
await session.commit()
invalidate_user_snapshot(tg_id)
await invalidate_keys_list(tg_id)
await invalidate_key_details(email)
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при сохранении ключа: {e}")
await session.rollback()
raise
invalidate_user_snapshot(uid)
if u.tg_id is not None:
invalidate_user_snapshot(u.tg_id)
await invalidate_keys_list(session, uid)
await invalidate_key_details(email)
def _key_to_cache_dict(k: Key) -> dict:
@@ -143,12 +156,16 @@ def _key_to_cache_dict(k: Key) -> dict:
}
async def get_keys(session: AsyncSession, tg_id: int):
ckey = cache_key("keys_list", tg_id)
async def get_keys(session: AsyncSession, legacy_user_ref: int):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return []
uid = u.id
ckey = cache_key("keys_list", uid)
cached = await cache_get(ckey)
if isinstance(cached, list):
return [SimpleNamespace(**d) for d in cached]
result = await session.execute(select(Key).where(Key.tg_id == tg_id))
result = await session.execute(select(Key).where(Key.user_id == uid))
rows = result.scalars().all()
serialized = [_key_to_cache_dict(k) for k in rows]
await cache_set(ckey, serialized, KEYS_LIST_CACHE_TTL_SEC)
@@ -160,24 +177,33 @@ async def get_all_keys(session: AsyncSession):
return result.scalars().all()
async def get_key_by_server(session: AsyncSession, tg_id: int, client_id: str):
stmt = select(Key).where(Key.tg_id == tg_id, Key.client_id == client_id)
async def get_key_by_server(session: AsyncSession, legacy_user_ref: int, client_id: str):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return None
stmt = select(Key).where(Key.user_id == u.id, Key.client_id == client_id)
result = await session.execute(stmt)
return result.scalar_one_or_none()
async def get_key_by_email(session: AsyncSession, email: str, tg_id: int | None = None) -> Key | None:
async def get_key_by_email(session: AsyncSession, email: str, legacy_user_ref: int | None = None) -> Key | None:
stmt = select(Key).where(Key.email == email)
if tg_id is not None:
stmt = stmt.where(Key.tg_id == tg_id)
if legacy_user_ref is not None:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return None
stmt = stmt.where(Key.user_id == u.id)
result = await session.execute(stmt.limit(1))
return result.scalar_one_or_none()
async def get_key_by_client_id(session: AsyncSession, client_id: str, tg_id: int | None = None) -> Key | None:
async def get_key_by_client_id(session: AsyncSession, client_id: str, legacy_user_ref: int | None = None) -> Key | None:
stmt = select(Key).where(Key.client_id == client_id)
if tg_id is not None:
stmt = stmt.where(Key.tg_id == tg_id)
if legacy_user_ref is not None:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return None
stmt = stmt.where(Key.user_id == u.id)
result = await session.execute(stmt.limit(1))
return result.scalar_one_or_none()
@@ -218,15 +244,15 @@ async def get_key_details(session: AsyncSession, email: str) -> dict | None:
if isinstance(cached, dict):
return cached
stmt = select(Key, User).join(User, Key.tg_id == User.tg_id).where(Key.email == email)
stmt = select(Key, User).join(User, Key.user_id == User.id).where(Key.email == email)
result = await session.execute(stmt)
row = result.first()
if not row:
return None
key, user = row
expiry_date = datetime.utcfromtimestamp(key.expiry_time / 1000)
current_date = datetime.utcnow()
expiry_date = datetime.fromtimestamp(key.expiry_time / 1000, UTC)
current_date = datetime.now(UTC)
time_left = expiry_date - current_date
if time_left.total_seconds() <= 0:
@@ -267,39 +293,155 @@ async def get_key_details(session: AsyncSession, email: str) -> dict | None:
return out
async def get_key_count(session: AsyncSession, tg_id: int) -> int:
cached = await cache_get(cache_key("key_count", tg_id))
async def get_key_count(session: AsyncSession, legacy_user_ref: int) -> int:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return 0
uid = u.id
cached = await cache_get(cache_key("key_count", uid))
if cached is not None:
try:
return int(cached)
except (TypeError, ValueError):
pass
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.user_id == uid))
count = result.scalar() or 0
await cache_set(cache_key("key_count", tg_id), count, KEY_COUNT_CACHE_TTL_SEC)
await cache_set(cache_key("key_count", uid), count, KEY_COUNT_CACHE_TTL_SEC)
return count
async def delete_key(session: AsyncSession, identifier: int | str, commit: bool = True):
tg_id_for_cache = None
async def get_key_by_user_and_email(session: AsyncSession, user_id: int, email: str) -> Key | None:
"""Возвращает ORM-объект Key по паре (users.id, email) или None."""
result = await session.execute(
select(Key).where(Key.user_id == int(user_id), Key.email == email)
)
return result.scalar_one_or_none()
async def delete_key_by_user_and_email(session: AsyncSession, user_id: int, email: str) -> None:
"""Удаляет ключ по паре (users.id, email). Commit — ответственность caller'а."""
await session.execute(
delete(Key).where(Key.user_id == int(user_id), Key.email == email)
)
async def get_user_keys_with_servers_by_email(
session: AsyncSession, user_id: int, email: str
) -> list[tuple[str, str, dict]]:
"""Возвращает ключи пользователя + инфо о серверах (join Key × Server).
Каждый элемент — ``(client_id, server_id, server_info_dict)``. Join
делается по (Key.server_id == Server.server_name OR Server.cluster_name),
чтобы поддержать и country-mode (server_id = cluster), и cluster-mode
(server_id = server_name).
Используется в ``services.operations.traffic.get_user_traffic``.
"""
from sqlalchemy import or_
from database.models import Server
join_cond = or_(
Key.server_id == Server.server_name,
Key.server_id == Server.cluster_name,
)
result = await session.execute(
select(Key.client_id, Key.server_id, Server)
.select_from(Key)
.join(Server, join_cond)
.where(Server.enabled.is_(True), Key.user_id == int(user_id), Key.email == email)
)
rows = []
for client_id, server_id, server in result.all():
rows.append((
client_id,
server_id,
{
"server_name": server.server_name,
"cluster_name": server.cluster_name,
"api_url": server.api_url,
"panel_type": server.panel_type,
},
))
return rows
async def get_key_client_id_by_email_and_server(
session: AsyncSession, email: str, server_id: str
) -> str | None:
"""Возвращает ``client_id`` первого ключа для пары (email, server_id).
Используется для remnawave traffic reset, где нам нужен только client_id,
без остальных полей ключа.
"""
result = await session.execute(
select(Key.client_id)
.where(Key.email == email, Key.server_id == server_id)
.limit(1)
)
return result.scalar()
async def count_keys_by_server_id(session: AsyncSession, server_id: str) -> int:
"""Сколько всего ключей привязано к указанному server_id (кластеру или серверу).
Используется для проверки max_keys лимита. ``server_id`` — строка
(у ``keys.server_id`` колонка типа String, содержит либо cluster_name,
либо server_name в зависимости от страны/кластера).
"""
result = await session.execute(
select(func.count()).select_from(Key).where(Key.server_id == server_id)
)
return int(result.scalar() or 0)
async def get_all_key_server_ids(session: AsyncSession) -> list[str]:
"""Список всех ``server_id`` из таблицы keys (с повторениями).
Используется в ``services.clusters.select_cluster`` для подсчёта загрузки
кластеров. Возвращаем только server_id строки без подгрузки остальных
полей, чтобы не тянуть сотни мегабайт для огромных deployments.
"""
result = await session.execute(select(Key.server_id))
return [row[0] for row in result.all() if row[0] is not None]
async def count_active_keys_for_user(session: AsyncSession, user_id: int) -> int:
"""Количество незамороженных ключей у пользователя (по internal users.id).
Отличается от `get_key_count`: не кэшируется и явно исключает замороженные.
Используется в проверке "новый пользователь" для купонных правил.
"""
result = await session.execute(
select(func.count())
.select_from(Key)
.where(Key.user_id == int(user_id), Key.is_frozen.is_(False))
)
return int(result.scalar() or 0)
async def delete_key(session: AsyncSession, identifier: int | str):
legacy_for_cache = None
email_for_cache = None
if isinstance(identifier, str):
res = await session.execute(
select(Key.tg_id, Key.email).where(Key.client_id == identifier).limit(1)
select(Key.user_id, Key.email).where(Key.client_id == identifier).limit(1)
)
row = res.first()
if row:
tg_id_for_cache, email_for_cache = row[0], row[1]
legacy_for_cache, email_for_cache = row[0], row[1]
await cache_delete(cache_key("key_email", identifier))
await session.execute(delete(Key).where(Key.client_id == identifier))
else:
tg_id_for_cache = identifier
stmt = delete(Key).where(Key.tg_id == identifier if isinstance(identifier, int) else Key.client_id == identifier)
await session.execute(stmt)
if commit:
await session.commit()
if tg_id_for_cache is not None:
invalidate_user_snapshot(tg_id_for_cache)
await invalidate_keys_list(tg_id_for_cache)
u = await resolve_user_optional(session, identifier)
if u is None:
logger.info(f"Ключ не удалён: пользователь {identifier} не найден")
return
legacy_for_cache = u.id
await session.execute(delete(Key).where(Key.user_id == u.id))
if legacy_for_cache is not None:
invalidate_user_snapshot(legacy_for_cache)
await invalidate_keys_list(session, legacy_for_cache)
if email_for_cache is not None:
await invalidate_key_details(str(email_for_cache))
logger.info(f"Ключ с идентификатором {identifier} удалён")
@@ -307,7 +449,6 @@ async def delete_key(session: AsyncSession, identifier: int | str, commit: bool
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).values(expiry_time=new_expiry_time))
await session.commit()
await invalidate_key_details_by_client_id(session, client_id)
logger.info(f"Срок действия ключа {client_id} обновлён до {new_expiry_time}")
@@ -317,59 +458,120 @@ async def get_client_id_by_email(session: AsyncSession, email: str):
return result.scalar_one_or_none()
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.commit()
await invalidate_keys_list(tg_id)
async def update_key_notified(session: AsyncSession, legacy_user_ref: int, client_id: str):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
await session.execute(update(Key).where(Key.user_id == u.id, Key.client_id == client_id).values(notified=True))
await invalidate_keys_list(session, u.id)
await invalidate_key_details_by_client_id(session, client_id)
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, legacy_user_ref: int, client_id: str, time_left: int):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
await session.execute(
text(
"""
UPDATE keys
SET expiry_time = :expiry,
is_frozen = TRUE
WHERE tg_id = :tg_id
WHERE user_id = :user_id
AND client_id = :client_id
"""
),
{"expiry": time_left, "tg_id": tg_id, "client_id": client_id},
{"expiry": time_left, "user_id": u.id, "client_id": client_id},
)
await invalidate_keys_list(tg_id)
await invalidate_keys_list(session, u.id)
await invalidate_key_details_by_client_id(session, client_id)
async def mark_key_as_unfrozen(
session: AsyncSession,
tg_id: int,
legacy_user_ref: int,
client_id: str,
new_expiry_time: int,
):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
await session.execute(
text(
"""
UPDATE keys
SET expiry_time = :expiry,
is_frozen = FALSE
WHERE tg_id = :tg_id
WHERE user_id = :user_id
AND client_id = :client_id
"""
),
{"expiry": new_expiry_time, "tg_id": tg_id, "client_id": client_id},
{"expiry": new_expiry_time, "user_id": u.id, "client_id": client_id},
)
await invalidate_keys_list(tg_id)
await invalidate_keys_list(session, u.id)
await invalidate_key_details_by_client_id(session, client_id)
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.commit()
await invalidate_key_details_by_client_id(session, client_id)
logger.info(f"Тариф ключа {client_id} обновлён на {tariff_id}")
async def update_key_renewal_snapshot(
session: AsyncSession,
email: str,
*,
tariff_id: int,
selected_device_limit: int | None = None,
current_device_limit: int | None = None,
selected_traffic_limit: int | None = None,
current_traffic_limit: int | None = None,
apply_limits: bool = True,
) -> None:
"""Обновляет tariff_id и (опционально) лимиты ключа после продления.
``apply_limits=True`` — выставить все четыре лимита (для non-configurable
тарифов). ``apply_limits=False`` — обновить только ``tariff_id``, лимиты
не трогать (configurable-тарифы обновляют их через `save_key_config_with_mode`).
"""
values: dict = {"tariff_id": tariff_id}
if apply_limits:
values["selected_device_limit"] = selected_device_limit
values["current_device_limit"] = current_device_limit
values["selected_traffic_limit"] = selected_traffic_limit
values["current_traffic_limit"] = current_traffic_limit
await session.execute(update(Key).where(Key.email == email).values(**values))
await invalidate_key_details(email)
async def update_key_post_creation_snapshot(
session: AsyncSession,
*,
user_id: int,
email: str,
selected_device_limit: int | None,
selected_traffic_limit: int | None,
selected_price_rub: int | None,
) -> None:
"""Дозаписывает выбранные пользователем параметры ключа сразу после создания.
Используется из `services.keys.create_vpn_key_headless` — тариф/лимиты не
всегда известны на момент `create_key_on_cluster`, поэтому после него
идёт snapshot-апдейт для полей, которые нужны для отображения в UI.
"""
await session.execute(
update(Key)
.where(Key.user_id == int(user_id), Key.email == email)
.values(
selected_device_limit=selected_device_limit,
selected_traffic_limit=selected_traffic_limit,
selected_price_rub=selected_price_rub,
)
)
await invalidate_key_details(email)
async def get_subscription_link(session: AsyncSession, email: str) -> str | None:
result = await session.execute(select(func.coalesce(Key.key, Key.remnawave_link)).where(Key.email == email))
return result.scalar_one_or_none()
@@ -377,7 +579,6 @@ async def get_subscription_link(session: AsyncSession, email: str) -> str | None
async def update_key_client_id(session: AsyncSession, email: str, new_client_id: str):
await session.execute(update(Key).where(Key.email == email).values(client_id=new_client_id))
await session.commit()
await invalidate_key_details(email)
logger.info(f"client_id обновлён для {email} -> {new_client_id}")
@@ -385,7 +586,6 @@ async def update_key_client_id(session: AsyncSession, email: str, new_client_id:
async def update_key_link(session: AsyncSession, email: str, link: str) -> bool:
q = update(Key).where(Key.email == email).values(key=link).returning(Key.client_id)
res = await session.execute(q)
await session.commit()
ok = res.scalar_one_or_none() is not None
if ok:
await invalidate_key_details(email)
@@ -403,7 +603,6 @@ async def update_key_subscription_links(session: AsyncSession, email: str, link:
.returning(Key.client_id)
)
res = await session.execute(stmt)
await session.commit()
ok = res.scalar_one_or_none() is not None
if ok:
await invalidate_key_details(email)
@@ -444,10 +643,13 @@ async def save_key_config_with_mode(
await invalidate_key_details(email)
async def reset_key_tariff_state(session: AsyncSession, tg_id: int, email: str, tariff_id: int) -> None:
async def reset_key_tariff_state(session: AsyncSession, legacy_user_ref: int, email: str, tariff_id: int) -> None:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
await session.execute(
update(Key)
.where(Key.tg_id == tg_id, Key.email == email)
.where(Key.user_id == u.id, Key.email == email)
.values(
tariff_id=tariff_id,
selected_device_limit=None,
@@ -457,25 +659,27 @@ async def reset_key_tariff_state(session: AsyncSession, tg_id: int, email: str,
selected_price_rub=None,
)
)
await session.commit()
await invalidate_keys_list(tg_id)
await invalidate_keys_list(session, u.id)
await invalidate_key_details(email)
async def save_key_tariff_selection(
session: AsyncSession,
tg_id: int,
legacy_user_ref: int,
email: str,
tariff_id: int,
selected_devices: int | None,
selected_traffic_gb: int | None,
) -> None:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
selected_devices_val = int(selected_devices) if selected_devices is not None else None
selected_traffic_val = int(selected_traffic_gb) if selected_traffic_gb is not None and int(selected_traffic_gb) > 0 else None
await session.execute(
update(Key)
.where(Key.tg_id == tg_id, Key.email == email)
.where(Key.user_id == u.id, Key.email == email)
.values(
tariff_id=tariff_id,
selected_device_limit=selected_devices_val,
@@ -485,8 +689,7 @@ async def save_key_tariff_selection(
selected_price_rub=None,
)
)
await session.commit()
await invalidate_keys_list(tg_id)
await invalidate_keys_list(session, u.id)
await invalidate_key_details(email)
@@ -510,7 +713,6 @@ async def save_admin_key_config(
selected_price_rub=selected_price,
)
)
await session.commit()
await invalidate_key_details(email)
@@ -527,6 +729,5 @@ async def reset_key_current_limits_to_selected(session: AsyncSession, client_id:
),
{"client_id": client_id},
)
await session.commit()
await invalidate_key_details_by_client_id(session, client_id)
logger.info(f"Текущие лимиты ключа {client_id} сброшены к выбранным")
+1
View File
@@ -0,0 +1 @@
from .schema_upgrade import *
File diff suppressed because it is too large Load Diff
-405
View File
@@ -1,405 +0,0 @@
import secrets
import uuid
from datetime import datetime
from sqlalchemy import (
JSON,
BigInteger,
Boolean,
Column,
DateTime,
Float,
ForeignKey,
Index,
Integer,
Numeric,
String,
Text,
UniqueConstraint,
text as sql_text,
)
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import Mapped, declarative_base, mapped_column, relationship
Base = declarative_base()
class DictLikeMixin:
def __getitem__(self, key):
return getattr(self, key)
def get(self, key, default=None):
return getattr(self, key, default)
def to_dict(self):
return {column.name: getattr(self, column.name) for column in self.__table__.columns}
class Identity(DictLikeMixin, Base):
"""Слой идентификации: к одному identity можно привязать email и/или Telegram (tg_id)."""
__tablename__ = "identities"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
email = Column(String(255), unique=True, nullable=True, index=True)
tg_id = Column(BigInteger, unique=True, nullable=True, index=True)
api_token_hash = Column(String(64), nullable=True, index=True)
token_issued_at = Column(DateTime, nullable=True)
password_hash = Column(String(64), nullable=True)
is_admin = Column(Boolean, nullable=False, server_default=sql_text("false"))
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
class User(DictLikeMixin, Base):
__tablename__ = "users"
tg_id = Column(BigInteger, primary_key=True)
identity_id = Column(
String(36),
ForeignKey("identities.id", ondelete="SET NULL", onupdate="CASCADE"),
nullable=True,
index=True,
)
username = Column(String)
first_name = Column(String)
last_name = Column(String)
language_code = Column(String)
is_bot = Column(Boolean, default=False)
balance = Column(Float, default=0.0)
trial = Column(Integer, default=0)
preferred_currency = Column(String(10), nullable=False, server_default="RUB", index=True)
source_code = Column(
String,
ForeignKey(
"tracking_sources.code",
ondelete="SET NULL",
onupdate="CASCADE",
),
nullable=True,
)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow)
class Key(DictLikeMixin, Base):
__tablename__ = "keys"
tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=False, index=True)
client_id = Column(String, primary_key=True)
email = Column(String, unique=True)
created_at = Column(BigInteger)
expiry_time = Column(BigInteger)
key = Column(String)
server_id = Column(String)
remnawave_link = Column(String)
tariff_id = Column(Integer, ForeignKey("tariffs.id", ondelete="SET NULL"))
is_frozen = Column(Boolean, default=False)
alias = Column(String)
notified = Column(Boolean, default=False)
notified_24h = Column(Boolean, default=False)
selected_device_limit = Column(Integer, nullable=True)
selected_traffic_limit = Column(BigInteger, nullable=True)
selected_price_rub = Column(Integer, nullable=True)
current_device_limit = Column(Integer, nullable=True)
current_traffic_limit = Column(BigInteger, nullable=True)
class Tariff(DictLikeMixin, Base):
__tablename__ = "tariffs"
id = Column(Integer, primary_key=True)
name = Column(String)
group_code = Column(String)
duration_days = Column(Integer)
price_rub = Column(Integer)
traffic_limit = Column(BigInteger, nullable=True)
device_limit = Column(Integer, nullable=True)
is_active = Column(Boolean, default=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow)
subgroup_title = Column(String, nullable=True)
sort_order = Column(Integer, nullable=True)
vless = Column(Boolean, default=False)
external_squad: Mapped[str | None] = mapped_column(String(64), nullable=True)
configurable = Column(Boolean, nullable=False, server_default="false")
device_options = Column(JSONB, nullable=True)
traffic_options_gb = Column(JSONB, nullable=True)
device_step_rub = Column(Integer, nullable=True)
device_overrides = Column(JSONB, nullable=True)
traffic_step_rub = Column(Integer, nullable=True)
traffic_overrides = Column(JSONB, nullable=True)
class Server(DictLikeMixin, Base):
__tablename__ = "servers"
id = Column(Integer, primary_key=True, autoincrement=True)
cluster_name = Column(String)
server_name = Column(String, unique=True)
api_url = Column(String)
subscription_url = Column(String)
inbound_id = Column(String)
panel_type = Column(String)
max_keys = Column(Integer)
tariff_group = Column(String)
enabled = Column(Boolean, default=True)
subgroups = relationship("ServerSubgroup", back_populates="server", cascade="all, delete-orphan")
groups = relationship("ServerSpecialgroup", back_populates="server", cascade="all, delete-orphan")
class ServerSubgroup(DictLikeMixin, Base):
__tablename__ = "server_subgroups"
id = Column(Integer, primary_key=True, autoincrement=True)
server_id = Column(Integer, ForeignKey("servers.id", ondelete="CASCADE"), index=True, nullable=False)
group_code = Column(String, nullable=False)
subgroup_title = Column(String, nullable=False)
server = relationship("Server", back_populates="subgroups")
__table_args__ = (UniqueConstraint("server_id", "subgroup_title", name="uq_server_subgroup"),)
class ServerSpecialgroup(DictLikeMixin, Base):
__tablename__ = "server_specialgroups"
id = Column(Integer, primary_key=True, autoincrement=True)
server_id = Column(Integer, ForeignKey("servers.id", ondelete="CASCADE"), index=True, nullable=False)
group_code = Column(String, nullable=False)
server = relationship("Server")
__table_args__ = (UniqueConstraint("server_id", "group_code", name="uq_server_group"),)
class Payment(DictLikeMixin, Base):
__tablename__ = "payments"
id = Column(Integer, primary_key=True, autoincrement=True)
tg_id = Column(BigInteger, ForeignKey("users.tg_id"))
amount = Column(Float)
payment_system = Column(String)
status = Column(String)
created_at = Column(DateTime, default=datetime.utcnow)
original_amount = Column(Numeric(18, 8), nullable=True)
currency = Column(String(10), nullable=False, server_default="RUB")
payment_id = Column(String(128), nullable=True, index=True)
metadata_ = Column("metadata", JSONB, nullable=True)
class Coupon(DictLikeMixin, Base):
__tablename__ = "coupons"
id = Column(Integer, primary_key=True)
code = Column(String, unique=True)
amount = Column(Integer)
usage_limit = Column(Integer)
usage_count = Column(Integer, default=0)
is_used = Column(Boolean, default=False)
days = Column(Integer, nullable=True)
new_users_only = Column(Boolean, nullable=False, server_default=sql_text("false"))
percent = Column(Integer, nullable=True)
max_discount_amount = Column(Integer, nullable=True)
min_order_amount = Column(Integer, nullable=True)
class CouponUsage(DictLikeMixin, Base):
__tablename__ = "coupon_usages"
coupon_id = Column(Integer, ForeignKey("coupons.id", ondelete="CASCADE"), primary_key=True)
user_id = Column(BigInteger, primary_key=True)
used_at = Column(DateTime, default=datetime.utcnow)
class Referral(DictLikeMixin, Base):
__tablename__ = "referrals"
referred_tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="CASCADE"), primary_key=True)
referrer_tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="CASCADE"), primary_key=True)
reward_issued = Column(Boolean, default=False)
class Notification(DictLikeMixin, Base):
__tablename__ = "notifications"
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)
class ScheduledBroadcast(DictLikeMixin, Base):
__tablename__ = "scheduled_broadcasts"
__table_args__ = (
Index("ix_scheduled_broadcasts_status_time", "status", "scheduled_for"),
Index("ix_scheduled_broadcasts_creator_time", "created_by_tg_id", "created_at"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
created_by_tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="SET NULL"), nullable=True, index=True)
status = Column(String(32), nullable=False, server_default=sql_text("'scheduled'"), index=True)
send_to = Column(String(32), nullable=False, index=True)
cluster_name = Column(String, nullable=True)
text = Column(Text, nullable=False)
photo = Column(String, nullable=True)
keyboard_json = Column(JSONB, nullable=True)
scheduled_for = Column(DateTime(timezone=True), nullable=False, index=True)
workers = Column(Integer, nullable=False, server_default=sql_text("5"))
messages_per_second = Column(Integer, nullable=False, server_default=sql_text("35"))
stats_json = Column(JSONB, nullable=True)
error_text = Column(Text, nullable=True)
started_at = Column(DateTime(timezone=True), nullable=True)
sent_at = Column(DateTime(timezone=True), nullable=True)
cancelled_at = Column(DateTime(timezone=True), nullable=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
class Gift(DictLikeMixin, Base):
__tablename__ = "gifts"
gift_id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
sender_tg_id = Column(BigInteger, ForeignKey("users.tg_id"))
recipient_tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=True)
selected_months = Column(Integer)
expiry_time = Column(DateTime)
gift_link = Column(String)
created_at = Column(DateTime, default=datetime.utcnow)
is_used = Column(Boolean, default=False)
is_unlimited = Column(Boolean, default=False)
max_usages = Column(Integer, nullable=True)
tariff_id: Mapped[int | None] = mapped_column(ForeignKey("tariffs.id"))
selected_device_limit = Column(Integer, nullable=True)
selected_traffic_gb = Column(Integer, nullable=True)
selected_price_rub = Column(Integer, nullable=True)
class GiftUsage(DictLikeMixin, Base):
__tablename__ = "gift_usages"
gift_id = Column(String, ForeignKey("gifts.gift_id"), primary_key=True)
tg_id = Column(BigInteger, primary_key=True)
used_at = Column(DateTime, default=datetime.utcnow)
class ManualBan(DictLikeMixin, Base):
__tablename__ = "manual_bans"
tg_id = Column(BigInteger, primary_key=True)
banned_at = Column(DateTime(timezone=True), default=datetime.utcnow)
reason = Column(Text)
banned_by = Column(BigInteger)
until = Column(DateTime(timezone=True), nullable=True)
class TemporaryData(DictLikeMixin, Base):
__tablename__ = "temporary_data"
tg_id = Column(BigInteger, primary_key=True)
state = Column(String)
data = Column(JSON)
updated_at = Column(DateTime, default=datetime.utcnow)
class BlockedUser(DictLikeMixin, Base):
__tablename__ = "blocked_users"
tg_id = Column(BigInteger, primary_key=True)
class TrackingSource(DictLikeMixin, Base):
__tablename__ = "tracking_sources"
id = Column(Integer, primary_key=True)
name = Column(String)
code = Column(String, unique=True)
type = Column(String)
created_by = Column(BigInteger)
created_at = Column(DateTime, default=datetime.utcnow)
class AuditEvent(DictLikeMixin, Base):
"""События аудита (флоу пользователя)."""
__tablename__ = "audit_events"
__table_args__ = (
Index("ix_audit_events_tg_created", "actor_tg_id", "created_at"),
Index("ix_audit_events_identity_created", "actor_identity_id", "created_at"),
)
id = Column(Integer, primary_key=True, autoincrement=True)
event_type = Column(String(64), nullable=False, index=True)
channel = Column(String(32), nullable=False, index=True)
actor_identity_id = Column(
String(36),
ForeignKey("identities.id", ondelete="SET NULL", onupdate="CASCADE"),
nullable=True,
index=True,
)
actor_tg_id = Column(BigInteger, nullable=True, index=True)
path_or_handler = Column(String(255), nullable=False)
entity_type = Column(String(64), nullable=True, index=True)
entity_id = Column(String(255), nullable=True, index=True)
result = Column(String(32), nullable=False, server_default=sql_text("'success'"))
reason = Column(Text, nullable=True)
metadata_ = Column("metadata", JSONB, nullable=True)
request_id = Column(String(64), nullable=True, index=True)
created_at = Column(DateTime, default=datetime.utcnow, index=True)
class Admin(Base):
__tablename__ = "admins"
tg_id = Column(BigInteger, primary_key=True)
token = Column(String, unique=True, nullable=True)
description = Column(String, nullable=True)
role = Column(String, nullable=False, default="admin")
added_at = Column(DateTime, default=datetime.utcnow)
@staticmethod
def generate_token() -> str:
return secrets.token_urlsafe(32)
class Setting(DictLikeMixin, Base):
__tablename__ = "settings"
key = Column(String, primary_key=True)
value = Column(JSONB, nullable=True)
description = Column(Text, nullable=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
class WebPage(DictLikeMixin, Base):
__tablename__ = "web_pages"
slug = Column(String(64), primary_key=True)
title = Column(String(255), nullable=True)
class WebTheme(DictLikeMixin, Base):
__tablename__ = "web_themes"
page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), primary_key=True)
tokens = Column(JSONB, nullable=False, default=dict)
class WebBlock(DictLikeMixin, Base):
__tablename__ = "web_blocks"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), index=True, nullable=False)
order = Column(Integer, nullable=False, default=0)
type = Column(String(64), nullable=False)
data = Column(JSONB, nullable=False, default=dict)
+61
View File
@@ -0,0 +1,61 @@
from ._base import Base, DictLikeMixin
from .admin import Admin, Setting
from .audit import AuditEvent
from .coupons import Coupon, CouponUsage
from .gifts import Gift, GiftUsage
from .identity import Identity
from .keys import Key
from .notifications import Notification, ScheduledBroadcast
from .payments import Payment
from .referrals import Referral
from .servers import Server, ServerSpecialgroup, ServerSubgroup
from .tariffs import Tariff
from .users import BlockedUser, ManualBan, TemporaryData, TrackingSource, User
from .web import (
WebBlock,
WebCustomElementBuild,
WebFlow,
WebFlowEvent,
WebNotification,
WebPage,
WebPageVariant,
WebPageVariantBlock,
WebPushSubscription,
WebTheme,
)
__all__ = [
"Base",
"DictLikeMixin",
"Identity",
"User",
"ManualBan",
"TemporaryData",
"BlockedUser",
"TrackingSource",
"Key",
"Tariff",
"Server",
"ServerSubgroup",
"ServerSpecialgroup",
"Payment",
"Coupon",
"CouponUsage",
"Referral",
"Notification",
"ScheduledBroadcast",
"Gift",
"GiftUsage",
"AuditEvent",
"Admin",
"Setting",
"WebPage",
"WebTheme",
"WebBlock",
"WebPageVariant",
"WebPageVariantBlock",
"WebPushSubscription",
"WebNotification",
"WebFlow",
]
+21
View File
@@ -0,0 +1,21 @@
from sqlalchemy.orm import declarative_base
Base = declarative_base()
class DictLikeMixin:
"""Позволяет обращаться к ORM-объектам как к словарю.
Используется legacy-кодом, который мигрировал с dict-результатов asyncpg
на ORM и не хочет переписывать все `row["field"]` / `row.get("field")`.
"""
def __getitem__(self, key):
return getattr(self, key)
def get(self, key, default=None):
return getattr(self, key, default)
def to_dict(self):
return {column.name: getattr(self, column.name) for column in self.__table__.columns}
+32
View File
@@ -0,0 +1,32 @@
import secrets
from datetime import datetime
from sqlalchemy import BigInteger, Column, DateTime, String, Text
from sqlalchemy.dialects.postgresql import JSONB
from ._base import Base, DictLikeMixin
class Admin(Base):
__tablename__ = "admins"
tg_id = Column(BigInteger, primary_key=True)
token = Column(String, unique=True, nullable=True)
description = Column(String, nullable=True)
role = Column(String, nullable=False, default="admin")
added_at = Column(DateTime, default=datetime.utcnow)
@staticmethod
def generate_token() -> str:
return secrets.token_urlsafe(32)
class Setting(DictLikeMixin, Base):
__tablename__ = "settings"
key = Column(String, primary_key=True)
value = Column(JSONB, nullable=True)
description = Column(Text, nullable=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
+44
View File
@@ -0,0 +1,44 @@
from datetime import datetime
from sqlalchemy import (
BigInteger,
Column,
DateTime,
ForeignKey,
Index,
Integer,
String,
Text,
text as sql_text,
)
from sqlalchemy.dialects.postgresql import JSONB
from ._base import Base, DictLikeMixin
class AuditEvent(DictLikeMixin, Base):
"""События аудита (флоу пользователя)."""
__tablename__ = "audit_events"
__table_args__ = (
Index("ix_audit_events_tg_created", "actor_tg_id", "created_at"),
Index("ix_audit_events_identity_created", "actor_identity_id", "created_at"),
)
id = Column(Integer, primary_key=True, autoincrement=True)
event_type = Column(String(64), nullable=False, index=True)
channel = Column(String(32), nullable=False, index=True)
actor_identity_id = Column(
String(36),
ForeignKey("identities.id", ondelete="SET NULL", onupdate="CASCADE"),
nullable=True,
index=True,
)
actor_tg_id = Column(BigInteger, nullable=True, index=True)
path_or_handler = Column(String(255), nullable=False)
entity_type = Column(String(64), nullable=True, index=True)
entity_id = Column(String(255), nullable=True, index=True)
result = Column(String(32), nullable=False, server_default=sql_text("'success'"))
reason = Column(Text, nullable=True)
metadata_ = Column("metadata", JSONB, nullable=True)
request_id = Column(String(64), nullable=True, index=True)
created_at = Column(DateTime, default=datetime.utcnow, index=True)
+40
View File
@@ -0,0 +1,40 @@
from datetime import datetime
from sqlalchemy import (
BigInteger,
Boolean,
Column,
DateTime,
ForeignKey,
Integer,
String,
text as sql_text,
)
from ._base import Base, DictLikeMixin
class Coupon(DictLikeMixin, Base):
__tablename__ = "coupons"
id = Column(Integer, primary_key=True)
code = Column(String, unique=True)
amount = Column(Integer)
usage_limit = Column(Integer)
usage_count = Column(Integer, default=0)
is_used = Column(Boolean, default=False)
days = Column(Integer, nullable=True)
new_users_only = Column(Boolean, nullable=False, server_default=sql_text("false"))
percent = Column(Integer, nullable=True)
max_discount_amount = Column(Integer, nullable=True)
min_order_amount = Column(Integer, nullable=True)
class CouponUsage(DictLikeMixin, Base):
__tablename__ = "coupon_usages"
coupon_id = Column(Integer, ForeignKey("coupons.id", ondelete="CASCADE"), primary_key=True)
user_id = Column(BigInteger, primary_key=True)
tg_id = Column(BigInteger, nullable=True, index=True)
used_at = Column(DateTime, default=datetime.utcnow)
+39
View File
@@ -0,0 +1,39 @@
import uuid
from datetime import datetime
from sqlalchemy import BigInteger, Boolean, Column, DateTime, ForeignKey, Integer, String
from sqlalchemy.orm import Mapped, mapped_column
from ._base import Base, DictLikeMixin
class Gift(DictLikeMixin, Base):
__tablename__ = "gifts"
gift_id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
sender_user_id = Column(BigInteger, nullable=True)
recipient_user_id = Column(BigInteger, nullable=True)
sender_tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=True, index=True)
recipient_tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=True, index=True)
selected_months = Column(Integer)
expiry_time = Column(DateTime)
gift_link = Column(String)
created_at = Column(DateTime, default=datetime.utcnow)
is_used = Column(Boolean, default=False)
is_unlimited = Column(Boolean, default=False)
max_usages = Column(Integer, nullable=True)
tariff_id: Mapped[int | None] = mapped_column(ForeignKey("tariffs.id"))
selected_device_limit = Column(Integer, nullable=True)
selected_traffic_gb = Column(Integer, nullable=True)
selected_price_rub = Column(Integer, nullable=True)
class GiftUsage(DictLikeMixin, Base):
__tablename__ = "gift_usages"
gift_id = Column(String, ForeignKey("gifts.gift_id"), primary_key=True)
user_id = Column(BigInteger, nullable=False, primary_key=True)
tg_id = Column(BigInteger, nullable=True, index=True)
used_at = Column(DateTime, default=datetime.utcnow)
+31
View File
@@ -0,0 +1,31 @@
import uuid
from datetime import datetime
from sqlalchemy import (
BigInteger,
Boolean,
Column,
DateTime,
String,
text as sql_text,
)
from ._base import Base, DictLikeMixin
class Identity(DictLikeMixin, Base):
"""Слой идентификации: к одному identity можно привязать email и/или Telegram (tg_id)."""
__tablename__ = "identities"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
email = Column(String(255), unique=True, nullable=True, index=True)
tg_id = Column(BigInteger, unique=True, nullable=True, index=True)
api_token_hash = Column(String(64), nullable=True, index=True)
token_issued_at = Column(DateTime, nullable=True)
password_hash = Column(String(64), nullable=True)
email_verified = Column(Boolean, nullable=False, server_default=sql_text("false"))
is_admin = Column(Boolean, nullable=False, server_default=sql_text("false"))
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
+29
View File
@@ -0,0 +1,29 @@
from sqlalchemy import BigInteger, Boolean, Column, ForeignKey, Integer, String
from ._base import Base, DictLikeMixin
class Key(DictLikeMixin, Base):
__tablename__ = "keys"
tg_id = Column(BigInteger, ForeignKey("users.tg_id"), primary_key=True, nullable=False, index=True)
user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=True, index=True)
client_id = Column(String, primary_key=True)
email = Column(String, unique=True)
created_at = Column(BigInteger)
expiry_time = Column(BigInteger)
key = Column(String)
server_id = Column(String)
remnawave_link = Column(String)
tariff_id = Column(Integer, ForeignKey("tariffs.id", ondelete="SET NULL"))
is_frozen = Column(Boolean, default=False)
alias = Column(String)
notified = Column(Boolean, default=False)
notified_24h = Column(Boolean, default=False)
selected_device_limit = Column(Integer, nullable=True)
selected_traffic_limit = Column(BigInteger, nullable=True)
selected_price_rub = Column(Integer, nullable=True)
current_device_limit = Column(Integer, nullable=True)
current_traffic_limit = Column(BigInteger, nullable=True)
+55
View File
@@ -0,0 +1,55 @@
import uuid
from datetime import UTC, datetime
from sqlalchemy import (
BigInteger,
Column,
DateTime,
ForeignKey,
Index,
Integer,
String,
Text,
text as sql_text,
)
from sqlalchemy.dialects.postgresql import JSONB
from ._base import Base, DictLikeMixin
class Notification(DictLikeMixin, Base):
__tablename__ = "notifications"
tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="CASCADE"), nullable=True, index=True)
user_id = Column(BigInteger, nullable=False, primary_key=True)
notification_type = Column(String, primary_key=True)
last_notification_time = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC))
class ScheduledBroadcast(DictLikeMixin, Base):
__tablename__ = "scheduled_broadcasts"
__table_args__ = (
Index("ix_scheduled_broadcasts_status_time", "status", "scheduled_for"),
Index("ix_scheduled_broadcasts_creator_time", "created_by_tg_id", "created_at"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
created_by_user_id = Column(BigInteger, nullable=True, index=True)
created_by_tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="SET NULL"), nullable=True, index=True)
status = Column(String(32), nullable=False, server_default=sql_text("'scheduled'"), index=True)
send_to = Column(String(32), nullable=False, index=True)
cluster_name = Column(String, nullable=True)
text = Column(Text, nullable=False)
photo = Column(String, nullable=True)
keyboard_json = Column(JSONB, nullable=True)
scheduled_for = Column(DateTime(timezone=True), nullable=False, index=True)
workers = Column(Integer, nullable=False, server_default=sql_text("5"))
messages_per_second = Column(Integer, nullable=False, server_default=sql_text("35"))
stats_json = Column(JSONB, nullable=True)
error_text = Column(Text, nullable=True)
started_at = Column(DateTime(timezone=True), nullable=True)
sent_at = Column(DateTime(timezone=True), nullable=True)
cancelled_at = Column(DateTime(timezone=True), nullable=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
+22
View File
@@ -0,0 +1,22 @@
from datetime import datetime
from sqlalchemy import BigInteger, Column, DateTime, Float, ForeignKey, Integer, Numeric, String
from sqlalchemy.dialects.postgresql import JSONB
from ._base import Base, DictLikeMixin
class Payment(DictLikeMixin, Base):
__tablename__ = "payments"
id = Column(Integer, primary_key=True, autoincrement=True)
user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=True, index=True)
tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=True, index=True)
amount = Column(Float)
payment_system = Column(String)
status = Column(String)
created_at = Column(DateTime, default=datetime.utcnow)
original_amount = Column(Numeric(18, 8), nullable=True)
currency = Column(String(10), nullable=False, server_default="RUB")
payment_id = Column(String(128), nullable=True, index=True)
metadata_ = Column("metadata", JSONB, nullable=True)
+13
View File
@@ -0,0 +1,13 @@
from sqlalchemy import BigInteger, Boolean, Column, ForeignKey
from ._base import Base, DictLikeMixin
class Referral(DictLikeMixin, Base):
__tablename__ = "referrals"
referred_user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), primary_key=True)
referrer_user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), primary_key=True)
referred_tg_id = Column(BigInteger, nullable=True, index=True)
referrer_tg_id = Column(BigInteger, nullable=True, index=True)
reward_issued = Column(Boolean, default=False)
+47
View File
@@ -0,0 +1,47 @@
from sqlalchemy import Boolean, Column, ForeignKey, Integer, String, UniqueConstraint
from sqlalchemy.orm import relationship
from ._base import Base, DictLikeMixin
class Server(DictLikeMixin, Base):
__tablename__ = "servers"
id = Column(Integer, primary_key=True, autoincrement=True)
cluster_name = Column(String)
server_name = Column(String, unique=True)
api_url = Column(String)
subscription_url = Column(String)
inbound_id = Column(String)
panel_type = Column(String)
max_keys = Column(Integer)
tariff_group = Column(String)
enabled = Column(Boolean, default=True)
subgroups = relationship("ServerSubgroup", back_populates="server", cascade="all, delete-orphan")
groups = relationship("ServerSpecialgroup", back_populates="server", cascade="all, delete-orphan")
class ServerSubgroup(DictLikeMixin, Base):
__tablename__ = "server_subgroups"
id = Column(Integer, primary_key=True, autoincrement=True)
server_id = Column(Integer, ForeignKey("servers.id", ondelete="CASCADE"), index=True, nullable=False)
group_code = Column(String, nullable=False)
subgroup_title = Column(String, nullable=False)
server = relationship("Server", back_populates="subgroups")
__table_args__ = (UniqueConstraint("server_id", "subgroup_title", name="uq_server_subgroup"),)
class ServerSpecialgroup(DictLikeMixin, Base):
__tablename__ = "server_specialgroups"
id = Column(Integer, primary_key=True, autoincrement=True)
server_id = Column(Integer, ForeignKey("servers.id", ondelete="CASCADE"), index=True, nullable=False)
group_code = Column(String, nullable=False)
server = relationship("Server")
__table_args__ = (UniqueConstraint("server_id", "group_code", name="uq_server_group"),)
+37
View File
@@ -0,0 +1,37 @@
from datetime import datetime
from sqlalchemy import BigInteger, Boolean, Column, DateTime, Integer, String
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import Mapped, mapped_column
from ._base import Base, DictLikeMixin
class Tariff(DictLikeMixin, Base):
__tablename__ = "tariffs"
id = Column(Integer, primary_key=True)
name = Column(String)
group_code = Column(String)
duration_days = Column(Integer)
price_rub = Column(Integer)
traffic_limit = Column(BigInteger, nullable=True)
device_limit = Column(Integer, nullable=True)
is_active = Column(Boolean, default=True)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow)
subgroup_title = Column(String, nullable=True)
sort_order = Column(Integer, nullable=True)
vless = Column(Boolean, default=False)
external_squad: Mapped[str | None] = mapped_column(String(64), nullable=True)
configurable = Column(Boolean, nullable=False, server_default="false")
device_options = Column(JSONB, nullable=True)
traffic_options_gb = Column(JSONB, nullable=True)
device_step_rub = Column(Integer, nullable=True)
device_overrides = Column(JSONB, nullable=True)
traffic_step_rub = Column(Integer, nullable=True)
traffic_overrides = Column(JSONB, nullable=True)
+88
View File
@@ -0,0 +1,88 @@
from datetime import datetime
from sqlalchemy import (
JSON,
BigInteger,
Boolean,
Column,
DateTime,
Float,
ForeignKey,
Identity as SAIdentity,
Integer,
String,
Text,
)
from ._base import Base, DictLikeMixin
class User(DictLikeMixin, Base):
__tablename__ = "users"
id = Column(BigInteger, SAIdentity(always=False), primary_key=True)
tg_id = Column(BigInteger, nullable=True, unique=True, index=True)
identity_id = Column(
String(36),
ForeignKey("identities.id", ondelete="SET NULL", onupdate="CASCADE"),
nullable=True,
index=True,
)
username = Column(String)
first_name = Column(String)
last_name = Column(String)
language_code = Column(String)
is_bot = Column(Boolean, default=False)
balance = Column(Float, default=0.0)
trial = Column(Integer, default=0)
preferred_currency = Column(String(10), nullable=False, server_default="RUB", index=True)
source_code = Column(
String,
ForeignKey(
"tracking_sources.code",
ondelete="SET NULL",
onupdate="CASCADE",
),
nullable=True,
)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow)
class ManualBan(DictLikeMixin, Base):
__tablename__ = "manual_bans"
user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=False, primary_key=True)
tg_id = Column(BigInteger, nullable=True, index=True)
banned_at = Column(DateTime(timezone=True), default=datetime.utcnow)
reason = Column(Text)
banned_by = Column(BigInteger)
until = Column(DateTime(timezone=True), nullable=True)
class TemporaryData(DictLikeMixin, Base):
__tablename__ = "temporary_data"
user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=False, primary_key=True)
tg_id = Column(BigInteger, nullable=True, index=True)
state = Column(String)
data = Column(JSON)
updated_at = Column(DateTime, default=datetime.utcnow)
class BlockedUser(DictLikeMixin, Base):
__tablename__ = "blocked_users"
user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=False, primary_key=True)
tg_id = Column(BigInteger, nullable=True, index=True)
class TrackingSource(DictLikeMixin, Base):
__tablename__ = "tracking_sources"
id = Column(Integer, primary_key=True)
name = Column(String)
code = Column(String, unique=True)
type = Column(String)
created_by = Column(BigInteger)
created_at = Column(DateTime, default=datetime.utcnow)
+166
View File
@@ -0,0 +1,166 @@
import uuid
from datetime import UTC, datetime
from sqlalchemy import (
BigInteger,
Boolean,
Column,
DateTime,
ForeignKey,
Index,
Integer,
String,
Text,
UniqueConstraint,
text as sql_text,
)
from sqlalchemy.dialects.postgresql import JSONB
from ._base import Base, DictLikeMixin
class WebPage(DictLikeMixin, Base):
__tablename__ = "web_pages"
slug = Column(String(64), primary_key=True)
title = Column(String(255), nullable=True)
class WebTheme(DictLikeMixin, Base):
__tablename__ = "web_themes"
page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), primary_key=True)
tokens = Column(JSONB, nullable=False, default=dict)
class WebBlock(DictLikeMixin, Base):
__tablename__ = "web_blocks"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), index=True, nullable=False)
order = Column(Integer, nullable=False, default=0)
type = Column(String(64), nullable=False)
data = Column(JSONB, nullable=False, default=dict)
class WebPageVariant(DictLikeMixin, Base):
__tablename__ = "web_page_variants"
__table_args__ = (
UniqueConstraint("page_slug", "variant_key", name="uq_web_page_variants_page_slug_variant_key"),
Index("ix_web_page_variants_page_slug_is_active", "page_slug", "is_active"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), index=True, nullable=False)
variant_key = Column(String(64), nullable=False)
name = Column(String(255), nullable=False, default="Default")
is_active = Column(Boolean, nullable=False, server_default=sql_text("false"))
theme_tokens = Column(JSONB, nullable=False, default=dict)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
class WebPageVariantBlock(DictLikeMixin, Base):
__tablename__ = "web_page_variant_blocks"
__table_args__ = (Index("ix_web_page_variant_blocks_variant_id_order", "variant_id", "order"),)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
variant_id = Column(String(36), ForeignKey("web_page_variants.id", ondelete="CASCADE"), index=True, nullable=False)
order = Column(Integer, nullable=False, default=0)
type = Column(String(64), nullable=False)
data = Column(JSONB, nullable=False, default=dict)
class WebPushSubscription(DictLikeMixin, Base):
__tablename__ = "web_push_subscriptions"
__table_args__ = (
Index("ix_web_push_subscriptions_user_id", "user_id"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
user_id = Column(BigInteger, nullable=False)
identity_id = Column(String(36), nullable=True, index=True)
endpoint = Column(Text, nullable=False, unique=True)
keys_json = Column(JSONB, nullable=False)
created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC))
class WebNotification(DictLikeMixin, Base):
__tablename__ = "web_notifications"
__table_args__ = (
Index("ix_web_notifications_user_read", "user_id", "read"),
Index("ix_web_notifications_created", "created_at"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
user_id = Column(BigInteger, nullable=False, index=True)
identity_id = Column(String(36), nullable=True, index=True)
type = Column(String(32), nullable=False, default="system")
title = Column(String(255), nullable=False)
message = Column(Text, nullable=False, default="")
read = Column(Boolean, nullable=False, default=False)
data = Column(JSONB, nullable=True)
created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC))
class WebFlowEvent(DictLikeMixin, Base):
__tablename__ = "web_flow_events"
__table_args__ = (
Index("ix_web_flow_events_flow_node", "flow_id", "node_id"),
Index("ix_web_flow_events_created", "created_at"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
flow_id = Column(String(64), nullable=False)
node_id = Column(String(64), nullable=False)
node_type = Column(String(32), nullable=False, default="")
event_type = Column(String(32), nullable=False)
ab_variant = Column(String(16), nullable=True)
device = Column(String(16), nullable=True)
locale = Column(String(8), nullable=True)
authenticated = Column(Boolean, nullable=True)
event_metadata = Column("metadata", JSONB, nullable=True)
created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC))
class WebCustomElementBuild(DictLikeMixin, Base):
__tablename__ = "web_custom_element_builds"
__table_args__ = (
Index("ix_web_custom_element_builds_status", "status"),
Index("ix_web_custom_element_builds_created", "created_at"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
label = Column(String(255), nullable=False, default="")
slug = Column(String(128), nullable=False, default="")
runtime = Column(String(32), nullable=False, default="react-component")
source_kind = Column(String(32), nullable=False, default="inline-code")
source_value = Column(Text, nullable=False, default="")
export_name = Column(String(128), nullable=False, default="default")
props_schema_text = Column(Text, nullable=False, default="")
sample_props_text = Column(Text, nullable=False, default="")
events_text = Column(Text, nullable=False, default="")
notes = Column(Text, nullable=False, default="")
status = Column(String(32), nullable=False, default="queued")
summary = Column(Text, nullable=False, default="")
next_steps = Column(JSONB, nullable=False, default=list)
artifact = Column(JSONB, nullable=True)
upload_meta = Column(JSONB, nullable=True)
worker_id = Column(String(64), nullable=True)
worker_claimed_at = Column(DateTime(timezone=True), nullable=True)
completed_at = Column(DateTime(timezone=True), nullable=True)
created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC))
updated_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC), onupdate=lambda: datetime.now(UTC))
class WebFlow(DictLikeMixin, Base):
__tablename__ = "web_flows"
id = Column(String(64), primary_key=True, default="default")
name = Column(String(255), nullable=False, default="Основной flow")
nodes = Column(JSONB, nullable=False, default=list)
edges = Column(JSONB, nullable=False, default=list)
entry_node_id = Column(String(64), nullable=True)
version = Column(Integer, nullable=False, default=1)
updated_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC), onupdate=lambda: datetime.now(UTC))
+360 -309
View File
@@ -1,111 +1,155 @@
from collections import defaultdict
from datetime import datetime, timedelta
from datetime import UTC, datetime, timedelta
from sqlalchemy import and_, delete, func, select, tuple_
from sqlalchemy.dialects.postgresql import insert
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from config import DISCOUNT_ACTIVE_HOURS
from core.bootstrap import NOTIFICATIONS_CONFIG
from database.models import Key, Notification, User
from database.models import BlockedUser, Key, Notification, User
from database.access.resolution import resolve_user_optional
from logger import logger
_NOTIFICATION_TIME_BATCH_SIZE = 300
_BULK_ADD_NOTIFICATIONS_BATCH_SIZE = 1000
async def add_notification(session: AsyncSession, tg_id: int, notification_type: str):
try:
stmt = (
insert(Notification)
.values(
tg_id=tg_id,
notification_type=notification_type,
last_notification_time=datetime.utcnow(),
)
.on_conflict_do_update(
index_elements=[Notification.tg_id, Notification.notification_type],
set_={"last_notification_time": datetime.utcnow()},
)
)
await session.execute(stmt)
await session.commit()
logger.info(f"✅ Добавлено уведомление {notification_type} для пользователя {tg_id}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при добавлении уведомления: {e}")
await session.rollback()
raise
def _utc_now() -> datetime:
return datetime.now(UTC)
async def delete_notification(session: AsyncSession, tg_id: int, notification_type: str):
def _as_utc(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=UTC)
return value.astimezone(UTC)
async def _map_legacy_refs_to_user_ids(session: AsyncSession, refs: list[int]) -> dict[int, int]:
if not refs:
return {}
from sqlalchemy import or_
uniq = list(dict.fromkeys(refs))
r = await session.execute(select(User.id, User.tg_id).where(or_(User.tg_id.in_(uniq), User.id.in_(uniq))))
m: dict[int, int] = {}
for uid, tgid in r.all():
m[int(uid)] = int(uid)
if tgid is not None:
m[int(tgid)] = int(uid)
return {ref: m[ref] for ref in uniq if ref in m}
async def add_notification(session: AsyncSession, legacy_user_ref: int, notification_type: str):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
ins = insert(Notification).values(
user_id=u.id,
tg_id=u.tg_id,
notification_type=notification_type,
last_notification_time=_utc_now(),
)
stmt = ins.on_conflict_do_update(
index_elements=[Notification.user_id, Notification.notification_type],
set_={
"last_notification_time": ins.excluded.last_notification_time,
"tg_id": ins.excluded.tg_id,
},
)
await session.execute(stmt)
logger.info(f"✅ Добавлено уведомление {notification_type} для пользователя {u.id}")
async def delete_notification(session: AsyncSession, legacy_user_ref: int, notification_type: str):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
uid = u.id
await session.execute(
delete(Notification).where(
Notification.tg_id == tg_id,
Notification.user_id == uid,
Notification.notification_type == notification_type,
)
)
await session.commit()
logger.debug(f"🗑 Уведомление {notification_type} для пользователя {tg_id} удалено")
logger.debug(f"🗑 Уведомление {notification_type} для пользователя {uid} удалено")
async def bulk_add_notifications(
session: AsyncSession, items: list[tuple[int, str]], *, commit: bool = False
) -> None:
"""Вставка/обновление многих (tg_id, notification_type) батчами (лимит параметров PostgreSQL). Без commit, если commit=False."""
async def bulk_add_notifications(session: AsyncSession, items: list[tuple[int, str]]) -> None:
"""Вставка/обновление многих (legacy_user_ref, notification_type) батчами (лимит параметров PostgreSQL)."""
if not items:
return
now = datetime.utcnow()
id_map = await _map_legacy_refs_to_user_ids(session, [p[0] for p in items])
mapped = [(id_map[r], n) for r, n in items if r in id_map]
if not mapped:
return
uids = list({uid for uid, _ in mapped})
tg_map_r = await session.execute(select(User.id, User.tg_id).where(User.id.in_(uids)))
tg_by_uid = {int(r.id): r.tg_id for r in tg_map_r.all()}
now = _utc_now()
total = 0
for i in range(0, len(items), _BULK_ADD_NOTIFICATIONS_BATCH_SIZE):
batch = items[i : i + _BULK_ADD_NOTIFICATIONS_BATCH_SIZE]
stmt = insert(Notification).values(
for i in range(0, len(mapped), _BULK_ADD_NOTIFICATIONS_BATCH_SIZE):
batch = mapped[i : i + _BULK_ADD_NOTIFICATIONS_BATCH_SIZE]
ins = insert(Notification).values(
[
{"tg_id": tg_id, "notification_type": ntype, "last_notification_time": now}
for tg_id, ntype in batch
{
"user_id": uid,
"tg_id": tg_by_uid.get(uid),
"notification_type": ntype,
"last_notification_time": now,
}
for uid, ntype in batch
]
).on_conflict_do_update(
index_elements=[Notification.tg_id, Notification.notification_type],
set_={"last_notification_time": now},
)
stmt = ins.on_conflict_do_update(
index_elements=[Notification.user_id, Notification.notification_type],
set_={
"last_notification_time": ins.excluded.last_notification_time,
"tg_id": ins.excluded.tg_id,
},
)
await session.execute(stmt)
total += len(batch)
if commit:
await session.commit()
logger.info(f"✅ Bulk: добавлено/обновлено {total} уведомлений")
INACTIVE_TRIAL_REGISTERED_TYPE = "inactive_trial_registered"
async def bulk_delete_notifications(
session: AsyncSession, items: list[tuple[int, str]], *, commit: bool = False
) -> None:
"""Удаление многих (tg_id, notification_type) батчами (лимит параметров PostgreSQL). Без commit, если commit=False."""
async def bulk_delete_notifications(session: AsyncSession, items: list[tuple[int, str]]) -> None:
"""Удаление многих (legacy_user_ref, notification_type) батчами (лимит параметров PostgreSQL)."""
if not items:
return
id_map = await _map_legacy_refs_to_user_ids(session, [p[0] for p in items])
mapped = [(id_map[r], n) for r, n in items if r in id_map]
if not mapped:
return
total = 0
for i in range(0, len(items), _BULK_ADD_NOTIFICATIONS_BATCH_SIZE):
batch = items[i : i + _BULK_ADD_NOTIFICATIONS_BATCH_SIZE]
for i in range(0, len(mapped), _BULK_ADD_NOTIFICATIONS_BATCH_SIZE):
batch = mapped[i : i + _BULK_ADD_NOTIFICATIONS_BATCH_SIZE]
stmt = delete(Notification).where(
tuple_(Notification.tg_id, Notification.notification_type).in_(batch)
tuple_(Notification.user_id, Notification.notification_type).in_(batch)
)
await session.execute(stmt)
total += len(batch)
if commit:
await session.commit()
logger.debug(f"🗑 Bulk: удалено {total} уведомлений")
async def check_notification_time(session: AsyncSession, tg_id: int, notification_type: str, hours: int = 12) -> bool:
async def check_notification_time(session: AsyncSession, legacy_user_ref: int, notification_type: str, hours: int = 12) -> bool:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return True
stmt = select(Notification.last_notification_time).where(
Notification.tg_id == tg_id, Notification.notification_type == notification_type
Notification.user_id == u.id, Notification.notification_type == notification_type
)
result = await session.execute(stmt)
last_time = result.scalar_one_or_none()
if not last_time:
return True
return datetime.utcnow() - last_time > timedelta(hours=hours)
return _utc_now() - _as_utc(last_time) > timedelta(hours=hours)
async def check_notification_time_bulk(
@@ -121,37 +165,46 @@ async def check_notification_time_bulk(
"""
if not items:
return set()
now = datetime.utcnow()
now = _utc_now()
threshold = now - timedelta(hours=hours)
can_notify = set()
found = set()
try:
for batch in (
items[i : i + _NOTIFICATION_TIME_BATCH_SIZE]
for i in range(0, len(items), _NOTIFICATION_TIME_BATCH_SIZE)
):
stmt = select(
Notification.tg_id,
Notification.notification_type,
Notification.last_notification_time,
).where(tuple_(Notification.tg_id, Notification.notification_type).in_(batch))
result = await session.execute(stmt)
for row in result:
found.add((row.tg_id, row.notification_type))
if row.last_notification_time is None or row.last_notification_time < threshold:
can_notify.add((row.tg_id, row.notification_type))
for pair in items:
if pair not in found:
can_notify.add(pair)
except SQLAlchemyError:
await session.rollback()
raise
for batch in (
items[i : i + _NOTIFICATION_TIME_BATCH_SIZE]
for i in range(0, len(items), _NOTIFICATION_TIME_BATCH_SIZE)
):
id_map = await _map_legacy_refs_to_user_ids(session, [p[0] for p in batch])
mapped_batch = [(id_map[r], n) for r, n in batch if r in id_map]
if not mapped_batch:
continue
stmt = select(
Notification.user_id,
Notification.notification_type,
Notification.last_notification_time,
).where(tuple_(Notification.user_id, Notification.notification_type).in_(mapped_batch))
result = await session.execute(stmt)
uid_to_ref: dict[int, int] = {}
for r, _n in batch:
if r in id_map:
uid_to_ref[id_map[r]] = r
for row in result:
ref = uid_to_ref.get(row.user_id, row.user_id)
found.add((ref, row.notification_type))
row_time = _as_utc(row.last_notification_time)
if row_time is None or row_time < threshold:
can_notify.add((ref, row.notification_type))
for pair in items:
if pair not in found:
can_notify.add(pair)
return can_notify
async def get_last_notification_time(session: AsyncSession, tg_id: int, notification_type: str) -> int | None:
async def get_last_notification_time(session: AsyncSession, legacy_user_ref: int, notification_type: str) -> int | None:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return None
stmt = select(Notification.last_notification_time).where(
Notification.tg_id == tg_id, Notification.notification_type == notification_type
Notification.user_id == u.id, Notification.notification_type == notification_type
)
result = await session.execute(stmt)
ts = result.scalar_one_or_none()
@@ -173,15 +226,24 @@ async def get_last_notification_times_bulk(
out = {}
for chunk in _batched_list(pairs, _BULK_NOTIFICATION_BATCH_SIZE):
id_map = await _map_legacy_refs_to_user_ids(session, [p[0] for p in chunk])
mapped = [(id_map[r], n) for r, n in chunk if r in id_map]
if not mapped:
continue
stmt = select(
Notification.tg_id,
Notification.user_id,
Notification.notification_type,
Notification.last_notification_time,
).where(tuple_(Notification.tg_id, Notification.notification_type).in_(chunk))
).where(tuple_(Notification.user_id, Notification.notification_type).in_(mapped))
result = await session.execute(stmt)
for tg_id, ntype, last_time in result.all():
uid_to_ref: dict[int, int] = {}
for r, _n in chunk:
if r in id_map:
uid_to_ref[id_map[r]] = r
for uid, ntype, last_time in result.all():
if last_time:
out[(tg_id, ntype)] = int(last_time.timestamp() * 1000)
ref = uid_to_ref.get(uid, uid)
out[(ref, ntype)] = int(last_time.timestamp() * 1000)
return out
@@ -202,54 +264,51 @@ async def get_hot_lead_notification_flags(
"""
if not tg_ids:
return {}
stmt = select(Notification.tg_id, Notification.notification_type).where(
Notification.tg_id.in_(tg_ids),
stmt = select(Notification.user_id, Notification.notification_type).where(
Notification.user_id.in_(tg_ids),
Notification.notification_type.in_(_HOT_LEAD_NOTIFICATION_TYPES),
)
result = await session.execute(stmt)
out = defaultdict(set)
for tg_id, ntype in result.all():
out[tg_id].add(ntype)
for uid, ntype in result.all():
out[uid].add(ntype)
return dict(out)
async def check_hot_lead_discount(session: AsyncSession, tg_id: int) -> dict:
try:
result = await session.execute(
select(Notification.notification_type, Notification.last_notification_time)
.where(Notification.tg_id == tg_id)
.where(Notification.notification_type.in_(["hot_lead_step_2", "hot_lead_step_3"]))
.order_by(Notification.last_notification_time.desc())
.limit(1)
)
row = result.first()
if not row:
return {"available": False}
notification_type, last_time = row
hours = int(NOTIFICATIONS_CONFIG.get("DISCOUNT_ACTIVE_HOURS", DISCOUNT_ACTIVE_HOURS))
expires_at = last_time + timedelta(hours=hours)
current_time = datetime.utcnow()
if current_time > expires_at:
return {"available": False}
tariff_group = "discounts" if notification_type == "hot_lead_step_2" else "discounts_max"
return {
"available": True,
"type": notification_type,
"tariff_group": tariff_group,
"expires_at": expires_at,
}
except Exception as e:
logger.error(f"❌ Ошибка при проверке скидки горячего лида для {tg_id}: {e}")
await session.rollback()
async def check_hot_lead_discount(session: AsyncSession, legacy_user_ref: int) -> dict:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return {"available": False}
result = await session.execute(
select(Notification.notification_type, Notification.last_notification_time)
.where(Notification.user_id == u.id)
.where(Notification.notification_type.in_(["hot_lead_step_2", "hot_lead_step_3"]))
.order_by(Notification.last_notification_time.desc())
.limit(1)
)
row = result.first()
if not row:
return {"available": False}
notification_type, last_time = row
hours = int(NOTIFICATIONS_CONFIG.get("DISCOUNT_ACTIVE_HOURS", DISCOUNT_ACTIVE_HOURS))
expires_at = last_time + timedelta(hours=hours)
current_time = _utc_now()
if current_time > expires_at:
return {"available": False}
tariff_group = "discounts" if notification_type == "hot_lead_step_2" else "discounts_max"
return {
"available": True,
"type": notification_type,
"tariff_group": tariff_group,
"expires_at": expires_at,
}
_BULK_NOTIFICATION_BATCH_SIZE = 250
@@ -274,144 +333,160 @@ async def check_notifications_bulk(
tg_ids: list[int] = None,
emails: list[str] = None,
) -> list[dict]:
from sqlalchemy import select
now = _utc_now()
from database.models import BlockedUser, Notification
try:
now = datetime.utcnow()
if notification_type == "inactive_trial":
stmt_inactive = (
select(User.tg_id)
.where(
and_(
User.trial.in_([0, -1]),
~User.tg_id.in_(select(BlockedUser.tg_id)),
~User.tg_id.in_(select(Key.tg_id.distinct())),
if notification_type == "inactive_trial":
stmt_inactive = (
select(User.id)
.where(
and_(
User.trial.in_([0, -1]),
User.tg_id.isnot(None),
~User.id.in_(select(BlockedUser.user_id)),
~User.id.in_(select(Key.user_id.distinct())),
)
)
)
result_inactive = await session.execute(stmt_inactive)
inactive_user_ids = [r[0] for r in result_inactive.all()]
if inactive_user_ids:
already = set()
for chunk in _batched_list(inactive_user_ids, _NOTIFICATION_TIME_BATCH_SIZE):
result_existing = await session.execute(
select(Notification.user_id).where(
Notification.notification_type == INACTIVE_TRIAL_REGISTERED_TYPE,
Notification.user_id.in_(chunk),
)
)
)
result_inactive = await session.execute(stmt_inactive)
inactive_tg_ids = [r[0] for r in result_inactive.all()]
if inactive_tg_ids:
already = set()
for chunk in _batched_list(inactive_tg_ids, _NOTIFICATION_TIME_BATCH_SIZE):
result_existing = await session.execute(
select(Notification.tg_id).where(
Notification.notification_type == INACTIVE_TRIAL_REGISTERED_TYPE,
Notification.tg_id.in_(chunk),
)
already.update(r[0] for r in result_existing.all())
to_register = [uid for uid in inactive_user_ids if uid not in already]
if to_register:
for batch in _batched_list(to_register, _BULK_ADD_NOTIFICATIONS_BATCH_SIZE):
await bulk_add_notifications(
session,
[(uid, INACTIVE_TRIAL_REGISTERED_TYPE) for uid in batch],
)
already.update(r[0] for r in result_existing.all())
to_register = [tid for tid in inactive_tg_ids if tid not in already]
if to_register:
for batch in _batched_list(to_register, _BULK_ADD_NOTIFICATIONS_BATCH_SIZE):
await bulk_add_notifications(
session,
[(tid, INACTIVE_TRIAL_REGISTERED_TYPE) for tid in batch],
commit=True,
)
logger.info(f"Зарегистрировано как неактивные (шаг 1): {len(to_register)} пользователей.")
logger.info(f"Зарегистрировано как неактивные (шаг 1): {len(to_register)} пользователей.")
subq_registered = (
select(
Notification.tg_id,
func.max(Notification.last_notification_time).label("registered_time"),
)
.where(Notification.notification_type == INACTIVE_TRIAL_REGISTERED_TYPE)
.group_by(Notification.tg_id)
.subquery()
subq_registered = (
select(
Notification.user_id,
func.max(Notification.last_notification_time).label("registered_time"),
)
subq_sent = (
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()
.where(Notification.notification_type == INACTIVE_TRIAL_REGISTERED_TYPE)
.group_by(Notification.user_id)
.subquery()
)
subq_sent = (
select(
Notification.user_id,
func.max(Notification.last_notification_time).label("last_notification_time"),
)
stmt = (
select(
User.tg_id,
Key.email,
User.username,
User.first_name,
User.last_name,
subq_registered.c.registered_time,
subq_sent.c.last_notification_time,
)
.select_from(User)
.outerjoin(Key, Key.tg_id == User.tg_id)
.outerjoin(subq_registered, subq_registered.c.tg_id == User.tg_id)
.outerjoin(subq_sent, subq_sent.c.tg_id == User.tg_id)
.where(
and_(
User.trial.in_([0, -1]),
~User.tg_id.in_(select(BlockedUser.tg_id)),
~User.tg_id.in_(select(Key.tg_id.distinct())),
)
.where(Notification.notification_type == notification_type)
.group_by(Notification.user_id)
.subquery()
)
stmt = (
select(
User.tg_id,
Key.email,
User.username,
User.first_name,
User.last_name,
subq_registered.c.registered_time,
subq_sent.c.last_notification_time,
)
.select_from(User)
.outerjoin(Key, Key.user_id == User.id)
.outerjoin(subq_registered, subq_registered.c.user_id == User.id)
.outerjoin(subq_sent, subq_sent.c.user_id == User.id)
.where(
and_(
User.trial.in_([0, -1]),
User.tg_id.isnot(None),
~User.id.in_(select(BlockedUser.user_id)),
~User.id.in_(select(Key.user_id.distinct())),
)
)
)
result = await session.execute(stmt)
users = []
for row in result:
registered_time = row.registered_time
last_sent_time = row.last_notification_time
first_ok = (
registered_time is not None
and (now - _as_utc(registered_time)) >= timedelta(hours=hours)
and last_sent_time is None
)
second_ok = last_sent_time is not None and (now - _as_utc(last_sent_time)) > timedelta(hours=hours)
if first_ok or second_ok:
users.append({
"tg_id": row.tg_id,
"email": row.email,
"username": row.username,
"first_name": row.first_name,
"last_name": row.last_name,
"last_notification_time": int(last_sent_time.timestamp() * 1000) if last_sent_time else None,
})
logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}")
return users
subq_last_notification = (
select(Notification.user_id, func.max(Notification.last_notification_time).label("last_notification_time"))
.where(Notification.notification_type == notification_type)
.group_by(Notification.user_id)
.subquery()
)
def make_stmt(tg_ids_batch: list[int] | None, emails_batch: list[str] | None):
stmt = (
select(
User.tg_id,
Key.email,
User.username,
User.first_name,
User.last_name,
subq_last_notification.c.last_notification_time,
)
.select_from(User)
.outerjoin(Key, Key.user_id == User.id)
.outerjoin(subq_last_notification, subq_last_notification.c.user_id == User.id)
)
if tg_ids_batch:
stmt = stmt.where(User.tg_id.in_(tg_ids_batch))
if emails_batch:
stmt = stmt.where(Key.email.in_(emails_batch))
return stmt
def _can_notify(last_time):
return last_time is None or (now - _as_utc(last_time)) > timedelta(hours=hours)
users: list[dict] = []
seen: set[tuple[int, str | None]] = set()
if tg_ids and emails and len(tg_ids) == len(emails):
for tg_ids_chunk, emails_chunk in _batched_pairs(tg_ids, emails, _BULK_NOTIFICATION_BATCH_SIZE):
stmt = make_stmt(tg_ids_chunk, emails_chunk)
result = await session.execute(stmt)
users = []
for row in result:
registered_time = row.registered_time
last_sent_time = row.last_notification_time
first_ok = (
registered_time is not None
and (now - registered_time) >= timedelta(hours=hours)
and last_sent_time is None
)
second_ok = last_sent_time is not None and (now - last_sent_time) > timedelta(hours=hours)
if first_ok or second_ok:
key = (row.tg_id, row.email)
if key in seen:
continue
seen.add(key)
last_time = row.last_notification_time
if _can_notify(last_time):
users.append({
"tg_id": row.tg_id,
"email": row.email,
"username": row.username,
"first_name": row.first_name,
"last_name": row.last_name,
"last_notification_time": int(last_sent_time.timestamp() * 1000) if last_sent_time else None,
"last_notification_time": int(last_time.timestamp() * 1000) if last_time else None,
})
logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}")
return users
subq_last_notification = (
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()
)
def make_stmt(tg_ids_batch: list[int] | None, emails_batch: list[str] | None):
stmt = (
select(
User.tg_id,
Key.email,
User.username,
User.first_name,
User.last_name,
subq_last_notification.c.last_notification_time,
)
.select_from(User)
.outerjoin(Key, Key.tg_id == User.tg_id)
.outerjoin(subq_last_notification, subq_last_notification.c.tg_id == User.tg_id)
)
if tg_ids_batch:
stmt = stmt.where(User.tg_id.in_(tg_ids_batch))
if emails_batch:
stmt = stmt.where(Key.email.in_(emails_batch))
return stmt
def _can_notify(last_time):
return last_time is None or (now - last_time) > timedelta(hours=hours)
users: list[dict] = []
seen: set[tuple[int, str | None]] = set()
if tg_ids and emails and len(tg_ids) == len(emails):
for tg_ids_chunk, emails_chunk in _batched_pairs(tg_ids, emails, _BULK_NOTIFICATION_BATCH_SIZE):
elif tg_ids and emails:
for tg_ids_chunk in _batched_list(tg_ids, _BULK_NOTIFICATION_BATCH_SIZE):
for emails_chunk in _batched_list(emails, _BULK_NOTIFICATION_BATCH_SIZE):
stmt = make_stmt(tg_ids_chunk, emails_chunk)
result = await session.execute(stmt)
for row in result:
@@ -429,68 +504,15 @@ async def check_notifications_bulk(
"last_name": row.last_name,
"last_notification_time": int(last_time.timestamp() * 1000) if last_time else None,
})
elif tg_ids and emails:
for tg_ids_chunk in _batched_list(tg_ids, _BULK_NOTIFICATION_BATCH_SIZE):
for emails_chunk in _batched_list(emails, _BULK_NOTIFICATION_BATCH_SIZE):
stmt = make_stmt(tg_ids_chunk, emails_chunk)
result = await session.execute(stmt)
for row in result:
key = (row.tg_id, row.email)
if key in seen:
continue
seen.add(key)
last_time = row.last_notification_time
if _can_notify(last_time):
users.append({
"tg_id": row.tg_id,
"email": row.email,
"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,
})
elif tg_ids:
for tg_ids_chunk in _batched_list(tg_ids, _BULK_NOTIFICATION_BATCH_SIZE):
stmt = make_stmt(tg_ids_chunk, None)
result = await session.execute(stmt)
for row in result:
key = (row.tg_id, row.email)
if key in seen:
continue
seen.add(key)
last_time = row.last_notification_time
if _can_notify(last_time):
users.append({
"tg_id": row.tg_id,
"email": row.email,
"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,
})
elif emails:
for emails_chunk in _batched_list(emails, _BULK_NOTIFICATION_BATCH_SIZE):
stmt = make_stmt(None, emails_chunk)
result = await session.execute(stmt)
for row in result:
key = (row.tg_id, row.email)
if key in seen:
continue
seen.add(key)
last_time = row.last_notification_time
if _can_notify(last_time):
users.append({
"tg_id": row.tg_id,
"email": row.email,
"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,
})
else:
stmt = make_stmt(None, None)
elif tg_ids:
for tg_ids_chunk in _batched_list(tg_ids, _BULK_NOTIFICATION_BATCH_SIZE):
stmt = make_stmt(tg_ids_chunk, None)
result = await session.execute(stmt)
for row in result:
key = (row.tg_id, row.email)
if key in seen:
continue
seen.add(key)
last_time = row.last_notification_time
if _can_notify(last_time):
users.append({
@@ -501,11 +523,40 @@ async def check_notifications_bulk(
"last_name": row.last_name,
"last_notification_time": int(last_time.timestamp() * 1000) if last_time else None,
})
elif emails:
for emails_chunk in _batched_list(emails, _BULK_NOTIFICATION_BATCH_SIZE):
stmt = make_stmt(None, emails_chunk)
result = await session.execute(stmt)
for row in result:
key = (row.tg_id, row.email)
if key in seen:
continue
seen.add(key)
last_time = row.last_notification_time
if _can_notify(last_time):
users.append({
"tg_id": row.tg_id,
"email": row.email,
"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,
})
else:
stmt = make_stmt(None, None)
result = await session.execute(stmt)
for row in result:
last_time = row.last_notification_time
if _can_notify(last_time):
users.append({
"tg_id": row.tg_id,
"email": row.email,
"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,
})
logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}")
return users
logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}")
return users
except Exception as e:
logger.error(f"Ошибка при массовой проверке уведомлений типа {notification_type}: {e}")
await session.rollback()
return []
+135 -102
View File
@@ -1,12 +1,12 @@
from datetime import datetime, timedelta
from pytz import timezone
from sqlalchemy import and_, insert, select, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy import and_, func, insert, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from core.cache_config import PAYMENT_PENDING_CACHE_TTL_SEC
from core.redis_cache import cache_delete, cache_get, cache_key, cache_set
from database.access.resolution import resolve_user_optional
from database.models import Payment
from logger import logger
@@ -52,52 +52,59 @@ async def invalidate_payment_cache(payment_id: str) -> None:
async def add_payment(
session: AsyncSession,
tg_id: int,
amount: float,
payment_system: str,
legacy_user_ref: int | None = None,
amount: float = 0,
payment_system: str = "",
*,
tg_id: int | None = None,
status: str = "success",
currency: str = "RUB",
payment_id: str | None = None,
metadata: dict | None = None,
original_amount: float | None = None,
) -> int:
try:
now_moscow = datetime.now(MOSCOW_TZ).replace(tzinfo=None)
stmt = (
insert(Payment)
.values(
tg_id=tg_id,
amount=amount,
payment_system=payment_system,
status=status,
created_at=now_moscow,
currency=currency,
payment_id=payment_id,
metadata_=metadata,
original_amount=original_amount,
)
.returning(Payment.id)
if legacy_user_ref is None:
legacy_user_ref = tg_id
if legacy_user_ref is None:
raise ValueError("legacy_user_ref is required for payment")
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
raise ValueError(f"user not found for payment: {legacy_user_ref}")
now_moscow = datetime.now(MOSCOW_TZ).replace(tzinfo=None)
stmt = (
insert(Payment)
.values(
user_id=u.id,
tg_id=u.tg_id,
amount=amount,
payment_system=payment_system,
status=status,
created_at=now_moscow,
currency=currency,
payment_id=payment_id,
metadata_=metadata,
original_amount=original_amount,
)
result = await session.execute(stmt)
internal_id = result.scalar_one()
logger.info(
f"Добавлен платёж id={internal_id}: tg_id={tg_id}, amount={amount}, system={payment_system}, status={status}"
)
return internal_id
except SQLAlchemyError as e:
await session.rollback()
logger.error(f"Ошибка при добавлении платежа: {e}")
raise
.returning(Payment.id)
)
result = await session.execute(stmt)
internal_id = result.scalar_one()
logger.info(
f"Добавлен платёж id={internal_id}: user_id={u.id}, amount={amount}, system={payment_system}, status={status}"
)
return internal_id
async def get_last_payments(
session: AsyncSession,
tg_id: int,
legacy_user_ref: int,
limit: int = 3,
statuses: list[str] | None = None,
):
query = select(Payment).where(Payment.tg_id == tg_id)
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return []
query = select(Payment).where(Payment.user_id == u.id)
if statuses:
query = query.where(Payment.status.in_(statuses))
@@ -109,7 +116,8 @@ async def get_last_payments(
return [
{
"id": p.id,
"tg_id": p.tg_id,
"tg_id": p.user_id,
"user_id": p.user_id,
"amount": p.amount,
"currency": p.currency,
"status": p.status,
@@ -124,27 +132,23 @@ async def get_last_payments(
async def get_payment_by_id(session: AsyncSession, internal_id: int) -> dict | None:
try:
result = await session.execute(select(Payment).where(Payment.id == internal_id).limit(1))
payment = result.scalar_one_or_none()
if not payment:
return None
return {
"id": payment.id,
"tg_id": payment.tg_id,
"amount": payment.amount,
"currency": payment.currency,
"status": payment.status,
"payment_system": payment.payment_system,
"payment_id": payment.payment_id,
"created_at": payment.created_at,
"metadata": payment.metadata_,
"original_amount": payment.original_amount,
}
except SQLAlchemyError as e:
logger.error(f"Ошибка при поиске платежа id={internal_id}: {e}")
await session.rollback()
result = await session.execute(select(Payment).where(Payment.id == internal_id).limit(1))
payment = result.scalar_one_or_none()
if not payment:
return None
return {
"id": payment.id,
"tg_id": payment.user_id,
"user_id": payment.user_id,
"amount": payment.amount,
"currency": payment.currency,
"status": payment.status,
"payment_system": payment.payment_system,
"payment_id": payment.payment_id,
"created_at": payment.created_at,
"metadata": payment.metadata_,
"original_amount": payment.original_amount,
}
async def update_payment_status(
@@ -155,32 +159,49 @@ async def update_payment_status(
payment_id: str | None = None,
metadata_patch: dict | None = None,
) -> bool:
try:
result = await session.execute(select(Payment).where(Payment.id == internal_id).limit(1))
payment = result.scalar_one_or_none()
if not payment:
logger.info(f"Не удалось сменить статус: платёж id={internal_id} не найден")
return False
payment.status = new_status
if payment_id is not None:
payment.payment_id = payment_id
base = payment.metadata_ or {}
if new_status == "success" and "status_changed_at" not in base:
base["status_changed_at"] = datetime.utcnow().replace(tzinfo=None).isoformat()
if metadata_patch:
base.update(metadata_patch)
if base:
payment.metadata_ = base
await session.commit()
logger.info(f"Статус платежа id={internal_id} изменён на {new_status}")
return True
except SQLAlchemyError as e:
await session.rollback()
logger.error(f"Ошибка при смене статуса платежа id={internal_id}: {e}")
result = await session.execute(select(Payment).where(Payment.id == internal_id).limit(1))
payment = result.scalar_one_or_none()
if not payment:
logger.info(f"Не удалось сменить статус: платёж id={internal_id} не найден")
return False
payment.status = new_status
if payment_id is not None:
payment.payment_id = payment_id
base = payment.metadata_ or {}
if new_status == "success" and "status_changed_at" not in base:
base["status_changed_at"] = datetime.utcnow().replace(tzinfo=None).isoformat()
if metadata_patch:
base.update(metadata_patch)
if base:
payment.metadata_ = base
await session.flush()
logger.info(f"Статус платежа id={internal_id} изменён на {new_status}")
return True
async def get_payment_from_db_by_payment_id(session: AsyncSession, pid: str) -> dict | None:
if not str(pid or "").strip():
return None
result = await session.execute(select(Payment).where(Payment.payment_id == pid).limit(1))
payment = result.scalar_one_or_none()
if not payment:
return None
return {
"id": payment.id,
"tg_id": payment.user_id,
"user_id": payment.user_id,
"amount": payment.amount,
"currency": payment.currency,
"status": payment.status,
"payment_system": payment.payment_system,
"payment_id": payment.payment_id,
"created_at": payment.created_at,
"metadata": payment.metadata_,
"original_amount": payment.original_amount,
}
async def get_payment_by_payment_id(session: AsyncSession, pid: str) -> dict | None:
"""Сначала Redis (pending), затем БД. Из кэша возвращается запись без id — вебхук делает add_payment."""
@@ -198,27 +219,36 @@ async def get_payment_by_payment_id(session: AsyncSession, pid: str) -> dict | N
"metadata": cached.get("metadata"),
"original_amount": cached.get("original_amount"),
}
try:
result = await session.execute(select(Payment).where(Payment.payment_id == pid).limit(1))
payment = result.scalar_one_or_none()
if not payment:
return None
return {
"id": payment.id,
"tg_id": payment.tg_id,
"amount": payment.amount,
"currency": payment.currency,
"status": payment.status,
"payment_system": payment.payment_system,
"payment_id": payment.payment_id,
"created_at": payment.created_at,
"metadata": payment.metadata_,
"original_amount": payment.original_amount,
}
except SQLAlchemyError as e:
logger.error(f"Ошибка при поиске платежа payment_id={pid}: {e}")
await session.rollback()
result = await session.execute(select(Payment).where(Payment.payment_id == pid).limit(1))
payment = result.scalar_one_or_none()
if not payment:
return None
return {
"id": payment.id,
"tg_id": payment.user_id,
"user_id": payment.user_id,
"amount": payment.amount,
"currency": payment.currency,
"status": payment.status,
"payment_system": payment.payment_system,
"payment_id": payment.payment_id,
"created_at": payment.created_at,
"metadata": payment.metadata_,
"original_amount": payment.original_amount,
}
async def count_successful_payments(session: AsyncSession, user_id: int) -> int:
"""Сколько успешных платежей у пользователя (по internal user id).
Используется для проверки "новый пользователь" в купонных правилах.
"""
result = await session.execute(
select(func.count())
.select_from(Payment)
.where(Payment.user_id == int(user_id), func.lower(Payment.status) == "success")
)
return int(result.scalar() or 0)
async def cancel_expired_pending_payments(session: AsyncSession) -> int:
@@ -234,17 +264,19 @@ async def cancel_expired_pending_payments(session: AsyncSession) -> int:
.values(status="cancelled")
)
res = await session.execute(stmt)
await session.commit()
affected = res.rowcount or 0
return affected
async def get_all_payments(
session: AsyncSession,
tg_id: int,
legacy_user_ref: int,
statuses: list[str] | None = None,
) -> list[dict]:
query = select(Payment).where(Payment.tg_id == tg_id)
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return []
query = select(Payment).where(Payment.user_id == u.id)
if statuses:
query = query.where(Payment.status.in_(statuses))
@@ -256,7 +288,8 @@ async def get_all_payments(
return [
{
"id": p.id,
"tg_id": p.tg_id,
"tg_id": p.user_id,
"user_id": p.user_id,
"amount": p.amount,
"currency": p.currency,
"status": p.status,
+94 -75
View File
@@ -1,50 +1,62 @@
from sqlalchemy import and_, desc, func, insert, select, text, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from config import CHECK_REFERRAL_REWARD_ISSUED, REFERRAL_BONUS_PERCENTAGES
from core.bootstrap import BUTTONS_CONFIG
from database.access.resolution import resolve_user_optional
from database.models import Referral
from logger import logger
async def add_referral(session: AsyncSession, referred_tg_id: int, referrer_tg_id: int):
try:
if referred_tg_id == referrer_tg_id:
logger.warning(f"⚠️ Попытка самореферала: {referred_tg_id}")
return
async def add_referral(session: AsyncSession, referred_legacy: int, referrer_legacy: int):
ru = await resolve_user_optional(session, referred_legacy)
rf = await resolve_user_optional(session, referrer_legacy)
if ru is None or rf is None:
return
if ru.id == rf.id:
logger.warning(f"⚠️ Попытка самореферала: {referred_legacy}")
return
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}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при добавлении реферала: {e}")
await session.rollback()
raise
stmt = insert(Referral).values(
referred_user_id=ru.id,
referrer_user_id=rf.id,
referred_tg_id=ru.tg_id,
referrer_tg_id=rf.tg_id,
)
await session.execute(stmt)
logger.info(f"✅ Добавлена реферальная связь: {ru.id}{rf.id}")
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)
async def get_referral_by_referred_id(session: AsyncSession, referred_legacy: int) -> dict | None:
ru = await resolve_user_optional(session, referred_legacy)
if ru is None:
return None
stmt = select(Referral).where(Referral.referred_user_id == ru.id)
result = await session.execute(stmt)
row = result.scalar_one_or_none()
return dict(row.__dict__) if row else None
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)
async def get_total_referrals(session: AsyncSession, referrer_legacy: int) -> int:
ru = await resolve_user_optional(session, referrer_legacy)
if ru is None:
return 0
stmt = select(func.count()).select_from(Referral).where(Referral.referrer_user_id == ru.id)
result = await session.execute(stmt)
return result.scalar()
async def get_active_referrals(session: AsyncSession, referrer_tg_id: int) -> int:
async def get_active_referrals(session: AsyncSession, referrer_legacy: int) -> int:
ru = await resolve_user_optional(session, referrer_legacy)
if ru is None:
return 0
stmt = (
select(func.count())
.select_from(Referral)
.where(
and_(
Referral.referrer_tg_id == referrer_tg_id,
Referral.reward_issued is True,
Referral.referrer_user_id == ru.id,
Referral.reward_issued.is_(True),
)
)
)
@@ -52,44 +64,51 @@ async def get_active_referrals(session: AsyncSession, referrer_tg_id: int) -> in
return result.scalar()
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.commit()
async def mark_referral_reward_issued(session: AsyncSession, referred_legacy: int):
ru = await resolve_user_optional(session, referred_legacy)
if ru is None:
return
await session.execute(update(Referral).where(Referral.referred_user_id == ru.id).values(reward_issued=True))
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_legacy: int, max_levels: int) -> float:
referral_enabled = bool(BUTTONS_CONFIG.get("REFERRAL_BUTTON_ENABLED", True))
if not referral_enabled:
logger.debug("Реферальная программа отключена, бонусы не начисляются")
return 0.0
ru = await resolve_user_optional(session, referrer_legacy)
if ru is None:
return 0.0
uid = ru.id
if CHECK_REFERRAL_REWARD_ISSUED:
bonus_cte = """
WITH RECURSIVE
referral_levels AS (
SELECT
referred_tg_id,
referrer_tg_id,
referred_user_id,
referrer_user_id,
1 AS level
FROM referrals
WHERE referrer_tg_id = :tg_id AND reward_issued = TRUE
WHERE referrer_user_id = :user_id AND reward_issued = TRUE
UNION
SELECT
r.referred_tg_id,
r.referrer_tg_id,
r.referred_user_id,
r.referrer_user_id,
rl.level + 1
FROM referrals r
JOIN referral_levels rl ON r.referrer_tg_id = rl.referred_tg_id
JOIN referral_levels rl ON r.referrer_user_id = rl.referred_user_id
WHERE rl.level < :max_levels AND r.reward_issued = TRUE
),
earliest_payments AS (
SELECT DISTINCT ON (tg_id) tg_id, amount, created_at
SELECT DISTINCT ON (user_id) user_id, amount, created_at
FROM payments
WHERE status = 'success'
AND payment_system NOT IN ('coupon', 'admin', 'referral')
ORDER BY tg_id, created_at
ORDER BY user_id, created_at
)
"""
bonus_query = (
@@ -110,7 +129,7 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m
END
), 0) AS total_bonus
FROM referral_levels rl
JOIN earliest_payments ep ON rl.referred_tg_id = ep.tg_id
JOIN earliest_payments ep ON rl.referred_user_id = ep.user_id
WHERE rl.level <= :max_levels
"""
)
@@ -119,20 +138,20 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m
WITH RECURSIVE
referral_levels AS (
SELECT
referred_tg_id,
referrer_tg_id,
referred_user_id,
referrer_user_id,
1 AS level
FROM referrals
WHERE referrer_tg_id = :tg_id
WHERE referrer_user_id = :user_id
UNION
SELECT
r.referred_tg_id,
r.referrer_tg_id,
r.referred_user_id,
r.referrer_user_id,
rl.level + 1
FROM referrals r
JOIN referral_levels rl ON r.referrer_tg_id = rl.referred_tg_id
JOIN referral_levels rl ON r.referrer_user_id = rl.referred_user_id
WHERE rl.level < :max_levels
)
"""
@@ -154,7 +173,7 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m
END
), 0) AS total_bonus
FROM referral_levels rl
JOIN payments p ON rl.referred_tg_id = p.tg_id
JOIN payments p ON rl.referred_user_id = p.user_id
WHERE p.status = 'success'
AND p.payment_system NOT IN ('coupon', 'admin', 'referral')
AND rl.level <= :max_levels
@@ -163,7 +182,7 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m
result = await session.execute(
text(bonus_query),
{"tg_id": referrer_tg_id, "max_levels": max_levels},
{"user_id": uid, "max_levels": max_levels},
)
total_bonus_raw = result.scalar()
total_bonus = round(float(total_bonus_raw or 0), 2)
@@ -172,29 +191,32 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m
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_legacy: int, max_levels: int) -> dict:
ru = await resolve_user_optional(session, referrer_legacy)
if ru is None:
return {}
query = """
WITH RECURSIVE referral_levels AS (
SELECT referred_tg_id, referrer_tg_id, 1 AS level
SELECT referred_user_id, referrer_user_id, 1 AS level
FROM referrals
WHERE referrer_tg_id = :referrer_tg_id
WHERE referrer_user_id = :referrer_user_id
UNION
SELECT r.referred_tg_id, r.referrer_tg_id, rl.level + 1
SELECT r.referred_user_id, r.referrer_user_id, rl.level + 1
FROM referrals r
JOIN referral_levels rl ON r.referrer_tg_id = rl.referred_tg_id
JOIN referral_levels rl ON r.referrer_user_id = rl.referred_user_id
WHERE rl.level < :max_levels
)
SELECT level,
COUNT(*) AS level_count,
COUNT(CASE WHEN reward_issued THEN 1 END) AS active_level_count
FROM referral_levels rl
JOIN referrals r ON rl.referred_tg_id = r.referred_tg_id
JOIN referrals r ON rl.referred_user_id = r.referred_user_id
GROUP BY level
ORDER BY level
"""
result = await session.execute(
text(query),
{"referrer_tg_id": referrer_tg_id, "max_levels": max_levels},
{"referrer_user_id": ru.id, "max_levels": max_levels},
)
return {
row["level"]: {
@@ -205,38 +227,35 @@ async def get_referrals_by_level(session: AsyncSession, referrer_tg_id: int, max
}
async def get_referral_stats(session: AsyncSession, referrer_tg_id: int):
try:
logger.info(f"[ReferralStats] Получение статистики для пользователя {referrer_tg_id}")
async def get_referral_stats(session: AsyncSession, referrer_legacy: int):
logger.info(f"[ReferralStats] Получение статистики для пользователя {referrer_legacy}")
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)
total_referrals = await get_total_referrals(session, referrer_legacy)
active_referrals = await get_active_referrals(session, referrer_legacy)
max_levels = len(REFERRAL_BONUS_PERCENTAGES)
referrals_by_level = await get_referrals_by_level(session, referrer_legacy, max_levels)
total_referral_bonus = await get_total_referral_bonus(session, referrer_legacy, max_levels)
return {
"total_referrals": total_referrals,
"active_referrals": active_referrals,
"referrals_by_level": referrals_by_level,
"total_referral_bonus": total_referral_bonus,
}
except Exception as e:
logger.error(f"[ReferralStats] Ошибка при получении статистики для пользователя {referrer_tg_id}: {e}")
await session.rollback()
raise
return {
"total_referrals": total_referrals,
"active_referrals": active_referrals,
"referrals_by_level": referrals_by_level,
"total_referral_bonus": total_referral_bonus,
}
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))
async def get_user_referral_count(session: AsyncSession, legacy: int) -> int:
ru = await resolve_user_optional(session, legacy)
if ru is None:
return 0
result = await session.execute(select(func.count()).select_from(Referral).where(Referral.referrer_user_id == ru.id))
return result.scalar_one() or 0
async def get_referral_position(session: AsyncSession, referral_count: int) -> int:
subq = (
select(Referral.referrer_tg_id)
.group_by(Referral.referrer_tg_id)
select(Referral.referrer_user_id)
.group_by(Referral.referrer_user_id)
.having(func.count() > referral_count)
.subquery()
)
@@ -248,10 +267,10 @@ async def get_referral_position(session: AsyncSession, referral_count: int) -> i
async def get_top_referrals(session: AsyncSession, limit: int = 5):
query = (
select(Referral.referrer_tg_id, func.count().label("referral_count"))
.group_by(Referral.referrer_tg_id)
select(Referral.referrer_user_id, func.count().label("referral_count"))
.group_by(Referral.referrer_user_id)
.order_by(desc("referral_count"))
.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_user_id": row.referrer_user_id, "referral_count": row.referral_count} for row in result.all()]
+11 -12
View File
@@ -3,6 +3,7 @@ from datetime import datetime, timezone
from sqlalchemy import select, update
from sqlalchemy.ext.asyncio import AsyncSession
from database.access.resolution import resolve_user_optional
from database.models import ScheduledBroadcast
@@ -34,8 +35,16 @@ async def create_scheduled_broadcast(
messages_per_second: int,
status: str = SCHEDULED_BROADCAST_STATUS_SCHEDULED,
) -> ScheduledBroadcast:
created_by_uid = None
mirror_tg = created_by_tg_id
if created_by_tg_id is not None:
cu = await resolve_user_optional(session, created_by_tg_id)
if cu is not None:
created_by_uid = cu.id
mirror_tg = cu.tg_id
broadcast = ScheduledBroadcast(
created_by_tg_id=created_by_tg_id,
created_by_user_id=created_by_uid,
created_by_tg_id=mirror_tg,
send_to=send_to,
cluster_name=cluster_name,
text=text,
@@ -47,7 +56,7 @@ async def create_scheduled_broadcast(
status=status,
)
session.add(broadcast)
await session.commit()
await session.flush()
await session.refresh(broadcast)
return broadcast
@@ -91,9 +100,7 @@ async def update_scheduled_broadcast(
.values(**values)
)
if not result.rowcount:
await session.rollback()
return None
await session.commit()
return await get_scheduled_broadcast(session, broadcast_id)
@@ -112,9 +119,7 @@ async def cancel_scheduled_broadcast(session: AsyncSession, broadcast_id: str) -
)
)
if not result.rowcount:
await session.rollback()
return None
await session.commit()
return await get_scheduled_broadcast(session, broadcast_id)
@@ -148,9 +153,7 @@ async def claim_due_scheduled_broadcasts(session: AsyncSession, limit: int = 10)
if claim_result.rowcount:
claimed_ids.append(broadcast_id)
if not claimed_ids:
await session.rollback()
return []
await session.commit()
result = await session.execute(
select(ScheduledBroadcast)
.where(ScheduledBroadcast.id.in_(claimed_ids))
@@ -177,9 +180,7 @@ async def start_scheduled_broadcast(session: AsyncSession, broadcast_id: str) ->
)
)
if not result.rowcount:
await session.rollback()
return None
await session.commit()
return await get_scheduled_broadcast(session, broadcast_id)
@@ -200,7 +201,6 @@ async def mark_scheduled_broadcast_sent(
updated_at=datetime.utcnow(),
)
)
await session.commit()
return await get_scheduled_broadcast(session, broadcast_id)
@@ -218,5 +218,4 @@ async def mark_scheduled_broadcast_failed(
updated_at=datetime.utcnow(),
)
)
await session.commit()
return await get_scheduled_broadcast(session, broadcast_id)
+183 -198
View File
@@ -1,5 +1,4 @@
from sqlalchemy import delete, func, insert, select, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from core.cache_config import SERVERS_CACHE_TTL_SEC
@@ -20,35 +19,23 @@ async def create_server(
subscription_url: str,
inbound_id: str,
):
try:
stmt = insert(Server).values(
cluster_name=cluster_name,
server_name=server_name,
api_url=api_url,
subscription_url=subscription_url,
inbound_id=inbound_id,
)
await session.execute(stmt)
await session.commit()
await _invalidate_servers_cache()
logger.info(f"✅ Сервер {server_name} добавлен в кластер {cluster_name}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при добавлении сервера {server_name}: {e}")
await session.rollback()
raise
stmt = insert(Server).values(
cluster_name=cluster_name,
server_name=server_name,
api_url=api_url,
subscription_url=subscription_url,
inbound_id=inbound_id,
)
await session.execute(stmt)
await _invalidate_servers_cache()
logger.info(f"✅ Сервер {server_name} добавлен в кластер {cluster_name}")
async def delete_server(session: AsyncSession, server_name: str):
try:
stmt = delete(Server).where(Server.server_name == server_name)
await session.execute(stmt)
await session.commit()
await _invalidate_servers_cache()
logger.info(f"🗑 Сервер {server_name} удалён")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при удалении сервера {server_name}: {e}")
await session.rollback()
raise
stmt = delete(Server).where(Server.server_name == server_name)
await session.execute(stmt)
await _invalidate_servers_cache()
logger.info(f"🗑 Сервер {server_name} удалён")
async def get_servers(session: AsyncSession, include_enabled: bool = False) -> dict:
@@ -59,63 +46,58 @@ async def get_servers(session: AsyncSession, include_enabled: bool = False) -> d
if isinstance(cached, dict):
return cached
try:
stmt = select(Server)
result = await session.execute(stmt)
servers = result.scalars().all()
stmt = select(Server)
result = await session.execute(stmt)
servers = result.scalars().all()
ids = [s.id for s in servers]
subs_map = {}
tariffs_map = {}
if ids:
r = await session.execute(
select(ServerSubgroup.server_id, ServerSubgroup.subgroup_title).where(ServerSubgroup.server_id.in_(ids))
ids = [s.id for s in servers]
subs_map = {}
tariffs_map = {}
if ids:
r = await session.execute(
select(ServerSubgroup.server_id, ServerSubgroup.subgroup_title).where(ServerSubgroup.server_id.in_(ids))
)
for sid, sg in r.all():
if sg and sg.isdigit():
tariffs_map.setdefault(sid, []).append(int(sg))
else:
subs_map.setdefault(sid, []).append(sg)
groups_map = {}
if ids:
r2 = await session.execute(
select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where(
ServerSpecialgroup.server_id.in_(ids)
)
for sid, sg in r.all():
if sg and sg.isdigit():
tariffs_map.setdefault(sid, []).append(int(sg))
else:
subs_map.setdefault(sid, []).append(sg)
)
for sid, gc in r2.all():
groups_map.setdefault(sid, []).append(gc)
groups_map = {}
if ids:
r2 = await session.execute(
select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where(
ServerSpecialgroup.server_id.in_(ids)
)
)
for sid, gc in r2.all():
groups_map.setdefault(sid, []).append(gc)
allowed = set(ALLOWED_GROUP_CODES)
allowed = set(ALLOWED_GROUP_CODES)
grouped = {}
for s in servers:
if not include_enabled and not s.enabled:
continue
cluster = s.cluster_name
special = sorted({g for g in groups_map.get(s.id, []) if g in allowed})
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,
"tariff_subgroups": subs_map.get(s.id, []),
"tariff_ids": tariffs_map.get(s.id, []),
"special_groups": special,
"cluster_name": cluster,
"server_id": s.id,
})
await cache_set(cache_key_servers, grouped, SERVERS_CACHE_TTL_SEC)
return grouped
except SQLAlchemyError as e:
logger.error(f"Ошибка при получении серверов: {e}")
await session.rollback()
return {}
grouped = {}
for s in servers:
if not include_enabled and not s.enabled:
continue
cluster = s.cluster_name
special = sorted({g for g in groups_map.get(s.id, []) if g in allowed})
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,
"tariff_subgroups": subs_map.get(s.id, []),
"tariff_ids": tariffs_map.get(s.id, []),
"special_groups": special,
"cluster_name": cluster,
"server_id": s.id,
})
await cache_set(cache_key_servers, grouped, SERVERS_CACHE_TTL_SEC)
return grouped
async def get_clusters(session: AsyncSession) -> list[str]:
@@ -139,14 +121,49 @@ async def check_unique_server_name(session: AsyncSession, server_name: str, clus
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))
row = result.first()
return {"cluster_name": row[0]} if row else None
except SQLAlchemyError as e:
logger.error(f"Ошибка при поиске кластера для сервера {server_name}: {e}")
await session.rollback()
return None
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
async def get_panel_types_for_cluster(session: AsyncSession, cluster_name: str) -> list[str]:
"""Список panel_type всех серверов кластера (для проверки "весь remnawave")."""
result = await session.execute(
select(Server.panel_type).where(Server.cluster_name == cluster_name)
)
return list(result.scalars().all())
async def get_panel_type_for_server(session: AsyncSession, server_name: str) -> str | None:
"""Возвращает panel_type конкретного сервера по его имени."""
result = await session.execute(
select(Server.panel_type).where(Server.server_name == server_name)
)
return result.scalar_one_or_none()
async def get_enabled_server_subscription_url(session: AsyncSession, server_name: str) -> str | None:
"""Возвращает ``subscription_url`` для включённого сервера по его имени."""
result = await session.execute(
select(Server.subscription_url).where(Server.server_name == server_name, Server.enabled.is_(True))
)
return result.scalar()
async def cluster_name_exists(session: AsyncSession, cluster_name: str) -> bool:
"""Есть ли хоть один сервер с таким cluster_name."""
result = await session.execute(
select(Server).where(Server.cluster_name == cluster_name).limit(1)
)
return result.scalars().first() is not None
async def get_cluster_name_for_server_name(session: AsyncSession, server_name: str) -> str | None:
"""Возвращает cluster_name для указанного server_name (строго по server_name)."""
result = await session.execute(
select(Server.cluster_name).where(Server.server_name == server_name).limit(1)
)
return result.scalar()
async def get_cluster_name_by_server(session: AsyncSession, server_id_or_name: str) -> str | None:
@@ -162,132 +179,100 @@ async def get_cluster_name_by_server(session: AsyncSession, server_id_or_name: s
async def get_server_by_name(session: AsyncSession, server_name: str) -> dict | None:
try:
stmt = select(Server).where(Server.server_name == server_name)
result = await session.execute(stmt)
server = result.scalar_one_or_none()
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,
"cluster_name": server.cluster_name,
"server_name": server.server_name,
"api_url": server.api_url,
"subscription_url": server.subscription_url,
"inbound_id": server.inbound_id,
"panel_type": server.panel_type,
"enabled": server.enabled,
"max_keys": server.max_keys,
"tariff_group": server.tariff_group,
}
return None
except SQLAlchemyError as e:
logger.error(f"Ошибка при получении сервера {server_name}: {e}")
await session.rollback()
return None
if server:
return {
"id": server.id,
"cluster_name": server.cluster_name,
"server_name": server.server_name,
"api_url": server.api_url,
"subscription_url": server.subscription_url,
"inbound_id": server.inbound_id,
"panel_type": server.panel_type,
"enabled": server.enabled,
"max_keys": server.max_keys,
"tariff_group": server.tariff_group,
}
return None
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)
await session.commit()
await _invalidate_servers_cache()
logger.info(f"✅ Поле {field} сервера {server_name} обновлено на {value}")
return True
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при обновлении поля {field} сервера {server_name}: {e}")
await session.rollback()
return False
stmt = update(Server).where(Server.server_name == server_name).values(**{field: value})
await session.execute(stmt)
await _invalidate_servers_cache()
logger.info(f"✅ Поле {field} сервера {server_name} обновлено на {value}")
return True
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):
logger.error(f"❌ Сервер с именем {new_name} уже существует")
return False
stmt_server = update(Server).where(Server.server_name == old_name).values(server_name=new_name)
await session.execute(stmt_server)
stmt_keys = update(Key).where(Key.server_id == old_name).values(server_id=new_name)
await session.execute(stmt_keys)
await session.commit()
await _invalidate_servers_cache()
logger.info(f"✅ Сервер переименован с {old_name} на {new_name}")
return True
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при переименовании сервера {old_name}: {e}")
await session.rollback()
if not await check_unique_server_name(session, new_name):
logger.error(f"❌ Сервер с именем {new_name} уже существует")
return False
stmt_server = update(Server).where(Server.server_name == old_name).values(server_name=new_name)
await session.execute(stmt_server)
stmt_keys = update(Key).where(Key.server_id == old_name).values(server_id=new_name)
await session.execute(stmt_keys)
await _invalidate_servers_cache()
logger.info(f"✅ Сервер переименован с {old_name} на {new_name}")
return True
async def get_available_clusters(session: AsyncSession) -> list[str]:
try:
stmt = select(Server.cluster_name).distinct().order_by(Server.cluster_name)
result = await session.execute(stmt)
return [row[0] for row in result.all()]
except SQLAlchemyError as e:
logger.error(f"Ошибка при получении списка кластеров: {e}")
await session.rollback()
return []
stmt = select(Server.cluster_name).distinct().order_by(Server.cluster_name)
result = await session.execute(stmt)
return [row[0] for row in result.all()]
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(
(Server.cluster_name == old_cluster) & (Server.server_name != server_name)
)
result = await session.execute(stmt_remaining)
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)
await session.execute(stmt_update_keys)
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()
await session.execute(
update(Server)
.where(Server.server_name == server_name)
.values(cluster_name=new_cluster, tariff_group=new_tariff_group)
)
if server_data.get("id") is None:
rid = await session.execute(select(Server.id).where(Server.server_name == server_name).limit(1))
server_id = rid.scalar_one_or_none()
else:
server_id = server_data["id"]
if server_id is not None and new_tariff_group is not None:
await session.execute(
update(ServerSubgroup).where(ServerSubgroup.server_id == server_id).values(group_code=new_tariff_group)
)
await session.commit()
await _invalidate_servers_cache()
logger.info(
f"✅ Сервер {server_name} перемещен в кластер {new_cluster} с обновлением тарифной группы и привязок подгрупп"
)
return True
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при обновлении кластера сервера {server_name}: {e}")
await session.rollback()
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(
(Server.cluster_name == old_cluster) & (Server.server_name != server_name)
)
result = await session.execute(stmt_remaining)
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)
await session.execute(stmt_update_keys)
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()
await session.execute(
update(Server)
.where(Server.server_name == server_name)
.values(cluster_name=new_cluster, tariff_group=new_tariff_group)
)
if server_data.get("id") is None:
rid = await session.execute(select(Server.id).where(Server.server_name == server_name).limit(1))
server_id = rid.scalar_one_or_none()
else:
server_id = server_data["id"]
if server_id is not None and new_tariff_group is not None:
await session.execute(
update(ServerSubgroup).where(ServerSubgroup.server_id == server_id).values(group_code=new_tariff_group)
)
await _invalidate_servers_cache()
logger.info(
f"✅ Сервер {server_name} перемещен в кластер {new_cluster} с обновлением тарифной группы и привязок подгрупп"
)
return True
async def resolve_device_limit_from_group(session: AsyncSession, server_id: str) -> int | None:
r = await session.execute(select(Server.tariff_group).where(Server.server_name == server_id))
+1
View File
@@ -0,0 +1 @@
from .init_db import *
@@ -4,12 +4,19 @@ from sqlalchemy import select
from config import ADMIN_ID
from database import db
from database.migrations.schema_upgrade import (
apply_all_migrations,
ensure_tg_mirror_columns_and_backfill,
)
from database.models import Admin, Base, User
async def init_db():
async with db.engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
await apply_all_migrations(conn)
async with db.engine.begin() as conn:
await ensure_tg_mirror_columns_and_backfill(conn)
async with db.async_session_maker() as session:
result = await session.execute(select(User).where(User.tg_id == 0))
+3 -3
View File
@@ -152,15 +152,15 @@ 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.user_id).where(Key.expiry_time > int(datetime.utcnow().timestamp() * 1000)).distinct()
)
stmt = (
select(Payment.tg_id)
select(Payment.user_id)
.where(Payment.amount > 0)
.where(Payment.status == "success")
.where(Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED))
.where(not_(exists(subquery_active_keys.where(Key.tg_id == Payment.tg_id))))
.where(not_(exists(subquery_active_keys.where(Key.user_id == Payment.user_id))))
.distinct()
)
+124 -185
View File
@@ -4,7 +4,6 @@ from collections import defaultdict
from datetime import datetime
from sqlalchemy import delete, func, insert, select, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from core.cache_config import TARIFF_BY_ID_CACHE_TTL_SEC, TARIFFS_FOR_CLUSTER_CACHE_TTL_SEC
@@ -31,6 +30,7 @@ async def _invalidate_tariff_cache(tariff_id: int | None = None) -> None:
if tariff_id is not None:
await cache_delete(cache_key("tariff", tariff_id))
await cache_delete_pattern("tariffs_cluster:*")
await cache_delete_pattern("tariffs_public:*")
def create_subgroup_hash(subgroup_title: str, group_code: str) -> str:
@@ -60,43 +60,37 @@ async def find_subgroup_by_hash(session: AsyncSession, subgroup_hash: str, group
async def get_tariffs(
session: AsyncSession, tariff_id: int = None, group_code: str = None, with_subgroup_weights: bool = False
):
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.sort_order, Tariff.id)
)
else:
result = await session.execute(select(Tariff).order_by(Tariff.sort_order, Tariff.id))
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.sort_order, Tariff.id)
)
else:
result = await session.execute(select(Tariff).order_by(Tariff.sort_order, Tariff.id))
tariffs = [dict(r.__dict__) for r in result.scalars().all()]
tariffs = [dict(r.__dict__) for r in result.scalars().all()]
if with_subgroup_weights and group_code:
tariffs_without_order = [t for t in tariffs if t.get("sort_order") is None]
if tariffs_without_order:
for tariff in tariffs_without_order:
tariff["sort_order"] = 1
await session.execute(update(Tariff).where(Tariff.id == tariff["id"]).values(sort_order=1))
await session.commit()
if with_subgroup_weights and group_code:
tariffs_without_order = [t for t in tariffs if t.get("sort_order") is None]
if tariffs_without_order:
for tariff in tariffs_without_order:
tariff["sort_order"] = 1
await session.execute(update(Tariff).where(Tariff.id == tariff["id"]).values(sort_order=1))
grouped = defaultdict(list)
for t in tariffs:
grouped[t.get("subgroup_title")].append(t)
grouped = defaultdict(list)
for t in tariffs:
grouped[t.get("subgroup_title")].append(t)
subgroup_weights = {}
for subgroup, tariffs_list in grouped.items():
if subgroup:
total_weight = sum(t.get("sort_order", 1) for t in tariffs_list)
subgroup_weights[subgroup] = total_weight
subgroup_weights = {}
for subgroup, tariffs_list in grouped.items():
if subgroup:
total_weight = sum(t.get("sort_order", 1) for t in tariffs_list)
subgroup_weights[subgroup] = total_weight
return {"tariffs": tariffs, "subgroup_weights": subgroup_weights}
return {"tariffs": tariffs, "subgroup_weights": subgroup_weights}
return tariffs
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при получении тарифов: {e}")
await session.rollback()
return []
return tariffs
async def get_tariff_names_groups_subgroups_durations(
@@ -133,18 +127,13 @@ async def get_tariff_by_id(session: AsyncSession, tariff_id: int):
cached = await cache_get(key)
if isinstance(cached, dict):
return cached
try:
result = await session.execute(select(Tariff).where(Tariff.id == tariff_id))
tariff = result.scalar_one_or_none()
if not tariff:
return None
row = _row_to_cache_dict(dict(tariff.__dict__))
await cache_set(key, row, TARIFF_BY_ID_CACHE_TTL_SEC)
return row
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при получении тарифа по ID {tariff_id}: {e}")
await session.rollback()
result = await session.execute(select(Tariff).where(Tariff.id == tariff_id))
tariff = result.scalar_one_or_none()
if not tariff:
return None
row = _row_to_cache_dict(dict(tariff.__dict__))
await cache_set(key, row, TARIFF_BY_ID_CACHE_TTL_SEC)
return row
async def get_tariff_group_codes(session: AsyncSession) -> list[str]:
@@ -152,6 +141,14 @@ async def get_tariff_group_codes(session: AsyncSession) -> list[str]:
return [row[0] for row in result.fetchall() if row[0]]
async def get_active_tariff_by_id(session: AsyncSession, tariff_id: int) -> Tariff | None:
"""Возвращает ORM-объект Tariff по id, если тариф активен (is_active=True)."""
result = await session.execute(
select(Tariff).where(Tariff.id == int(tariff_id), Tariff.is_active.is_(True))
)
return result.scalar_one_or_none()
async def get_active_tariffs_by_group_code(session: AsyncSession, group_code: str) -> list[Tariff]:
result = await session.execute(
select(Tariff).where(Tariff.group_code == group_code, Tariff.is_active.is_(True)).order_by(Tariff.id)
@@ -164,105 +161,78 @@ async def get_tariffs_for_cluster(session: AsyncSession, cluster_name: str):
cached = await cache_get(key)
if isinstance(cached, list):
return cached
try:
server_row = await session.execute(
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.cluster_name == cluster_name).limit(1)
select(Server.tariff_group).where(Server.server_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)
)
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.sort_order, Tariff.id)
)
rows = [_row_to_cache_dict(dict(r.__dict__)) for r in result.scalars().all()]
await cache_set(key, rows, TARIFFS_FOR_CLUSTER_CACHE_TTL_SEC)
return rows
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при получении тарифов для кластера {cluster_name}: {e}")
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.sort_order, Tariff.id)
)
rows = [_row_to_cache_dict(dict(r.__dict__)) for r in result.scalars().all()]
await cache_set(key, rows, TARIFFS_FOR_CLUSTER_CACHE_TTL_SEC)
return rows
async def create_tariff(session: AsyncSession, data: dict):
try:
data["created_at"] = datetime.utcnow()
data["updated_at"] = datetime.utcnow()
data["created_at"] = datetime.utcnow()
data["updated_at"] = datetime.utcnow()
if "sort_order" not in data:
group_code = data.get("group_code")
if group_code:
result = await session.execute(
select(func.max(Tariff.sort_order)).where(
Tariff.group_code == group_code, Tariff.sort_order.isnot(None)
)
if "sort_order" not in data:
group_code = data.get("group_code")
if group_code:
result = await session.execute(
select(func.max(Tariff.sort_order)).where(
Tariff.group_code == group_code, Tariff.sort_order.isnot(None)
)
max_order = result.scalar() or 0
else:
result = await session.execute(select(func.max(Tariff.sort_order)).where(Tariff.sort_order.isnot(None)))
max_order = result.scalar() or 0
)
max_order = result.scalar() or 0
else:
result = await session.execute(select(func.max(Tariff.sort_order)).where(Tariff.sort_order.isnot(None)))
max_order = result.scalar() or 0
data["sort_order"] = max_order + 1
data["sort_order"] = max_order + 1
stmt = insert(Tariff).values(**data).returning(Tariff)
result = await session.execute(stmt)
await session.commit()
await _invalidate_tariff_cache()
return result.scalar_one()
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при создании тарифа: {e}")
await session.rollback()
return None
stmt = insert(Tariff).values(**data).returning(Tariff)
result = await session.execute(stmt)
await _invalidate_tariff_cache()
return result.scalar_one()
async def update_tariff(session: AsyncSession, tariff_id: int, updates: dict):
if not updates:
return False
try:
updates["updated_at"] = datetime.utcnow()
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(**updates))
await session.commit()
await _invalidate_tariff_cache(tariff_id)
return True
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при обновлении тарифа ID={tariff_id}: {e}")
await session.rollback()
return False
updates["updated_at"] = datetime.utcnow()
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(**updates))
await _invalidate_tariff_cache(tariff_id)
return True
async def delete_tariff(session: AsyncSession, tariff_id: int):
try:
await session.execute(delete(Tariff).where(Tariff.id == tariff_id))
await session.commit()
await _invalidate_tariff_cache(tariff_id)
return True
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при удалении тарифа ID={tariff_id}: {e}")
await session.rollback()
return False
await session.execute(delete(Tariff).where(Tariff.id == tariff_id))
await _invalidate_tariff_cache(tariff_id)
return True
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)))
tariff = result.scalar_one_or_none()
if tariff:
return True
logger.warning(f"[TARIFF] Тариф {tariff_id} не найден в БД")
return False
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при проверке тарифа {tariff_id}: {e}")
await session.rollback()
return False
result = await session.execute(select(Tariff).where(Tariff.id == tariff_id, Tariff.is_active.is_(True)))
tariff = result.scalar_one_or_none()
if tariff:
return True
logger.warning(f"[TARIFF] Тариф {tariff_id} не найден в БД")
return False
async def get_vless_enabled(session: AsyncSession, tariff_id: int | None) -> bool:
@@ -292,90 +262,59 @@ async def get_vless_enabled_batch(
async def get_tariff_sort_order(session: AsyncSession, tariff_id: int) -> int:
try:
result = await session.execute(select(Tariff.sort_order).where(Tariff.id == tariff_id))
sort_order = result.scalar_one_or_none()
result = await session.execute(select(Tariff.sort_order).where(Tariff.id == tariff_id))
sort_order = result.scalar_one_or_none()
if sort_order is None:
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=1))
await session.commit()
await _invalidate_tariff_cache(tariff_id)
return 1
if sort_order is None:
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=1))
await _invalidate_tariff_cache(tariff_id)
return 1
return sort_order
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при получении sort_order для тарифа {tariff_id}: {e}")
await session.rollback()
return None
return sort_order
async def move_tariff_up(session: AsyncSession, tariff_id: int) -> bool:
try:
current_order = await get_tariff_sort_order(session, tariff_id)
new_order = max(1, current_order - 1)
current_order = await get_tariff_sort_order(session, tariff_id)
new_order = max(1, current_order - 1)
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=new_order))
await session.commit()
await _invalidate_tariff_cache(tariff_id)
return True
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при перемещении тарифа {tariff_id} вверх: {e}")
await session.rollback()
return False
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=new_order))
await _invalidate_tariff_cache(tariff_id)
return True
async def move_tariff_down(session: AsyncSession, tariff_id: int) -> bool:
try:
current_order = await get_tariff_sort_order(session, tariff_id)
new_order = current_order + 1
current_order = await get_tariff_sort_order(session, tariff_id)
new_order = current_order + 1
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=new_order))
await session.commit()
await _invalidate_tariff_cache(tariff_id)
return True
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при перемещении тарифа {tariff_id} вниз: {e}")
await session.rollback()
return False
await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=new_order))
await _invalidate_tariff_cache(tariff_id)
return True
async def initialize_tariff_sort_orders(session: AsyncSession, group_code: str) -> bool:
try:
result = await session.execute(select(Tariff).where(Tariff.group_code == group_code).order_by(Tariff.id))
tariffs = result.scalars().all()
result = await session.execute(select(Tariff).where(Tariff.group_code == group_code).order_by(Tariff.id))
tariffs = result.scalars().all()
if not tariffs:
return True
for i, tariff in enumerate(tariffs):
new_sort_order = 1 + i
await session.execute(update(Tariff).where(Tariff.id == tariff.id).values(sort_order=new_sort_order))
await session.commit()
await _invalidate_tariff_cache()
if not tariffs:
return True
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при инициализации sort_order для группы {group_code}: {e}")
await session.rollback()
return False
for i, tariff in enumerate(tariffs):
new_sort_order = 1 + i
await session.execute(update(Tariff).where(Tariff.id == tariff.id).values(sort_order=new_sort_order))
await _invalidate_tariff_cache()
return True
async def initialize_all_tariff_weights(session: AsyncSession) -> bool:
try:
result = await session.execute(select(Tariff).where(Tariff.sort_order.is_(None)))
tariffs_without_weight = result.scalars().all()
result = await session.execute(select(Tariff).where(Tariff.sort_order.is_(None)))
tariffs_without_weight = result.scalars().all()
if not tariffs_without_weight:
return True
for tariff in tariffs_without_weight:
await session.execute(update(Tariff).where(Tariff.id == tariff.id).values(sort_order=1))
await session.commit()
await _invalidate_tariff_cache()
if not tariffs_without_weight:
return True
except SQLAlchemyError as e:
logger.error(f"[TARIFF] Ошибка при инициализации весов тарифов: {e}")
await session.rollback()
return False
for tariff in tariffs_without_weight:
await session.execute(update(Tariff).where(Tariff.id == tariff.id).values(sort_order=1))
await _invalidate_tariff_cache()
return True
+49 -25
View File
@@ -1,35 +1,49 @@
from datetime import datetime
from sqlalchemy import delete, select
from sqlalchemy import delete, or_, select
from sqlalchemy.dialects.postgresql import insert
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from database.access.resolution import resolve_user_optional
from database.models import TemporaryData
from logger import logger
async def create_temporary_data(session: AsyncSession, tg_id: int, state: str, data: dict):
try:
stmt = (
insert(TemporaryData)
.values(tg_id=tg_id, state=state, data=data, updated_at=datetime.utcnow())
.on_conflict_do_update(
index_elements=[TemporaryData.tg_id],
set_={"state": state, "data": data, "updated_at": datetime.utcnow()},
async def create_temporary_data(session: AsyncSession, legacy_user_ref: int, state: str, data: dict):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
raise ValueError(f"user not found for temporary data: {legacy_user_ref}")
ins = insert(TemporaryData).values(
user_id=u.id,
tg_id=u.tg_id,
state=state,
data=data,
updated_at=datetime.utcnow(),
)
stmt = ins.on_conflict_do_update(
index_elements=[TemporaryData.user_id],
set_={
"state": state,
"data": data,
"updated_at": datetime.utcnow(),
"tg_id": ins.excluded.tg_id,
},
)
await session.execute(stmt)
logger.info(f"📝 Временные данные сохранены для user_id={u.id}")
async def get_temporary_data(session: AsyncSession, legacy_user_ref: int) -> dict | None:
u = await resolve_user_optional(session, legacy_user_ref)
if u is not None:
stmt = select(TemporaryData).where(
or_(
TemporaryData.user_id == u.id,
TemporaryData.tg_id == u.tg_id,
)
)
await session.execute(stmt)
await session.commit()
logger.info(f"📝 Временные данные сохранены для {tg_id}")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при сохранении временных данных для {tg_id}: {e}")
await session.rollback()
raise
async def get_temporary_data(session: AsyncSession, tg_id: int) -> dict | None:
stmt = select(TemporaryData).where(TemporaryData.tg_id == tg_id)
else:
stmt = select(TemporaryData).where(TemporaryData.tg_id == legacy_user_ref)
result = await session.execute(stmt)
row = result.scalar_one_or_none()
if row:
@@ -37,7 +51,17 @@ async def get_temporary_data(session: AsyncSession, tg_id: int) -> dict | None:
return None
async def clear_temporary_data(session: AsyncSession, tg_id: int):
await session.execute(delete(TemporaryData).where(TemporaryData.tg_id == tg_id))
await session.commit()
logger.info(f"🗑 Временные данные очищены для {tg_id}")
async def clear_temporary_data(session: AsyncSession, legacy_user_ref: int):
u = await resolve_user_optional(session, legacy_user_ref)
if u is not None:
await session.execute(
delete(TemporaryData).where(
or_(
TemporaryData.user_id == u.id,
TemporaryData.tg_id == u.tg_id,
)
)
)
else:
await session.execute(delete(TemporaryData).where(TemporaryData.tg_id == legacy_user_ref))
logger.info(f"🗑 Временные данные очищены для {legacy_user_ref}")
+20 -27
View File
@@ -1,5 +1,4 @@
from sqlalchemy import and_, func, insert, not_, select
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy import and_, func, insert, select
from sqlalchemy.ext.asyncio import AsyncSession
from core.constants import PAYMENT_SYSTEMS_EXCLUDED
@@ -8,40 +7,34 @@ from logger import logger
async def create_tracking_source(session: AsyncSession, name: str, code: str, type_: str, created_by: int):
try:
stmt = insert(TrackingSource).values(
name=name,
code=code,
type=type_,
created_by=created_by,
)
await session.execute(stmt)
await session.commit()
logger.info(f"🆕 Источник трафика {code} создан")
except SQLAlchemyError as e:
logger.error(f"❌ Ошибка при создании источника {code}: {e}")
await session.rollback()
raise
stmt = insert(TrackingSource).values(
name=name,
code=code,
type=type_,
created_by=created_by,
)
await session.execute(stmt)
logger.info(f"🆕 Источник трафика {code} создан")
async def get_all_tracking_sources(session: AsyncSession) -> list[dict]:
registrations_subq = (
select(func.count(func.distinct(User.tg_id)))
select(func.count(func.distinct(User.id)))
.where(User.source_code == TrackingSource.code)
.correlate(TrackingSource)
.scalar_subquery()
)
trials_subq = (
select(func.count(func.distinct(User.tg_id)))
select(func.count(func.distinct(User.id)))
.where((User.source_code == TrackingSource.code) & (User.trial == 1))
.correlate(TrackingSource)
.scalar_subquery()
)
payments_subq = (
select(func.count(func.distinct(Payment.tg_id)))
.join(User, Payment.tg_id == User.tg_id)
select(func.count(func.distinct(Payment.user_id)))
.join(User, Payment.user_id == User.id)
.where(
(User.source_code == TrackingSource.code)
& (Payment.status == "success")
@@ -89,20 +82,20 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict |
_src_name, _src_code, created_at = src
reg_subq = (
select(func.count(func.distinct(User.tg_id)))
select(func.count(func.distinct(User.id)))
.where((User.source_code == code) & (User.created_at >= created_at))
.scalar_subquery()
)
trial_subq = (
select(func.count(func.distinct(User.tg_id)))
select(func.count(func.distinct(User.id)))
.where((User.source_code == code) & (User.trial == 1) & (User.created_at >= created_at))
.scalar_subquery()
)
payments_subq = (
select(func.count(func.distinct(Payment.tg_id)))
.join(User, Payment.tg_id == User.tg_id)
select(func.count(func.distinct(Payment.user_id)))
.join(User, Payment.user_id == User.id)
.where(
(User.source_code == code)
& (Payment.status == "success")
@@ -114,7 +107,7 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict |
amount_subq = (
select(func.coalesce(func.sum(Payment.amount), 0.0))
.join(User, Payment.tg_id == User.tg_id)
.join(User, Payment.user_id == User.id)
.where(
(User.source_code == code)
& (Payment.status == "success")
@@ -141,11 +134,11 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict |
payments_base = (
select(
Payment.tg_id.label("tg_id"),
Payment.user_id.label("tg_id"),
Payment.amount.label("amount"),
Payment.created_at.label("dt"),
)
.join(User, Payment.tg_id == User.tg_id)
.join(User, Payment.user_id == User.id)
.where(
(User.source_code == code)
& (Payment.status == "success")
+217 -192
View File
@@ -1,8 +1,7 @@
from datetime import datetime
from sqlalchemy import delete, exists, func, or_, select, update
from sqlalchemy import delete, func, or_, select, update
from sqlalchemy.dialects.postgresql import insert
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.ext.asyncio import AsyncSession
from core.cache_config import (
@@ -11,6 +10,7 @@ from core.cache_config import (
USER_SNAPSHOT_CACHE_TTL_SEC,
)
from core.redis_cache import cache_delete, cache_get, cache_key, cache_set
from database.access.resolution import resolve_user_optional
from database.models import (
BlockedUser,
CouponUsage,
@@ -45,36 +45,28 @@ async def add_user(
language_code: str = None,
is_bot: bool = False,
source_code: str = None,
commit: bool = True,
) -> bool:
try:
stmt = (
insert(User)
.values(
tg_id=tg_id,
username=username,
first_name=first_name,
last_name=last_name,
language_code=language_code,
is_bot=is_bot,
source_code=source_code,
)
.on_conflict_do_nothing(index_elements=["tg_id"])
.returning(User.tg_id)
) -> int | None:
stmt = (
insert(User)
.values(
tg_id=tg_id,
username=username,
first_name=first_name,
last_name=last_name,
language_code=language_code,
is_bot=is_bot,
source_code=source_code,
)
res = await session.execute(stmt)
inserted_tg_id = res.scalar_one_or_none()
if inserted_tg_id is None:
return False
if commit:
await session.commit()
await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC)
logger.info(f"[DB] Новый пользователь добавлен: {tg_id} (source: {source_code})")
return True
except SQLAlchemyError as e:
logger.error(f"[DB] Ошибка при добавлении пользователя {tg_id}: {e}")
await session.rollback()
raise
.on_conflict_do_nothing(index_elements=["tg_id"])
.returning(User.id)
)
res = await session.execute(stmt)
inserted_id = res.scalar_one_or_none()
if inserted_id is None:
return None
await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC)
logger.info(f"[DB] Новый пользователь добавлен: tg_id={tg_id} id={inserted_id} (source: {source_code})")
return int(inserted_id)
async def invalidate_balance_cache(tg_id: int) -> None:
@@ -87,109 +79,130 @@ async def invalidate_profile_cache(tg_id: int) -> None:
async def update_balance(
session: AsyncSession,
tg_id: int,
legacy_user_ref: int,
amount: float,
) -> None:
try:
amount = float(amount)
res = await session.execute(
update(User)
.where(User.tg_id == tg_id)
.values(balance=func.coalesce(User.balance, 0) + amount)
.returning(User.balance)
)
new_balance = res.scalar_one_or_none()
if new_balance is not None:
old_balance = new_balance - amount
await session.commit()
if new_balance is not None:
logger.info(f"[DB] Баланс пользователя {tg_id} обновлён: {old_balance}{new_balance}")
else:
logger.info(f"[DB] Баланс пользователя {tg_id} не изменён: пользователь не найден")
await invalidate_balance_cache(tg_id)
await invalidate_profile_cache(tg_id)
except SQLAlchemyError as e:
logger.error(f"[DB] Ошибка при обновлении баланса пользователя {tg_id}: {e}")
await session.rollback()
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
logger.info(f"[DB] Баланс не изменён: пользователь {legacy_user_ref} не найден")
return
uid = u.id
amount = float(amount)
res = await session.execute(
update(User)
.where(User.id == uid)
.values(balance=func.coalesce(User.balance, 0) + amount)
.returning(User.balance)
)
new_balance = res.scalar_one_or_none()
if new_balance is not None:
old_balance = new_balance - amount
logger.info(f"[DB] Баланс пользователя id={uid} обновлён: {old_balance}{new_balance}")
else:
logger.info(f"[DB] Баланс пользователя id={uid} не изменён: пользователь не найден")
await invalidate_balance_cache(uid)
await invalidate_profile_cache(uid)
if u.tg_id is not None:
await invalidate_balance_cache(u.tg_id)
await invalidate_profile_cache(u.tg_id)
async def check_user_exists(session: AsyncSession, tg_id: int) -> bool:
cached = await cache_get(cache_key("user_exists", tg_id))
async def check_user_exists(session: AsyncSession, legacy_user_ref: int) -> bool:
cached = await cache_get(cache_key("user_exists", legacy_user_ref))
if isinstance(cached, bool):
return cached
stmt = select(exists().where(User.tg_id == tg_id))
result = await session.execute(stmt)
value = result.scalar()
await cache_set(cache_key("user_exists", tg_id), bool(value), USER_EXISTS_CACHE_TTL_SEC)
u = await resolve_user_optional(session, legacy_user_ref)
value = u is not None
await cache_set(cache_key("user_exists", legacy_user_ref), value, USER_EXISTS_CACHE_TTL_SEC)
return value
async def get_balance(session: AsyncSession, tg_id: int) -> float:
cached = await cache_get(cache_key("balance", tg_id))
async def get_balance(session: AsyncSession, legacy_user_ref: int) -> float:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return 0.0
uid = u.id
cached = await cache_get(cache_key("balance", uid))
if cached is not None:
try:
return round(float(cached), 1)
except (TypeError, ValueError):
pass
result = await session.execute(select(func.coalesce(User.balance, 0.0)).where(User.tg_id == tg_id))
result = await session.execute(select(func.coalesce(User.balance, 0.0)).where(User.id == uid))
balance = result.scalar_one_or_none()
value = round(float(balance or 0.0), 1)
await cache_set(cache_key("balance", tg_id), value, BALANCE_CACHE_TTL_SEC)
await cache_set(cache_key("balance", uid), value, BALANCE_CACHE_TTL_SEC)
return value
async def set_user_balance(
session: AsyncSession,
tg_id: int,
legacy_user_ref: int,
balance: float,
) -> None:
try:
old_balance_result = await session.execute(select(func.coalesce(User.balance, 0.0)).where(User.tg_id == tg_id))
old_balance = old_balance_result.scalar_one_or_none()
if old_balance is None:
await session.execute(update(User).where(User.tg_id == tg_id).values(balance=balance))
await session.commit()
await invalidate_balance_cache(tg_id)
await invalidate_profile_cache(tg_id)
return
balance = float(balance)
await session.execute(update(User).where(User.tg_id == tg_id).values(balance=balance))
await session.commit()
await invalidate_balance_cache(tg_id)
await invalidate_profile_cache(tg_id)
except SQLAlchemyError as e:
logger.error(f"Ошибка при установке баланса для пользователя {tg_id}: {e}")
await session.rollback()
raise
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
uid = u.id
balance = float(balance)
await session.execute(update(User).where(User.id == uid).values(balance=balance))
await invalidate_balance_cache(uid)
await invalidate_profile_cache(uid)
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.commit()
await invalidate_profile_cache(tg_id)
invalidate_user_snapshot(tg_id)
logger.info(f"[DB] Триал статус обновлён для пользователя {tg_id}: {status}")
except SQLAlchemyError as e:
logger.error(f"[DB] Ошибка при обновлении триала пользователя {tg_id}: {e}")
await session.rollback()
raise
async def get_user_preferred_currency(session: AsyncSession, tg_id: int) -> str | None:
"""Предпочитаемая валюта пользователя по ``tg_id``, если установлена."""
result = await session.execute(
select(User.preferred_currency).where(User.tg_id == int(tg_id))
)
return result.scalar()
async def get_trial(session: AsyncSession, tg_id: int) -> int:
result = await session.execute(select(func.coalesce(User.trial, 0)).where(User.tg_id == tg_id))
async def mark_trial_started_if_eligible(session: AsyncSession, tg_id: int) -> None:
"""Переводит `trial` в 1, только если текущее значение in [0, -1] (пользователь
ещё не использовал триал). Условный update без пред-чтения — атомарно на уровне БД.
Используется в `services.operations.creation.create_key_on_cluster` после
успешного создания ключа.
"""
await session.execute(
update(User).where(User.tg_id == tg_id, User.trial.in_([0, -1])).values(trial=1)
)
async def update_trial(session: AsyncSession, legacy_user_ref: int, status: int):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
uid = u.id
await session.execute(update(User).where(User.id == uid).values(trial=status))
await invalidate_profile_cache(uid)
invalidate_user_snapshot(uid)
if u.tg_id is not None:
await invalidate_profile_cache(u.tg_id)
invalidate_user_snapshot(u.tg_id)
logger.info(f"[DB] Триал статус обновлён для пользователя id={uid}: {status}")
async def get_trial(session: AsyncSession, legacy_user_ref: int) -> int:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return 0
result = await session.execute(select(func.coalesce(User.trial, 0)).where(User.id == u.id))
trial = result.scalar_one_or_none()
return int(trial or 0)
async def get_balance_and_trial(session: AsyncSession, tg_id: int) -> tuple[float, int]:
async def get_balance_and_trial(session: AsyncSession, legacy_user_ref: int) -> tuple[float, int]:
"""Один запрос к БД для баланса и триала (профиль при промахе кэша)."""
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return 0.0, 0
result = await session.execute(
select(
func.coalesce(User.balance, 0.0),
func.coalesce(User.trial, 0),
).where(User.tg_id == tg_id)
).where(User.id == u.id)
)
row = result.one_or_none()
if row is None:
@@ -198,18 +211,21 @@ async def get_balance_and_trial(session: AsyncSession, tg_id: int) -> tuple[floa
return round(float(balance or 0.0), 1), int(trial or 0)
async def get_balance_trial_key_count(session: AsyncSession, tg_id: int) -> tuple[float, int, int]:
async def get_balance_trial_key_count(session: AsyncSession, legacy_user_ref: int) -> tuple[float, int, int]:
"""
Один запрос: баланс, триал и число ключей пользователя (для профиля при промахе кэша).
Возвращает (balance_rub, trial_status, key_count).
"""
key_count_subq = select(func.count()).select_from(Key).where(Key.tg_id == User.tg_id).scalar_subquery()
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return 0.0, 0, 0
key_count_subq = select(func.count()).select_from(Key).where(Key.user_id == User.id).scalar_subquery()
result = await session.execute(
select(
func.coalesce(User.balance, 0.0),
func.coalesce(User.trial, 0),
key_count_subq,
).where(User.tg_id == tg_id)
).where(User.id == u.id)
)
row = result.one_or_none()
if row is None:
@@ -233,117 +249,131 @@ async def upsert_user(
only_if_exists: bool = False,
) -> dict | None:
"""Создаёт пользователя или обновляет поля профиля."""
try:
now = datetime.utcnow()
returning_cols = list(User.__table__.c)
now = datetime.utcnow()
returning_cols = list(User.__table__.c)
if only_if_exists:
username_value = username if username else User.username
first_name_value = first_name if first_name else User.first_name
last_name_value = last_name if last_name else User.last_name
language_code_value = language_code if language_code else User.language_code
res = await session.execute(
update(User)
.where(User.tg_id == tg_id)
.values(
username=username_value,
first_name=first_name_value,
last_name=last_name_value,
language_code=language_code_value,
is_bot=is_bot,
updated_at=now,
)
.returning(*returning_cols)
)
row = res.mappings().one_or_none()
if row is None:
return None
await session.commit()
await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC)
return dict(row)
if only_if_exists:
username_value = username if username else User.username
first_name_value = first_name if first_name else User.first_name
last_name_value = last_name if last_name else User.last_name
language_code_value = language_code if language_code else User.language_code
res = await session.execute(
insert(User)
update(User)
.where(User.tg_id == tg_id)
.values(
tg_id=tg_id,
username=username,
first_name=first_name,
last_name=last_name,
language_code=language_code,
username=username_value,
first_name=first_name_value,
last_name=last_name_value,
language_code=language_code_value,
is_bot=is_bot,
created_at=now,
updated_at=now,
)
.on_conflict_do_update(
index_elements=[User.tg_id],
set_={
"username": username,
"first_name": first_name,
"last_name": last_name,
"language_code": language_code,
"is_bot": is_bot,
"updated_at": now,
},
)
.returning(*returning_cols)
)
row = res.mappings().one()
await session.commit()
row = res.mappings().one_or_none()
if row is None:
return None
await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC)
return dict(row)
except SQLAlchemyError as e:
logger.error(f"[DB] Ошибка при UPSERT пользователя {tg_id}: {e}")
await session.rollback()
raise
async def delete_user_data(session: AsyncSession, tg_id: int):
try:
from database.keys import delete_key
await session.execute(delete(Notification).where(Notification.tg_id == tg_id))
await session.execute(
delete(GiftUsage).where(GiftUsage.gift_id.in_(select(Gift.gift_id).where(Gift.sender_tg_id == tg_id)))
res = await session.execute(
insert(User)
.values(
tg_id=tg_id,
username=username,
first_name=first_name,
last_name=last_name,
language_code=language_code,
is_bot=is_bot,
created_at=now,
updated_at=now,
)
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(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))
.on_conflict_do_update(
index_elements=[User.tg_id],
set_={
"username": username,
"first_name": first_name,
"last_name": last_name,
"language_code": language_code,
"is_bot": is_bot,
"updated_at": now,
},
)
await session.execute(delete(CouponUsage).where(CouponUsage.user_id == tg_id))
await delete_key(session, tg_id, commit=False)
await session.execute(delete(TemporaryData).where(TemporaryData.tg_id == tg_id))
await session.execute(delete(BlockedUser).where(BlockedUser.tg_id == tg_id))
await session.execute(delete(User).where(User.tg_id == tg_id))
await session.commit()
logger.info(f"[DB] Данные пользователя {tg_id} полностью удалены")
except SQLAlchemyError as e:
await session.rollback()
logger.error(f"[DB] Ошибка при удалении данных пользователя {tg_id}: {e}")
raise
.returning(*returning_cols)
)
row = res.mappings().one()
await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC)
return dict(row)
async def mark_trial_extended(tg_id: int, session: AsyncSession):
await session.execute(update(User).where(User.tg_id == tg_id).values(trial=-1))
await session.commit()
invalidate_user_snapshot(tg_id)
async def delete_user_data(session: AsyncSession, legacy_user_ref: int):
from database.keys import delete_key
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
uid = u.id
await session.execute(delete(Notification).where(Notification.user_id == uid))
await session.execute(
delete(GiftUsage).where(GiftUsage.gift_id.in_(select(Gift.gift_id).where(Gift.sender_user_id == uid)))
)
await session.execute(delete(Gift).where(Gift.sender_user_id == uid))
await session.execute(update(Gift).where(Gift.recipient_user_id == uid).values(recipient_user_id=None))
await session.execute(delete(Payment).where(Payment.user_id == uid))
await session.execute(
delete(Referral).where(or_(Referral.referrer_user_id == uid, Referral.referred_user_id == uid))
)
await session.execute(delete(CouponUsage).where(CouponUsage.user_id == uid))
await delete_key(session, uid)
await session.execute(
delete(TemporaryData).where(
or_(
TemporaryData.user_id == uid,
TemporaryData.tg_id == u.tg_id,
)
)
)
await session.execute(
delete(BlockedUser).where(
or_(
BlockedUser.user_id == uid,
BlockedUser.tg_id == u.tg_id,
)
)
)
await session.execute(delete(User).where(User.id == uid))
logger.info(f"[DB] Данные пользователя id={uid} полностью удалены")
async def get_user_snapshot(session: AsyncSession, tg_id: int) -> tuple[int, int] | None:
cached = await cache_get(cache_key("user_snapshot", tg_id))
async def mark_trial_extended(legacy_user_ref: int, session: AsyncSession):
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return
await session.execute(update(User).where(User.id == u.id).values(trial=-1))
invalidate_user_snapshot(u.id)
if u.tg_id is not None:
invalidate_user_snapshot(u.tg_id)
async def get_user_snapshot(session: AsyncSession, legacy_user_ref: int) -> tuple[int, int] | None:
u = await resolve_user_optional(session, legacy_user_ref)
if u is None:
return None
uid = u.id
cached = await cache_get(cache_key("user_snapshot", uid))
if isinstance(cached, list) and len(cached) == 2:
return (int(cached[0]), int(cached[1]))
if isinstance(cached, tuple) and len(cached) == 2:
return (int(cached[0]), int(cached[1]))
keys_count_sq = select(func.count(Key.client_id)).where(Key.tg_id == tg_id).scalar_subquery()
res = await session.execute(select(func.coalesce(User.trial, 0), keys_count_sq).where(User.tg_id == tg_id))
keys_count_sq = select(func.count(Key.client_id)).where(Key.user_id == uid).scalar_subquery()
res = await session.execute(select(func.coalesce(User.trial, 0), keys_count_sq).where(User.id == uid))
row = res.first()
if row is None:
return None
value = (int(row[0]), int(row[1]))
await cache_set(cache_key("user_snapshot", tg_id), [value[0], value[1]], USER_SNAPSHOT_CACHE_TTL_SEC)
await cache_set(cache_key("user_snapshot", uid), [value[0], value[1]], USER_SNAPSHOT_CACHE_TTL_SEC)
return value
@@ -351,7 +381,6 @@ async def upsert_source_if_empty(
session: AsyncSession,
tg_id: int,
source_code: str,
commit: bool = True,
) -> bool:
if not source_code:
return False
@@ -367,8 +396,4 @@ async def upsert_source_if_empty(
)
res = await session.execute(stmt)
changed_tg_id = res.scalar_one_or_none()
if changed_tg_id is None:
return False
if commit:
await session.commit()
return True
return changed_tg_id is not None
+234
View File
@@ -0,0 +1,234 @@
from datetime import UTC, datetime
from sqlalchemy import delete, func, select, update
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.ext.asyncio import AsyncSession
from database.models import User, WebNotification, WebPushSubscription
from logger import logger
async def upsert_push_subscription(
session: AsyncSession,
*,
user_id: int,
identity_id: str | None,
endpoint: str,
keys_json: dict,
) -> WebPushSubscription:
"""Upsert push subscription by endpoint (unique)."""
stmt = pg_insert(WebPushSubscription).values(
user_id=user_id,
identity_id=identity_id,
endpoint=endpoint,
keys_json=keys_json,
created_at=datetime.now(UTC),
).on_conflict_do_update(
index_elements=["endpoint"],
set_={
"user_id": user_id,
"identity_id": identity_id,
"keys_json": keys_json,
"created_at": datetime.now(UTC),
},
).returning(WebPushSubscription)
result = await session.execute(stmt)
return result.scalar_one()
async def get_push_subscriptions_by_user(
session: AsyncSession, user_id: int,
) -> list[WebPushSubscription]:
result = await session.execute(
select(WebPushSubscription).where(WebPushSubscription.user_id == user_id)
)
return list(result.scalars().all())
async def get_push_subscriptions_by_identity(
session: AsyncSession, identity_id: str,
) -> list[WebPushSubscription]:
result = await session.execute(
select(WebPushSubscription).where(WebPushSubscription.identity_id == identity_id)
)
return list(result.scalars().all())
async def delete_push_subscription_by_endpoint(
session: AsyncSession, endpoint: str,
) -> None:
await session.execute(
delete(WebPushSubscription).where(WebPushSubscription.endpoint == endpoint)
)
async def get_notifications_for_identity(
session: AsyncSession,
identity_id: str,
limit: int = 20,
offset: int = 0,
) -> list[WebNotification]:
result = await session.execute(
select(WebNotification)
.where(WebNotification.identity_id == identity_id)
.order_by(WebNotification.created_at.desc())
.limit(limit)
.offset(offset)
)
return list(result.scalars().all())
async def count_unread_for_identity(
session: AsyncSession, identity_id: str,
) -> int:
result = await session.execute(
select(func.count())
.select_from(WebNotification)
.where(
WebNotification.identity_id == identity_id,
WebNotification.read is False,
)
)
return result.scalar() or 0
async def mark_all_read_for_identity(
session: AsyncSession, identity_id: str,
) -> int:
result = await session.execute(
update(WebNotification)
.where(
WebNotification.identity_id == identity_id,
WebNotification.read is False,
)
.values(read=True)
)
return result.rowcount
async def resolve_identity_id_by_tg_id(
session: AsyncSession, tg_id: int,
) -> str | None:
"""Resolve identity_id from user's tg_id."""
result = await session.execute(
select(User.identity_id).where(User.tg_id == tg_id)
)
return result.scalar_one_or_none()
async def create_notification(
session: AsyncSession,
*,
user_id: int,
identity_id: str | None,
type: str = "system",
title: str,
message: str = "",
data: dict | None = None,
) -> WebNotification:
notif = WebNotification(
user_id=user_id,
identity_id=identity_id,
type=type,
title=title,
message=message,
data=data,
)
session.add(notif)
await session.flush()
return notif
def _render_template(template: str, **kwargs: object) -> str:
"""Safe format — unknown placeholders stay as-is."""
try:
return template.format_map(
{k: str(v) for k, v in kwargs.items() if v is not None}
| type("_Defaults", (), {"__missing__": lambda self, k: f"{{{k}}}"})()
)
except Exception:
return template
def _get_web_config_str(key: str, default: str) -> str:
try:
from core.settings.web_config import WEB_CONFIG
val = WEB_CONFIG.get(key)
return str(val).strip() if val else default
except Exception:
return default
async def notify_web(
session: AsyncSession,
*,
tg_id: int,
type: str = "system",
title: str | None = None,
message: str | None = None,
data: dict | None = None,
template_vars: dict | None = None,
) -> WebNotification | None:
"""Создаёт web-уведомление по tg_id.
title/message — если None, берутся из WEB_CONFIG шаблонов по type.
template_vars — подстановки в шаблон ({email}, {amount}, {name}, {duration}).
"""
try:
identity_id = await resolve_identity_id_by_tg_id(session, tg_id)
if not identity_id:
return None
vars_ = template_vars or {}
type_key_map = {
"payment": ("WEB_NOTIFY_PAYMENT_TITLE", "WEB_NOTIFY_PAYMENT_MESSAGE"),
"key_created": ("WEB_NOTIFY_KEY_CREATED_TITLE", "WEB_NOTIFY_KEY_CREATED_MESSAGE"),
"key_expiry": ("WEB_NOTIFY_KEY_EXPIRY_TITLE", "WEB_NOTIFY_KEY_EXPIRY_MESSAGE"),
"gift_received": ("WEB_NOTIFY_GIFT_TITLE", "WEB_NOTIFY_GIFT_MESSAGE"),
}
title_key, msg_key = type_key_map.get(type, (None, None))
resolved_title = title
if resolved_title is None and title_key:
resolved_title = _render_template(_get_web_config_str(title_key, ""), **vars_)
resolved_title = resolved_title or type
resolved_message = message
if resolved_message is None and msg_key:
resolved_message = _render_template(_get_web_config_str(msg_key, ""), **vars_)
resolved_message = resolved_message or ""
notif = await create_notification(
session,
user_id=tg_id,
identity_id=identity_id,
type=type,
title=resolved_title,
message=resolved_message,
data=data,
)
try:
from services.web_push import push_enabled, send_push_to_many
if push_enabled():
subs = await get_push_subscriptions_by_identity(session, identity_id)
if subs:
sub_infos = [
{"endpoint": s.endpoint, "keys": s.keys_json}
for s in subs
]
sent = await send_push_to_many(
sub_infos,
title=resolved_title,
body=resolved_message,
url="/dashboard/notifications",
)
logger.debug("[notify_web] push sent to {}/{} subscriptions", sent, len(sub_infos))
except Exception as push_err:
logger.warning("[notify_web] push delivery failed: {}", push_err)
return notif
except Exception:
return None