import json, re
from services.pdf_extractor.config import groq_client, EXTRACT_PROMPT

MODELS = [
      "llama-3.3-70b-versatile",
    "qwen/qwen3-32b",
    "openai/gpt-oss-120b",
    "openai/gpt-oss-20b",
]

SKIP_KEYWORDS = [
    'Freight Payment', 'Ship Method', 'Gross Weight',
    'Net Weight', 'No. of Pack', 'Customer pick-up',
    'CERTIFIED TRUE', 'RECEIVED BY','Export HTS', 'Import HTS', 'US ECCN', 
    'HAZ CLASS', 'R/N NBR', 'EAR99',  
]

def smart_trim(text: str, max_chars: int = 12000) -> str:
    lines = text.split('\n')
    useful_lines = []
    for line in lines:
        stripped = line.strip()
        if not stripped:
            continue
        if any(kw in stripped for kw in SKIP_KEYWORDS):
            continue
        if len(stripped) > 400:  # long legal lines skip
            continue
        useful_lines.append(stripped)
    result = '\n'.join(useful_lines)
    print(f"Smart trim: {len(text)} → {len(result)} chars")
    return result[:max_chars]

def fix_incomplete_json(raw: str) -> dict:
    try:
        return json.loads(raw)
    except json.JSONDecodeError:
        print("JSON incomplete — trimming to last complete item...")
        last_brace = raw.rfind('},')
        if last_brace != -1:
            raw = raw[:last_brace+1] + '\n  ]\n}'
            return json.loads(raw)
        raise
def clean_number(val):
    if not val:
        return ''
    val = str(val).strip().replace(' ', '')
    if ',' in val and '.' in val:
        if val.index(',') < val.index('.'):
            val = val.replace(',', '')        # 1,704.00 → 1704.00
        else:
            val = val.replace('.', '').replace(',', '.')  # 1.704,00 → 1704.00
    elif ',' in val:
        parts = val.split(',')
        if len(parts[-1]) == 3:
            val = val.replace(',', '')        # 1,704 → 1704
        else:
            val = val.replace(',', '.')       # 1,70 → 1.70
    return val

def validate_and_fix_rows(items: list) -> list:
    fixed = []
    for row in items:
        if not isinstance(row, dict):
            continue
        if not row.get('description') and not row.get('model_item_no') and not row.get('qty'):
            continue

        desc = row.get('description', '') or ''
        for marker in ['Customer Part#:', 'Customer Part:', 'Cust Part#:']:
            if marker in desc:
                desc = desc[:desc.index(marker)].strip()
        row['description'] = desc

        try:
            qty = float(clean_number(str(row.get('qty', '') or '')))
            up  = float(clean_number(str(row.get('unit_price', '') or '')))
            amt = float(clean_number(str(row.get('amount', '') or '')))
            if qty > 0 and up > 0 and amt > 0:
                calc = round(qty * up, 2)
                if abs(calc - amt) > 1:
                    calc2 = round(qty * amt, 2)
                    if abs(calc2 - up) < 1:
                        row['unit_price'], row['amount'] = row['amount'], row['unit_price']
                        print(f"Fixed swap — line {row.get('line_no')}")
            row['unit_price'] = clean_number(str(row.get('unit_price', '') or ''))
            row['amount'] = clean_number(str(row.get('amount', '') or ''))
        except:
            pass

        fixed.append(row)
    return fixed

def call_ai(text: str) -> dict:
    # ── Page-aware chunking ──────────────────────────────
    # Pages ko split karo
    pages = re.split(r'--- PAGE \d+ ---', text)
    pages = [p.strip() for p in pages if p.strip()]

    if not pages:
        pages = [text]

    print(f"Total pages found: {len(pages)}")

    # 4 pages per chunk, overlap ke liye 1 page repeat
    PAGES_PER_CHUNK = 4
    chunks = []
    i = 0
    while i < len(pages):
        chunk_pages = pages[i:i+PAGES_PER_CHUNK]
        chunks.append('\n\n'.join(chunk_pages))
        i += PAGES_PER_CHUNK - 1  # 1 page overlap
        if i >= len(pages):
            break

    print(f"Total chunks: {len(chunks)}")
    # ────────────────────────────────────────────────────

    all_items = []
    vendor_name = "Unknown"
    invoice_no = ""
    invoice_date = ""
    po_no = ""

    for idx, chunk in enumerate(chunks):
        useful = smart_trim(chunk, max_chars=12000)

        # Chunk 2+ mein header inject karo
        if idx > 0 and invoice_no:
            header = f"[HEADER INFO] Invoice No: {invoice_no}, Invoice Date: {invoice_date}, PO No: {po_no}\n\n"
            useful = header + useful

        prompt = EXTRACT_PROMPT.replace("{text}", useful)
        print(f"\nChunk {idx+1}/{len(chunks)} — chars: {len(useful)}")

        last_error = None
        success = False

        for model in MODELS:
            try:
                print(f"Trying model: {model}")
                response = groq_client.chat.completions.create(
                    model=model,
                    messages=[{"role": "user", "content": prompt}],
                    max_tokens=4000
                )
                raw = response.choices[0].message.content.strip()
                if not raw:
                    print(f"Empty response — trying next model")
                    continue

                raw = re.sub(r'^```[\w]*\s*', '', raw)
                raw = re.sub(r'\s*```$', '', raw)
                data = fix_incomplete_json(raw.strip())

                items = data.get('line_items', [])
                items = [item for item in items if isinstance(item, dict)]

                if idx == 0:
                    vendor_name = data.get('vendor_name', 'Unknown') or 'Unknown'
                    if items:
                        invoice_no = items[0].get('invoice_no', '')
                        invoice_date = items[0].get('invoice_date', '')
                        po_no = items[0].get('po_no', '')

                items = validate_and_fix_rows(items)
                all_items.extend(items)
                print(f"Chunk {idx+1} rows: {len(items)}, Running total: {len(all_items)}")
                success = True
                break

            except Exception as e:
                err = str(e)
                if "429" in err:
                    print(f"RATE LIMIT: {model} — skipping")
                elif "413" in err:
                    print(f"TOO LARGE: {model} — skipping")
                else:
                    print(f"ERROR: {model} — {err[:80]}")
                last_error = e
                continue

        if not success and last_error and len(all_items) == 0:
            raise last_error

    # Deduplicate
    seen = set()
    unique_items = []
    for item in all_items:
        key = (item.get('line_no', ''), item.get('model_item_no', ''))
        if key not in seen:
            seen.add(key)
            unique_items.append(item)

    print(f"\nTOTAL unique rows: {len(unique_items)}")
    return {"vendor_name": vendor_name, "line_items": unique_items}