"""Invisible Playwright driver: operates the muse.ai web client using invisible-playwright
(C++ patched anti-detect Firefox) to bypass Cloudflare Turnstile, anti-bot, and fingerprint checks.
"""

from __future__ import annotations

import asyncio
import base64
import contextlib
import logging
import time
from collections.abc import AsyncIterator, Callable
from dataclasses import dataclass, field
from typing import Any
from urllib.parse import urlsplit

from ..accounts.model import Account
from ..config import Settings
from ..core.prompt import followup_text
from ..errors import (
    UpstreamAuthError,
    UpstreamError,
    UpstreamQuotaError,
    UpstreamRefused,
    UpstreamTimeout,
)
from ..upstream import muse
from .base import (
    ChatRequest,
    DriverCapabilities,
    ImageRequest,
    InputImage,
    MediaResult,
    MuseDriver,
    SessionInfo,
    VideoRequest,
)
from .browser import dom
from .browser.driver import (
    HOT_BUBBLE_LIMIT,
    POLL_INTERVAL,
    STABLE_POLLS_DONE,
    STABLE_POLLS_FORCE,
    is_thread_url,
)

log = logging.getLogger(__name__)


@dataclass
class _PlaywrightTab:
    account_id: str
    context: Any
    page: Any
    hint: str | None = None
    turns: list[tuple[str, str]] = field(default_factory=list)
    image_count: int = 0
    thread_url: str = ""
    quota_note: str = ""
    busy: bool = False
    stale: bool = False
    closed: bool = False

    @property
    def has_state(self) -> bool:
        return bool(self.turns or self.thread_url)


class InvisiblePlaywrightDriver(MuseDriver):
    """Driver using invisible-playwright (anti-detect Firefox) to bypass bot blocking."""

    name = "invisible"
    capabilities = DriverCapabilities(
        chat=True, chat_images=True, image=True, image_edit=True, video=True, renew_session=True
    )

    def __init__(self, settings: Settings) -> None:
        self.settings = settings
        self._ip: Any | None = None
        self._browser: Any | None = None
        self._contexts: dict[str, Any] = {}
        self._tabs: dict[str, list[_PlaywrightTab]] = {}
        self._opening: dict[str, int] = {}
        self._tab_cond = asyncio.Condition()
        self._ctx_lock = asyncio.Lock()

    # ------------------------------------------------------------ lifecycle
    async def startup(self) -> None:
        from invisible_playwright.async_api import InvisiblePlaywright

        proxy_dict = None
        if self.settings.browser_proxy:
            proxy_dict = {"server": self.settings.browser_proxy}

        log.info("launching invisible-playwright anti-detect browser (headless=%s)", self.settings.headless)
        self._ip = InvisiblePlaywright(
            headless=self.settings.headless,
            proxy=proxy_dict,
            humanize=True,
        )
        self._browser = await self._ip.__aenter__()
        log.info("invisible-playwright driver ready")

    async def shutdown(self) -> None:
        for tabs in list(self._tabs.values()):
            for tab in list(tabs):
                await self._close_tab(tab)
        for ctx in list(self._contexts.values()):
            with contextlib.suppress(Exception):
                await ctx.close()
        self._contexts.clear()
        if self._browser:
            with contextlib.suppress(Exception):
                await self._browser.close()
        if self._ip:
            with contextlib.suppress(Exception):
                await self._ip.__aexit__(None, None, None)
        self._browser = None
        self._ip = None
        log.info("invisible-playwright driver stopped")

    async def health(self) -> dict[str, Any]:
        ok = self._browser is not None and self._browser.is_connected()
        tabs = [t for ts in self._tabs.values() for t in ts]
        return {
            "driver": self.name,
            "ok": ok,
            "tabs": len(tabs),
            "busy_tabs": sum(t.busy for t in tabs),
        }

    # ------------------------------------------------------------ tabs
    async def _checkout(
        self, account: Account, prefer: Callable[[_PlaywrightTab], bool] | None = None
    ) -> _PlaywrightTab:
        limit = max(1, self.settings.account_max_concurrency)
        dead: list[_PlaywrightTab] = []
        async with self._tab_cond:
            while True:
                tabs = self._tabs.setdefault(account.id, [])
                for t in [t for t in tabs if t.closed or (not t.busy and (t.stale or t.page.is_closed()))]:
                    tabs.remove(t)
                    dead.append(t)
                idle = [t for t in tabs if not t.busy]
                can_open = len(tabs) + self._opening.get(account.id, 0) < limit
                match = next((t for t in idle if prefer is None or prefer(t)), None)
                if match or (idle and not can_open):
                    tab = match or next((t for t in idle if not t.has_state), idle[0])
                    tab.busy = True
                    break
                if can_open:
                    self._opening[account.id] = self._opening.get(account.id, 0) + 1
                    tab = None
                    break
                await self._tab_cond.wait()
        for t in dead:
            await self._close_tab(t)
        if tab:
            return tab
        try:
            tab = await self._open_tab(account)
        finally:
            async with self._tab_cond:
                self._opening[account.id] -= 1
                if tab:
                    tab.busy = True
                    self._tabs.setdefault(account.id, []).append(tab)
                self._tab_cond.notify_all()
        log.info("opened invisible tab %d for account %s", len(self._tabs[account.id]), account.id)
        return tab

    async def _checkin(self, tab: _PlaywrightTab) -> None:
        tab.busy = False
        await asyncio.shield(self._notify_tabs())
        if tab.stale:
            await self._close_tab(tab)

    async def _notify_tabs(self) -> None:
        async with self._tab_cond:
            self._tab_cond.notify_all()

    async def _context(self, account: Account) -> Any:
        async with self._ctx_lock:
            if ctx := self._contexts.get(account.id):
                return ctx
            assert self._browser, "Browser not initialized"
            ctx = await self._browser.new_context(
                viewport={"width": 1440, "height": 900},
            )
            cookies = []
            floor = time.time() + 3600
            for name, value in account.cookies.items():
                if not value:
                    continue
                exp = account.cookie_expires.get(name) or 0
                cookies.append({
                    "name": name,
                    "value": value,
                    "domain": ".muse.ai",
                    "path": "/",
                    "secure": True,
                    "expires": exp if exp > floor else int(time.time() + 7 * 86400),
                })
            if cookies:
                await ctx.add_cookies(cookies)
            self._contexts[account.id] = ctx
            return ctx

    async def _open_tab(self, account: Account) -> _PlaywrightTab:
        if not account.cookies.get("hatch_sess"):
            raise UpstreamAuthError("account is missing session cookie: hatch_sess")
        ctx = await self._context(account)
        page = await ctx.new_page()
        return _PlaywrightTab(account.id, ctx, page)

    async def _close_tab(self, tab: _PlaywrightTab) -> None:
        if tab.closed:
            return
        tab.closed, tab.busy = True, False
        await asyncio.shield(self._dispose_tab(tab))

    async def _dispose_tab(self, tab: _PlaywrightTab) -> None:
        async with self._tab_cond:
            tabs = self._tabs.get(tab.account_id, [])
            if tab in tabs:
                tabs.remove(tab)
            last = not tabs and not self._opening.get(tab.account_id)
            context = self._contexts.pop(tab.account_id, None) if last else None
            self._tab_cond.notify_all()
        with contextlib.suppress(Exception):
            await tab.page.close()
        if context:
            with contextlib.suppress(Exception):
                await context.close()

    @staticmethod
    def _forget(tab: _PlaywrightTab) -> None:
        tab.hint = None
        tab.turns = []
        tab.image_count = 0
        tab.thread_url = ""

    def _load_hot(self, tab: _PlaywrightTab, account: Account) -> None:
        if tab.turns or tab.thread_url:
            return
        hot = account.meta.get("hot_page")
        if not isinstance(hot, dict):
            return
        turns = []
        for item in hot.get("turns") or []:
            if isinstance(item, (list, tuple)) and len(item) == 2:
                turns.append((str(item[0]), str(item[1])))
        tab.turns = turns
        tab.hint = hot.get("hint") or None
        tab.thread_url = str(hot.get("url") or "")
        try:
            tab.image_count = int(hot.get("image_count") or 0)
        except (TypeError, ValueError):
            tab.image_count = 0

    def _save_hot(self, tab: _PlaywrightTab, account: Account, req: ChatRequest, url: str) -> None:
        if url:
            tab.thread_url = url
        tab.hint = req.conversation_hint
        tab.turns = list(req.turns)
        tab.image_count = len(req.images)
        account.meta["hot_page"] = {
            "hint": req.conversation_hint or "",
            "turns": [list(turn) for turn in req.turns],
            "url": tab.thread_url,
            "image_count": tab.image_count,
        }

    async def _href(self, tab: _PlaywrightTab) -> str:
        try:
            return tab.page.url or ""
        except Exception:
            return ""

    async def _on_thread(self, tab: _PlaywrightTab) -> bool:
        href = await self._href(tab)
        if not is_thread_url(href):
            return False
        if tab.thread_url and urlsplit(href).path.rstrip("/") != urlsplit(tab.thread_url).path.rstrip("/"):
            return False
        try:
            page = await tab.page.evaluate(dom.PAGE_STATE) or {}
            if not page.get("hasInput"):
                return False
            if "Connecting..." in (page.get("head") or ""):
                return False
            state = await self._state(tab)
        except Exception:
            return False
        if state.get("generating"):
            return False
        if any(hint in state.get("tail", "") for hint in dom.STALL_HINTS):
            return False
        return (state.get("agentCount") or 0) < HOT_BUBBLE_LIMIT

    async def _new_thread(self, tab: _PlaywrightTab) -> None:
        self._forget(tab)
        await tab.page.goto(muse.NEW_THREAD_URL, wait_until="domcontentloaded", timeout=self.settings.page_ready_timeout * 1000)
        deadline = time.monotonic() + self.settings.page_ready_timeout
        state: dict = {}
        while time.monotonic() < deadline:
            await asyncio.sleep(0.25)
            with contextlib.suppress(Exception):
                state = await tab.page.evaluate(dom.PAGE_STATE) or {}
                if state.get("ready"):
                    return
        head = (state.get("head") or "").lower()
        if any(h in head for h in dom.LOGIN_HINTS):
            raise UpstreamAuthError("muse.ai redirected to login; cookies are no longer valid")
        raise UpstreamTimeout("muse.ai page did not become ready")

    # ------------------------------------------------------------ input & dom
    async def _attach(self, tab: _PlaywrightTab, images: list[InputImage]) -> None:
        for idx, img in enumerate(images):
            ext = img.mime.split("/")[-1].replace("jpeg", "jpg")
            res = await tab.page.evaluate(
                dom.attach_file(base64.b64encode(img.data).decode(), img.mime, f"input_{idx}.{ext}")
            )
            if not (res or {}).get("ok"):
                raise UpstreamError(f"failed to attach image: {(res or {}).get('err')}")
        if images:
            await asyncio.sleep(1.0)

    async def _send(self, tab: _PlaywrightTab, prompt: str) -> None:
        res = await tab.page.evaluate(dom.fill_input(prompt))
        if not (res or {}).get("ok"):
            raise UpstreamError(f"cannot fill chat input: {(res or {}).get('err')}")
        deadline = time.monotonic() + 5
        while time.monotonic() < deadline:
            if await tab.page.evaluate(dom.CLICK_SEND) == "clicked":
                return
            await asyncio.sleep(0.1)
        await tab.page.keyboard.press("Enter")
        await asyncio.sleep(0.5)
        if not await tab.page.evaluate(dom.INPUT_EMPTY):
            raise UpstreamError("send button not found and Enter did not submit")

    async def _state(self, tab: _PlaywrightTab) -> dict:
        state = await tab.page.evaluate(dom.CHAT_STATE) or {}
        tail = state.get("tail", "")
        low = tail.lower()
        if hit := next((h for h in dom.QUOTA_HINTS if h in low), None):
            i = low.find(hit)
            snippet = " ".join(tail[max(0, i - 120):i + 120].split())
            reply = state.get("lastText", "")
            if not state.get("generating") and hit in reply.lower():
                log.warning("quota hint %r in reply on %s: %s", hit, tab.thread_url or "new thread", snippet)
                raise UpstreamQuotaError(f"muse.ai reports the account is out of quota: …{snippet}…")
            if snippet != tab.quota_note:
                tab.quota_note = snippet
                log.info("ignoring quota hint %r outside the reply on %s: %s", hit, tab.thread_url or "new thread", snippet)
        return state

    async def _prepare(self, tab: _PlaywrightTab, prompt: str, images: list[InputImage]) -> dict:
        await self._new_thread(tab)
        base = await self._state(tab)
        await self._attach(tab, images)
        await self._send(tab, prompt)
        return base

    async def _deny_approvals(self, tab: _PlaywrightTab) -> None:
        denied = await tab.page.evaluate(dom.DENY_APPROVALS) or []
        for what in denied:
            log.info("denied muse.ai %s", what.lower())

    # ------------------------------------------------------------ chat
    async def _live_continuation(self, tab: _PlaywrightTab) -> bool:
        last_user = next((text for role, text in reversed(tab.turns) if role == "user"), "")
        if not last_user:
            return False
        try:
            page = await tab.page.evaluate(dom.PAGE_STATE) or {}
            if not page.get("hasInput"):
                return False
            state = await self._state(tab)
            present = await tab.page.evaluate(
                f"!!(document.body && document.body.innerText.includes({dom.q(last_user[:80])}))"
            )
        except Exception:
            return False
        if state.get("generating"):
            return False
        if any(hint in state.get("tail", "") for hint in dom.STALL_HINTS):
            return False
        if (state.get("agentCount") or 0) >= HOT_BUBBLE_LIMIT:
            return False
        return bool(present)

    async def _open_saved(self, tab: _PlaywrightTab) -> None:
        await tab.page.goto(tab.thread_url, wait_until="domcontentloaded", timeout=self.settings.page_ready_timeout * 1000)
        deadline = time.monotonic() + self.settings.page_ready_timeout
        while time.monotonic() < deadline:
            await asyncio.sleep(0.25)
            if await self._on_thread(tab):
                return
        raise UpstreamTimeout("saved conversation did not become ready")

    async def _begin_chat(self, account: Account, req: ChatRequest):
        hint = req.conversation_hint or ""
        tab = await self._checkout(account, prefer=lambda t: t.has_state and (t.hint or "") == hint)
        prompt, images = req.prompt, req.images
        try:
            if not any(t.has_state for t in self._tabs.get(account.id, []) if t is not tab):
                self._load_hot(tab, account)
            same_owner = (tab.hint or "") == (req.conversation_hint or "")
            follow = followup_text(tab.turns, req.turns) if tab.turns and same_owner else None
            if follow and len(req.images) >= tab.image_count and (
                tab.thread_url or await self._live_continuation(tab)
            ):
                if await self._on_thread(tab) or await self._live_continuation(tab):
                    log.info("reusing hot page %s", tab.thread_url or "open page")
                elif tab.thread_url:
                    log.info("reopening %s", tab.thread_url)
                    await self._open_saved(tab)
                prompt = follow
                images = req.images[tab.image_count:]
            else:
                log.info("opening a new thread")
                saved = account.meta.get("hot_page")
                if isinstance(saved, dict) and (saved.get("hint") or "") == hint:
                    account.meta.pop("hot_page", None)
                await self._new_thread(tab)
            base = await self._state(tab)
            await self._attach(tab, images)
            await self._send(tab, prompt)
            return tab, base
        except Exception as exc:
            await self._close_tab(tab)
            await self._checkin(tab)
            raise UpstreamError(f"browser error: {exc}") from exc

    async def chat_stream(self, account: Account, req: ChatRequest) -> AsyncIterator[str]:
        tab, base = await self._begin_chat(account, req)
        base_count, base_text = base.get("agentCount", 0), base.get("lastText", "")
        started = time.monotonic()
        emitted, last, stable, got_first = "", None, 0, False
        finished = False
        try:
            while True:
                if req.cancel and req.cancel.is_set():
                    return
                elapsed = time.monotonic() - started
                if elapsed > req.timeout:
                    raise UpstreamTimeout("assistant reply timed out")
                if not got_first and elapsed > req.first_token_timeout:
                    raise UpstreamTimeout("no first token from assistant")
                await asyncio.sleep(POLL_INTERVAL)
                st = await self._state(tab)
                if st.get("approval"):
                    await self._deny_approvals(tab)
                    continue
                text = st.get("lastText", "")
                is_new = st.get("agentCount", 0) > base_count or (text and text != base_text)
                if not is_new or not text:
                    if elapsed > 15 and any(h in st.get("tail", "") for h in dom.STALL_HINTS):
                        raise UpstreamTimeout("upstream workspace is stuck connecting")
                    continue
                got_first = True
                if text != last:
                    delta = text[len(emitted):] if text.startswith(emitted) else text
                    if delta:
                        emitted = text
                        yield delta
                    last, stable = text, 0
                else:
                    stable += 1
                    if (not st.get("generating") and stable >= STABLE_POLLS_DONE) or stable >= STABLE_POLLS_FORCE:
                        finished = True
                        return
        except Exception as exc:
            await self._close_tab(tab)
            raise UpstreamError(f"browser error: {exc}") from exc
        finally:
            if finished:
                url = ""
                for _ in range(10):
                    url = await self._href(tab)
                    if is_thread_url(url):
                        break
                    await asyncio.sleep(0.2)
                self._save_hot(tab, account, req, url if is_thread_url(url) else "")
            else:
                self._forget(tab)
            await self._checkin(tab)

    # ------------------------------------------------------------ media
    @staticmethod
    def _media_prompt(prompt: str, kind: str, size: str | None, duration: int | None = None) -> str:
        hints = []
        if size:
            hints.append(f"aspect ratio {size}")
        if duration:
            hints.append(f"{duration} seconds long")
        verb = "Generate an image" if kind == "image" else "Generate a video"
        suffix = f" ({', '.join(hints)})" if hints else ""
        return f"{verb}{suffix}: {prompt}"

    _MEDIA_GRACE = {"video": 90.0, "image": 8.0}
    _OVERRUN_FACTOR = 2.5

    async def _wait_media(
        self, tab: _PlaywrightTab, base: dict, kind: str, timeout: float, on_progress, cancel: asyncio.Event | None
    ) -> dict:
        base_atts = len(base.get("attachments", []))
        base_count = base.get("agentCount", 0)
        grace = self._MEDIA_GRACE.get(kind, 8.0)
        started = time.monotonic()
        text_done_at: float | None = None
        denied = asked_to_attach = busy = False
        while True:
            elapsed = time.monotonic() - started
            if elapsed >= timeout and not (busy and elapsed < timeout * self._OVERRUN_FACTOR):
                break
            if cancel and cancel.is_set():
                raise asyncio.CancelledError
            await asyncio.sleep(0.6)
            st = await self._state(tab)
            if st.get("approval"):
                await self._deny_approvals(tab)
                denied, text_done_at = True, None
                continue
            new_atts = st.get("attachments", [])[base_atts:]
            fresh = [a for a in new_atts if a.get("src") and a.get("kind") == kind]
            if fresh:
                return fresh[-1]
            busy = bool(st.get("generating")) or any(kind in (a.get("tid") or "") for a in new_atts)
            elapsed = time.monotonic() - started
            if on_progress:
                on_progress(min(95, int(elapsed / timeout * 100)))
            if any(kind in (a.get("tid") or "") for a in new_atts):
                text_done_at = None
                continue
            text = st.get("lastText", "")
            if st.get("agentCount", 0) > base_count and text and not st.get("generating"):
                if text_done_at is None:
                    text_done_at = time.monotonic()
                elif time.monotonic() - text_done_at >= grace:
                    if not asked_to_attach:
                        await self._send(
                            tab,
                            f"Please attach the final {kind} file here in the chat"
                            + (" instead of uploading it." if denied else ", as a single attachment."),
                        )
                        asked_to_attach, text_done_at = True, None
                        base_count = st.get("agentCount", 0)
                        continue
                    probe = await tab.page.evaluate(dom.LAST_BUBBLE_MEDIA) or {}
                    url = self._pick_media_url(probe.get("links", []), kind)
                    if url:
                        return {"src": url, "kind": kind}
                    raise UpstreamRefused(f"upstream replied with text only: {text[:200]}")
            else:
                text_done_at = None
        if busy:
            log.warning("%s still generating after %.0fs; giving up", kind, elapsed)
        raise UpstreamTimeout(f"{kind} generation timed out")

    _VIDEO_EXT = (".mp4", ".webm", ".mov", ".m4v")
    _IMAGE_EXT = (".png", ".jpg", ".jpeg", ".webp", ".gif", ".avif")
    _MEDIA_HOSTS = (
        "videodelivery.net",
        "cloudflarestream.com",
        "cloudflarestorage.com",
        "amazonaws.com",
        "storage.googleapis.com",
        "blob.core.windows.net",
        "musecdn",
        "cdn.muse",
        "muse-cdn",
    )
    _MEDIA_PATHS = ("/media/", "/download", "/dl/", "/video", "/attachment", "/files/")

    @classmethod
    def _pick_media_url(cls, links: list[str], kind: str) -> str | None:
        exts = cls._VIDEO_EXT if kind == "video" else cls._IMAGE_EXT

        def score(u: str) -> int:
            low = u.lower()
            path = low.split("?", 1)[0].split("#", 1)[0]
            if u.startswith("blob:") or path.endswith(exts):
                return 3
            if any(h in low for h in cls._MEDIA_HOSTS):
                return 2
            if any(seg in low for seg in cls._MEDIA_PATHS):
                return 1
            return 0

        cands = [u for u in links if u and not u.startswith(("mailto:", "javascript:"))]
        cands.sort(key=score, reverse=True)
        return cands[0] if cands and score(cands[0]) > 0 else None

    async def _download(self, tab: _PlaywrightTab, att: dict, kind: str) -> MediaResult:
        res = await tab.page.evaluate(dom.fetch_as_base64(att["src"]))
        if not (res or {}).get("ok"):
            raise UpstreamError(f"failed to download generated {kind}: {(res or {}).get('err')}")
        mime = res.get("mime") or ("video/mp4" if kind == "video" else "image/png")
        return MediaResult(
            data=base64.b64decode(res["b64"]),
            mime=mime,
            kind=kind,
            width=att.get("w") or None,
            height=att.get("h") or None,
        )

    async def _generate(
        self,
        account: Account,
        prompt: str,
        images: list[InputImage],
        kind: str,
        timeout: float,
        on_progress,
        cancel,
    ) -> MediaResult:
        tab = await self._checkout(account, prefer=lambda t: not t.has_state)
        try:
            base = await self._prepare(tab, prompt, images)
            att = await self._wait_media(tab, base, kind, timeout, on_progress, cancel)
            return await self._download(tab, att, kind)
        except Exception as exc:
            await self._close_tab(tab)
            raise UpstreamError(f"browser error: {exc}") from exc
        finally:
            await self._checkin(tab)

    async def generate_image(self, account: Account, req: ImageRequest) -> list[MediaResult]:
        prompt = self._media_prompt(req.prompt, "image", req.size)
        results = []
        for _ in range(max(1, req.n)):
            r = await self._generate(
                account, prompt, req.reference_images, "image", req.timeout, req.on_progress, req.cancel
            )
            r.revised_prompt = req.prompt
            results.append(r)
        return results

    async def generate_video(self, account: Account, req: VideoRequest) -> MediaResult:
        prompt = self._media_prompt(req.prompt, "video", req.size, req.duration)
        images = [req.first_frame] if req.first_frame else []
        return await self._generate(
            account, prompt, images, "video", req.timeout, req.on_progress, req.cancel
        )

    async def renew_session(self, account: Account) -> SessionInfo:
        info = await muse.renew_session(account.cookies)
        async with self._tab_cond:
            tabs = list(self._tabs.get(account.id, []))
            for tab in tabs:
                tab.stale = True
            idle = [t for t in tabs if not t.busy]
        for tab in idle:
            await self._close_tab(tab)
        return info
