"""Public contracts sync — direct from IMPIC dumps on dados.gov.pt.

Replaces the old ptdata.org pipeline. The dataset
`Contratos Públicos - Portal Base - IMPIC` publishes one ZIP per year
(2012-current) biweekly, each containing a single JSON array of contracts.

Upside over ptdata:
  * Same canonical source — IMPIC feeds both ptdata and this dump
  * No mystery 504s — dados.gov.pt is served via CDN, predictable
  * Full 14-year archive available on demand
  * No account / API key / rate-limit friction

Downside: each year is ~13-60 MB compressed (100-500 MB uncompressed JSON).
We stream-parse with `ijson` so memory stays bounded (~1 MB window).

Data flow:
  1. GET dataset metadata from dados.gov.pt API
  2. For each year ZIP whose last_modified > our `contract_sync_state`:
        - Download ZIP → keep in memory (≤60 MB)
        - Unzip → stream-parse the JSON array
        - For every contract object, check if any NIF in `adjudicatarios`
          or `adjudicante` is monitored
        - Upsert matches into `public_contracts`
  3. Refresh `companies.contracts_total` cache from local aggregates
"""
import io
import json
import logging
import re
import zipfile
from datetime import date, datetime
from typing import Any, Iterator

import httpx
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession

logger = logging.getLogger(__name__)

DATASET_ID = "66d72d488ca4b7cb2de28712"
DATASET_URL = f"https://dados.gov.pt/api/1/datasets/{DATASET_ID}/"
HTTP_TIMEOUT = httpx.Timeout(120.0, connect=30.0)  # ZIPs can be 60 MB+

# "504595067 - Escola Profissional Amar Terra Verde, L.da"
_NIF_NAME_RE = re.compile(r"^\s*(\d{9})\s*[-—]\s*(.+?)\s*$")


def _parse_nif_name(s: str) -> tuple[str, str] | None:
    """Entity strings come as `'NIF - Name'`. Split out the 9-digit NIF."""
    if not s:
        return None
    m = _NIF_NAME_RE.match(s)
    if not m:
        return None
    return m.group(1), m.group(2)


def _parse_pt_date(s: Any) -> date | None:
    if not s or not isinstance(s, str):
        return None
    for fmt in ("%d/%m/%Y", "%Y-%m-%d"):
        try:
            return datetime.strptime(s.strip(), fmt).date()
        except ValueError:
            continue
    return None


def _extract_matches(
    contract: dict, monitored: dict[str, str],
) -> list[tuple[str, str]]:
    """Return list of (company_id, role) tuples for every monitored NIF that
    appears as adjudicatário. Private companies (our entire monitoring set)
    are never `adjudicantes` — those are public entities only, so we skip
    that side of the contract for efficiency."""
    out: list[tuple[str, str]] = []
    seen_companies: set[str] = set()
    for entity in contract.get("adjudicatarios") or []:
        parsed = _parse_nif_name(str(entity))
        if not parsed:
            continue
        nif, _ = parsed
        cid = monitored.get(nif)
        if cid and cid not in seen_companies:
            out.append((cid, "supplier"))
            seen_companies.add(cid)
    return out


def _supplier_cols(contract: dict, monitored_cid: str, role: str) -> dict[str, Any]:
    """Build kwargs for INSERT into public_contracts, from one IMPIC contract."""
    # Extract counterpart(s) for display: when the monitored company is
    # supplier, store the awarding entity; when it's awarding, store the suppliers.
    ae_name = ae_nif = None
    sup_names: list[str] = []
    sup_nifs: list[str] = []
    for entity in contract.get("adjudicante") or []:
        parsed = _parse_nif_name(str(entity))
        if parsed:
            ae_nif, ae_name = parsed
            break  # awarding entity is typically a single item
    for entity in contract.get("adjudicatarios") or []:
        parsed = _parse_nif_name(str(entity))
        if parsed:
            sup_nifs.append(parsed[0])
            sup_names.append(parsed[1])
    # CPV codes come as ["31720000-9 - Equipamento electromecânico"] — keep only the code
    cpv_codes: list[str] = []
    for raw in contract.get("cpv") or []:
        code = str(raw).split(" ", 1)[0].strip()
        if code:
            cpv_codes.append(code)
    return {
        "cid": monitored_cid,
        "ext": str(contract.get("idcontrato") or ""),
        "role": role,
        "title": contract.get("objectoContrato") or contract.get("descContrato"),
        "ae_name": ae_name,
        "ae_nif": ae_nif,
        "sup_names": sup_names,
        "sup_nifs": sup_nifs,
        "cp": contract.get("precoContratual"),
        "bp": contract.get("precoBaseProcedimento"),
        "ap": contract.get("PrecoTotalEfetivo"),
        "pd": _parse_pt_date(contract.get("dataPublicacao")),
        "sd": _parse_pt_date(contract.get("dataCelebracaoContrato")),
        "cld": _parse_pt_date(contract.get("dataFechoContrato")),
        "ed": contract.get("prazoExecucao"),
        "proc": contract.get("tipoprocedimento"),
        "cpv": cpv_codes,
        "dc": None,  # not in IMPIC dump (would need geocode from localExecucao)
        "mc": None,
        "raw": json.dumps(contract, ensure_ascii=False, default=str),
    }


async def _monitored_nifs(session: AsyncSession) -> dict[str, str]:
    """Return {nif: company_id} for every monitored company (internal +
    competitor + analysis). Cheap lookup table used against each contract."""
    rows = (
        await session.execute(
            text(
                """
                SELECT nif, id::text
                FROM companies
                WHERE monitored = TRUE
                  -- clientes incluídos: saber que contratos públicos um cliente
                  -- ganha é contexto de solvência, e custa zero (é o mesmo
                  -- dataset, só se acrescentam NIFs ao mapa)
                  AND monitoring_type IN ('internal','competitor','analysis','client')
                """
            )
        )
    ).all()
    return {nif: cid for nif, cid in rows}


def _iter_contracts_from_zip(zip_bytes: bytes) -> Iterator[dict]:
    """Stream-parse the JSON array inside the ZIP. ijson keeps memory flat
    even for 500 MB uncompressed JSON. `use_float=True` avoids Decimal
    objects that would not survive `json.dumps` later."""
    import ijson  # lazy import — ijson is a new dep
    zf = zipfile.ZipFile(io.BytesIO(zip_bytes))
    inner = zf.namelist()[0]
    with zf.open(inner) as f:
        for obj in ijson.items(f, "item", use_float=True):
            yield obj


async def _fetch_dataset_metadata(client: httpx.AsyncClient) -> dict:
    r = await client.get(DATASET_URL)
    r.raise_for_status()
    return r.json()


async def _already_synced(
    session: AsyncSession, resource_id: str, last_modified: datetime,
) -> bool:
    row = (
        await session.execute(
            text(
                """
                SELECT last_modified FROM contract_sync_state
                WHERE resource_id = :rid
                """
            ),
            {"rid": resource_id},
        )
    ).first()
    if not row:
        return False
    return row[0] >= last_modified


async def _record_sync(
    session: AsyncSession,
    resource_id: str, resource_title: str, resource_url: str,
    last_modified: datetime, rows_kept: int, rows_inserted: int,
) -> None:
    await session.execute(
        text(
            """
            INSERT INTO contract_sync_state
                (resource_id, resource_title, resource_url, last_modified,
                 synced_at, rows_kept, rows_inserted)
            VALUES (:rid, :t, :u, :lm, now(), :kept, :ins)
            ON CONFLICT (resource_id) DO UPDATE SET
                resource_title = EXCLUDED.resource_title,
                resource_url   = EXCLUDED.resource_url,
                last_modified  = EXCLUDED.last_modified,
                synced_at      = EXCLUDED.synced_at,
                rows_kept      = EXCLUDED.rows_kept,
                rows_inserted  = EXCLUDED.rows_inserted
            """
        ),
        {
            "rid": resource_id, "t": resource_title, "u": resource_url,
            "lm": last_modified, "kept": rows_kept, "ins": rows_inserted,
        },
    )
    await session.commit()


async def _upsert_one(session: AsyncSession, cols: dict) -> bool:
    if not cols["ext"]:
        return False
    result = await session.execute(
        text(
            """
            INSERT INTO public_contracts (
              company_id, external_id, role, title,
              awarding_entity_name, awarding_entity_nif,
              supplier_names, supplier_nifs,
              contract_price, base_price, actual_price,
              publication_date, signing_date, close_date,
              execution_days, procedure_type, cpv_codes,
              district_code, municipality_code, raw_json,
              source
            )
            VALUES (
              :cid, :ext, :role, :title,
              :ae_name, :ae_nif,
              :sup_names, :sup_nifs,
              :cp, :bp, :ap,
              :pd, :sd, :cld,
              :ed, :proc, :cpv,
              :dc, :mc, CAST(:raw AS JSONB),
              'impic'
            )
            ON CONFLICT (company_id, external_id) DO NOTHING
            RETURNING id
            """
        ),
        cols,
    )
    return result.first() is not None


async def _refresh_company_totals(session: AsyncSession) -> None:
    """Sync `companies.contracts_total` + `contracts_total_value` from our
    local `public_contracts` aggregate. Called at the end of a sync so the
    Intelligence Summary + ContractsSection KPIs stay correct without a
    round-trip to ptdata."""
    await session.execute(text("""
        UPDATE companies c SET
          contracts_total = COALESCE(sub.n, 0),
          contracts_total_value = COALESCE(sub.val, 0),
          contracts_fetched_at = now()
        FROM (
          SELECT
            company_id,
            count(*) FILTER (WHERE role = 'supplier')::int AS n,
            coalesce(sum(contract_price) FILTER (WHERE role = 'supplier'), 0) AS val
          FROM public_contracts
          GROUP BY company_id
        ) sub
        WHERE c.id = sub.company_id
    """))
    # Companies with zero contracts still need their cache zeroed (for the
    # case where ptdata had inflated numbers that don't match local reality).
    await session.execute(text("""
        UPDATE companies c SET
          contracts_total = 0,
          contracts_total_value = 0,
          contracts_fetched_at = now()
        WHERE NOT EXISTS (
          SELECT 1 FROM public_contracts pc WHERE pc.company_id = c.id
        )
    """))
    await session.commit()


async def sync_impic_contracts(
    session: AsyncSession, force: bool = False,
) -> dict[str, Any]:
    """Main entry point. Scans the IMPIC dataset, downloads each pending
    year-file, filters to monitored companies, upserts, records sync state."""
    monitored = await _monitored_nifs(session)
    if not monitored:
        return {"error": "no monitored companies"}
    logger.info("impic sync starting — monitored NIFs: %d", len(monitored))

    stats: dict[str, Any] = {
        "total_resources": 0, "resources_processed": 0, "resources_skipped": 0,
        "rows_inserted": 0, "rows_kept_by_year": {},
    }

    async with httpx.AsyncClient(timeout=HTTP_TIMEOUT, follow_redirects=True) as client:
        meta = await _fetch_dataset_metadata(client)
        zip_resources = [
            r for r in meta.get("resources", [])
            if (r.get("format") or "").lower() == "zip"
        ]
        zip_resources.sort(key=lambda r: r.get("title") or "")
        stats["total_resources"] = len(zip_resources)

        for res in zip_resources:
            rid = res["id"]
            title = res.get("title") or rid
            url = res["url"]
            last_mod = datetime.fromisoformat(
                (res.get("last_modified") or "").replace("Z", "+00:00")
            )
            if not force and await _already_synced(session, rid, last_mod):
                stats["resources_skipped"] += 1
                continue

            logger.info("impic sync: downloading %s (%s bytes)", title, res.get("filesize"))
            r = await client.get(url)
            r.raise_for_status()

            kept = 0
            inserted = 0
            for contract in _iter_contracts_from_zip(r.content):
                matches = _extract_matches(contract, monitored)
                if not matches:
                    continue
                for cid, role in matches:
                    kept += 1
                    cols = _supplier_cols(contract, cid, role)
                    if await _upsert_one(session, cols):
                        inserted += 1
                # Commit per contract keeps lock contention low but adds round-trips;
                # commit every 200 rows instead.
                if kept % 200 == 0:
                    await session.commit()
            await session.commit()

            await _record_sync(
                session, rid, title, url, last_mod, kept, inserted,
            )
            stats["resources_processed"] += 1
            stats["rows_inserted"] += inserted
            stats["rows_kept_by_year"][title] = kept
            logger.info(
                "impic sync: %s → kept=%d new=%d", title, kept, inserted,
            )

    await _refresh_company_totals(session)
    logger.info("impic sync complete: %s", stats)
    return stats
