import io

from pydantic import ValidationError
from PIL import Image

from src.core.llm import invoke_json_extraction
from src.prompts.invoice import invoice_system_prompt
from src.prompts.movement_out import movement_out_system_prompt
from src.prompts.movement_out_pdf import movement_out_pdf_system_prompt
from src.prompts.sales_order import sales_order_system_prompt
from src.prompts.strategic_stock import (
    strategic_stock_chunk_prompt,
    strategic_stock_document_prompt,
    strategic_stock_system_prompt,
    strategic_stock_parse_compact,
)
from src.schemas.response import ExtractedFields, ExtractionMetadata, ExtractionResponse


def extract_document(
    document_type: str,
    image_bytes: bytes,
) -> ExtractionResponse:
    system_prompt = (
        invoice_system_prompt()
        if document_type == "incoming_invoice"
        else movement_out_system_prompt()
    )
    parsed, metadata = invoke_json_extraction(system_prompt, image_bytes)

    try:
        fields = ExtractedFields(**parsed)
    except ValidationError as exc:
        raise ValueError(f"Model returned unexpected shape: {exc}. Raw: {parsed}") from exc

    validation_flags = _validation_flags(fields)

    return ExtractionResponse(
        document_type=document_type,
        extracted_fields=fields,
        validation_flags=validation_flags,
        confidence=0.99 if not validation_flags else 0.75,
        raw_model_output=parsed,
        metadata=ExtractionMetadata(**metadata),
    )


def extract_sales_order(image_bytes: bytes) -> ExtractionResponse:
    """Extract fields from a Sales Order PDF."""
    parsed, metadata = invoke_json_extraction(sales_order_system_prompt(), image_bytes)

    try:
        fields = ExtractedFields(**parsed)
    except ValidationError as exc:
        raise ValueError(f"Model returned unexpected shape: {exc}. Raw: {parsed}") from exc

    # Map invoice_date → document_date if needed
    raw = parsed if isinstance(parsed, dict) else {}
    if not fields.document_date and raw.get("invoice_date"):
        fields.document_date = raw["invoice_date"]

    # Store po_number and bill_to in reference_number if not already set
    if not fields.referenceNumber if hasattr(fields, "referenceNumber") else not fields.reference_number:
        po = raw.get("po_number")
        bill = raw.get("bill_to")
        if po:
            fields.reference_number = po  # type: ignore[attr-defined]

    validation_flags = _sales_order_validation_flags(fields)

    return ExtractionResponse(
        document_type="incoming_invoice",
        extracted_fields=fields,
        validation_flags=validation_flags,
        confidence=0.99 if not validation_flags else 0.75,
        raw_model_output=parsed,
        metadata=ExtractionMetadata(**metadata),
    )


def extract_strategic_stock(image_bytes: bytes) -> ExtractionResponse:
    """Extract fields from a Strategic Stock Movement Sheet PDF using chunked row extraction."""
    parsed, metadata = _extract_strategic_stock_chunked(image_bytes)
    parsed = _retry_sparse_strategic_stock(parsed, image_bytes)

    # Expand compact keys (n, ic, in, bn, etc.) back to full names
    parsed = strategic_stock_parse_compact(parsed)
    parsed = _normalize_strategic_stock_numbers(parsed)

    try:
        fields = ExtractedFields(**parsed)
    except ValidationError as exc:
        raise ValueError(f"Model returned unexpected shape: {exc}. Raw: {parsed}") from exc

    # Signature validation flag
    raw = parsed if isinstance(parsed, dict) else {}
    sig_count = raw.get("signature_count", 0)
    try:
        sig_count = int(sig_count)
    except (TypeError, ValueError):
        sig_count = 0

    validation_flags = _validation_flags(fields)
    if sig_count < 3:
        validation_flags.append(f"insufficient_signatures:{sig_count}")

    confidence = 0.99 if not [f for f in validation_flags if not f.startswith("insufficient")] else 0.75

    return ExtractionResponse(
        document_type="incoming_invoice",
        extracted_fields=fields,
        validation_flags=validation_flags,
        confidence=confidence,
        raw_model_output=parsed,
        metadata=ExtractionMetadata(**metadata),
    )


def _extract_strategic_stock_chunked(image_bytes: bytes) -> tuple[dict, dict]:
    # Pass 1 — extract header, warehouse list, totals, signatures
    # Use a taller header crop (25%) to ensure column headers are fully visible
    doc_image_bytes = _strategic_stock_summary_image(image_bytes)
    doc_parsed, metadata = invoke_json_extraction(
        strategic_stock_document_prompt(),
        doc_image_bytes,
        max_tokens_override=4096,
    )
    warehouse_headers = _normalize_warehouse_headers(doc_parsed.get("warehouses") if isinstance(doc_parsed, dict) else None)

    # Pass 2 — extract rows in chunks, injecting exact column order
    merged_items: dict[int, dict] = {}
    chunk_count = 0
    for chunk_bytes in _strategic_stock_row_chunks(image_bytes):
        chunk_count += 1
        chunk_parsed, _ = invoke_json_extraction(
            strategic_stock_chunk_prompt(warehouse_headers),
            chunk_bytes,
            max_tokens_override=12000,
        )
        for li in chunk_parsed.get("line_items", []) if isinstance(chunk_parsed, dict) else []:
            line_number = li.get("n") or li.get("line_number")
            if not isinstance(line_number, int):
                try:
                    line_number = int(line_number)
                except (TypeError, ValueError):
                    continue
            existing = merged_items.get(line_number)
            if existing is None or _compact_item_score(li) > _compact_item_score(existing):
                merged_items[line_number] = li

    merged = dict(doc_parsed) if isinstance(doc_parsed, dict) else {}
    merged["line_items"] = [merged_items[k] for k in sorted(merged_items)]
    merged = _enrich_strategic_stock_from_rows(merged)
    metadata["chunk_count"] = chunk_count
    return merged, metadata


def _retry_sparse_strategic_stock(parsed: dict, image_bytes: bytes) -> dict:
    """Retry strategic stock extraction when the model returns an obviously incomplete row set."""
    line_items = parsed.get("line_items", []) if isinstance(parsed, dict) else []
    if not isinstance(line_items, list) or len(line_items) >= 10:
        return parsed

    retry_prompt = (
        strategic_stock_system_prompt()
        + "\n\nFINAL CHECK BEFORE YOU ANSWER:\n"
        + "- Recount the SL/serial-number rows from top to bottom.\n"
        + "- This sheet is often 40+ rows and may not fit if you stop early.\n"
        + "- If you currently have fewer than 10 line_items, your extraction is incomplete.\n"
        + "- Return EVERY numbered row before the totals row.\n"
        + "- Keep JSON compact and prioritize completing line_items."
    )
    retried, _ = invoke_json_extraction(
        retry_prompt,
        image_bytes,
        max_tokens_override=16384,
    )
    retried_items = retried.get("line_items", []) if isinstance(retried, dict) else []
    if isinstance(retried_items, list) and len(retried_items) > len(line_items):
        return retried
    return parsed


def _strategic_stock_row_chunks(image_bytes: bytes) -> list[bytes]:
    img = Image.open(io.BytesIO(image_bytes))
    if img.mode in ("RGBA", "P"):
        img = img.convert("RGB")

    width, height = img.size
    # Use 25% for header to ensure column headers are always included in chunks
    header_bottom = int(height * 0.25)
    table_top = int(height * 0.25)
    table_bottom = int(height * 0.90)
    table_height = max(table_bottom - table_top, height // 2)

    target_chunks = 4
    overlap = int(table_height * 0.08)  # 8% overlap to avoid cutting rows
    step = max(1, table_height // target_chunks)

    chunks: list[bytes] = []
    header_crop = img.crop((0, 0, width, header_bottom))
    start = table_top
    while start < table_bottom:
        end = min(table_bottom, start + step + overlap)
        row_crop = img.crop((0, max(0, start - overlap), width, end))
        crop = _stack_with_header_context(header_crop, row_crop)
        buf = io.BytesIO()
        crop.save(buf, format="PNG")
        chunks.append(buf.getvalue())
        if end >= table_bottom:
            break
        start += step

    return chunks


def _strategic_stock_summary_image(image_bytes: bytes) -> bytes:
    img = Image.open(io.BytesIO(image_bytes))
    if img.mode in ("RGBA", "P"):
        img = img.convert("RGB")

    width, height = img.size
    # Take top 25% (header + column names) and bottom 16% (totals + signatures)
    top = img.crop((0, 0, width, int(height * 0.25)))
    bottom = img.crop((0, int(height * 0.84), width, height))
    summary = _stack_with_header_context(top, bottom)
    buf = io.BytesIO()
    summary.save(buf, format="PNG")
    return buf.getvalue()


def _compact_item_score(li: dict) -> int:
    return sum(
        1
        for value in li.values()
        if value not in (None, "", {}, [])
    )


def _normalize_strategic_stock_numbers(parsed: dict) -> dict:
    if not isinstance(parsed, dict):
        return parsed

    for key in ("quantity_mt", "total_bags_all", "total_mt_all", "signature_count"):
        if key in parsed:
            parsed[key] = _coerce_number(parsed.get(key))

    normalized_items = []
    for li in parsed.get("line_items", []) if isinstance(parsed.get("line_items", []), list) else []:
        if not isinstance(li, dict):
            normalized_items.append(li)
            continue
        item = dict(li)
        for key in ("line_number", "quantity", "quantity_mt", "unit_price", "amount", "vat_pct", "vat_amount", "amount_with_vat", "total_bags", "total_mt"):
            if key in item:
                item[key] = _coerce_number(item.get(key))
        if "unit" in item and item["unit"] is not None and not isinstance(item["unit"], str):
            item["unit"] = str(item["unit"])
        wb = item.get("warehouse_bags")
        if isinstance(wb, dict):
            item["warehouse_bags"] = {
                wh: (_coerce_number(bags) if bags is not None else None)
                for wh, bags in wb.items()
            }
            item["warehouse_bags"] = _clean_warehouse_bags(item["warehouse_bags"])
        normalized_items.append(item)
    parsed["line_items"] = normalized_items
    return parsed


def _coerce_number(value):
    if isinstance(value, (int, float)) or value is None:
        return value
    if isinstance(value, str):
        cleaned = value.replace(",", "").strip()
        if cleaned == "":
            return None
        try:
            if "." in cleaned:
                return float(cleaned)
            return int(cleaned)
        except ValueError:
            return value
    return value


def _stack_with_header_context(header_crop: Image.Image, row_crop: Image.Image) -> Image.Image:
    gap = 10
    canvas = Image.new(
        "RGB",
        (max(header_crop.width, row_crop.width), header_crop.height + gap + row_crop.height),
        color="#ffffff",
    )
    canvas.paste(header_crop, (0, 0))
    canvas.paste(row_crop, (0, header_crop.height + gap))
    return canvas


def _enrich_strategic_stock_from_rows(parsed: dict) -> dict:
    if not isinstance(parsed, dict):
        return parsed

    line_items = parsed.get("line_items", [])
    if not isinstance(line_items, list):
        return parsed

    warehouse_keys: list[str] = _normalize_warehouse_headers(parsed.get("warehouses"))
    total_bags_sum = 0.0
    total_mt_sum = 0.0
    has_total_bags = False
    has_total_mt = False

    for li in line_items:
        if not isinstance(li, dict):
            continue
        wb = li.get("wb") or li.get("warehouse_bags")
        if isinstance(wb, dict):
            cleaned_wb = _clean_warehouse_bags(wb)
            if "warehouse_bags" in li:
                li["warehouse_bags"] = cleaned_wb
            if "wb" in li:
                li["wb"] = cleaned_wb
            for key in cleaned_wb.keys():
                if isinstance(key, str) and key not in warehouse_keys:
                    warehouse_keys.append(key)
        else:
            cleaned_wb = {}

        warehouse_total = sum(
            float(value) for value in cleaned_wb.values()
            if isinstance(value, (int, float))
        )

        tb = _coerce_number(li.get("tb") if "tb" in li else li.get("total_bags"))
        if warehouse_total > 0 and (not isinstance(tb, (int, float)) or abs(float(tb) - warehouse_total) > 0.5):
            tb = warehouse_total
            if "tb" in li:
                li["tb"] = tb
            else:
                li["total_bags"] = tb

        tm = _coerce_number(li.get("tm") if "tm" in li else li.get("total_mt"))
        computed_mt = _calculate_total_mt(tb, li.get("u") if "u" in li else li.get("unit"))
        if computed_mt is not None and (not isinstance(tm, (int, float)) or abs(float(tm) - computed_mt) > 0.051):
            tm = computed_mt
            if "tm" in li:
                li["tm"] = tm
            else:
                li["total_mt"] = tm
            if "quantity_mt" in li:
                li["quantity_mt"] = tm

        if isinstance(tb, (int, float)):
            total_bags_sum += float(tb)
            has_total_bags = True
        if isinstance(tm, (int, float)):
            total_mt_sum += float(tm)
            has_total_mt = True

    if warehouse_keys:
        parsed["warehouses"] = warehouse_keys

    if parsed.get("total_bags_all") in (None, "", 0) and has_total_bags:
        parsed["total_bags_all"] = round(total_bags_sum, 3)

    if parsed.get("total_mt_all") in (None, "", 0) and has_total_mt:
        parsed["total_mt_all"] = round(total_mt_sum, 3)

    if parsed.get("quantity_mt") in (None, "", 0) and has_total_mt:
        parsed["quantity_mt"] = round(total_mt_sum, 3)

    return parsed


def _normalize_warehouse_headers(warehouses) -> list[str]:
    if not isinstance(warehouses, list):
        return []
    normalized: list[str] = []
    for warehouse in warehouses:
        if not isinstance(warehouse, str):
            continue
        cleaned = " ".join(warehouse.replace("\n", " ").split())
        if cleaned and cleaned not in normalized:
            normalized.append(cleaned)
    return normalized


def _clean_warehouse_bags(warehouse_bags: dict) -> dict:
    cleaned: dict = {}
    for raw_key, raw_value in warehouse_bags.items():
        if not isinstance(raw_key, str):
            continue
        key = " ".join(raw_key.replace("\n", " ").split())
        if not key:
            continue
        lower = key.lower()
        if "total bag" in lower or "total mt" in lower or lower in {"unit", "shipment no.", "shipment no", "quality report no.", "quality report no"}:
            continue
        value = _coerce_number(raw_value)
        if value in (None, "", 0):
            continue
        cleaned[key] = value
    return cleaned


def _calculate_total_mt(total_bags, unit) -> float | None:
    if not isinstance(total_bags, (int, float)):
        return None
    if unit is None:
        return None
    unit_text = str(unit).strip().lower()
    digits = []
    current = ""
    for char in unit_text:
        if char.isdigit() or char == ".":
            current += char
        elif current:
            digits.append(current)
            current = ""
    if current:
        digits.append(current)
    if not digits:
        return None
    try:
        unit_kg = float(digits[0])
    except ValueError:
        return None
    if unit_kg <= 0:
        return None
    return round((float(total_bags) * unit_kg) / 1000.0, 3)


def extract_movement_out_pdf(image_bytes: bytes) -> ExtractionResponse:
    """Extract fields from a Movement Out PDF (generated document with signatures)."""
    parsed, metadata = invoke_json_extraction(movement_out_pdf_system_prompt(), image_bytes)

    try:
        fields = ExtractedFields(**parsed)
    except ValidationError as exc:
        raise ValueError(f"Model returned unexpected shape: {exc}. Raw: {parsed}") from exc

    raw = parsed if isinstance(parsed, dict) else {}
    sig_count = raw.get("signature_count", 0)
    try:
        sig_count = int(sig_count)
    except (TypeError, ValueError):
        sig_count = 0

    validation_flags = _validation_flags(fields)
    if sig_count < 3:
        validation_flags.append(f"insufficient_signatures:{sig_count}")

    confidence = 0.99 if not [f for f in validation_flags if not f.startswith("insufficient")] else 0.75

    return ExtractionResponse(
        document_type="movement_out",
        extracted_fields=fields,
        validation_flags=validation_flags,
        confidence=confidence,
        raw_model_output=parsed,
        metadata=ExtractionMetadata(**metadata),
    )


def _validation_flags(fields: ExtractedFields) -> list[str]:
    flags: list[str] = []
    if not fields.warehouse_name:
        flags.append("missing_warehouse_name")
    if not fields.item_code:
        flags.append("missing_item_code")
    if not fields.batch_number:
        flags.append("missing_batch_number")
    # expiry_date not flagged — absent on most invoices
    if fields.quantity_mt is None:
        flags.append("missing_quantity_mt")
    return flags


def _sales_order_validation_flags(fields: ExtractedFields) -> list[str]:
    flags: list[str] = []
    if not fields.document_number:
        flags.append("missing_invoice_number")
    if not fields.document_date:
        flags.append("missing_invoice_date")
    return flags
