# tests/test_broadcast.py
from __future__ import annotations
import pytest
import aiosqlite
from unittest.mock import AsyncMock, patch, MagicMock
from aiogram.exceptions import TelegramForbiddenError, TelegramRetryAfter
from db import init_db, record_start, record_subscription, mark_blocked
from services.broadcast import broadcast, broadcast_unsubscribed, broadcast_to_all
from aiogram.types import InlineKeyboardMarkup, InlineKeyboardButton
from aiogram.utils.keyboard import InlineKeyboardBuilder
from config import Speaker

SPEAKERS = [Speaker(name="S1", channel="@chan1")]


@pytest.fixture
async def db():
    async with aiosqlite.connect(":memory:") as conn:
        await init_db(conn)
        yield conn


async def test_broadcast_delivers_to_verified_user(db):
    await record_start(db, user_id=1, conference_id="conf-1")
    await record_subscription(db, user_id=1, conference_id="conf-1")

    bot = AsyncMock()
    bot.get_chat_member.return_value = MagicMock(status="member")

    report = await broadcast(bot, db, conference_id="conf-1", speakers=SPEAKERS, text="Hello")
    assert report["delivered"] == 1
    assert report["skipped"] == 0
    bot.send_message.assert_called_once_with(chat_id=1, text="Hello", parse_mode="HTML", reply_markup=None)


async def test_broadcast_marks_blocked_on_403(db):
    await record_start(db, user_id=2, conference_id="conf-1")
    await record_subscription(db, user_id=2, conference_id="conf-1")

    bot = AsyncMock()
    bot.get_chat_member.return_value = MagicMock(status="member")
    error = TelegramForbiddenError(method=MagicMock(), message="bot was blocked by the user")
    bot.send_message.side_effect = error

    report = await broadcast(bot, db, conference_id="conf-1", speakers=SPEAKERS, text="Hello")
    assert report["skipped"] == 1

    async with db.execute("SELECT user_id FROM blocked_users WHERE user_id=2") as cur:
        row = await cur.fetchone()
    assert row is not None


async def test_broadcast_skips_unsubscribed_after_verification(db):
    await record_start(db, user_id=3, conference_id="conf-1")
    await record_subscription(db, user_id=3, conference_id="conf-1")

    bot = AsyncMock()
    bot.get_chat_member.return_value = MagicMock(status="left")

    report = await broadcast(bot, db, conference_id="conf-1", speakers=SPEAKERS, text="Hello")
    assert report["skipped"] == 1
    bot.send_message.assert_not_called()


async def test_broadcast_retries_on_429(db):
    await record_start(db, user_id=4, conference_id="conf-1")
    await record_subscription(db, user_id=4, conference_id="conf-1")

    bot = AsyncMock()
    bot.get_chat_member.return_value = MagicMock(status="member")
    retry_error = TelegramRetryAfter(method=MagicMock(), message="Too Many Requests", retry_after=0)
    bot.send_message.side_effect = [retry_error, None]

    with patch("services.broadcast.asyncio.sleep", new=AsyncMock()):
        report = await broadcast(bot, db, conference_id="conf-1", speakers=SPEAKERS, text="Hello")
    assert report["delivered"] == 1


async def test_broadcast_returns_error_when_429_exhausts_retries(db):
    await record_start(db, user_id=5, conference_id="conf-1")
    await record_subscription(db, user_id=5, conference_id="conf-1")

    bot = AsyncMock()
    bot.get_chat_member.return_value = MagicMock(status="member")
    retry_error = TelegramRetryAfter(method=MagicMock(), message="Too Many Requests", retry_after=0)
    bot.send_message.side_effect = retry_error  # always raises — never succeeds

    with patch("services.broadcast.asyncio.sleep", new=AsyncMock()):
        report = await broadcast(bot, db, conference_id="conf-1", speakers=SPEAKERS, text="Hello")
    assert report["errors"] == 1
    assert report["delivered"] == 0


def _make_keyboard() -> InlineKeyboardMarkup:
    builder = InlineKeyboardBuilder()
    builder.row(InlineKeyboardButton(text="✅", callback_data="check_sub:conf-1"))
    return builder.as_markup()


async def test_broadcast_unsubscribed_delivers_text_message(db):
    await record_start(db, user_id=10, conference_id="conf-1")
    bot = AsyncMock()
    with patch("services.broadcast.asyncio.sleep", new=AsyncMock()):
        report = await broadcast_unsubscribed(bot, db, "conf-1", SPEAKERS, "Привет", None, _make_keyboard())
    assert report["delivered"] == 1
    assert report["skipped"] == 0
    bot.send_message.assert_called_once()


async def test_broadcast_unsubscribed_sends_photo_then_text(db):
    await record_start(db, user_id=11, conference_id="conf-1")
    bot = AsyncMock()
    with patch("services.broadcast.asyncio.sleep", new=AsyncMock()):
        report = await broadcast_unsubscribed(bot, db, "conf-1", SPEAKERS, "Привет", "photo123", _make_keyboard())
    assert report["delivered"] == 1
    bot.send_photo.assert_called_once_with(chat_id=11, photo="photo123")
    bot.send_message.assert_called_once()


async def test_broadcast_unsubscribed_photo_403_marks_blocked(db):
    await record_start(db, user_id=12, conference_id="conf-1")
    bot = AsyncMock()
    bot.send_photo.side_effect = TelegramForbiddenError(method=MagicMock(), message="blocked")
    with patch("services.broadcast.asyncio.sleep", new=AsyncMock()):
        report = await broadcast_unsubscribed(bot, db, "conf-1", SPEAKERS, "Привет", "photo123", _make_keyboard())
    assert report["skipped"] == 1
    async with db.execute("SELECT user_id FROM blocked_users WHERE user_id=12") as cur:
        row = await cur.fetchone()
    assert row is not None


async def test_broadcast_unsubscribed_step2_failure_is_error(db):
    await record_start(db, user_id=13, conference_id="conf-1")
    bot = AsyncMock()
    bot.send_photo.return_value = None
    bot.send_message.side_effect = Exception("network error")
    with patch("services.broadcast.asyncio.sleep", new=AsyncMock()):
        report = await broadcast_unsubscribed(bot, db, "conf-1", SPEAKERS, "Привет", "photo123", _make_keyboard())
    assert report["errors"] == 1
    assert report["delivered"] == 0


async def test_broadcast_unsubscribed_skips_already_subscribed(db):
    await record_start(db, user_id=14, conference_id="conf-1")
    await record_subscription(db, user_id=14, conference_id="conf-1")
    bot = AsyncMock()
    with patch("services.broadcast.asyncio.sleep", new=AsyncMock()):
        report = await broadcast_unsubscribed(bot, db, "conf-1", SPEAKERS, "Привет", None, _make_keyboard())
    assert report["delivered"] == 0
    bot.send_message.assert_not_called()


async def test_broadcast_to_all_text_only_with_keyboard(db):
    await record_start(db, user_id=10, conference_id="conf-1")
    bot = AsyncMock()
    kb = _make_keyboard()
    report = await broadcast_to_all(bot, db, text="Hello", photo_file_id=None, keyboard=kb, exclude_user_ids=set())
    assert report["delivered"] == 1
    assert report["skipped"] == 0
    assert report["errors"] == 0
    bot.send_message.assert_called_once()
    kwargs = bot.send_message.call_args.kwargs
    assert kwargs["chat_id"] == 10
    assert kwargs["text"] == "Hello"
    assert kwargs["reply_markup"] is kb


async def test_broadcast_to_all_photo_short_caption_single_message(db):
    await record_start(db, user_id=11, conference_id="conf-1")
    bot = AsyncMock()
    kb = _make_keyboard()
    report = await broadcast_to_all(bot, db, text="short", photo_file_id="photo123", keyboard=kb, exclude_user_ids=set())
    assert report["delivered"] == 1
    bot.send_photo.assert_called_once()
    kwargs = bot.send_photo.call_args.kwargs
    assert kwargs["photo"] == "photo123"
    assert kwargs["caption"] == "short"
    assert kwargs["reply_markup"] is kb
    bot.send_message.assert_not_called()


async def test_broadcast_to_all_photo_long_caption_two_messages(db):
    await record_start(db, user_id=12, conference_id="conf-1")
    long_text = "x" * 1500  # > 1024
    bot = AsyncMock()
    kb = _make_keyboard()
    report = await broadcast_to_all(bot, db, text=long_text, photo_file_id="photo123", keyboard=kb, exclude_user_ids=set())
    assert report["delivered"] == 1
    # photo sent without caption, no keyboard
    bot.send_photo.assert_called_once()
    photo_kwargs = bot.send_photo.call_args.kwargs
    assert "caption" not in photo_kwargs or photo_kwargs.get("caption") in (None, "")
    assert "reply_markup" not in photo_kwargs or photo_kwargs.get("reply_markup") is None
    # text sent separately with the keyboard
    bot.send_message.assert_called_once()
    msg_kwargs = bot.send_message.call_args.kwargs
    assert msg_kwargs["text"] == long_text
    assert msg_kwargs["reply_markup"] is kb


async def test_broadcast_to_all_excludes_admin(db):
    await record_start(db, user_id=20, conference_id="conf-1")
    await record_start(db, user_id=99, conference_id="conf-1")  # admin
    bot = AsyncMock()
    report = await broadcast_to_all(bot, db, text="Hi", photo_file_id=None, keyboard=None, exclude_user_ids={99})
    assert report["delivered"] == 1
    # Only one send_message call
    assert bot.send_message.call_count == 1
    assert bot.send_message.call_args.kwargs["chat_id"] == 20


async def test_broadcast_to_all_forbidden_marks_blocked(db):
    await record_start(db, user_id=30, conference_id="conf-1")
    bot = AsyncMock()
    bot.send_message.side_effect = TelegramForbiddenError(method=MagicMock(), message="blocked")
    report = await broadcast_to_all(bot, db, text="Hi", photo_file_id=None, keyboard=None, exclude_user_ids=set())
    assert report["skipped"] == 1
    assert report["delivered"] == 0
    async with db.execute("SELECT user_id FROM blocked_users WHERE user_id=30") as cur:
        row = await cur.fetchone()
    assert row is not None


async def test_broadcast_to_all_photo_403_marks_blocked(db):
    await record_start(db, user_id=31, conference_id="conf-1")
    bot = AsyncMock()
    bot.send_photo.side_effect = TelegramForbiddenError(method=MagicMock(), message="blocked")
    report = await broadcast_to_all(bot, db, text="Hi", photo_file_id="p", keyboard=None, exclude_user_ids=set())
    assert report["skipped"] == 1
    bot.send_message.assert_not_called()


async def test_broadcast_unsubscribed_escapes_speaker_name(db):
    await record_start(db, user_id=40, conference_id="conf-1")
    bot = AsyncMock()
    evil_speakers = [Speaker(name="<i>Evil</i>", channel="@evilchan")]
    with patch("services.broadcast.asyncio.sleep", new=AsyncMock()):
        report = await broadcast_unsubscribed(bot, db, "conf-1", evil_speakers, "Привет", None, _make_keyboard())
    assert report["delivered"] == 1
    sent_text = bot.send_message.call_args.kwargs["text"]
    assert "<i>Evil</i>" not in sent_text
    assert "&lt;i&gt;Evil&lt;/i&gt;" in sent_text


async def test_broadcast_to_all_logs_summary(db, caplog):
    await record_start(db, user_id=41, conference_id="conf-1")
    bot = AsyncMock()
    with caplog.at_level("INFO", logger="services.broadcast"):
        report = await broadcast_to_all(bot, db, text="Hi", photo_file_id=None, keyboard=None, exclude_user_ids=set())
    assert report["delivered"] == 1
    assert any("finished: sent=1 blocked=0 failed=0" in r.message for r in caplog.records)


async def test_broadcast_to_all_logs_progress_every_25(db, caplog):
    for uid in range(1000, 1030):  # 30 users
        await record_start(db, user_id=uid, conference_id="conf-1")
    bot = AsyncMock()
    with caplog.at_level("INFO", logger="services.broadcast"):
        await broadcast_to_all(bot, db, text="Hi", photo_file_id=None, keyboard=None, exclude_user_ids=set())
    progress_lines = [r.message for r in caplog.records if "progress:" in r.message]
    assert any("25/30" in line for line in progress_lines)
    # Only one progress line for 30 users (25 is the only multiple of 25 <= 30).
    assert len(progress_lines) == 1
