"""Message forwarding engine extracted from core.py for maintainability.

Uses `SignalParser` to:
- Skip promo / advert messages unconditionally.
- Only trigger price-augmentation on high-confidence entry signals.
"""

from __future__ import annotations

import asyncio
import logging
import time
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Optional

from telethon import events
from telethon.errors.rpcerrorlist import FloodWaitError

import config
from signal_parser import SignalParser, SignalType
from signal_processor import SignalProcessor

if TYPE_CHECKING:
    from core import ForwarderCore

logger = logging.getLogger(__name__)


class MessageForwarder:
    """Handles Telegram new-message / edit / delete events and forwarding logic."""

    def __init__(self, core: ForwarderCore) -> None:
        self.core = core
        self._parser = SignalParser()
        self._handlers_installed = False
        # id(client) -> account_id, so an event can be routed to the account
        # that owns the source (and deduped when multiple clients see it).
        self._client_accounts: dict = {}
        # Auto-pair dedupe: {(source_id, msg_id): timestamp}. Auto pairs
        # (account_id NULL) may fire on several connected clients — the first
        # one wins for a short window so the message isn't forwarded N times.
        self._auto_claims: dict = {}

    # ── Handler wiring ──────────────────────────────────────────────────────

    def install(self, client, account_id: int) -> None:
        client.add_event_handler(self._on_new_message, events.NewMessage())
        client.add_event_handler(self._on_edit, events.MessageEdited())
        client.add_event_handler(self._on_delete, events.MessageDeleted())
        self._client_accounts[id(client)] = account_id
        self._handlers_installed = True

    def remove(self, client) -> None:
        client.remove_event_handler(self._on_new_message)
        client.remove_event_handler(self._on_edit)
        client.remove_event_handler(self._on_delete)
        self._client_accounts.pop(id(client), None)
        self._handlers_installed = False

    @staticmethod
    def _get_client(event) -> Any:
        """Extract the Telethon client from an event."""
        return getattr(event, "_client", None) or getattr(event, "client", None)

    def _lookup_pairs(self, event, source_id: int, claim: bool = True) -> list:
        """Pairs for this event, scoped to the account whose client fired.

        Bound pairs (account_id set) are only handled by their own account's
        client. Auto pairs (account_id NULL) are claimed by the first client
        that sees the message — later clients skip to avoid duplicates.
        """
        client = self._get_client(event)
        acct_id = self._client_accounts.get(id(client)) if client else None
        pairs: list = []
        if acct_id is not None:
            pairs += self.core._pairs.get(acct_id, {}).get(source_id, [])
        auto_pairs = self.core._pairs.get(None, {}).get(source_id, [])
        if auto_pairs and claim:
            msg_id = getattr(event, "id", None) or getattr(
                getattr(event, "message", None), "id", None
            )
            key = (source_id, msg_id)
            now = time.time()
            if len(self._auto_claims) > 500:
                cutoff = now - 10
                self._auto_claims = {
                    k: v for k, v in self._auto_claims.items() if v >= cutoff
                }
            if key in self._auto_claims:
                return []
            self._auto_claims[key] = now
            pairs += auto_pairs
        return pairs

    def _claim_auto(self, key) -> bool:
        """Claim a dedupe key for an auto pair. Returns False if already claimed."""
        now = time.time()
        if len(self._auto_claims) > 500:
            cutoff = now - 10
            self._auto_claims = {k: v for k, v in self._auto_claims.items() if v >= cutoff}
        if key in self._auto_claims:
            return False
        self._auto_claims[key] = now
        return True

    # ── Core event handlers ─────────────────────────────────────────────────

    async def _on_new_message(self, event) -> None:
        if not self.core._forwarding:
            return
        source_id = event.chat_id
        pairs = self._lookup_pairs(event, source_id)
        if not pairs:
            return

        msg = event.message
        text = msg.message or ""
        signal_type = self._parser.classify(text)
        is_entry = signal_type in (SignalType.ENTRY, SignalType.ENTRY_PENDING)
        is_promo = signal_type == SignalType.PROMO
        if is_promo:
            return  # silently drop promos — no webhook, no bot msg, no forward

        # ── Per-message actions (sent once, not duplicated per pair) ───────
        # Only fire if at least one pair for this source uses price_augment
        # or forward_via_bot (the flags that justify a SignalProcessor).
        needs_signal_proc = any(
            p.get("price_augment") or p.get("forward_via_bot") for p in pairs
        )

        if self.core._signal_processor and needs_signal_proc:
            # Verbatim webhook — every signal type, sent once per message
            try:
                await self.core._signal_processor.send_signal_webhook(
                    source_chat_id=source_id,
                    message_id=msg.id,
                    text=text,
                    date=int(msg.date.timestamp()) if msg.date else None,
                    reply_to_message_id=msg.reply_to_msg_id,
                )
            except Exception as exc:
                logger.warning(
                    "Signal webhook failed for msg %s: %s", msg.id, exc
                )

            # Signal notification to admin chat — once per message.
            # cTrader client may not be running (forward_via_bot only, no
            # price_augment), in which case the price line is simply omitted.
            try:
                # Use the symbol from the first price_augment pair, or default
                sym = next(
                    (p.get("augment_symbol") or config.CTRADER_SYMBOL
                     for p in pairs if p.get("price_augment")),
                    config.CTRADER_SYMBOL,
                )
                await self.core._signal_processor.send_signal_notification(
                    self.core._ctrader_client,
                    source_chat_id=source_id,
                    message_id=msg.id,
                    text=text,
                    signal_type=signal_type.value,
                    date=int(msg.date.timestamp()) if msg.date else None,
                    reply_to_message_id=msg.reply_to_msg_id,
                    symbol=sym,
                )
            except Exception as exc:
                logger.warning(
                    "Signal notification failed for msg %s: %s", msg.id, exc
                )

        for pair in pairs:
            # ── Optional LLM classification (OpenAI-compatible) ───────────
            # Opt-in per pair + global master switch. Runs for *every*
            # non-promo message (entry / manage / result / unknown) so the
            # model can confirm entries and detect updates to previous
            # entries (reply_to_message_id is passed for that reason).
            # Fire-and-forget: it never blocks or delays forwarding.
            if (
                pair.get("llm_enabled")
                and self.core._llm_engine
                and self.core._llm_engine.available
            ):
                try:
                    asyncio.create_task(
                        self.core._llm_engine.classify_and_post(
                            self.core,
                            source_chat_id=source_id,
                            message_id=msg.id,
                            text=text,
                            date=int(msg.date.timestamp()) if msg.date else None,
                            reply_to_message_id=msg.reply_to_msg_id,
                            source_title=pair.get("source_title") or "",
                            symbol=pair.get("augment_symbol") or config.CTRADER_SYMBOL,
                        )
                    )
                except Exception as exc:
                    logger.warning(
                        "LLM classify schedule failed for msg %s: %s", msg.id, exc
                    )

            # ── Price-augmented admin message ────────────────────────────
            # Detailed snapshot with live bid/ask — sent to admin chat only.
            if (
                pair.get("price_augment")
                and is_entry
                and self.core._signal_processor
                and self.core._ctrader_client
            ):
                try:
                    await self.core._signal_processor.send_augmented_snapshot(
                        self.core._ctrader_client,
                        source_chat_id=source_id,
                        message_id=msg.id,
                        text=text,
                        date=int(msg.date.timestamp()) if msg.date else None,
                        reply_to_message_id=msg.reply_to_msg_id,
                        symbol=pair.get("augment_symbol") or config.CTRADER_SYMBOL,
                    )
                except Exception as exc:
                    logger.warning(
                        "Price augmentation failed for msg %s: %s", msg.id, exc
                    )

            if self._should_forward(msg, pair):
                await self._forward_message(msg, pair, self._get_client(event))

    async def _on_edit(self, event) -> None:
        if not self.core._forwarding:
            return
        source_id = event.chat_id
        pairs = self._lookup_pairs(event, source_id)
        if not pairs:
            return

        client = self._get_client(event)
        if not client:
            return
        msg = event.message
        text = msg.message or ""
        for pair in pairs:
            if pair.get("forward_as_link"):
                continue
            if not self._should_forward(msg, pair):
                continue
            dest_id = pair["dest_chat_id"]
            mapping = await self.core._run_db(
                self.core.db.get_message_mapping, source_id, dest_id, msg.id
            )
            if not mapping:
                continue
            dest_msg_id = mapping["dest_msg_id"]
            via_bot = bool(mapping.get("via_bot"))
            try:
                if via_bot:
                    bot_token = await self._resolve_bot_token(pair)
                    await self._edit_via_bot(msg, dest_id, dest_msg_id, bot_token=bot_token)
                else:
                    await self._send_with_retry(
                        lambda: client.edit_message(
                            dest_id,
                            dest_msg_id,
                            text,
                            formatting_entities=msg.entities,
                            link_preview=False,
                        )
                    )
                logger.info("Edited %s -> %s/%s", msg.id, dest_id, dest_msg_id)
            except Exception as e:
                logger.error("Edit error (dest=%s): %s", dest_id, e, exc_info=False)

    async def _on_delete(self, event) -> None:
        if not self.core._forwarding:
            return
        source_id = getattr(event, "chat_id", None)
        pairs = self._lookup_pairs(event, source_id, claim=False)
        auto_pairs = self.core._pairs.get(None, {}).get(source_id, [])
        if not pairs and not auto_pairs:
            return

        client = self._get_client(event)
        if not client:
            return
        deleted_ids = getattr(event, "deleted_ids", [])
        if not deleted_ids:
            return
        has_auto = bool(auto_pairs)
        for src_id in deleted_ids:
            cur_pairs = pairs
            if has_auto:
                # Auto-pair deletes fire on every client that sees the event —
                # the first client to claim a deleted id handles them.
                if not self._claim_auto((source_id, src_id)):
                    continue
                cur_pairs = pairs + auto_pairs
            for pair in cur_pairs:
                dest_id = pair["dest_chat_id"]
                mapping = await self.core._run_db(
                    self.core.db.get_message_mapping, source_id, dest_id, src_id
                )
                if not mapping:
                    continue
                dest_msg_id = mapping["dest_msg_id"]
                via_bot = bool(mapping.get("via_bot"))
                try:
                    if via_bot and self.core._signal_processor:
                        bot_token = await self._resolve_bot_token(pair)
                        await self.core._signal_processor.delete_bot_message(
                            dest_id, dest_msg_id, bot_token=bot_token
                        )
                    else:
                        await client.delete_messages(dest_id, [dest_msg_id])
                    logger.info("Deleted %s -> %s/%s", src_id, dest_id, dest_msg_id)
                except Exception as e:
                    logger.error(
                        "Delete error (dest=%s): %s", dest_id, e, exc_info=False
                    )

    # ── Forwarding helpers ──────────────────────────────────────────────────

    def _should_forward(self, msg, pair: dict) -> bool:
        text = msg.message or ""
        # Hard-filter promos — they never belong in a signal channel.
        if self._parser.is_promo(text):
            return False

        has_text = bool(text.strip())
        has_media = bool(msg.media)
        ftype = pair.get("filter_type") or "none"
        if ftype == "text_only" and not has_text:
            return False
        if ftype == "media_only" and not has_media:
            return False
        if not pair.get("include_text", True) and not has_media:
            return False
        if not pair.get("include_media", True) and has_media:
            return False
        if pair.get("skip_standalone_media") and has_media and not has_text:
            return False
        return True

    async def _resolve_bot_token(self, pair: dict) -> Optional[str]:
        """Load the per-pair bot token from the DB, falling back to the global token."""
        bot_token_id = pair.get("bot_token_id")
        if not bot_token_id:
            return None
        try:
            return await self.core._run_db(
                self.core.db.get_decrypted_bot_token, bot_token_id
            )
        except Exception as exc:
            logger.warning(
                "Could not load bot token id=%s for pair %s: %s",
                bot_token_id,
                pair.get("id"),
                exc,
            )
            return None

    async def _send_with_retry(self, coro_factory, max_retries=3):
        last_exc = None
        for attempt in range(max_retries):
            try:
                return await coro_factory()
            except FloodWaitError as e:
                wait = min(e.seconds or 5, 300)
                logger.warning(
                    "FloodWait on attempt %s: sleeping %ss", attempt + 1, wait
                )
                await asyncio.sleep(wait)
            except Exception as e:
                last_exc = e
                raise
        if last_exc:
            raise last_exc
        raise RuntimeError("Max retries exceeded")

    async def _forward_message(self, msg, pair: dict, client=None) -> None:
        source_id = pair["source_chat_id"]
        dest_id = pair["dest_chat_id"]
        text = msg.message or ""
        via_bot = bool(pair.get("forward_via_bot"))

        reply_to = None
        if msg.reply_to_msg_id:
            reply_to = await self.core._run_db(
                self.core.db.get_mapped_message, source_id, dest_id, msg.reply_to_msg_id
            )

        bot_token = await self._resolve_bot_token(pair) if via_bot else None

        try:
            if via_bot and not msg.media:
                sent_id = await self._forward_via_bot(
                    msg, pair, dest_id, reply_to=reply_to, bot_token=bot_token
                )
            elif pair.get("forward_as_link"):
                sent = await self._send_with_retry(
                    lambda: client.forward_messages(dest_id, msg)
                )
                sent_id = getattr(sent, "id", None)
                if isinstance(sent, list) and sent:
                    sent_id = sent[0].id
            elif msg.media:
                sent = await self._send_with_retry(
                    lambda: client.send_file(
                        dest_id,
                        file=msg.media,
                        caption=text,
                        reply_to=reply_to,
                        formatting_entities=msg.entities,
                        link_preview=False,
                    )
                )
                sent_id = getattr(sent, "id", None)
                if isinstance(sent, list) and sent:
                    sent_id = sent[0].id
            else:
                sent = await self._send_with_retry(
                    lambda: client.send_message(
                        dest_id,
                        text,
                        reply_to=reply_to,
                        formatting_entities=msg.entities,
                        link_preview=False,
                    )
                )
                sent_id = getattr(sent, "id", None)
                if isinstance(sent, list) and sent:
                    sent_id = sent[0].id

            if sent_id:
                await self.core._run_db(
                    self.core.db.save_message_mapping,
                    source_id,
                    msg.id,
                    dest_id,
                    sent_id,
                    via_bot=via_bot,
                )
                logger.info(
                    "Forwarded %s -> %s (dest=%s, via_bot=%s)",
                    msg.id,
                    sent_id,
                    dest_id,
                    via_bot,
                )
        except Exception as e:
            logger.error("Forward error (dest=%s): %s", dest_id, e, exc_info=False)

    async def _forward_via_bot(
        self,
        msg,
        pair: dict,
        dest_id: int,
        reply_to: Optional[int] = None,
        bot_token: Optional[str] = None,
    ) -> Optional[int]:
        if not self.core._signal_processor:
            logger.warning("No signal processor available for bot forwarding")
            return None

        entities = None
        if msg.entities:
            try:
                entities = [
                    {
                        "type": e.__class__.__name__.lower().replace(
                            "messageentity", ""
                        ),
                        "offset": e.offset,
                        "length": e.length,
                    }
                    for e in msg.entities
                ]
            except Exception as exc:
                logger.debug("Could not serialize entities: %s", exc)

        return await self.core._signal_processor.send_bot_text(
            dest_chat_id=dest_id,
            text=msg.message or "",
            entities=entities or None,
            reply_to_message_id=reply_to,
            bot_token=bot_token,
        )

    async def _edit_via_bot(
        self,
        msg,
        dest_id: int,
        dest_msg_id: int,
        bot_token: Optional[str] = None,
    ) -> None:
        if not self.core._signal_processor:
            return
        entities = None
        if msg.entities:
            try:
                entities = [
                    {
                        "type": e.__class__.__name__.lower().replace(
                            "messageentity", ""
                        ),
                        "offset": e.offset,
                        "length": e.length,
                    }
                    for e in msg.entities
                ]
            except Exception as exc:
                logger.debug("Could not serialize entities: %s", exc)
        await self.core._signal_processor.edit_bot_text(
            dest_chat_id=dest_id,
            dest_message_id=dest_msg_id,
            text=msg.message or "",
            entities=entities or None,
            bot_token=bot_token,
        )

    # ── Test signal emission ────────────────────────────────────────────────

    async def emit_test_signal(self, pair: dict, text: str) -> dict:
        """Imitate a source-channel message for a pair without posting to the source."""
        source_id = pair["source_chat_id"]
        dest_id = pair["dest_chat_id"]

        class FakeMsg:
            id = int(time.time() * 1000)
            message = text
            media = None
            entities = None
            reply_to_msg_id = None
            date = datetime.fromtimestamp(time.time(), tz=timezone.utc)

        msg = FakeMsg()
        is_entry = self._parser.is_high_confidence_entry(text)

        if self.core._signal_processor:
            try:
                await self.core._signal_processor.send_signal_webhook(
                    source_chat_id=source_id,
                    message_id=msg.id,
                    text=text,
                    date=int(msg.date.timestamp()),
                )
            except Exception as exc:
                logger.warning("Test webhook failed: %s", exc)

        if (
            pair.get("llm_enabled")
            and self.core._llm_engine
            and self.core._llm_engine.available
        ):
            try:
                await self.core._llm_engine.classify_and_post(
                    self.core,
                    source_chat_id=source_id,
                    message_id=msg.id,
                    text=text,
                    date=int(msg.date.timestamp()),
                    reply_to_message_id=None,
                    source_title=pair.get("source_title") or "",
                    symbol=pair.get("augment_symbol") or config.CTRADER_SYMBOL,
                )
            except Exception as exc:
                logger.warning("Test LLM classify failed: %s", exc)

        if (
            pair.get("price_augment")
            and is_entry
            and self.core._signal_processor
            and self.core._ctrader_client
        ):
            try:
                await self.core._signal_processor.send_augmented_snapshot(
                    self.core._ctrader_client,
                    source_chat_id=source_id,
                    message_id=msg.id,
                    text=text,
                    date=int(msg.date.timestamp()),
                    symbol=pair.get("augment_symbol") or config.CTRADER_SYMBOL,
                )
            except Exception as exc:
                logger.warning("Test price augmentation failed: %s", exc)

        if self._should_forward(msg, pair):
            client = self.core._get_client(pair.get("account_id"))
            await self._forward_message(msg, pair, client)

        return {
            "status": "ok",
            "pair_id": pair["id"],
            "source_chat_id": source_id,
            "dest_chat_id": dest_id,
            "via_bot": bool(pair.get("forward_via_bot")),
            "fake_message_id": msg.id,
        }
