import uuid
from typing import Any

from sqlalchemy import or_, select

from app.models import Admin, Transaction, User, Wallet
from app.repositories.base import BaseRepository


class UserRepository(BaseRepository[User]):
    model = User

    async def get_by_telegram_id(self, telegram_id: int) -> User | None:
        return await self.get_by(telegram_id=telegram_id)

    async def get_by_referral_code(self, code: str) -> User | None:
        return await self.get_by(referral_code=code)

    async def search(self, page: int, size: int, search: str | None = None) -> tuple[list[User], int]:
        stmt = select(User).order_by(User.created_at.desc())
        if search:
            pattern = f"%{search}%"
            conditions = [
                User.username.ilike(pattern),
                User.first_name.ilike(pattern),
                User.last_name.ilike(pattern),
            ]
            if search.strip().isdigit():
                conditions.append(User.telegram_id == int(search.strip()))
            stmt = stmt.where(or_(*conditions))
        return await self.paginate(stmt, page, size)

    async def count_referrals(self, user_id: uuid.UUID) -> int:
        stmt = select(User).where(User.referred_by == user_id)
        return await self.count(stmt)


class WalletRepository(BaseRepository[Wallet]):
    model = Wallet

    async def get_by_user(self, user_id: uuid.UUID) -> Wallet | None:
        return await self.get_by(user_id=user_id)

    async def get_for_update(self, user_id: uuid.UUID) -> Wallet | None:
        """Row-level lock so concurrent balance changes cannot interleave."""
        stmt = select(Wallet).where(Wallet.user_id == user_id).with_for_update()
        return (await self.session.execute(stmt)).scalar_one_or_none()


class TransactionRepository(BaseRepository[Transaction]):
    model = Transaction

    async def history(
        self, wallet_id: uuid.UUID, page: int, size: int, **filters: Any
    ) -> tuple[list[Transaction], int]:
        stmt = (
            select(Transaction)
            .where(Transaction.wallet_id == wallet_id)
            .order_by(Transaction.created_at.desc())
        )
        tx_type = filters.get("type")
        if tx_type:
            stmt = stmt.where(Transaction.type == tx_type)
        return await self.paginate(stmt, page, size)

    async def get_by_reference(self, reference: str) -> Transaction | None:
        return await self.get_by(reference=reference)


class AdminRepository(BaseRepository[Admin]):
    model = Admin

    async def get_by_username(self, username: str) -> Admin | None:
        return await self.get_by(username=username)
