"""Optional LLM signal classification via any OpenAI-compatible provider.

This module is entirely additive and opt-in. When disabled (the default) it
adds *zero* overhead to the forwarder: no client is constructed and no network
calls are made.

Design guardrails for shared hosting (~256 MB RAM, single process):
- Remote API only — no local models, no torch/numpy/transformers. Only the
  Python standard library (urllib) is used; any provider exposing the OpenAI
  ``/chat/completions`` API works (OpenAI, Cerebras, OpenRouter, Groq,
  Together, local Ollama, ...).
- A ``Semaphore(1)`` guarantees at most **one** in-flight LLM call (RPM + RAM).
- A tiny TTL dedup map stops message edits from double-billing the same signal.
- The outbound call is always executed through ``run_in_executor`` from the
  asyncio loop, so a slow provider never blocks message forwarding.
- Hard request timeout (default 12s), small ``max_tokens`` (default 120) and
  input truncation (default 1500 chars) bound both cost and latency.

The engine never raises: any failure is returned as an ``LlmVerdict`` with
``error`` set so the rest of the pipeline is unaffected.
"""

from __future__ import annotations

import asyncio
import hashlib
import hmac
import json
import logging
import re
import threading
import time
import urllib.error
import urllib.request
from dataclasses import dataclass, field
from typing import Any, Optional

logger = logging.getLogger(__name__)

DEFAULT_BASE_URL = "https://api.cerebras.ai/v1"


def _err_detail(exc: Exception) -> str:
    """Human-readable provider error, preferring the JSON body when present."""
    detail = str(exc)[:200]
    if isinstance(exc, urllib.error.HTTPError):
        try:  # provider error bodies carry the real reason (bad key, bad model…)
            detail = exc.read().decode("utf-8", "replace")[:200]
        except Exception:
            pass
        return f"HTTP {exc.code}: {detail}"
    return detail


# ---------------------------------------------------------------------------
# Result type
# ---------------------------------------------------------------------------


@dataclass(slots=True)
class LlmVerdict:
    """Classification result for a single source-channel message."""

    label: str = "other"  # "entry" | "update" | "other"
    confidence: float = 0.0  # 0.0 .. 1.0
    reason: str = ""
    model: str = ""
    latency_ms: int = 0
    error: str = ""

    @property
    def is_signal(self) -> bool:
        return self.label in ("entry", "update") and not self.error

    def as_dict(self) -> dict[str, Any]:
        return {
            "label": self.label,
            "confidence": round(float(self.confidence), 3),
            "reason": self.reason,
            "model": self.model,
            "latency_ms": self.latency_ms,
            "error": self.error,
        }


# ---------------------------------------------------------------------------
# Prompting
# ---------------------------------------------------------------------------

_SYSTEM_PROMPT = (
    "You are a strict classifier for trading-signal Telegram messages. "
    "Decide whether the message is a NEW trade ENTRY, an UPDATE to a previous "
    "entry (e.g. move SL, partial close, cancel pending, breakeven, TP/SL hit, "
    "or a reply that modifies an earlier signal), or OTHER (promo, chat, "
    "noise). Reply with ONLY a compact JSON object, no markdown, no prose: "
    '{"label":"entry|update|other","confidence":0.0-1.0,"reason":"short"}. '
    "Use high confidence only when the intent is unambiguous."
)

_FENCE_RE = re.compile(r"```(?:json)?\s*(.*?)```", re.I | re.S)
_JSON_RE = re.compile(r"\{.*\}", re.S)


def _parse_verdict_text(raw: str) -> tuple[str, float, str]:
    """Extract (label, confidence, reason) from a model's free-form reply."""
    text = (raw or "").strip()
    m = _FENCE_RE.search(text)
    if m:
        text = m.group(1).strip()
    m = _JSON_RE.search(text)
    candidate = m.group(0) if m else text
    try:
        obj = json.loads(candidate)
    except Exception:
        return "other", 0.0, "unparseable model reply"
    label = str(obj.get("label", "other")).strip().lower()
    if label not in ("entry", "update", "other"):
        label = "other"
    try:
        conf = float(obj.get("confidence", 0.0))
    except (TypeError, ValueError):
        conf = 0.0
    conf = max(0.0, min(1.0, conf))
    reason = str(obj.get("reason", "")).strip()[:300]
    return label, conf, reason


# ---------------------------------------------------------------------------
# Engine
# ---------------------------------------------------------------------------


@dataclass
class _Dedup:
    _ttl: float = 60.0
    _max: int = 512
    _seen: dict = field(default_factory=dict)  # key -> ts

    def is_duplicate(self, key: str) -> bool:
        now = time.time()
        # prune opportunistically to keep the map tiny
        if len(self._seen) > self._max:
            self._seen = {
                k: t for k, t in self._seen.items() if now - t < self._ttl
            }
        ts = self._seen.get(key)
        if ts is not None and now - ts < self._ttl:
            return True
        self._seen[key] = now
        return False


class LlmEngine:
    """Holds live LLM config and the classify+post flow."""

    USER_AGENT = "ssfx-signal-forwarder/llm-1.0"

    def __init__(self, db):
        self.db = db
        self._lock = threading.Lock()
        self._semaphore = threading.Semaphore(1)
        self._dedup = _Dedup()

        # Live configuration (overwritten by reload_from_db / update_config).
        self.enabled: bool = False
        self.provider: str = "openai-compatible"
        self.model: str = "llama3.1-8b"
        self.base_url: str = DEFAULT_BASE_URL
        self.api_key: str = ""
        self.min_confidence: float = 0.80
        self.webhook_url: str = ""
        self.webhook_secret: str = ""
        self.timeout: float = 12.0
        self.max_tokens: int = 120
        self.max_input_chars: int = 1500

    # ── Config ────────────────────────────────────────────────────────────

    @staticmethod
    def _bool(val: Any) -> bool:
        if isinstance(val, bool):
            return val
        return str(val or "").strip().lower() in ("1", "true", "yes", "on")

    def reload_from_db(self) -> None:
        """Load global LLM config from config_vars (called at startup / on save).

        DB values win; when a DB value is absent, the corresponding ``LLM_*``
        env var (see config.py) is used as a first-boot seed so a fresh deploy
        can enable the feature without touching the UI.
        """
        import config as _cfg  # lazy: avoids a circular import at module load

        db = self.db

        def _db_or_env(key: str, env_val: str) -> str:
            v = db.get_config(key)
            return v if v not in (None, "") else env_val

        enabled = self._bool(_db_or_env("LLM_ENABLED", "1" if _cfg.LLM_ENABLED else ""))
        provider = (_db_or_env("LLM_PROVIDER", _cfg.LLM_PROVIDER) or "openai-compatible").strip().lower()
        model = (_db_or_env("LLM_MODEL", _cfg.LLM_MODEL) or "llama3.1-8b").strip()
        base_url = (_db_or_env("LLM_BASE_URL", _cfg.LLM_BASE_URL) or DEFAULT_BASE_URL).strip()
        api_key = db.get_secret_config("LLM_API_KEY") or _cfg.LLM_API_KEY or ""
        webhook_url = (_db_or_env("LLM_WEBHOOK_URL", _cfg.LLM_WEBHOOK_URL) or "").strip()
        webhook_secret = db.get_secret_config("LLM_WEBHOOK_SECRET") or _cfg.LLM_WEBHOOK_SECRET or ""
        try:
            min_confidence = float(_db_or_env("LLM_MIN_CONFIDENCE", str(_cfg.LLM_MIN_CONFIDENCE)))
        except (TypeError, ValueError):
            min_confidence = 0.80
        try:
            timeout = float(_db_or_env("LLM_TIMEOUT", str(_cfg.LLM_TIMEOUT)))
        except (TypeError, ValueError):
            timeout = 12.0
        try:
            max_tokens = int(_db_or_env("LLM_MAX_TOKENS", str(_cfg.LLM_MAX_TOKENS)))
        except (TypeError, ValueError):
            max_tokens = 120
        self.update_config(
            enabled=enabled,
            provider=provider,
            model=model,
            base_url=base_url,
            api_key=api_key,
            min_confidence=min_confidence,
            webhook_url=webhook_url,
            webhook_secret=webhook_secret,
            timeout=timeout,
            max_tokens=max_tokens,
        )

    def update_config(self, **kw: Any) -> None:
        with self._lock:
            for key in (
                "enabled",
                "provider",
                "model",
                "base_url",
                "api_key",
                "min_confidence",
                "webhook_url",
                "webhook_secret",
                "timeout",
                "max_tokens",
                "max_input_chars",
            ):
                if key in kw and kw[key] is not None:
                    setattr(self, key, kw[key])
            self.provider = (self.provider or "openai-compatible").strip().lower()
            self.model = (self.model or "").strip()
            self.base_url = (self.base_url or DEFAULT_BASE_URL).strip().rstrip("/")
            self.webhook_url = (self.webhook_url or "").strip()
            self.min_confidence = max(0.0, min(1.0, float(self.min_confidence)))
            self.max_tokens = max(16, min(512, int(self.max_tokens)))
            self.timeout = max(3.0, min(60.0, float(self.timeout)))

    @property
    def client_available(self) -> bool:
        return True  # stdlib HTTP client — nothing extra to install

    @property
    def available(self) -> bool:
        return bool(self.enabled and self.api_key and self.model and self.base_url)

    def status(self) -> dict[str, Any]:
        return {
            "enabled": bool(self.enabled),
            "provider": self.provider or "openai-compatible",
            "model": self.model,
            "base_url": self.base_url or DEFAULT_BASE_URL,
            "api_key_set": bool(self.api_key),
            "client_available": self.client_available,
            "available": self.available,
            "min_confidence": self.min_confidence,
            "webhook_url_set": bool(self.webhook_url),
            "timeout": self.timeout,
            "max_tokens": self.max_tokens,
        }

    # ── Classification (sync, time-boxed, semaphore-guarded) ──────────────

    def classify(
        self,
        text: str,
        reply_to_message_id: Optional[int] = None,
        source_title: str = "",
    ) -> LlmVerdict:
        if not self.available:
            return LlmVerdict(error="llm_disabled")
        if not (text or "").strip():
            return LlmVerdict(error="empty")

        snippet = text.strip()
        if len(snippet) > self.max_input_chars:
            snippet = snippet[: self.max_input_chars]

        user_bits = []
        if source_title:
            user_bits.append(f"Channel: {source_title}")
        if reply_to_message_id:
            user_bits.append(
                f"This message is a reply to message id {reply_to_message_id} "
                "(likely an update to that earlier signal)."
            )
        user_bits.append("Message:\n" + snippet)
        user_content = "\n".join(user_bits)

        # OpenAI-compatible endpoint: POST {base_url}/chat/completions.
        url = f"{self.base_url.rstrip('/')}/chat/completions"
        body = json.dumps(
            {
                "model": self.model,
                "messages": [
                    {"role": "system", "content": _SYSTEM_PROMPT},
                    {"role": "user", "content": user_content},
                ],
                "max_tokens": self.max_tokens,
                "temperature": 0.0,
            }
        ).encode("utf-8")
        headers = {
            "Content-Type": "application/json",
            "Authorization": f"Bearer {self.api_key}",
            "User-Agent": self.USER_AGENT,
        }

        # Only one LLM call in flight; if busy, skip rather than queue.
        acquired = self._semaphore.acquire(blocking=False)
        if not acquired:
            return LlmVerdict(model=self.model, error="busy")
        start = time.time()
        try:
            req = urllib.request.Request(url, data=body, headers=headers, method="POST")
            with urllib.request.urlopen(req, timeout=self.timeout) as resp:
                resp_data = resp.read()
            latency_ms = int((time.time() - start) * 1000)
            raw = ""
            try:
                obj = json.loads(resp_data)
                raw = obj["choices"][0]["message"]["content"] or ""
            except Exception:
                raw = ""
            label, conf, reason = _parse_verdict_text(raw)
            return LlmVerdict(
                label=label,
                confidence=conf,
                reason=reason,
                model=self.model,
                latency_ms=latency_ms,
            )
        except Exception as exc:  # noqa: BLE001 - never let LLM break forwarding
            latency_ms = int((time.time() - start) * 1000)
            detail = _err_detail(exc)
            logger.warning("LLM classify failed (%sms): %s", latency_ms, detail)
            return LlmVerdict(
                model=self.model, latency_ms=latency_ms, error=detail
            )
        finally:
            self._semaphore.release()

    # ── Interactive diagnostics (web UI: model autocomplete + test button) ──
    # Not semaphore-guarded: these are explicit web actions, not forwarder hot
    # paths — a slow click must not block classification, but does no harm.

    def _effective(self, base_url: str, api_key: str) -> tuple[str, str]:
        """Resolve overrides (empty = keep live/saved config)."""
        url = (base_url or self.base_url or DEFAULT_BASE_URL).strip().rstrip("/")
        key = (api_key or self.api_key or "").strip()
        return url, key

    def list_models(self, base_url: str = "", api_key: str = "") -> dict[str, Any]:
        """GET {base_url}/models → sorted model ids. Never raises."""
        url, key = self._effective(base_url, api_key)
        if not key:
            return {"ok": False, "models": [], "error": "no API key configured"}
        req = urllib.request.Request(
            url + "/models",
            headers={"Authorization": f"Bearer {key}", "User-Agent": self.USER_AGENT},
        )
        try:
            with urllib.request.urlopen(req, timeout=min(self.timeout, 10.0)) as resp:
                data = json.loads(resp.read())
            ids = sorted(
                {str(m.get("id", "")) for m in data.get("data", []) if m.get("id")}
            )
            return {"ok": True, "models": ids, "error": ""}
        except Exception as exc:
            return {"ok": False, "models": [], "error": _err_detail(exc)}

    def test_connection(
        self, base_url: str = "", api_key: str = "", model: str = ""
    ) -> dict[str, Any]:
        """Ping {base_url}/chat/completions with a tiny prompt. Never raises."""
        url, key = self._effective(base_url, api_key)
        model = (model or self.model or "").strip()
        if not key:
            return {"ok": False, "latency_ms": 0, "model": model, "error": "no API key configured"}
        if not model:
            return {"ok": False, "latency_ms": 0, "model": model, "error": "no model set"}
        body = json.dumps(
            {
                "model": model,
                "messages": [{"role": "user", "content": "Reply with exactly: pong"}],
                "max_tokens": 5,
                "temperature": 0.0,
            }
        ).encode("utf-8")
        headers = {
            "Content-Type": "application/json",
            "Authorization": f"Bearer {key}",
            "User-Agent": self.USER_AGENT,
        }
        start = time.time()
        try:
            req = urllib.request.Request(
                url + "/chat/completions", data=body, headers=headers, method="POST"
            )
            with urllib.request.urlopen(req, timeout=self.timeout) as resp:
                resp.read()
            return {
                "ok": True,
                "latency_ms": int((time.time() - start) * 1000),
                "model": model,
                "error": "",
            }
        except Exception as exc:
            return {
                "ok": False,
                "latency_ms": int((time.time() - start) * 1000),
                "model": model,
                "error": _err_detail(exc),
            }

    # ── Async orchestration used by the forwarder ─────────────────────────

    async def classify_and_post(
        self,
        core,
        source_chat_id: int,
        message_id: int,
        text: str,
        date: Optional[int] = None,
        reply_to_message_id: Optional[int] = None,
        source_title: str = "",
        symbol: Optional[str] = None,
    ) -> LlmVerdict:
        """Classify a message, log latency, and optionally POST the result."""
        if not self.available:
            return LlmVerdict(error="llm_disabled")

        dedup_key = f"{source_chat_id}:{message_id}"
        if self._dedup.is_duplicate(dedup_key):
            return LlmVerdict(model=self.model, error="dedup")

        loop = asyncio.get_running_loop()
        verdict: LlmVerdict = await loop.run_in_executor(
            None, self.classify, text, reply_to_message_id, source_title
        )

        logger.info(
            "LLM classify src=%s msg=%s label=%s conf=%.2f latency=%sms%s",
            source_chat_id,
            message_id,
            verdict.label,
            verdict.confidence,
            verdict.latency_ms,
            f" err={verdict.error}" if verdict.error else "",
        )

        # Price snapshot: fast cached tick only (no waiting) — zero added latency.
        price = self._price_snapshot(core, symbol)

        if (
            self.webhook_url
            and verdict.is_signal
            and verdict.confidence >= self.min_confidence
        ):
            await self._post_llm_webhook(
                source_chat_id=source_chat_id,
                message_id=message_id,
                text=text,
                date=date,
                reply_to_message_id=reply_to_message_id,
                source_title=source_title,
                verdict=verdict,
                price=price,
            )
        return verdict

    @staticmethod
    def _price_snapshot(core, symbol: Optional[str]) -> dict[str, Any]:
        snap: dict[str, Any] = {"price_available": False}
        client = getattr(core, "_ctrader_client", None)
        if not client:
            return snap
        sym = (symbol or getattr(client, "symbol", "") or "").upper()
        if not sym:
            # fall back to first configured symbol if any
            sym = ""
        try:
            tick = client.get_tick(sym) if sym else None
        except Exception:
            tick = None
        if tick:
            snap.update(
                {
                    "price_available": True,
                    "symbol": sym,
                    "bid": tick.bid,
                    "ask": tick.ask,
                    "spread": round(tick.ask - tick.bid, 5),
                    "tick_timestamp_ms": tick.timestamp_ms,
                }
            )
        return snap

    async def _post_llm_webhook(
        self,
        source_chat_id: int,
        message_id: int,
        text: str,
        date: Optional[int],
        reply_to_message_id: Optional[int],
        source_title: str,
        verdict: LlmVerdict,
        price: dict[str, Any],
    ) -> None:
        payload = json.dumps(
            {
                "event": "llm_signal",
                "received_at_ms": int(time.time() * 1000),
                "snapshot": {
                    "source_chat_id": source_chat_id,
                    "source_message_id": message_id,
                    "source_title": source_title,
                    "signal_text": text,
                    "signal_date": date,
                    "reply_to_message_id": reply_to_message_id,
                    "price": price,
                },
                "llm": verdict.as_dict(),
            },
            default=str,
        ).encode("utf-8")

        headers = {
            "Content-Type": "application/json",
            "User-Agent": self.USER_AGENT,
        }
        if self.webhook_secret:
            sig = hmac.new(
                self.webhook_secret.encode("utf-8"), payload, hashlib.sha256
            ).hexdigest()
            headers["X-Signal-Signature"] = f"sha256={sig}"

        url = self.webhook_url

        def _request():
            req = urllib.request.Request(
                url, data=payload, headers=headers, method="POST"
            )
            with urllib.request.urlopen(req, timeout=15) as resp:
                return resp.status

        try:
            loop = asyncio.get_running_loop()
            status = await loop.run_in_executor(None, _request)
            logger.info("LLM webhook returned HTTP %s (msg=%s)", status, message_id)
        except Exception as exc:
            logger.warning("LLM webhook failed (msg=%s): %s", message_id, exc)


if __name__ == "__main__":
    import sys

    assert _parse_verdict_text(
        '{"label":"entry","confidence":0.9,"reason":"ok"}'
    ) == ("entry", 0.9, "ok")
    assert _parse_verdict_text(
        '```json\n{"label":"update","confidence":1.0}\n```'
    )[0] == "update"
    assert _parse_verdict_text("just some text")[2] == "unparseable model reply"
    label, conf, _ = _parse_verdict_text('{"label":"weird","confidence":5}')
    assert (label, conf) == ("other", 1.0)  # clamped to [0,1], bad label -> other
    assert DEFAULT_BASE_URL.rstrip("/") + "/chat/completions" == (
        "https://api.cerebras.ai/v1/chat/completions"
    )
    print("llm_classifier self-check OK")
    sys.exit(0)
