
"""
database/database.py
--------------------
اتصال و مدیریت دیتابیس SQLite
"""

from __future__ import annotations

import logging

from contextlib import asynccontextmanager
from typing import AsyncIterator

from sqlalchemy import text
from sqlalchemy.ext.asyncio import (
    AsyncSession,
    async_sessionmaker,
    create_async_engine,
)

import config

from database.models import Base


logger = logging.getLogger(__name__)


# ============================================================
# ENGINE
# ============================================================

engine = create_async_engine(
    config.DATABASE_URL,
    echo=False,
    future=True,
)


# ============================================================
# SESSION
# ============================================================

async_session_maker = async_sessionmaker(
    bind=engine,
    class_=AsyncSession,
    expire_on_commit=False,
)

AsyncSessionLocal = async_session_maker


# ============================================================
# TABLE COLUMNS
# ============================================================

async def get_table_columns(
    conn,
    table_name: str,
) -> set[str]:

    result = await conn.execute(
        text(
            f"PRAGMA table_info({table_name})"
        )
    )

    return {
        row[1]
        for row in result.fetchall()
    }


# ============================================================
# ADD COLUMN IF MISSING
# ============================================================

async def add_column_if_missing(
    conn,
    table_name: str,
    column_name: str,
    column_definition: str,
) -> None:

    columns = await get_table_columns(
        conn,
        table_name,
    )

    if column_name in columns:
        return

    await conn.execute(
        text(
            f"""
            ALTER TABLE {table_name}
            ADD COLUMN {column_name} {column_definition}
            """
        )
    )

    logger.info(
        "ستون %s.%s اضافه شد.",
        table_name,
        column_name,
    )


# ============================================================
# MIGRATION
# ============================================================

async def migrate_database() -> None:

    async with engine.begin() as conn:

        # ====================================================
        # USERS
        # ====================================================

        await add_column_if_missing(
            conn,
            "users",
            "last_name",
            "VARCHAR(255)",
        )

        await add_column_if_missing(
            conn,
            "users",
            "referred_by_id",
            "INTEGER",
        )

        await add_column_if_missing(
            conn,
            "users",
            "test_account_used_count",
            "INTEGER NOT NULL DEFAULT 0",
        )

        await add_column_if_missing(
            conn,
            "users",
            "updated_at",
            "DATETIME",
        )

        await add_column_if_missing(
            conn,
            "users",
            "referral_balance",
            "INTEGER NOT NULL DEFAULT 0",
        )        # ====================================================
        # PLANS
        # ====================================================

        await add_column_if_missing(
            conn,
            "plans",
            "service_type",
            "VARCHAR(100)",
        )

        await add_column_if_missing(
            conn,
            "plans",
            "region",
            "VARCHAR(100)",
        )

        await add_column_if_missing(
            conn,
            "plans",
            "description",
            "TEXT",
        )

        await add_column_if_missing(
            conn,
            "plans",
            "is_active",
            "BOOLEAN NOT NULL DEFAULT 1",
        )

        # ====================================================
        # ORDERS
        # ====================================================

        await add_column_if_missing(
            conn,
            "orders",
            "subscription_text",
            "TEXT",
        )

        await add_column_if_missing(
            conn,
            "orders",
            "config_text",
            "TEXT",
        )

        await add_column_if_missing(
            conn,
            "orders",
            "receipt_file_id",
            "TEXT",
        )

        await add_column_if_missing(
            conn,
            "orders",
            "admin_note",
            "TEXT",
        )

        await add_column_if_missing(
            conn,
            "orders",
            "order_type",
            "VARCHAR(50) DEFAULT 'new'",
        )

        await add_column_if_missing(
            conn,
            "orders",
            "subscription_id",
            "INTEGER",
        )

        await add_column_if_missing(
            conn,
            "orders",
            "service_name",
            "VARCHAR(255)",
        )

        # ====================================================
        # SUBSCRIPTIONS
        # ====================================================

        await add_column_if_missing(
            conn,
            "subscriptions",
            "service_name",
            "VARCHAR(255)",
        )

        await add_column_if_missing(
            conn,
            "subscriptions",
            "subscription_text",
            "TEXT",
        )

        # ====================================================
        # TUTORIALS
        # ====================================================

        await add_column_if_missing(
            conn,
            "tutorials",
            "order_index",
            "INTEGER NOT NULL DEFAULT 0",
        )

        await add_column_if_missing(
            conn,
            "tutorials",
            "is_active",
            "BOOLEAN NOT NULL DEFAULT 1",
        )


# ============================================================
# INIT DATABASE
# ============================================================

async def init_db() -> None:

    async with engine.begin() as conn:

        await conn.run_sync(
            Base.metadata.create_all
        )

    await migrate_database()

    logger.info(
        "دیتابیس با موفقیت مقداردهی اولیه شد."
    )


# ============================================================
# SESSION CONTEXT
# ============================================================

@asynccontextmanager
async def get_session() -> AsyncIterator[AsyncSession]:

    async with async_session_maker() as session:

        try:

            yield session

        except Exception:

            await session.rollback()

            raise


# ============================================================
# CLOSE DATABASE
# ============================================================

async def close_db() -> None:

    await engine.dispose()

    logger.info(
        "اتصال دیتابیس بسته شد."
    )
