Files
Solo_bot/database/users.py
T

355 lines
13 KiB
Python

from datetime import datetime
from sqlalchemy import delete, exists, 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 (
BALANCE_CACHE_TTL_SEC,
USER_EXISTS_CACHE_TTL_SEC,
USER_SNAPSHOT_CACHE_TTL_SEC,
)
from core.redis_cache import cache_delete, cache_get, cache_key, cache_set
from database.models import (
BlockedUser,
CouponUsage,
Gift,
GiftUsage,
Key,
Notification,
Payment,
Referral,
TemporaryData,
User,
)
from logger import logger
def invalidate_user_snapshot(tg_id: int) -> None:
import asyncio
try:
loop = asyncio.get_running_loop()
loop.create_task(cache_delete(cache_key("user_snapshot", tg_id)))
except RuntimeError:
return
async def add_user(
session: AsyncSession,
tg_id: int,
username: str = None,
first_name: str = None,
last_name: str = None,
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)
)
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
async def invalidate_balance_cache(tg_id: int) -> None:
await cache_delete(cache_key("balance", tg_id))
async def invalidate_profile_cache(tg_id: int) -> None:
await cache_delete(cache_key("profile_data", tg_id))
async def update_balance(session: AsyncSession, tg_id: int, amount: float) -> None:
try:
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()
await session.commit()
if new_balance is not None:
old_balance = new_balance - amount
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()
async def check_user_exists(session: AsyncSession, tg_id: int) -> bool:
cached = await cache_get(cache_key("user_exists", tg_id))
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)
return value
async def get_balance(session: AsyncSession, tg_id: int) -> float:
cached = await cache_get(cache_key("balance", tg_id))
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))
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)
return value
async def set_user_balance(session: AsyncSession, tg_id: int, balance: float) -> None:
try:
await session.execute(update(User).where(User.tg_id == tg_id).values(balance=balance))
await session.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
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_trial(session: AsyncSession, tg_id: int) -> int:
result = await session.execute(select(func.coalesce(User.trial, 0)).where(User.tg_id == tg_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]:
"""Один запрос к БД для баланса и триала (профиль при промахе кэша)."""
result = await session.execute(
select(
func.coalesce(User.balance, 0.0),
func.coalesce(User.trial, 0),
).where(User.tg_id == tg_id)
)
row = result.one_or_none()
if row is None:
return 0.0, 0
balance, trial = row
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]:
"""
Один запрос: баланс, триал и число ключей пользователя (для профиля при промахе кэша).
Возвращает (balance_rub, trial_status, key_count).
"""
key_count_subq = select(func.count()).select_from(Key).where(Key.tg_id == User.tg_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)
)
row = result.one_or_none()
if row is None:
return 0.0, 0, 0
balance, trial, key_count = row
return (
round(float(balance or 0.0), 1),
int(trial or 0),
int(key_count or 0),
)
async def upsert_user(
session: AsyncSession,
tg_id: int,
username: str = None,
first_name: str = None,
last_name: str = None,
language_code: str = None,
is_bot: bool = False,
only_if_exists: bool = False,
) -> dict | None:
"""Создаёт пользователя или обновляет поля профиля."""
try:
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)
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,
)
.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()
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)))
)
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))
)
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
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 get_user_snapshot(session: AsyncSession, tg_id: int) -> tuple[int, int] | None:
cached = await cache_get(cache_key("user_snapshot", tg_id))
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))
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)
return value
async def upsert_source_if_empty(
session: AsyncSession,
tg_id: int,
source_code: str,
commit: bool = True,
) -> bool:
if not source_code:
return False
stmt = (
insert(User)
.values(tg_id=tg_id, source_code=source_code)
.on_conflict_do_update(
index_elements=["tg_id"],
set_={"source_code": insert(User).excluded.source_code},
where=(User.source_code.is_(None)),
)
.returning(User.tg_id)
)
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