from app.services.content.browser import get_browser
import logging
import json
from urllib.parse import urljoin, urlparse
import re
from typing import List, Dict, Set, Any, Optional
from bs4 import BeautifulSoup
import asyncio

from app.core.config import settings
from app.core.llm_client import client as openai_client
from app.core.prompts import PRODUCT_LINK_CLASSIFICATION_PROMPT
from app.services.infra.database import insert_llm_usage

logger = logging.getLogger(__name__)

def extract_jsonld_blocks(html: str) -> list:
    """Every JSON-LD object on a page, with @graph wrappers flattened.

    Must run before the crawler decomposes <script> tags, which would otherwise
    delete exactly the data this reads.
    """
    blocks = []
    soup = BeautifulSoup(html, "html.parser")
    for tag in soup.find_all("script", attrs={"type": "application/ld+json"}):
        raw = tag.string or tag.get_text() or ""
        try:
            parsed = json.loads(raw)
        except (json.JSONDecodeError, TypeError):
            # One malformed block is common and must not lose the others.
            logger.debug("Skipping malformed JSON-LD block")
            continue

        candidates = parsed if isinstance(parsed, list) else [parsed]
        for item in candidates:
            if not isinstance(item, dict):
                continue
            if isinstance(item.get("@graph"), list):
                blocks.extend(g for g in item["@graph"] if isinstance(g, dict))
            else:
                blocks.append(item)
    return blocks

def is_subpath(base_url: str, target_url: str) -> bool:
    """
    Checks if target_url is a sub-path of base_url.
    """
    try:
        base = urlparse(base_url)
        target = urlparse(target_url)

        if base.netloc != target.netloc:
            return False

        # Standardize paths to have trailing slashes for prefix matching
        b_path = base.path if base.path.endswith('/') else base.path + '/'
        t_path = target.path if target.path.endswith('/') else target.path + '/'

        # If base is root, anything on the same domain is a subpath
        if base.path in ['', '/']:
            return True

        return t_path.startswith(b_path)
    except Exception:
        return False

_PRODUCT_URL_PATTERN = re.compile(r"/(products?|item|p)/[^/?#]+", re.IGNORECASE)


def _looks_like_product_url(url: str) -> bool:
    """Free heuristic for the common case. Returns False for anything that
    doesn't match a known pattern -- callers fall back to the LLM classifier
    for sites this misses, they don't treat False as "not a product"."""
    return bool(_PRODUCT_URL_PATTERN.search(urlparse(url).path))


def _dedup_key(url: str) -> str:
    """The identity a URL is deduped on while crawling -- deliberately looser
    than the URL itself.

    The fragment is always dropped: it never reaches the server, so two URLs
    differing only by fragment (.../shirt and .../shirt#main-content) are
    always the same page. The query string is dropped too, but only for
    product-detail URLs -- a storefront's "related products" widget appends a
    unique tracking query string (pr_rec_id, pr_ref_pid, ...) to every link to
    a product, which without this would make the same product page look like
    a new one on every referring page. Non-product URLs keep their query
    string, since this function is shared with the general site-scrape path
    where a query string can mean genuinely different content (pagination).
    """
    parsed = urlparse(url)
    query = "" if _looks_like_product_url(url) else parsed.query
    return parsed._replace(query=query, fragment="").geturl()


async def _classify_product_links(candidates: list, tenant_id: str) -> set:
    """One LLM call, run in a thread since the OpenAI client here is sync
    and this must not block the event loop mid-crawl. Never raises -- a
    failed classification just means the crawl falls back to insertion
    order, which is today's behavior."""
    if not candidates:
        return set()

    numbered = "\n".join(
        f"{i}. {url} -- \"{text}\"" for i, (url, text) in enumerate(candidates)
    )
    prompt = PRODUCT_LINK_CLASSIFICATION_PROMPT.format(links=numbered)

    def _call():
        return openai_client.chat.completions.create(
            model=settings.LLM_MODEL,
            messages=[{"role": "user", "content": prompt}],
            max_completion_tokens=500,
            timeout=30.0,
            response_format={"type": "json_object"},
        )

    try:
        response = await asyncio.to_thread(_call)
        payload = json.loads(response.choices[0].message.content)
    except Exception as ex:
        logger.warning(f"Product-link classification failed: {ex}")
        return set()

    try:
        usage = getattr(response, "usage", None)
        if usage:
            insert_llm_usage(tenant_id, "Product Link Classification", settings.LLM_MODEL,
                             usage.prompt_tokens, usage.completion_tokens, usage.total_tokens)
    except Exception as ex:
        logger.warning(f"Could not log product-link classification usage: {ex}")

    indices = payload.get("product_indices") if isinstance(payload, dict) else None
    indices = indices or []
    return {candidates[i][0] for i in indices if isinstance(i, int) and 0 <= i < len(candidates)}

async def extract_text_from_url(url: str, max_depth: int = 1, max_pages: int = 20,
                                  visited: Set[str] = None, tenant_id: Optional[str] = None,
                                  seed_urls: list = None, progress=None) -> Dict[str, Any]:
    """
    Crawls the website starting from URL up to max_depth and max_pages using a headless browser.
    Returns a dictionary containing aggregated text, metadata, and color palette.

    Visits product-shaped URLs before plain nav/collection links -- a strict
    depth-by-depth BFS can exhaust the whole page budget on a site's nav
    surface before ever reaching a product page (see the crawl-prioritization
    spec). A regex heuristic (_looks_like_product_url) handles the common
    case for free; if it finds nothing after CRAWL_PRODUCT_FALLBACK_AFTER_PAGES
    pages, one LLM call (_classify_product_links) re-ranks the frontier for
    sites with non-standard URL conventions.

    seed_urls, when given, replaces the single-URL frontier with these --
    e.g. a sitemap's already-known product URLs.

    progress, if given, is called as progress(percent, step) after every
    batch -- best-effort, never raises.
    """
    if visited is None:
        visited = set()

    base_domain = urlparse(url).netloc
    results = {
        "text": "",
        "pages": [],
        "metadata": [],
        "colors": set(),
        "images": {},
        "urls_visited": []
    }

    browser = await get_browser()

    # Each frontier entry: (candidate_url, depth). product_priority holds
    # URLs known (by heuristic or LLM verdict) to be product pages -- sorted
    # to the front of every batch regardless of discovery order.
    frontier = [(u, 0) for u in seed_urls] if seed_urls else [(url, 0)]
    # Mirrors every URL ever placed on the frontier (independent of visited),
    # so newly discovered links can be deduped with an O(1) set lookup instead
    # of rescanning the whole frontier list on every link found. Keyed by
    # _dedup_key(), not the raw URL -- see that function for why.
    queued = {_dedup_key(u) for u, _ in frontier}
    # Same key space as queued, populated as pages are actually crawled --
    # separate from `visited` (raw URLs, used for the urls_visited report and
    # the max_pages budget count) so a fragment/tracking-param variant of an
    # already-visited product is recognised as a repeat, not a new page.
    visited_keys = set()
    product_priority = set()
    fallback_fired = False
    any_product_pattern_seen = False

    async def crawl_one(current_url, depth):
        visited.add(current_url)
        visited_keys.add(_dedup_key(current_url))
        logger.info(f"Crawling URL (Headless): {current_url} (depth {depth}, count {len(visited)})")

        new_links = []
        try:
            content, links, pw_meta = await browser.get_rendered_content(current_url)

            if not content:
                logger.warning(f"No content rendered for {current_url}")
                return new_links

            soup = BeautifulSoup(content, 'html.parser')
            results["urls_visited"].append(current_url)

            page_meta = {
                "url": current_url,
                "title": pw_meta.get("title") or (soup.title.string if soup.title else ""),
                "description": ""
            }
            desc_tag = soup.find("meta", attrs={"name": "description"})
            if desc_tag:
                page_meta["description"] = desc_tag.get("content", "")
            results["metadata"].append(page_meta)

            styles = soup.find_all("style")
            for style in styles:
                hex_colors = re.findall(r'#[0-9a-fA-F]{3,6}', style.string or "")
                results["colors"].update(hex_colors)
                rgb_colors = re.findall(r'rgb\(.*?\)', style.string or "")
                results["colors"].update(rgb_colors)

            page_jsonld = extract_jsonld_blocks(content)

            for element in soup(["script", "style", "nav", "footer", "header", "aside"]):
                element.decompose()

            page_images = {}
            img_tags = soup.find_all("img")
            junk_keywords = ["facebook", "twitter", "linkedin", "instagram", "tiktok", "social", "icon", "arrow", "chevron"]

            for img in img_tags:
                img_url = img.get("src")
                if img_url:
                    full_img_url = urljoin(current_url, img_url)
                    url_lower = full_img_url.lower()
                    if any(key in url_lower for key in junk_keywords):
                        continue

                    label = img.get("alt") or img.get("title")
                    if not label:
                        filename = full_img_url.split("/")[-1].split("?")[0]
                        label = filename.rsplit(".", 1)[0].replace("-", " ").replace("_", " ").capitalize()

                    label = label.strip()[:100]

                    if full_img_url not in results["images"]:
                        results["images"][full_img_url] = label
                    page_images[full_img_url] = label

            raw_text = soup.get_text()
            lines = (line.strip() for line in raw_text.splitlines())
            chunks = (phrase.strip() for line in lines for phrase in line.split("  "))
            cleaned_text = '\n'.join(chunk for chunk in chunks if chunk)

            results["pages"].append({
                "url": current_url,
                "text": cleaned_text,
                "jsonld": page_jsonld,
                "html": content,
                "title": pw_meta.get("title") or (soup.title.string if soup.title else ""),
                "images": [{"url": u, "label": l} for u, l in list(page_images.items())[:10]]
            })

            results["text"] += f"\n--- Content from {current_url} ---\n{cleaned_text}\n"

            if depth < max_depth:
                for link in links:
                    next_url = link["href"]
                    next_text = link.get("text", "")
                    if is_subpath(url, next_url) and _dedup_key(next_url) not in visited_keys:
                        if not any(next_url.lower().endswith(ext) for ext in ['.pdf', '.jpg', '.png', '.zip', '.gif']):
                            new_links.append((next_url, next_text, depth + 1))

        except Exception as e:
            logger.error(f"Failed to crawl {current_url} with Playwright: {e}", exc_info=True)

        return new_links

    # (url, text) pairs discovered but not yet queued -- used only to build
    # the LLM fallback's candidate list; the frontier itself only needs urls.
    pending_text_by_url = {}

    while frontier and len(visited) < max_pages:
        frontier.sort(key=lambda entry: 0 if entry[0] in product_priority or _looks_like_product_url(entry[0]) else 1)

        batch = []
        while (frontier and len(batch) < settings.CRAWL_BATCH_SIZE
               and len(visited) + len(batch) < max_pages):
            candidate_url, depth = frontier.pop(0)
            candidate_key = _dedup_key(candidate_url)
            if candidate_key in visited_keys or candidate_key in [_dedup_key(b[0]) for b in batch]:
                continue
            batch.append((candidate_url, depth))

        if not batch:
            break

        results_batch = await asyncio.gather(*[crawl_one(u, d) for u, d in batch])

        if progress is not None:
            try:
                pct = min(99, int(len(visited) * 100 / max_pages)) if max_pages else 0
                progress(pct, f"crawled {len(visited)} of {max_pages} pages")
            except Exception:
                logger.warning("Crawl progress callback failed", exc_info=True)

        for new_links in results_batch:
            for next_url, next_text, next_depth in new_links:
                next_key = _dedup_key(next_url)
                if next_key not in visited_keys and next_key not in queued:
                    frontier.append((next_url, next_depth))
                    queued.add(next_key)
                    pending_text_by_url[next_url] = next_text
                    if _looks_like_product_url(next_url):
                        any_product_pattern_seen = True

        if (not fallback_fired and not any_product_pattern_seen
                and len(visited) >= settings.CRAWL_PRODUCT_FALLBACK_AFTER_PAGES
                and frontier):
            fallback_fired = True
            candidates = [(f_url, pending_text_by_url.get(f_url, "")) for f_url, _ in frontier[:100]]
            try:
                verdict = await _classify_product_links(candidates, tenant_id)
            except Exception as ex:
                logger.warning(f"Product-link classification fallback failed: {ex}")
                verdict = set()
            product_priority |= verdict

    results["colors"] = sorted(list(results["colors"]))
    results["images"] = [{"url": u, "label": l} for u, l in list(results["images"].items())[:20]]

    return results
