# app/db/message_counter_dal.py
from sqlalchemy import select, update
from .models import MessageCounter, MessageTotals
from .engine import async_session


# ────────────────────────────────────────────────────────────────
# Получить глобальный счётчик (создать, если ещё нет)
# ────────────────────────────────────────────────────────────────
async def get_global_counter() -> MessageCounter:
    """
    Гарантирует, что в message_counter есть ровно одна строка.
    Возвращает эту строку.
    """
    async with async_session() as session:
        async with session.begin():
            # SELECT ... FOR UPDATE — блокируем строку/таблицу на время транзакции
            counter = await session.scalar(
                select(MessageCounter)
                .limit(1)
                .with_for_update()
            )

            if counter is None:                 # таблица пуста → создаём
                counter = MessageCounter(successful=0, failed=0)
                session.add(counter)

        # выход из session.begin() = COMMIT
        return counter


# ────────────────────────────────────────────────────────────────
# Инкрементировать глобальный счётчик
# ────────────────────────────────────────────────────────────────
async def increment_global_counter(ok: bool = True, step: int = 1) -> None:
    """
    Увеличивает successful или failed на step для глобальной строки.
    """
    async with async_session() as session:
        async with session.begin():
            counter = await session.scalar(
                select(MessageCounter)
                .limit(1)
                .with_for_update()
            )
            if counter is None:                 # первый вызов → создаём строку
                counter = MessageCounter(successful=0, failed=0)
                session.add(counter)

            # Убеждаемся, что значения не None
            if counter.successful is None:
                counter.successful = 0
            if counter.failed is None:
                counter.failed = 0

            if ok:
                counter.successful += step
            else:
                counter.failed += step


# ────────────────────────────────────────────────────────────────
# 3. Получить один счётчик по ID
# ────────────────────────────────────────────────────────────────
async def get_global_counter_readonly() -> MessageCounter | None:
    """
    Возвращает глобальную строку без блокировки (только чтение).
    """
    async with async_session() as session:
        result = await session.execute(select(MessageCounter).limit(1))
        return result.scalars().first()


# ────────────────────────────────────────────────────────────────
# 5. Получить сумму успешных / неуспешных
# ────────────────────────────────────────────────────────────────
async def get_totals() -> tuple[int, int]:
    """
    Возвращает (successful, failed) глобального счётчика.
    Если строки ещё нет – вернёт (0, 0).
    """
    async with async_session() as session:
        result = await session.execute(select(MessageCounter).limit(1))
        counter = result.scalars().first()
        if counter is None:
            return 0, 0
        
        # Убеждаемся, что значения не None
        successful = counter.successful if counter.successful is not None else 0
        failed = counter.failed if counter.failed is not None else 0
        
        return successful, failed


# ────────────────────────────────────────────────────────────────
# Перенести локальные счетчики в глобальные итоги
# ────────────────────────────────────────────────────────────────
async def finalize_counters() -> tuple[int, int]:
    """
    Переносит значения из MessageCounter в MessageTotals и обнуляет MessageCounter.
    Возвращает (перенесенные_успешные, перенесенные_неуспешные).
    """
    async with async_session() as session:
        async with session.begin():
            # Получаем текущие значения локального счетчика
            counter = await session.scalar(
                select(MessageCounter).limit(1).with_for_update()
            )
            if counter is None:
                return 0, 0
            
            # Получаем/создаем глобальный счетчик
            totals = await session.scalar(
                select(MessageTotals).limit(1).with_for_update()
            )
            if totals is None:
                totals = MessageTotals(sent_total=0, failed_total=0)
                session.add(totals)
            
            # Убеждаемся, что значения не None
            successful_to_move = counter.successful if counter.successful is not None else 0
            failed_to_move = counter.failed if counter.failed is not None else 0
            
            # Убеждаемся, что значения в totals не None
            if totals.sent_total is None:
                totals.sent_total = 0
            if totals.failed_total is None:
                totals.failed_total = 0
            
            totals.sent_total += successful_to_move
            totals.failed_total += failed_to_move
            
            # Обнуляем локальный счетчик
            counter.successful = 0
            counter.failed = 0
            
            return successful_to_move, failed_to_move


# ────────────────────────────────────────────────────────────────
# 6. Полный сброс счётчика
# ────────────────────────────────────────────────────────────────
async def reset_global_counter() -> None:
    """
    Обнуляет successful и failed в глобальном счётчике.
    Если строки нет – создаёт её с нулями.
    """
    async with async_session() as session:
        async with session.begin():
            counter = await session.scalar(
                select(MessageCounter)
                .limit(1)
                .with_for_update()
            )
            if counter is None:            # таблица пуста – создаём
                counter = MessageCounter(successful=0, failed=0)
                session.add(counter)
            else:                          # обнуляем существующую
                counter.successful = 0
                counter.failed = 0


# получить или создать единственную строку
async def get_totals_row() -> MessageTotals:
    async with async_session() as session:
        async with session.begin():
            row = await session.scalar(
                select(MessageTotals).limit(1).with_for_update()
            )
            if row is None:
                row = MessageTotals(sent_total=0, failed_total=0)
                session.add(row)
        return row

# инкрементировать (sent_total или failed_total)
async def inc_totals(sent: int = 0, failed: int = 0) -> None:
    async with async_session() as session:
        async with session.begin():
            row = await session.scalar(
                select(MessageTotals).limit(1).with_for_update()
            )
            if row is None:
                row = MessageTotals(sent_total=0, failed_total=0)
                session.add(row)
            
            # Убеждаемся, что значения не None
            if row.sent_total is None:
                row.sent_total = 0
            if row.failed_total is None:
                row.failed_total = 0
                
            row.sent_total += sent
            row.failed_total += failed

# чтение без блокировки
async def read_totals() -> tuple[int, int]:
    async with async_session() as session:
        result = await session.execute(select(MessageTotals).limit(1))
        row = result.scalars().first()
        if row is None:
            return 0, 0
        
        # Убеждаемся, что значения не None
        sent_total = row.sent_total if row.sent_total is not None else 0
        failed_total = row.failed_total if row.failed_total is not None else 0
        
        return sent_total, failed_total

# обнулить
async def reset_totals() -> None:
    async with async_session() as session:
        async with session.begin():
            row = await session.scalar(
                select(MessageTotals).limit(1).with_for_update()
            )
            if row is None:
                row = MessageTotals(sent_total=0, failed_total=0)
                session.add(row)
            else:
                row.sent_total = 0
                row.failed_total = 0