from __future__ import annotations

import hashlib
import html
import json
import re
from dataclasses import dataclass
from datetime import datetime, timezone, timedelta
from decimal import Decimal, InvalidOperation
from sqlite3 import Connection
from typing import Any, Callable
from urllib.parse import quote_plus, urlencode, urljoin
from urllib.request import Request, urlopen

from jarvis_finance.services.budget_common import new_id, row_to_dict
from jarvis_finance.services.grocery_optimizer import _items_for_receipt, upsert_product_match, known_mappings_for_item

SAFE_MATCH_FLAGS = {"exact_match", "close_match", "cheaper_private_label"}
UNSAFE_FLAGS = {"needs_review", "no_price", "no_unit_price", "different_pack_size", "different_quality", "organic_mismatch"}


def _now() -> str:
    return datetime.now(timezone.utc).isoformat()


def _decimal_text(value: Any) -> str | None:
    if value is None:
        return None
    text = str(value).replace("CHF", "").replace("'", "").replace("’", "").strip().replace(",", ".")
    match = re.search(r"-?\d+(?:\.\d+)?", text)
    if not match:
        return None
    try:
        return f"{Decimal(match.group(0)).quantize(Decimal('0.01')):.2f}"
    except InvalidOperation:
        return None


def _money(value: Any) -> Decimal:
    text = _decimal_text(value)
    return Decimal(text or "0.00")


def _abs_url(base: str, url: str) -> str:
    return urljoin(base, html.unescape(url or ""))


class GroceryProviderRateLimitError(RuntimeError):
    pass


@dataclass
class GroceryProviderResult:
    retailer: str
    product_name: str
    product_url: str
    price_text: str | None = None
    unit_price_text: str | None = None
    brand: str | None = None
    image_url: str | None = None
    package_size: str | None = None
    unit: str | None = None
    currency: str = "CHF"
    availability_status: str = "unknown"
    promotion_text: str | None = None
    source: str = "web_fetch"
    confidence: str = "0.60"
    quality_flags: list[str] | None = None
    status: str = "suggested"
    fetched_at: str | None = None
    raw_result: dict[str, Any] | None = None

    def as_dict(self) -> dict[str, Any]:
        price_decimal = _decimal_text(self.price_text)
        unit_decimal = _decimal_text(self.unit_price_text)
        flags = list(self.quality_flags or [])
        if not price_decimal and "no_price" not in flags:
            flags.append("no_price")
        if not unit_decimal and "no_unit_price" not in flags:
            flags.append("no_unit_price")
        if not flags:
            flags = ["needs_review"]
        return {
            "retailer": self.retailer,
            "product_name": self.product_name,
            "brand": self.brand,
            "product_url": self.product_url,
            "image_url": self.image_url,
            "price_text": price_decimal or self.price_text,
            "price_decimal_text": price_decimal,
            "currency": self.currency,
            "package_size": self.package_size,
            "unit": self.unit,
            "unit_price_text": self.unit_price_text,
            "unit_price_decimal_text": unit_decimal,
            "availability_status": self.availability_status,
            "promotion_text": self.promotion_text,
            "fetched_at": self.fetched_at or _now(),
            "source": self.source,
            "confidence": self.confidence,
            "quality_flags": flags,
            "status": self.status if self.status else ("suggested" if set(flags) & SAFE_MATCH_FLAGS else "needs_review"),
            "raw_result": self.raw_result or {},
        }


class GroceryProductProvider:
    retailer = "generic"
    base_url = "https://example.invalid/"
    future_status: str | None = None

    def __init__(self, fetcher: Callable[[str], str] | None = None, max_requests: int = 5):
        self.fetcher = fetcher or self._default_fetcher
        self.max_requests = max_requests
        self.request_count = 0

    def search_url(self, query: str, locale: str = "de-CH") -> str:
        return self.base_url + "search?q=" + quote_plus(query)

    def _default_fetcher(self, url: str) -> str:
        if self.future_status:
            raise NotImplementedError(self.future_status)
        req = Request(url, headers={"User-Agent": "JarvisFinanceBot/1.0 low-volume product price lookup"})
        with urlopen(req, timeout=8) as response:  # noqa: S310 - explicit user-triggered low-volume public lookup
            if response.status == 429:
                raise GroceryProviderRateLimitError("rate_limited")
            return response.read().decode("utf-8", errors="ignore")

    def _fetch(self, url: str) -> str:
        self.request_count += 1
        if self.request_count > self.max_requests:
            raise GroceryProviderRateLimitError("rate_limited")
        return self.fetcher(url)

    def normalize_result(self, raw_result: dict[str, Any]) -> GroceryProviderResult:
        return GroceryProviderResult(retailer=self.retailer, **raw_result)

    def parse_search_results(self, html_text: str) -> list[dict[str, Any]]:
        chunks = re.findall(r"<article[^>]*data-product[^>]*>(.*?)</article>", html_text, flags=re.I | re.S) or [html_text]
        results: list[dict[str, Any]] = []
        for chunk in chunks[:5]:
            link = re.search(r"<a[^>]+href=[\"']([^\"']+)[\"'][^>]*>(.*?)</a>", chunk, flags=re.I | re.S)
            price = re.search(r"class=[\"'][^\"']*price[^\"']*[\"'][^>]*>(.*?)</", chunk, flags=re.I | re.S)
            unit_price = re.search(r"class=[\"'][^\"']*unit-price[^\"']*[\"'][^>]*>(.*?)</", chunk, flags=re.I | re.S)
            brand = re.search(r"class=[\"'][^\"']*brand[^\"']*[\"'][^>]*>(.*?)</", chunk, flags=re.I | re.S)
            if not link:
                continue
            name = re.sub(r"<[^>]+>", "", link.group(2)).strip()
            results.append({
                "product_name": html.unescape(name),
                "product_url": _abs_url(self.base_url, link.group(1)),
                "price_text": re.sub(r"<[^>]+>", "", price.group(1)).strip() if price else None,
                "unit_price_text": re.sub(r"<[^>]+>", "", unit_price.group(1)).strip() if unit_price else None,
                "brand": re.sub(r"<[^>]+>", "", brand.group(1)).strip() if brand else None,
                "raw_result": {"source": "html_article"},
            })
        return results

    def _cache_key(self, query: str) -> str:
        return f"provider-search://{self.retailer}/{quote_plus(query.lower().strip())}"

    def get_cached_result(self, conn: Connection, query: str, max_results: int = 5, max_age_seconds: int = 86400) -> list[dict[str, Any]]:
        rows = conn.execute("SELECT * FROM grocery_product_details_cache WHERE retailer=? AND source_hash=? AND cache_status='cached' ORDER BY fetched_at DESC LIMIT ?", (self.retailer, self._cache_key(query), max_results)).fetchall()
        output = []
        cutoff = datetime.now(timezone.utc) - timedelta(seconds=max_age_seconds)
        for row in rows:
            d = row_to_dict(row)
            try:
                fetched_dt = datetime.fromisoformat(str(d.get("fetched_at")))
                if fetched_dt.tzinfo is None:
                    fetched_dt = fetched_dt.replace(tzinfo=timezone.utc)
                if fetched_dt < cutoff:
                    continue
            except Exception:
                continue
            output.append({"retailer": d["retailer"], "product_name": d.get("product_name"), "brand": d.get("brand"), "product_url": d.get("product_url"), "image_url": d.get("image_url"), "price_text": d.get("price_text"), "price_decimal_text": d.get("price_decimal_text"), "currency": d.get("currency") or "CHF", "package_size": d.get("package_size"), "unit": d.get("unit"), "unit_price_text": d.get("unit_price_text"), "unit_price_decimal_text": d.get("unit_price_decimal_text"), "availability_status": d.get("availability_status"), "promotion_text": d.get("promotion_text"), "fetched_at": d.get("fetched_at"), "source": "cache", "confidence": d.get("confidence") or "0", "quality_flags": json.loads(d.get("quality_flags_json") or "[]"), "status": "suggested"})
        return output

    def store_cache(self, conn: Connection, query: str, result: dict[str, Any]) -> None:
        detail_id = new_id("gdetail")
        conn.execute("""INSERT OR REPLACE INTO grocery_product_details_cache(detail_id,retailer,product_url,product_name,price_text,unit_price_text,package_size,ingredients_text,nutrition_json,fetched_at,cache_status,source_hash,brand,image_url,price_decimal_text,currency,unit,unit_price_decimal_text,availability_status,promotion_text,source,confidence,quality_flags_json,raw_result_json)
            VALUES (COALESCE((SELECT detail_id FROM grocery_product_details_cache WHERE retailer=? AND product_url=?), ?), ?, ?, ?, ?, ?, ?, NULL, '{}', ?, 'cached', ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
            (self.retailer, result["product_url"], detail_id, self.retailer, result["product_url"], result.get("product_name"), result.get("price_text"), result.get("unit_price_text"), result.get("package_size"), result.get("fetched_at") or _now(), self._cache_key(query), result.get("brand"), result.get("image_url"), result.get("price_decimal_text"), result.get("currency") or "CHF", result.get("unit"), result.get("unit_price_decimal_text"), result.get("availability_status"), result.get("promotion_text"), result.get("source"), result.get("confidence"), json.dumps(result.get("quality_flags") or []), json.dumps(result.get("raw_result") or {})))
        conn.commit()

    def search_products(self, conn: Connection, query: str, retailer: str | None = None, locale: str = "de-CH", use_cache: bool = True, max_results: int = 5, max_age_seconds: int = 86400) -> list[dict[str, Any]]:
        if self.future_status:
            return [{"retailer": self.retailer, "query": query, "status": "needs_review", "future_status": self.future_status, "quality_flags": ["needs_review"], "source": "skeleton"}]
        if use_cache:
            cached = self.get_cached_result(conn, query, max_results=max_results, max_age_seconds=max_age_seconds)
            if cached:
                return cached
        try:
            html_text = self._fetch(self.search_url(query, locale=locale))
            raw_rows = self.parse_search_results(html_text)[:max_results]
            results = [self.normalize_result(raw).as_dict() for raw in raw_rows]
            for result in results:
                self.store_cache(conn, query, result)
            return results
        except GroceryProviderRateLimitError:
            return [{"retailer": self.retailer, "query": query, "status": "needs_review", "quality_flags": ["rate_limited", "needs_review"], "source": "web_fetch", "fetched_at": _now()}]
        except Exception as exc:
            return [{"retailer": self.retailer, "query": query, "status": "needs_review", "quality_flags": ["provider_error", "needs_review"], "error": type(exc).__name__, "source": "web_fetch", "fetched_at": _now()}]

    def fetch_product_detail(self, url: str) -> dict[str, Any]:
        html_text = self._fetch(url)
        parsed = self.parse_search_results(html_text)
        return parsed[0] if parsed else {"product_url": url, "status": "needs_review"}


class MigrosProvider(GroceryProductProvider):
    retailer = "Migros"
    base_url = "https://www.migros.ch/de/search?query="

    def search_url(self, query: str, locale: str = "de-CH") -> str:
        return self.base_url + quote_plus(query)


class CoopProvider(GroceryProductProvider):
    retailer = "Coop"
    base_url = "https://www.coop.ch/de/search/?text="

    def search_url(self, query: str, locale: str = "de-CH") -> str:
        return self.base_url + quote_plus(query)


class AldiSuisseProvider(GroceryProductProvider):
    retailer = "Aldi Suisse"
    future_status = "skeleton_provider_not_live"


class LidlSchweizProvider(GroceryProductProvider):
    retailer = "Lidl"
    future_status = "skeleton_provider_not_live"


class DennerProvider(GroceryProductProvider):
    retailer = "Denner"
    future_status = "skeleton_provider_not_live"


class OttosProvider(GroceryProductProvider):
    retailer = "Otto's"
    future_status = "skeleton_provider_not_live"


def search_open_prices_spike(query: str | None = None, barcode: str | None = None, *, fetcher: Callable[[str], str] | None = None, size: int = 5) -> dict[str, Any]:
    """Low-volume explicit Open Prices spike. Returns candidates only; never marks secure savings by itself."""
    params: dict[str, Any] = {"size": min(max(int(size or 5), 1), 10)}
    if barcode:
        params["code"] = barcode
    if query:
        params["product_name__like"] = query
    url = "https://prices.openfoodfacts.org/api/v1/prices?" + urlencode(params)
    fetch_impl: Callable[[str], str]
    if fetcher is None:
        def default_fetch_impl(target: str) -> str:
            req = Request(target, headers={"User-Agent": "JarvisFinanceOpenPricesSpike/1.0 explicit-user-action"})
            with urlopen(req, timeout=8) as response:  # noqa: S310 - official public API, explicit user action
                return response.read().decode("utf-8", errors="ignore")
        fetch_impl = default_fetch_impl
    else:
        fetch_impl = fetcher
    try:
        payload = json.loads(fetch_impl(url))
    except Exception as exc:
        return {"source": "open_prices", "official_api": True, "url": url, "candidates": [], "status": "needs_review", "error": type(exc).__name__}
    candidates = []
    for item in payload.get("items", [])[: params["size"]]:
        product = item.get("product") or {}
        location = item.get("location") or {}
        proof = item.get("proof") or {}
        price = item.get("price") or item.get("price_is_discounted")
        date = item.get("date") or proof.get("date")
        source_url = f"https://prices.openfoodfacts.org/prices/{item.get('id')}" if item.get("id") else "https://prices.openfoodfacts.org"
        flags = ["needs_review"]
        if not price:
            flags.append("no_price")
        if not (product.get("product_quantity") or product.get("quantity")):
            flags.append("different_pack_size")
        candidates.append({"provider": "Open Prices", "product_name": product.get("product_name") or query or barcode, "barcode": product.get("code") or barcode, "retailer": location.get("osm_name") or location.get("osm_brand") or "unknown", "price_text": _decimal_text(price), "currency": item.get("currency") or proof.get("currency"), "price_date": date, "source_url": source_url, "quality_flags": flags, "status": "needs_review", "can_count_as_secure_saving": False, "attribution": "Open Prices / Open Food Facts contributors"})
    return {"source": "open_prices", "official_api": True, "url": url, "license_note": "Open Prices/Open Food Facts attribution required", "candidates": candidates, "status": "ok"}


def provider_for_retailer(retailer: str) -> GroceryProductProvider:
    mapping: dict[str, type[GroceryProductProvider]] = {"Migros": MigrosProvider, "Coop": CoopProvider, "Aldi Suisse": AldiSuisseProvider, "Lidl": LidlSchweizProvider, "Denner": DennerProvider, "Otto's": OttosProvider}
    return mapping.get(retailer, GroceryProductProvider)()


def quality_flags_for_match(original_name: str, result: dict[str, Any]) -> list[str]:
    original = original_name.lower()
    candidate = str(result.get("product_name") or "").lower()
    flags: list[str] = []
    if not result.get("price_decimal_text"):
        flags.append("no_price")
    if not result.get("unit_price_decimal_text"):
        flags.append("no_unit_price")
    if original and original == candidate:
        flags.append("exact_match")
    elif any(token in candidate for token in original.split()[:2]):
        flags.append("close_match")
    else:
        flags.append("needs_review")
    if ("bio" in original) != ("bio" in candidate):
        flags.append("organic_mismatch")
    if result.get("brand") and str(result.get("brand")).lower() in {"prix garantie", "m-budget", "aldi", "lidl"}:
        flags.append("cheaper_private_label")
    if "no_price" in flags or "needs_review" in flags:
        if "needs_review" not in flags:
            flags.append("needs_review")
    return flags


def compare_product_prices(*, original_price_text: str, candidate_price_text: str | None, original_unit_price_text: str | None = None, candidate_unit_price_text: str | None = None, currency: str = "CHF", quality_flags: list[str] | None = None) -> dict[str, Any]:
    flags = list(quality_flags or [])
    if currency != "CHF" or not candidate_price_text or "needs_review" in flags or "no_price" in flags:
        return {"can_calculate_savings": False, "savings_text": "0.00", "basis": "none", "quality_flags": flags}
    if original_unit_price_text and candidate_unit_price_text and _decimal_text(original_unit_price_text) and _decimal_text(candidate_unit_price_text):
        raw_savings = _money(original_unit_price_text) - _money(candidate_unit_price_text)
        if raw_savings <= Decimal("0.00"):
            if "not_cheaper" not in flags:
                flags.append("not_cheaper")
            return {"can_calculate_savings": False, "savings_text": "0.00", "basis": "unit_price", "quality_flags": flags}
        return {"can_calculate_savings": True, "savings_text": f"{raw_savings:.2f}", "basis": "unit_price", "quality_flags": flags}
    if not any(f in flags for f in ["different_pack_size", "different_quality", "organic_mismatch"]):
        raw_savings = _money(original_price_text) - _money(candidate_price_text)
        if raw_savings <= Decimal("0.00"):
            if "not_cheaper" not in flags:
                flags.append("not_cheaper")
            return {"can_calculate_savings": False, "savings_text": "0.00", "basis": "package_price", "quality_flags": flags}
        if "package_price_fallback" not in flags:
            flags.append("package_price_fallback")
        return {"can_calculate_savings": True, "savings_text": f"{raw_savings:.2f}", "basis": "package_price", "quality_flags": flags}
    return {"can_calculate_savings": False, "savings_text": "0.00", "basis": "none", "quality_flags": flags}


def search_and_store_product_matches(conn: Connection, receipt_id: str, *, providers: list[GroceryProductProvider], included_product_item_ids: list[str] | None = None, max_products: int = 5, use_cache: bool = True, only_safe_matches: bool = True, prefer_known_mappings: bool = True) -> dict[str, Any]:
    items = _items_for_receipt(conn, receipt_id)
    if included_product_item_ids is not None:
        allowed = set(included_product_item_ids)
        items = [item for item in items if item["product_item_id"] in allowed]
    items = items[:max_products]
    matches: list[dict[str, Any]] = []
    provider_calls = 0
    cache_hits = 0
    known_mapping_hits = 0
    for item in items:
        query = item.get("normalized_product_name") or item.get("raw_product_name")
        if prefer_known_mappings:
            accepted = [m for m in known_mappings_for_item(conn, item, include_rejected=False) if m.get("status") == "accepted"]
            accepted_with_sourced_price = False
            for m in accepted:
                flags = set(json.loads(m.get("quality_flags_json") or "[]"))
                has_sourced_price = bool(m.get("target_price_text") and m.get("last_price_checked_at") and "manual_price" not in flags)
                matches.append({"mapping_id": m["mapping_id"], "product_item_id": item["product_item_id"], "retailer": m["target_retailer"], "product_name": m["target_product_name"], "product_url": m["target_product_url"], "price_text": m.get("target_price_text"), "unit_price_text": m.get("target_unit_price_text"), "source_fetched_at": m.get("last_price_checked_at"), "fetched_at": m.get("updated_at"), "quality_flags": json.loads(m.get("quality_flags_json") or "[]"), "confidence": m.get("confidence"), "status": "suggested" if has_sourced_price else "needs_review", "source": "known_mapping", "cache_status": m.get("cache_status")})
                known_mapping_hits += 1
                accepted_with_sourced_price = accepted_with_sourced_price or has_sourced_price
            if accepted_with_sourced_price:
                continue
            rejected_targets = [row_to_dict(r) for r in conn.execute("SELECT target_retailer, target_product_url, target_product_name FROM grocery_product_mappings WHERE source_product_normalized_name=? AND status='rejected'", (item.get("normalized_product_name"),)).fetchall()]
        else:
            rejected_targets = []
        for provider in providers:
            before = provider.request_count
            results = provider.search_products(conn, query, use_cache=use_cache, max_results=3)
            provider_calls += max(0, provider.request_count - before)
            for result in results:
                if result.get("source") == "cache":
                    cache_hits += 1
                if not result.get("product_url") or not result.get("product_name"):
                    matches.append({"product_item_id": item["product_item_id"], **result})
                    continue
                flags = quality_flags_for_match(item.get("normalized_product_name") or item.get("raw_product_name"), result)
                if any(rt.get("target_retailer") == result.get("retailer") and (rt.get("target_product_url") == result.get("product_url") or rt.get("target_product_name") == result.get("product_name")) for rt in rejected_targets):
                    stored = upsert_product_match(conn, item["product_item_id"], retailer=result["retailer"], candidate_product_name=result["product_name"], candidate_url=result["product_url"], candidate_price_text=result.get("price_decimal_text") or result.get("price_text"), candidate_unit_price_text=result.get("unit_price_text"), quality_flags=[*flags, "rejected"], match_confidence=result.get("confidence") or "0.40", status="rejected", source=result.get("source") or "web_fetch", candidate_brand=result.get("brand"), candidate_package_size=result.get("package_size"), candidate_unit=result.get("unit"))
                    matches.append({"match_id": stored.get("match_id"), "product_item_id": item["product_item_id"], "retailer": result.get("retailer"), "product_name": result.get("product_name"), "product_url": result.get("product_url"), "price_text": result.get("price_decimal_text") or result.get("price_text"), "unit_price_text": result.get("unit_price_text"), "source_fetched_at": result.get("fetched_at"), "fetched_at": stored.get("fetched_at"), "quality_flags": [*flags, "rejected"], "confidence": result.get("confidence"), "status": "rejected", "source": result.get("source")})
                    continue
                comparison = compare_product_prices(original_price_text=str(item.get("total_price_text") or "0"), original_unit_price_text=item.get("unit_price_text"), candidate_price_text=result.get("price_decimal_text") or result.get("price_text"), candidate_unit_price_text=result.get("unit_price_decimal_text") or result.get("unit_price_text"), currency=result.get("currency") or "CHF", quality_flags=flags)
                status = "suggested" if comparison["can_calculate_savings"] and (not only_safe_matches or not (set(flags) & UNSAFE_FLAGS)) else "needs_review"
                if not comparison["can_calculate_savings"] or not result.get("price_decimal_text"):
                    stored = upsert_product_match(conn, item["product_item_id"], retailer=result["retailer"], candidate_product_name=result["product_name"], candidate_url=result["product_url"], candidate_price_text=result.get("price_decimal_text") or result.get("price_text"), candidate_unit_price_text=result.get("unit_price_text"), quality_flags=comparison["quality_flags"], match_confidence=result.get("confidence") or "0.40", status="needs_review", source=result.get("source") or "web_fetch", candidate_brand=result.get("brand"), candidate_package_size=result.get("package_size"), candidate_unit=result.get("unit"))
                    matches.append({"match_id": stored.get("match_id"), "product_item_id": item["product_item_id"], "retailer": result.get("retailer"), "product_name": result.get("product_name"), "product_url": result.get("product_url"), "price_text": result.get("price_decimal_text") or result.get("price_text"), "unit_price_text": result.get("unit_price_text"), "source_fetched_at": result.get("fetched_at"), "fetched_at": stored.get("fetched_at"), "quality_flags": comparison["quality_flags"], "confidence": result.get("confidence"), "status": "needs_review", "source": result.get("source")})
                    continue
                stored = upsert_product_match(conn, item["product_item_id"], retailer=result["retailer"], candidate_product_name=result["product_name"], candidate_url=result["product_url"], candidate_price_text=result.get("price_decimal_text") or result.get("price_text") or "0", candidate_unit_price_text=result.get("unit_price_text"), quality_flags=comparison["quality_flags"], match_confidence=result.get("confidence") or "0.60", status=status, source=result.get("source") or "web_fetch", candidate_brand=result.get("brand"), candidate_package_size=result.get("package_size"), candidate_unit=result.get("unit"))
                matches.append({"match_id": stored.get("match_id"), "product_item_id": item["product_item_id"], "retailer": result["retailer"], "product_name": result["product_name"], "product_url": result["product_url"], "price_text": result.get("price_decimal_text") or result.get("price_text"), "unit_price_text": result.get("unit_price_text"), "source_fetched_at": result.get("fetched_at"), "fetched_at": stored.get("fetched_at"), "quality_flags": comparison["quality_flags"], "confidence": result.get("confidence"), "status": status, "source": result.get("source")})
    return {"receipt_id": receipt_id, "provider_calls": provider_calls, "cache_hits": cache_hits, "known_mapping_hits": known_mapping_hits, "matches": matches}
