"""
database/crud.py
----------------
توابع CRUD دیتابیس ربات Vasl Bot
"""

from __future__ import annotations

import random
from datetime import datetime, timedelta

from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload

from database.models import (
    Order,
    OrderStatus,
    Plan,
    Referral,
    Subscription,
    SubscriptionStatus,
    SupportMessage,
    SupportTicket,
    SupportTicketStatus,
    SenderRole,
    TestConfig,
    TestRequest,
    TestRequestStatus,
    Tutorial,
    User,
)


# ============================================================
# USERS
# ============================================================

async def get_user_by_telegram_id(
    session: AsyncSession,
    telegram_id: int,
) -> User | None:

    result = await session.execute(
        select(User).where(
            User.telegram_id == telegram_id
        )
    )

    return result.scalar_one_or_none()


async def get_user_by_id(
    session: AsyncSession,
    user_id: int,
) -> User | None:

    result = await session.execute(
        select(User).where(
            User.id == user_id
        )
    )

    return result.scalar_one_or_none()


async def get_or_create_user(
    session: AsyncSession,
    telegram_id: int,
    username: str | None = None,
    first_name: str | None = None,
    last_name: str | None = None,
    referred_by_telegram_id: int | None = None,
) -> tuple[User, bool]:

    user = await get_user_by_telegram_id(
        session,
        telegram_id,
    )

    if user:

        user.username = username
        user.first_name = first_name
        user.last_name = last_name
        user.updated_at = datetime.utcnow()

        await session.commit()
        await session.refresh(user)

        return user, False

    referred_by_id = None

    if (
        referred_by_telegram_id
        and referred_by_telegram_id != telegram_id
    ):

        referrer = await get_user_by_telegram_id(
            session,
            referred_by_telegram_id,
        )

        if referrer:
            referred_by_id = referrer.id

    user = User(
        telegram_id=telegram_id,
        username=username,
        first_name=first_name,
        last_name=last_name,
        referred_by_id=referred_by_id,
        referral_balance=0,
        test_account_used_count=0,
        is_blocked=False,
    )

    session.add(user)

    await session.flush()

    if referred_by_id:

        referral = Referral(
            referrer_id=referred_by_id,
            referred_user_id=user.id,
            reward_amount=0,
            reward_paid=False,
        )

        session.add(referral)

    await session.commit()
    await session.refresh(user)

    return user, True


async def increment_test_usage(
    session: AsyncSession,
    user_id: int,
) -> User | None:

    user = await get_user_by_id(
        session,
        user_id,
    )

    if not user:
        return None

    user.test_account_used_count += 1

    await session.commit()
    await session.refresh(user)

    return user


# ============================================================
# REFERRALS
# ============================================================

REFERRAL_DISCOUNT_PER_SUCCESSFUL_REFERRAL = 0.5
REFERRAL_MAX_DISCOUNT = 5.0


async def get_successful_referral_count(
    session: AsyncSession,
    user_id: int,
) -> int:

    result = await session.execute(
        select(
            func.count(
                func.distinct(
                    Referral.referred_user_id
                )
            )
        )
        .where(
            Referral.referrer_id == user_id
        )
    )

    return int(
        result.scalar() or 0
    )


async def get_referral_discount_percent(
    session: AsyncSession,
    user_id: int,
) -> float:

    successful_count = (
        await get_successful_referral_count(
            session,
            user_id,
        )
    )

    discount = (
        successful_count
        * REFERRAL_DISCOUNT_PER_SUCCESSFUL_REFERRAL
    )

    return min(
        discount,
        REFERRAL_MAX_DISCOUNT,
    )


async def get_referral_stats(
    session: AsyncSession,
    user_id: int,
) -> tuple[int, int]:

    total_result = await session.execute(
        select(
            func.count(Referral.id)
        ).where(
            Referral.referrer_id == user_id
        )
    )

    total_referrals = int(
        total_result.scalar() or 0
    )

    successful_referrals = (
        await get_successful_referral_count(
            session,
            user_id,
        )
    )

    return (
        total_referrals,
        successful_referrals,
    )


# ============================================================
# PLANS
# ============================================================

async def get_plan_by_id(
    session: AsyncSession,
    plan_id: int,
) -> Plan | None:

    result = await session.execute(
        select(Plan).where(
            Plan.id == plan_id
        )
    )

    return result.scalar_one_or_none()


async def get_plan_by_name(
    session: AsyncSession,
    name: str,
) -> Plan | None:

    result = await session.execute(
        select(Plan).where(
            Plan.name == name
        )
    )

    return result.scalar_one_or_none()


async def list_active_plans(
    session: AsyncSession,
) -> list[Plan]:

    result = await session.execute(
        select(Plan)
        .where(
            Plan.is_active.is_(True)
        )
        .order_by(
            Plan.price.asc()
        )
    )

    return list(
        result.scalars().all()
    )


async def list_active_plans_for_service(
    session: AsyncSession,
    service_type: str | None = None,
    region: str | None = None,
) -> list[Plan]:

    query = (
        select(Plan)
        .where(
            Plan.is_active.is_(True)
        )
    )

    if service_type is not None:
        query = query.where(
            Plan.service_type == service_type
        )

    if region is not None:
        query = query.where(
            Plan.region == region
        )

    query = query.order_by(
        Plan.price.asc()
    )

    result = await session.execute(
        query
    )

    return list(
        result.scalars().all()
    )


async def list_all_plans(
    session: AsyncSession,
) -> list[Plan]:

    result = await session.execute(
        select(Plan)
        .order_by(
            Plan.id.asc()
        )
    )

    return list(
        result.scalars().all()
    )


async def create_plan(
    session: AsyncSession,
    name: str,
    service_type: str | None,
    region: str | None,
    volume_gb: int,
    duration_days: int,
    price: int,
    description: str | None = None,
    is_active: bool = True,
) -> Plan:

    plan = Plan(
        name=name,
        service_type=service_type,
        region=region,
        volume_gb=volume_gb,
        duration_days=duration_days,
        price=price,
        description=description,
        is_active=is_active,
    )

    session.add(plan)

    await session.commit()
    await session.refresh(plan)

    return plan


async def update_plan(
    session: AsyncSession,
    plan_id: int,
    name: str | None = None,
    service_type: str | None = None,
    region: str | None = None,
    volume_gb: int | None = None,
    duration_days: int | None = None,
    price: int | None = None,
    description: str | None = None,
    is_active: bool | None = None,
) -> Plan | None:

    plan = await get_plan_by_id(
        session,
        plan_id,
    )

    if not plan:
        return None

    if name is not None:
        plan.name = name

    if service_type is not None:
        plan.service_type = service_type

    if region is not None:
        plan.region = region

    if volume_gb is not None:
        plan.volume_gb = volume_gb

    if duration_days is not None:
        plan.duration_days = duration_days

    if price is not None:
        plan.price = price

    if description is not None:
        plan.description = description

    if is_active is not None:
        plan.is_active = is_active

    await session.commit()
    await session.refresh(plan)

    return plan


# ============================================================
# ORDERS
# ============================================================

async def generate_unique_receipt_number(
    session: AsyncSession,
) -> int:
    """
    تولید شماره رسید ۵ رقمی و تصادفی
    از 10000 تا 99999.

    شماره رسید از ID داخلی سفارش کاملاً جداست.
    """

    for _ in range(1000):

        receipt_number = random.randint(
            10000,
            99999,
        )

        result = await session.execute(
            select(Order.id).where(
                Order.receipt_number == receipt_number
            )
        )

        exists = result.scalar_one_or_none()

        if exists is None:
            return receipt_number

    raise RuntimeError(
        "امکان تولید شماره رسید یکتا وجود ندارد."
    )


async def create_order(
    session: AsyncSession,
    user_id: int,
    plan_id: int,
    amount: int,
    order_type: str = "new",
    subscription_id: int | None = None,
    service_name: str | None = None,
) -> Order:

    receipt_number = (
        await generate_unique_receipt_number(
            session
        )
    )

    order = Order(
        user_id=user_id,
        plan_id=plan_id,
        amount=amount,
        order_type=order_type,
        subscription_id=subscription_id,
        service_name=service_name,
        receipt_number=receipt_number,
        status=OrderStatus.WAITING_RECEIPT.value,
    )

    session.add(order)

    await session.commit()
    await session.refresh(order)

    return order


async def get_order(
    session: AsyncSession,
    order_id: int,
) -> Order | None:

    result = await session.execute(
        select(Order)
        .options(
            selectinload(Order.user),
            selectinload(Order.plan),
            selectinload(Order.subscription),
        )
        .where(
            Order.id == order_id
        )
    )

    return result.scalar_one_or_none()


async def get_order_by_receipt_number(
    session: AsyncSession,
    receipt_number: int,
) -> Order | None:

    result = await session.execute(
        select(Order)
        .options(
            selectinload(Order.user),
            selectinload(Order.plan),
            selectinload(Order.subscription),
        )
        .where(
            Order.receipt_number == receipt_number
        )
    )

    return result.scalar_one_or_none()


async def set_order_receipt(
    session: AsyncSession,
    order_id: int,
    receipt_file_id: str,
) -> Order | None:

    order = await get_order(
        session,
        order_id,
    )

    if not order:
        return None

    order.receipt_file_id = receipt_file_id
    order.status = OrderStatus.WAITING_ADMIN.value

    await session.commit()
    await session.refresh(order)

    return order


async def update_order_status(
    session: AsyncSession,
    order_id: int,
    status: str,
    admin_note: str | None = None,
) -> Order | None:

    order = await get_order(
        session,
        order_id,
    )

    if not order:
        return None

    order.status = status

    if admin_note is not None:
        order.admin_note = admin_note

    await session.commit()
    await session.refresh(order)

    return order


async def set_order_config(
    session: AsyncSession,
    order_id: int,
    config_text: str,
) -> Order | None:

    order = await get_order(
        session,
        order_id,
    )

    if not order:
        return None

    order.config_text = config_text
    order.status = OrderStatus.WAITING_SUBSCRIPTION.value

    await session.commit()
    await session.refresh(order)

    return order


async def set_order_subscription(
    session: AsyncSession,
    order_id: int,
    subscription_text: str,
) -> Order | None:

    order = await get_order(
        session,
        order_id,
    )

    if not order:
        return None

    order.subscription_text = subscription_text
    order.status = OrderStatus.COMPLETED.value

    await session.commit()
    await session.refresh(order)

    return order


async def set_order_subscription_id(
    session: AsyncSession,
    order_id: int,
    subscription_id: int,
) -> Order | None:

    order = await get_order(
        session,
        order_id,
    )

    if not order:
        return None

    order.subscription_id = subscription_id

    await session.commit()
    await session.refresh(order)

    return order


async def list_orders(
    session: AsyncSession,
) -> list[Order]:

    result = await session.execute(
        select(Order)
        .options(
            selectinload(Order.user),
            selectinload(Order.plan),
            selectinload(Order.subscription),
        )
        .order_by(
            Order.id.desc()
        )
    )

    return list(
        result.scalars().all()
    )


# ============================================================
# SUBSCRIPTIONS
# ============================================================

async def create_subscription(
    session: AsyncSession,
    user_id: int,
    plan_id: int,
    duration_days: int,
    config_text: str | None = None,
    service_name: str | None = None,
    subscription_text: str | None = None,
) -> Subscription:

    now = datetime.utcnow()

    expire_date = (
        now
        + timedelta(days=duration_days)
    )

    subscription = Subscription(
        user_id=user_id,
        plan_id=plan_id,
        service_name=service_name,
        subscription_text=subscription_text,
        config_text=config_text,
        start_date=now,
        expire_date=expire_date,
        status=SubscriptionStatus.ACTIVE.value,
    )

    session.add(subscription)

    await session.commit()
    await session.refresh(subscription)

    return subscription


async def get_subscription(
    session: AsyncSession,
    subscription_id: int,
) -> Subscription | None:

    result = await session.execute(
        select(Subscription)
        .options(
            selectinload(Subscription.user),
            selectinload(Subscription.plan),
        )
        .where(
            Subscription.id == subscription_id
        )
    )

    return result.scalar_one_or_none()


async def list_user_subscriptions(
    session: AsyncSession,
    user_id: int,
) -> list[Subscription]:

    result = await session.execute(
        select(Subscription)
        .options(
            selectinload(Subscription.plan),
        )
        .where(
            Subscription.user_id == user_id
        )
        .order_by(
            Subscription.id.desc()
        )
    )

    return list(
        result.scalars().all()
    )


async def extend_subscription(
    session: AsyncSession,
    subscription_id: int,
    duration_days: int,
    new_config_text: str | None = None,
    new_subscription_text: str | None = None,
) -> Subscription | None:

    subscription = await get_subscription(
        session,
        subscription_id,
    )

    if not subscription:
        return None

    now = datetime.utcnow()

    if (
        subscription.expire_date
        and subscription.expire_date > now
    ):
        base_date = subscription.expire_date
    else:
        base_date = now

    subscription.expire_date = (
        base_date
        + timedelta(days=duration_days)
    )

    subscription.status = (
        SubscriptionStatus.ACTIVE.value
    )

    if new_config_text is not None:
        subscription.config_text = new_config_text

    if new_subscription_text is not None:
        subscription.subscription_text = (
            new_subscription_text
        )

    await session.commit()
    await session.refresh(subscription)

    return subscription


# ============================================================
# TEST REQUESTS
# ============================================================

async def get_test_request(
    session: AsyncSession,
    user_id: int,
) -> TestRequest | None:

    result = await session.execute(
        select(TestRequest).where(
            TestRequest.user_id == user_id
        )
        .order_by(
            TestRequest.id.desc()
        )
    )

    return result.scalars().first()


async def create_test_request(
    session: AsyncSession,
    user_id: int,
) -> TestRequest:

    request = TestRequest(
        user_id=user_id,
        status=TestRequestStatus.PENDING.value,
    )

    session.add(request)

    await session.commit()
    await session.refresh(request)

    return request


async def set_test_request_config(
    session: AsyncSession,
    request_id: int,
    config_text: str,
) -> TestRequest | None:

    result = await session.execute(
        select(TestRequest).where(
            TestRequest.id == request_id
        )
    )

    request = result.scalar_one_or_none()

    if not request:
        return None

    request.config_text = config_text
    request.status = TestRequestStatus.APPROVED.value

    await session.commit()
    await session.refresh(request)

    return request


async def reject_test_request(
    session: AsyncSession,
    request_id: int,
) -> TestRequest | None:

    result = await session.execute(
        select(TestRequest).where(
            TestRequest.id == request_id
        )
    )

    request = result.scalar_one_or_none()

    if not request:
        return None

    request.status = TestRequestStatus.REJECTED.value

    await session.commit()
    await session.refresh(request)

    return request


# ============================================================
# TEST CONFIG POOL
# ============================================================

async def add_test_config(
    session: AsyncSession,
    config_text: str,
) -> TestConfig:

    test_config = TestConfig(
        config_text=config_text,
        used=False,
        used_by_user_id=None,
        used_at=None,
    )

    session.add(test_config)

    await session.commit()
    await session.refresh(test_config)

    return test_config


async def get_available_test_config_count(
    session: AsyncSession,
) -> int:

    result = await session.execute(
        select(func.count(TestConfig.id))
        .where(
            TestConfig.used.is_(False)
        )
    )

    return int(
        result.scalar() or 0
    )


async def get_available_test_config(
    session: AsyncSession,
) -> TestConfig | None:

    result = await session.execute(
        select(TestConfig)
        .where(
            TestConfig.used.is_(False)
        )
        .order_by(
            TestConfig.id.asc()
        )
        .limit(1)
    )

    return result.scalar_one_or_none()


async def use_test_config(
    session: AsyncSession,
    user_id: int,
) -> TestConfig | None:

    user = await get_user_by_id(
        session,
        user_id,
    )

    if not user:
        return None

    if user.test_account_used_count > 0:
        return None

    result = await session.execute(
        select(TestConfig)
        .where(
            TestConfig.used.is_(False)
        )
        .order_by(
            TestConfig.id.asc()
        )
        .limit(1)
    )

    test_config = (
        result.scalar_one_or_none()
    )

    if not test_config:
        return None

    test_config.used = True
    test_config.used_by_user_id = user_id
    test_config.used_at = datetime.utcnow()

    user.test_account_used_count += 1

    await session.commit()

    await session.refresh(
        test_config
    )

    return test_config


async def list_test_configs(
    session: AsyncSession,
) -> list[TestConfig]:

    result = await session.execute(
        select(TestConfig)
        .order_by(
            TestConfig.id.asc()
        )
    )

    return list(
        result.scalars().all()
    )


async def delete_available_test_config(
    session: AsyncSession,
    test_config_id: int,
) -> bool:

    result = await session.execute(
        select(TestConfig)
        .where(
            TestConfig.id == test_config_id,
            TestConfig.used.is_(False),
        )
    )

    test_config = (
        result.scalar_one_or_none()
    )

    if not test_config:
        return False

    await session.delete(
        test_config
    )

    await session.commit()

    return True


async def clear_used_test_configs(
    session: AsyncSession,
) -> int:

    result = await session.execute(
        select(TestConfig)
        .where(
            TestConfig.used.is_(True)
        )
    )

    configs = list(
        result.scalars().all()
    )

    count = len(configs)

    for test_config in configs:
        await session.delete(
            test_config
        )

    await session.commit()

    return count


# ============================================================
# TUTORIALS
# ============================================================

async def list_active_tutorials(
    session: AsyncSession,
) -> list[Tutorial]:

    result = await session.execute(
        select(Tutorial)
        .where(
            Tutorial.is_active.is_(True)
        )
        .order_by(
            Tutorial.order_index.asc(),
            Tutorial.id.asc(),
        )
    )

    return list(
        result.scalars().all()
    )


async def list_all_tutorials(
    session: AsyncSession,
) -> list[Tutorial]:

    result = await session.execute(
        select(Tutorial)
        .order_by(
            Tutorial.order_index.asc(),
            Tutorial.id.asc()
        )
    )

    return list(
        result.scalars().all()
    )


async def get_tutorial(
    session: AsyncSession,
    tutorial_id: int,
) -> Tutorial | None:

    result = await session.execute(
        select(Tutorial).where(
            Tutorial.id == tutorial_id
        )
    )

    return result.scalar_one_or_none()


async def create_tutorial(
    session: AsyncSession,
    title: str,
    content: str,
    order_index: int = 0,
    is_active: bool = True,
) -> Tutorial:

    tutorial = Tutorial(
        title=title,
        content=content,
        order_index=order_index,
        is_active=is_active,
    )

    session.add(tutorial)

    await session.commit()
    await session.refresh(tutorial)

    return tutorial


async def update_tutorial(
    session: AsyncSession,
    tutorial_id: int,
    title: str | None = None,
    content: str | None = None,
    order_index: int | None = None,
    is_active: bool | None = None,
) -> Tutorial | None:

    tutorial = await get_tutorial(
        session,
        tutorial_id,
    )

    if not tutorial:
        return None

    if title is not None:
        tutorial.title = title

    if content is not None:
        tutorial.content = content

    if order_index is not None:
        tutorial.order_index = order_index

    if is_active is not None:
        tutorial.is_active = is_active

    await session.commit()
    await session.refresh(tutorial)

    return tutorial


async def delete_tutorial(
    session: AsyncSession,
    tutorial_id: int,
) -> bool:

    tutorial = await get_tutorial(
        session,
        tutorial_id,
    )

    if not tutorial:
        return False

    await session.delete(tutorial)

    await session.commit()

    return True


# ============================================================
# SUPPORT
# ============================================================

async def get_or_create_open_ticket(
    session: AsyncSession,
    user_id: int,
) -> SupportTicket:

    result = await session.execute(
        select(SupportTicket)
        .where(
            SupportTicket.user_id == user_id,
            SupportTicket.status
            != SupportTicketStatus.CLOSED.value,
        )
        .order_by(
            SupportTicket.id.desc()
        )
    )

    ticket = result.scalars().first()

    if ticket:
        return ticket

    ticket = SupportTicket(
        user_id=user_id,
        status=SupportTicketStatus.OPEN.value,
    )

    session.add(ticket)

    await session.commit()
    await session.refresh(ticket)

    return ticket


async def get_ticket(
    session: AsyncSession,
    ticket_id: int,
) -> SupportTicket | None:

    result = await session.execute(
        select(SupportTicket)
        .options(
            selectinload(SupportTicket.user),
            selectinload(SupportTicket.messages),
        )
        .where(
            SupportTicket.id == ticket_id
        )
    )

    return result.scalar_one_or_none()


async def add_support_message(
    session: AsyncSession,
    ticket_id: int,
    sender_role: str,
    text: str | None = None,
    photo_file_id: str | None = None,
) -> SupportMessage | None:

    ticket = await get_ticket(
        session,
        ticket_id,
    )

    if not ticket:
        return None

    message = SupportMessage(
        ticket_id=ticket_id,
        sender_role=sender_role,
        text=text,
        photo_file_id=photo_file_id,
    )

    session.add(message)

    ticket.status = (
        SupportTicketStatus.OPEN.value
    )

    await session.commit()
    await session.refresh(message)

    return message


async def list_open_tickets(
    session: AsyncSession,
) -> list[SupportTicket]:

    result = await session.execute(
        select(SupportTicket)
        .options(
            selectinload(SupportTicket.user),
        )
        .where(
            SupportTicket.status
            == SupportTicketStatus.OPEN.value
        )
        .order_by(
            SupportTicket.id.asc()
        )
    )

    return list(
        result.scalars().all()
    )