"""Minimal async cTrader Open API spot-price client.

Designed for memory-constrained shared hosting (~250 MB budget).
Imports only the protobuf runtime and avoids numpy / heavy execution layers.
Keeps a single persistent TCP+SSL connection to cTrader and caches the latest
XAUUSD (or configured symbol) tick for instant lookup when a Telegram signal
arrives.
"""

from __future__ import annotations

import asyncio
import logging
import os
import ssl
import struct
import time
from dataclasses import dataclass
from typing import Optional

from ctrader_open_api.messages.OpenApiCommonMessages_pb2 import (
    ProtoHeartbeatEvent,
    ProtoMessage,
)
from ctrader_open_api.messages.OpenApiMessages_pb2 import (
    ProtoOAAccountAuthReq,
    ProtoOAApplicationAuthReq,
    ProtoOASpotEvent,
    ProtoOASubscribeSpotsReq,
    ProtoOASymbolsListReq,
)
from ctrader_open_api.protobuf import Protobuf

logger = logging.getLogger(__name__)

_FRAME_HEADER = struct.Struct(">I")
DEMO_HOST = "demo.ctraderapi.com"
LIVE_HOST = "live.ctraderapi.com"
CTRADER_PORT = 5035

# Rough symbol-name → decimal-digit mapping used by cTrader for relative prices.
_GUESS_DIGITS_OVERRIDES: dict[str, int] = {"XAUUSD": 3, "XAGUSD": 3}


def _guess_digits(name: str) -> int:
    """Estimate price digits from symbol name convention."""
    upper = name.upper()
    if upper in _GUESS_DIGITS_OVERRIDES:
        return _GUESS_DIGITS_OVERRIDES[upper]
    if any(tok in upper for tok in ("BTC", "ETH")):
        return 2
    if "XAU" in upper or "XAG" in upper:
        return 3
    if len(upper) == 6 and not any(c.isdigit() for c in upper):
        return 5
    return 1


def _price_from_relative(raw: int, digits: int) -> float:
    return round(raw / 1e5, digits)


@dataclass(slots=True)
class SpotTick:
    symbol_id: int
    symbol_name: str
    bid: float
    ask: float
    timestamp_ms: int


class CTraderSpotClient:
    """Lightweight cTrader client that subscribes to one or more spot feeds."""

    def __init__(
        self,
        client_id: str,
        client_secret: str,
        access_token: str,
        account_id: int,
        use_live: bool = False,
        symbols: Optional[list[str]] = None,
        symbol_ids: Optional[list[int]] = None,
    ):
        self.client_id = client_id
        self.client_secret = client_secret
        self.access_token = access_token
        self.account_id = account_id
        self.use_live = use_live
        self.symbols = [s.upper() for s in (symbols or ["XAUUSD"])]
        self.symbol_ids = list(symbol_ids or [])

        self._host = LIVE_HOST if use_live else DEMO_HOST
        self._reader: Optional[asyncio.StreamReader] = None
        self._writer: Optional[asyncio.StreamWriter] = None
        self._send_queue: asyncio.Queue[bytes] = asyncio.Queue()
        self._tasks: list[asyncio.Task[None]] = []
        self._pending: dict[str, asyncio.Future[Any]] = {}
        self._msg_counter = 0
        self._last_ticks: dict[int, SpotTick] = {}
        self._symbol_ids: dict[int, str] = {}  # symbol_id -> name
        self._symbol_digits: dict[int, int] = {}
        self._subscribed: set[int] = set()
        self._connected = False
        self._stopping = False
        self._reconnect_delay = 1.0
        self._max_reconnect_delay = 60.0
        self._price_events: dict[int, asyncio.Event] = {}
        self._start_lock = asyncio.Lock()
        self._ready_event = asyncio.Event()

    # ── Public API ───────────────────────────────────────────────────────────

    async def start(self) -> None:
        """Start the persistent connection loop (non-blocking).

        The actual connect/auth/subscribe work happens in a background task so
        the caller can create the processor immediately. The loop retries with
        exponential backoff until a connection is established or stop() is called.
        """
        async with self._start_lock:
            if self._tasks or self._connected:
                return
            self._stopping = False
            self._ready_event.clear()
            self._tasks.append(asyncio.create_task(self._startup_loop()))

    async def wait_ready(self, timeout: float = 60.0) -> bool:
        """Wait until the client is connected and subscribed."""
        try:
            await asyncio.wait_for(self._ready_event.wait(), timeout=timeout)
            return True
        except asyncio.TimeoutError:
            return False

    async def stop(self) -> None:
        """Disconnect and cancel background tasks."""
        self._stopping = True
        self._connected = False
        self._ready_event.clear()
        await self._send_queue.put(b"")  # sentinel to stop write loop
        for task in self._tasks:
            task.cancel()
        if self._tasks:
            await asyncio.gather(*self._tasks, return_exceptions=True)
            self._tasks.clear()
        if self._writer:
            self._writer.close()
            try:
                await self._writer.wait_closed()
            except Exception:
                pass
            self._writer = None
        logger.info("cTrader spot client stopped")

    def get_tick(self, symbol_name: str) -> Optional[SpotTick]:
        """Return the latest cached tick for a symbol, or None."""
        symbol_name = symbol_name.upper()
        for sid, name in self._symbol_ids.items():
            if name == symbol_name:
                return self._last_ticks.get(sid)
        return None

    async def wait_for_tick(
        self, symbol_name: str, timeout: float = 10.0
    ) -> Optional[SpotTick]:
        """Wait until a tick arrives for symbol_name (or timeout)."""
        symbol_name = symbol_name.upper()
        existing = self.get_tick(symbol_name)
        if existing:
            return existing

        sid = next(
            (sid for sid, name in self._symbol_ids.items() if name == symbol_name), None
        )
        if sid is None:
            # Symbol id not resolved yet; poll briefly.
            deadline = time.time() + timeout
            while time.time() < deadline:
                tick = self.get_tick(symbol_name)
                if tick:
                    return tick
                await asyncio.sleep(0.1)
            return self.get_tick(symbol_name)

        event = self._price_events.setdefault(sid, asyncio.Event())
        try:
            await asyncio.wait_for(event.wait(), timeout=timeout)
        except asyncio.TimeoutError:
            pass
        return self.get_tick(symbol_name)

    # ── Internal lifecycle ───────────────────────────────────────────────────

    async def _startup_loop(self) -> None:
        """Keep trying to connect until ready or stopped."""
        while not self._stopping:
            try:
                await self._connect_and_auth()
                self._connected = True
                self._ready_event.set()
                self._reconnect_delay = 1.0
                logger.info("cTrader spot client ready")
                return
            except asyncio.CancelledError:
                raise
            except Exception as exc:
                if self._stopping:
                    return
                logger.warning(
                    "cTrader startup failed: %s; retrying in %.1fs",
                    exc,
                    self._reconnect_delay,
                )
                await self._cleanup()
                await asyncio.sleep(self._reconnect_delay)
                self._reconnect_delay = min(
                    self._reconnect_delay * 2, self._max_reconnect_delay
                )

    async def _cleanup(self) -> None:
        """Cancel transport tasks and close the socket without touching the startup loop."""
        current = asyncio.current_task()
        tasks_to_cancel = [t for t in self._tasks if t is not current]
        for task in tasks_to_cancel:
            task.cancel()
        if tasks_to_cancel:
            await asyncio.gather(*tasks_to_cancel, return_exceptions=True)
            self._tasks = [t for t in self._tasks if t is current]
        if self._writer:
            try:
                self._writer.close()
                await self._writer.wait_closed()
            except Exception:
                pass
            self._writer = None
        self._reader = None

    async def _connect_and_auth(self) -> None:
        await self._connect()
        await self._authenticate_app()

        if not self.access_token:
            raise RuntimeError(
                "No cTrader access token provided.  "
                "Authenticate via /callback first and store the token in the vault."
            )

        await self._authenticate_account()
        await self._resolve_and_subscribe_symbols()
        self._tasks.append(asyncio.create_task(self._heartbeat_loop()))

    async def _connect(self) -> None:
        context = ssl.create_default_context()
        self._reader, self._writer = await asyncio.open_connection(
            self._host, CTRADER_PORT, ssl=context
        )
        self._tasks.append(asyncio.create_task(self._read_loop()))
        self._tasks.append(asyncio.create_task(self._write_loop()))
        logger.info("Connected to cTrader %s:%s", self._host, CTRADER_PORT)

    async def _reconnect(self) -> None:
        if self._stopping:
            return
        logger.warning("cTrader connection lost; reconnecting...")
        self._connected = False
        self._ready_event.clear()
        await self._cleanup()

        while not self._stopping:
            try:
                await asyncio.sleep(self._reconnect_delay)
                self._reconnect_delay = min(
                    self._reconnect_delay * 2, self._max_reconnect_delay
                )
                await self._connect_and_auth()
                self._connected = True
                self._ready_event.set()
                self._reconnect_delay = 1.0
                return
            except asyncio.CancelledError:
                raise
            except Exception as exc:
                logger.warning("cTrader reconnect failed: %s", exc)

    async def _authenticate_app(self) -> Any:
        req = ProtoOAApplicationAuthReq()
        req.clientId = self.client_id
        req.clientSecret = self.client_secret
        res = await self._send_request(req)
        logger.info("cTrader application authenticated")
        return res

    async def _authenticate_account(self) -> Any:
        req = ProtoOAAccountAuthReq()
        req.ctidTraderAccountId = self.account_id
        req.accessToken = self.access_token
        res = await self._send_request(req)
        error_code = getattr(res, "errorCode", "")
        if error_code:
            description = getattr(res, "description", "")
            raise RuntimeError(
                f"cTrader account auth failed: {error_code} ({description})"
            )
        logger.info("cTrader account %d authenticated", self.account_id)
        return res

    async def _resolve_and_subscribe_symbols(self) -> None:
        # Always query the broker's symbol list so explicit symbol ids are mapped
        # to their real names. Blindly mapping all ids to symbols[0] caused a
        # single id (e.g. 41) to be mis-labeled as the wrong symbol.
        req = ProtoOASymbolsListReq()
        req.ctidTraderAccountId = self.account_id
        res = await self._send_request(req)

        error_code = getattr(res, "errorCode", "")
        if error_code:
            description = getattr(res, "description", "")
            raise RuntimeError(
                f"cTrader symbol list failed: {error_code} ({description})"
            )

        raw_symbols = getattr(res, "symbol", None) or getattr(res, "symbols", [])
        symbol_by_name: dict[str, int] = {}
        symbol_by_id: dict[int, str] = {}
        for sym in raw_symbols:
            name = getattr(sym, "symbolName", None) or getattr(sym, "name", None)
            sid = getattr(sym, "symbolId", None) or getattr(sym, "symbol_id", None)
            if name and sid is not None:
                name_up = name.upper()
                symbol_by_name[name_up] = sid
                symbol_by_id[sid] = name_up

        ids_to_subscribe: list[int] = []

        if self.symbol_ids:
            for sid in self.symbol_ids:
                sym_name = symbol_by_id.get(sid)
                if sym_name is None:
                    logger.warning(
                        "Explicit symbol id %d not found on account; skipping", sid
                    )
                    continue
                self._symbol_ids[sid] = sym_name
                self._symbol_digits[sid] = _guess_digits(sym_name)
                ids_to_subscribe.append(sid)
        else:
            for sym_name in self.symbols:
                sid = symbol_by_name.get(sym_name)
                if sid is None:
                    logger.warning("Symbol %s not found on account", sym_name)
                    continue
                self._symbol_ids[sid] = sym_name
                self._symbol_digits[sid] = _guess_digits(sym_name)
                ids_to_subscribe.append(sid)

        if ids_to_subscribe:
            logger.info(
                "Resolved cTrader symbol ids: %s",
                ", ".join(
                    f"{sid}:{self._symbol_ids[sid]}" for sid in ids_to_subscribe
                ),
            )
            await self._subscribe_spots(ids_to_subscribe)

    async def _subscribe_spots(self, symbol_ids: list[int]) -> None:
        new_ids = [sid for sid in symbol_ids if sid not in self._subscribed]
        if not new_ids:
            return
        req = ProtoOASubscribeSpotsReq()
        req.ctidTraderAccountId = self.account_id
        req.symbolId.extend(new_ids)
        req.subscribeToSpotTimestamp = True
        res = await self._send_request(req)
        error_code = getattr(res, "errorCode", "")
        if error_code:
            description = getattr(res, "description", "")
            raise RuntimeError(
                f"cTrader spot subscription failed: {error_code} ({description})"
            )
        self._subscribed.update(new_ids)
        logger.info("Subscribed to cTrader spots: %s", new_ids)

    # ── Transport / protocol ─────────────────────────────────────────────────

    async def _send_request(self, message: Any, timeout: float = 10.0) -> Any:
        self._msg_counter += 1
        client_msg_id = str(self._msg_counter)
        if hasattr(message, "clientMsgId"):
            message.clientMsgId = client_msg_id

        envelope = ProtoMessage()
        envelope.payloadType = message.payloadType
        envelope.payload = message.SerializeToString()
        envelope.clientMsgId = client_msg_id

        loop = asyncio.get_running_loop()
        fut: asyncio.Future[Any] = loop.create_future()
        self._pending[client_msg_id] = fut
        await self._send_queue.put(envelope.SerializeToString())
        try:
            return await asyncio.wait_for(fut, timeout=timeout)
        except asyncio.TimeoutError:
            self._pending.pop(client_msg_id, None)
            raise

    async def _write_loop(self) -> None:
        try:
            while True:
                data = await self._send_queue.get()
                if data == b"":
                    return
                frame = _FRAME_HEADER.pack(len(data)) + data
                if self._writer and not self._writer.is_closing():
                    self._writer.write(frame)
                    await self._writer.drain()
        except asyncio.CancelledError:
            raise
        except (ConnectionError, OSError) as exc:
            logger.warning("cTrader write error: %s", exc)
            asyncio.create_task(self._reconnect())

    async def _read_loop(self) -> None:
        try:
            while True:
                header = await self._reader.readexactly(4)
                length = _FRAME_HEADER.unpack(header)[0]
                payload = await self._reader.readexactly(length)
                await self._on_frame(payload)
        except asyncio.CancelledError:
            raise
        except (asyncio.IncompleteReadError, ConnectionResetError, OSError) as exc:
            logger.warning("cTrader read error: %s", exc)
            asyncio.create_task(self._reconnect())

    async def _on_frame(self, raw: bytes) -> None:
        proto_msg = ProtoMessage()
        proto_msg.ParseFromString(raw)

        if proto_msg.payloadType == ProtoHeartbeatEvent().payloadType:
            return

        try:
            domain_obj = Protobuf.extract(proto_msg)
        except Exception as exc:
            logger.warning("Failed to decode cTrader frame: %s", exc)
            return

        if proto_msg.clientMsgId and proto_msg.clientMsgId in self._pending:
            fut = self._pending.pop(proto_msg.clientMsgId)
            if not fut.done():
                fut.set_result(domain_obj)

        if proto_msg.payloadType == ProtoOASpotEvent().payloadType:
            await self._on_spot_event(domain_obj)

    async def _on_spot_event(self, msg: Any) -> None:
        sid = int(msg.symbolId)
        name = self._symbol_ids.get(sid, str(sid))
        digits = self._symbol_digits.get(sid, _guess_digits(name))

        if not (msg.HasField("bid") and msg.HasField("ask")):
            return

        ts = int(getattr(msg, "timestamp", time.time() * 1000))
        tick = SpotTick(
            symbol_id=sid,
            symbol_name=name,
            bid=_price_from_relative(msg.bid, digits),
            ask=_price_from_relative(msg.ask, digits),
            timestamp_ms=ts,
        )
        self._last_ticks[sid] = tick
        event = self._price_events.setdefault(sid, asyncio.Event())
        event.set()
        logger.debug("Spot tick %s bid=%s ask=%s", name, tick.bid, tick.ask)

    async def _heartbeat_loop(self) -> None:
        try:
            while True:
                await asyncio.sleep(10)
                hb = ProtoHeartbeatEvent()
                envelope = ProtoMessage()
                envelope.payloadType = hb.payloadType
                envelope.payload = hb.SerializeToString()
                try:
                    await self._send_queue.put(envelope.SerializeToString())
                except Exception as exc:
                    logger.warning("Heartbeat send failed: %s", exc)
                    break
        except asyncio.CancelledError:
            pass
