WEB-APP/ Optimization/ Build fix/ Hotkey edit mode/ Log rotation/ Form a11y/ E2E non-blocking
This commit is contained in:
@@ -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 *
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
from .resolution import *
|
||||
from .tg_mirror import *
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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
@@ -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
@@ -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]:
|
||||
|
||||
@@ -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
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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} сброшены к выбранным")
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
from .schema_upgrade import *
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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}
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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"),)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
@@ -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
@@ -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()]
|
||||
|
||||
@@ -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
@@ -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))
|
||||
|
||||
@@ -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))
|
||||
@@ -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
@@ -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
@@ -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}")
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user