import asyncio
import json
import os
import re
import uuid
import requests
from datetime import datetime
from dotenv import load_dotenv
from playwright.sync_api import sync_playwright

load_dotenv(os.path.join(os.path.dirname(__file__), '../../.env'))
import mysql.connector

# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------
SOURCE = 'www.publicprocurement.be'
API_BASE_URL = 'https://www.publicprocurement.be/api/sea/search/publications'
ITEMS_PER_PAGE = 25
MAX_PAGES = 200


def get_token():
    """Launch a headless browser, visit the BDA page and intercept the Bearer token."""
    token_holder = []

    def handle_request(request):
        auth = request.headers.get('authorization', '')
        if auth.startswith('Bearer ') and not token_holder:
            token_holder.append(auth.split(' ', 1)[1])

    with sync_playwright() as p:
        browser = p.chromium.launch(headless=True)
        context = browser.new_context(
            user_agent='Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 '
                       '(KHTML, like Gecko) Chrome/146.0.0.0 Safari/537.36'
        )
        page = context.new_page()
        page.on('request', handle_request)
        page.goto('https://www.publicprocurement.be/bda', wait_until='networkidle', timeout=60000)
        browser.close()

    if not token_holder:
        raise RuntimeError("Could not capture Bearer token from browser session.")

    print("  Token captured from browser.")
    return token_holder[0]


def build_headers(token):
    return {
        'Accept':           'application/json',
        'Content-Type':     'application/json',
        'Authorization':    f'Bearer {token}',
        'account-type':     'public',
        'belgov-trace-id':  str(uuid.uuid4()),
        'User-Agent':       'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 '
                            '(KHTML, like Gecko) Chrome/146.0.0.0 Safari/537.36',
        'origin':           'https://www.publicprocurement.be',
        'referer':          'https://www.publicprocurement.be/bda',
    }


# ---------------------------------------------------------------------------
# DB
# ---------------------------------------------------------------------------
def get_db_connection():
    return mysql.connector.connect(
        host=os.getenv('DB_HOST'),
        port=int(os.getenv('DB_PORT', 3306)),
        user=os.getenv('DB_USER'),
        password=os.getenv('DB_PASSWORD'),
        database=os.getenv('DB_NAME'),
    )


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def load_cpv_codes():
    """Parse CPV_CODES from .env — stored as JSON dict {code: description}."""
    raw = os.getenv('CPV_CODES', '{}')
    try:
        mapping = json.loads(raw)
        codes = list(mapping.keys())
        print(f"Loaded {len(codes)} CPV codes from .env")
        return codes
    except json.JSONDecodeError:
        # Fallback: comma-separated plain list
        codes = [c.strip() for c in raw.split(',') if c.strip()]
        print(f"Loaded {len(codes)} CPV codes (plain list) from .env")
        return codes


def get_text(multilingual_list, lang='EN'):
    """Extract text for a given language from a [{language, text}] list."""
    if not multilingual_list:
        return None
    for item in multilingual_list:
        if item.get('language') == lang:
            return item.get('text')
    # Fallback: first available
    return multilingual_list[0].get('text') if multilingual_list else None


def parse_date(value):
    if not value:
        return None
    for fmt in ('%Y-%m-%d', '%Y-%m-%dT%H:%M:%S', '%Y-%m-%dT%H:%M:%S.%f', '%d/%m/%Y'):
        try:
            return datetime.strptime(value[:19], fmt).strftime('%Y-%m-%d')
        except ValueError:
            continue
    return value[:10] if value else None


def parse_datetime(value):
    if not value:
        return None
    for fmt in ('%Y-%m-%dT%H:%M:%S.%f', '%Y-%m-%dT%H:%M:%S', '%Y-%m-%d %H:%M:%S', '%Y-%m-%d'):
        try:
            return datetime.strptime(value[:26], fmt).strftime('%Y-%m-%d %H:%M:%S')
        except ValueError:
            continue
    return None


# ---------------------------------------------------------------------------
# API
# ---------------------------------------------------------------------------
def fetch_publications(cpv_code, page, token):
    """POST search for one page of publications filtered by CPV code."""
    body = {
        "cpvCodes":                    [cpv_code],
        "includeOrganisationChildren": True,
        "page":                        page,
        "pageSize":                    ITEMS_PER_PAGE,
    }
    try:
        resp = requests.post(API_BASE_URL, json=body, headers=build_headers(token), timeout=30)
        resp.raise_for_status()
        return resp.json()
    except requests.RequestException as e:
        print(f"  API error (CPV={cpv_code}, page={page}): {e}")
        return None


# ---------------------------------------------------------------------------
# DB inserts
# ---------------------------------------------------------------------------
def insert_tender(cursor, tender):
    cursor.execute(
        "SELECT id FROM tenders WHERE source = %s AND reference_number = %s",
        (tender['source'], tender['reference_number'])
    )
    row = cursor.fetchone()
    if row:
        return None  # duplicate

    sql = """
        INSERT INTO tenders
            (source, source_id, title, reference_number, description,
             organization, url, status, publication_type, closing_date,
             keyword, detail, created_at, updated_at)
        VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, 0,
                CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
    """
    cursor.execute(sql, (
        tender['source'],
        tender['reference_number'],
        tender['title'],
        tender['reference_number'],
        tender['description'],
        tender['organization'],
        tender['url'],
        tender['status'],
        tender['publication_type'],
        tender['closing_date'],
        tender['cpv_code'],
    ))
    return cursor.lastrowid


def upsert_tender_detail(cursor, tender_id, d):
    cursor.execute("SELECT id FROM tender_details WHERE tender_id = %s", (tender_id,))
    existing = cursor.fetchone()

    fields = {
        'reference_number':          d['reference_number'],
        'contracting_authority_name':d['contracting_authority_name'],
        'authority_id':              d['authority_id'],
        'category':                  d['category'],
        'cpv_code':                  d['cpv_code'],
        'cpv_description':           d['cpv_description'],
        'notice_type':               d['notice_type'],
        'procedure_type':            d['procedure_type'],
        'specific_procedure':        d['specific_procedure'],
        'framework_agreement':       d['framework_agreement'],
        'trade_agreements':          d['trade_agreements'],
        'nuts_place':                d['nuts_place'],
        'languages':                 d['languages'],
        'procedure_id':              d['procedure_id'],
        'publication_date':          d['publication_date'],
        'open_date':                 d['open_date'],
        'deadline':                  d['deadline'],
        'divided_into_lots':         d['divided_into_lots'],
        'eu_identifier':             d['eu_identifier'],
        'notice_id':                 d['notice_id'],
        'summary':                   d['summary'],
        'notice':                    d['notice'],
    }

    if existing:
        set_clause = ', '.join(f"`{col}` = %s" for col in fields)
        sql = f"UPDATE tender_details SET {set_clause}, updated_at = CURRENT_TIMESTAMP WHERE tender_id = %s"
        cursor.execute(sql, list(fields.values()) + [tender_id])
    else:
        cols = ', '.join(f"`{col}`" for col in fields)
        placeholders = ', '.join('%s' for _ in fields)
        sql = f"""
            INSERT INTO tender_details (tender_id, {cols}, created_at, updated_at)
            VALUES (%s, {placeholders}, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
        """
        cursor.execute(sql, [tender_id] + list(fields.values()))


# ---------------------------------------------------------------------------
# Parse publication → tender + detail dicts
# ---------------------------------------------------------------------------
def parse_publication(pub):
    dossier = pub.get('dossier', {})
    org = pub.get('organisation', {})

    # --- tenders table ---
    title = get_text(dossier.get('titles', []))
    description = get_text(dossier.get('descriptions', []))
    organization = get_text(org.get('organisationNames', []))
    reference_number = pub.get('referenceNumber') or dossier.get('referenceNumber')
    workspace_id = pub.get('publicationWorkspaceId', '')
    url = f"https://www.publicprocurement.be/publication-workspaces/{workspace_id}/general" if workspace_id else ''
    status = pub.get('publicationType')
    publication_type = parse_datetime(pub.get('publicationDate'))
    closing_date = parse_date(pub.get('vaultSubmissionDeadline'))

    tender = {
        'source':           SOURCE,
        'reference_number': reference_number,
        'title':            title,
        'description':      description,
        'organization':     organization,
        'url':              url,
        'status':           status,
        'publication_type': publication_type,
        'closing_date':     closing_date,
        'cpv_code':         pub.get('cpvMainCode', {}).get('code', ''),
    }

    # --- tender_details table ---
    cpv_main = pub.get('cpvMainCode', {})
    cpv_additional = pub.get('cpvAdditionalCodes', [])
    cpv_additional_codes = ', '.join(c.get('code', '') for c in cpv_additional) if cpv_additional else None

    natures = pub.get('natures', [])
    nuts = pub.get('nutsCodes', [])
    langs = pub.get('publicationLanguages', [])

    ted_refs = pub.get('publicationReferenceNumbersTED', [])
    eu_identifier = ted_refs[0] if ted_refs else None

    notice_ids = pub.get('noticeIds', [])
    notice_id = notice_ids[0] if notice_ids else None

    # Lots → store as JSON in notice column
    lots = pub.get('lots', [])
    lots_data = [
        {
            'title': get_text(lot.get('titles', [])),
            'description': get_text(lot.get('descriptions', [])),
        }
        for lot in lots
    ]
    notice_json = json.dumps(lots_data) if lots_data else None

    detail = {
        'reference_number':          dossier.get('referenceNumber') or dossier.get('number'),
        'contracting_authority_name':organization,
        'authority_id':              str(org.get('organisationId', '')) or None,
        'category':                  f"{cpv_main.get('code')} - {get_text(cpv_main.get('descriptions', []))}" if cpv_main else None,
        'cpv_code':                  cpv_main.get('code') + (f', {cpv_additional_codes}' if cpv_additional_codes else ''),
        'cpv_description':           get_text(cpv_main.get('descriptions', [])),
        'notice_type':               ', '.join(natures) if natures else None,
        'procedure_type':            dossier.get('procurementProcedureType'),
        'specific_procedure':        dossier.get('procurementProcedureType'),
        'framework_agreement':       dossier.get('specialPurchasingTechnique'),
        'trade_agreements':          dossier.get('legalBasis'),
        'nuts_place':                ', '.join(nuts) if nuts else None,
        'languages':                 ', '.join(langs) if langs else None,
        'procedure_id':              pub.get('procedureId'),
        'publication_date':          parse_datetime(pub.get('insertionDate') or pub.get('publicationDate')),
        'open_date':                 parse_date(pub.get('dispatchDate')),
        'deadline':                  parse_datetime(pub.get('vaultSubmissionDeadline')),
        'divided_into_lots':         str(len(lots)) if lots else '0',
        'eu_identifier':             eu_identifier,
        'notice_id':                 notice_id,
        'summary':                   description,
        'notice':                    notice_json,
    }

    return tender, detail


# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main():
    cpv_codes = load_cpv_codes()
    if not cpv_codes:
        print("No CPV codes found in .env (CPV_CODES). Exiting.")
        return

    conn = get_db_connection()
    cursor = conn.cursor()

    print("Fetching API token...")
    token = get_token()

    total_inserted = 0
    total_skipped = 0

    for idx, cpv_code in enumerate(cpv_codes, 1):
        print(f"\n[{idx}/{len(cpv_codes)}] CPV: {cpv_code}")

        for page in range(1, MAX_PAGES + 1):
            print(f"  Page {page}...", end=' ', flush=True)
            data = fetch_publications(cpv_code, page, token)

            if data is None:
                print("API error, skipping.")
                break

            publications = data.get('publications', [])
            total_count = data.get('totalCount', 0)

            if not publications:
                print("No results.")
                break

            print(f"{len(publications)} results (total={total_count})")

            for pub in publications:
                tender, detail = parse_publication(pub)

                if not tender['reference_number']:
                    continue

                try:
                    tender_id = insert_tender(cursor, tender)
                    if tender_id:
                        upsert_tender_detail(cursor, tender_id, detail)
                        total_inserted += 1
                    else:
                        total_skipped += 1
                except mysql.connector.Error:
                    conn = get_db_connection()
                    cursor = conn.cursor()
                    tender_id = insert_tender(cursor, tender)
                    if tender_id:
                        upsert_tender_detail(cursor, tender_id, detail)
                        total_inserted += 1
                    else:
                        total_skipped += 1

            conn.commit()

            # Stop if we've fetched all pages
            if page * ITEMS_PER_PAGE >= total_count or len(publications) < ITEMS_PER_PAGE:
                break

    cursor.close()
    conn.close()
    print(f"\nDone. Inserted: {total_inserted}, Skipped (duplicates): {total_skipped}")


if __name__ == '__main__':
    main()
