# %% [code]
from __future__ import annotations

import os
import time
from pathlib import Path
from typing import Iterator
from urllib.parse import urlencode

import pandas as pd
import requests

# Western Ghats bounding box (decimal degrees). Slightly conservative on the
# north end to avoid the Vindhya Ranges being mis-included; conservative on the
# east to exclude the Deccan Plateau.
WG_BBOX = {"lat_min": 8.0, "lat_max": 22.0, "lon_min": 72.0, "lon_max": 78.0}

# GBIF dataset keys for the four sources that are train-set-clean for our
# 3-model ensemble. iNat is the bulk of the records (~4475 India total before
# filtering); the other three combined contribute ~50 India records.
DATASETS: dict[str, str] = {
    "iNaturalist Research-grade": "50c9509d-22c7-4a22-a47d-8c48425ef4a7",
    # Other clean sources are smaller; they're included optionally and queried
    # only if iNat returns < MIN_RECORDS or the user sets INCLUDE_SMALL_SOURCES.
    # Resolve their dataset keys via:
    # GET https://api.gbif.org/v1/dataset/search?q=<name>&type=OCCURRENCE
    # and look at the result with the matching publishingOrganization.
}

# Optional smaller sources (~50 India records combined). Enable only if needed.
INCLUDE_SMALL_SOURCES: bool = False

GBIF_SEARCH = "https://api.gbif.org/v1/occurrence/search"
BIRD_TAXON_KEY = 212 # GBIF taxon key for class Aves
PAGE_SIZE = 300

# Rate Limiting (be polite; GBIF and iNat CDNs are free).
GBIF_REQ_DELAY_SEC = 0.3
DOWNLOAD_DELAY_SEC = 0.2
HTTP_TIMEOUT_SEC = 60

# Endemics common-name list, reused from src/species.py and inlined here so the
# script is self-contained (paste-into-Kaggle-notebook usage).
ENDEMIC_COMMON_NAMES: set[str] = {
    "Malabar Whistling-Thrush", "Malabar Grey Hornbill", "Nilgiri Flycatcher",
    "Black-and-orange Flycatcher", "White-bellied Sholakili", "Nilgiri Laughingthrush",
    "Wayanad Laughingthrush", "Malabar Trogon", "Malabar Parakeet",
    "Crimson-backed Sunbird", "Nilgiri Pipit", "Nilgiri Flowerpecker",
    "White-bellied Blue Flycatcher", "Rufous Babbler", "Grey-headed Bulbul",
    "Yellow-throated Bulbul", "Flame-throated Bulbul", "Malabar Lark",
    "Broad-tailed Grassbird", "Nilgiri Wood-Pigeon", "Painted Bush-Quail",
    "Black-throated Munia", "Sri Lanka Bay-Owl", "Malabar Barbed Owlet",
    "Malabar Starling", "Malabar Woodshrike", "Wynaad Laughingthrush",
    "Nilgiri Sholakili", "Indian Yellow Tit", "Vigors's Sunbird",
}

# Taxonomy-split synonyms. iNat sometimes uses newer (split) genus names than
# BirdCLEF's Clements 2021 baseline. Map both directions so we catch both.
TAXONOMY_SYNONYMS: dict[str, list[str]] = {
    # Clements : iNat / IOC / split alternatives
    "Pterorhinus delesserti": ["Garrulax delesserti"],
    "Trochalopteron cachinnans": ["Pterorhinus cachinnans", "Garrulax cachinnans"],
    "Sholicola major": ["Brachypteryx major"],
    "Sholicola albiventris": ["Brachypteryx albiventris"],
    "Brachypodius priocephalus": ["Pycnonotus priocephalus"],
    "Rubigula gularis": ["Pycnonotus gularis", "Pycnonotus melanicterus gularis"],
    "Argya subrufa": ["Turdoides subrufa"],
    "Iole indica": ["Acritillas indica"],
}

def resolve_paths() -> tuple[Path, Path]:
    """Return (birdclef_dir, output_dir), Kaggle-aware."""
    if Path("/kaggle/input").exists() and Path("/kaggle/working").exists():
        bc = Path("/kaggle/input/competitions/birdclef-2024")
        out = Path("/kaggle/working")
    else:
        repo = Path(__file__).resolve().parents[1]
        bc = Path(os.environ.get("WG_BIRDCLEF_DIR", repo / "data" / "raw" / "birdclef-2024"))
        out = Path(os.environ.get("WG_GBIF_OUT_DIR", repo / "data" / "external" / "gbif"))
    out.mkdir(parents=True, exist_ok=True)
    return bc, out

def load_birdclef_species(birdclef_dir: Path) -> pd.DataFrame:
    """Return one row per BirdCLEF 2024 species: primary_label, scientific_name, common_name."""
    meta_path = birdclef_dir / "train_metadata.csv"
    if not meta_path.exists():
        raise FileNotFoundError(
            f"Missing {meta_path}. On Kaggle, attach the birdclef-2024 dataset. "
            "Locally, set WG_BIRDCLEF_DIR or download train_metadata.csv."
        )
    meta = pd.read_csv(meta_path)
    cols = ["primary_label", "scientific_name", "common_name"]
    return meta[cols].drop_duplicates().reset_index(drop=True)

def build_species_lookup(species_df: pd.DataFrame) -> dict[str, dict]:
    """Lower-cased scientific name -> {primary_label, scientific_name, common_name, is_endemic, source_match}."""
    # Includes synonym entries so iNat records using newer (split) genus names
    # still resolve to the BirdCLEF eBird code.
    lookup: dict[str, dict] = {}
    for r in species_df.itertuples(index=False):
        info = {
            "primary_label": r.primary_label,
            "scientific_name": r.scientific_name,
            "common_name": r.common_name,
            "is_endemic": r.common_name in ENDEMIC_COMMON_NAMES,
            "match_via": "primary",
        }
        lookup[r.scientific_name.strip().lower()] = info
        for syn in TAXONOMY_SYNONYMS.get(r.scientific_name, []):
            lookup[syn.strip().lower()] = {**info, "match_via": f"synonym:{syn}"}
    return lookup

def gbif_paged_search(params: dict) -> Iterator[dict]:
    """Iterate occurrence records across GBIF pagination."""
    offset = 0
    while True:
        q = dict(params, limit=PAGE_SIZE, offset=offset)
        url = f"{GBIF_SEARCH}?{urlencode(q, doseq=True)}"
        r = requests.get(url, timeout=HTTP_TIMEOUT_SEC)
        r.raise_for_status()
        j = r.json()
        for rec in j.get("results", []):
            yield rec
        if j.get("endOfRecords", False) or not j.get("results"):
            break
        offset += PAGE_SIZE
        time.sleep(GBIF_REQ_DELAY_SEC)

def query_clean_audio_records(species_lookup: dict[str, dict]) -> list[dict]:
    """Pull all WG-bbox + clean-source + Aves audio records, retain only species in BirdCLEF 2024."""
    base = {
        "country": "IN",
        "mediaType": "Sound",
        "taxonKey": BIRD_TAXON_KEY,
        "decimalLatitude": f"{WG_BBOX['lat_min']},{WG_BBOX['lat_max']}",
        "decimalLongitude": f"{WG_BBOX['lon_min']},{WG_BBOX['lon_max']}",
    }
    out: list[dict] = []
    seen_keys: set[str] = set()
    for ds_name, ds_key in DATASETS.items():
        params = dict(base, datasetKey=ds_key)
        n_total = 0
        n_kept = 0
        print(f"[GBIF] querying {ds_name}...")
        for rec in gbif_paged_search(params):
            n_total += 1
            sci = (rec.get("species") or rec.get("scientificName") or "").strip()
            if not sci:
                continue
            info = species_lookup.get(sci.lower())
            if info is None:
                # Try genus + first epithet only (drops subspecies tail).
                parts = sci.split()
                if len(parts) >= 2:
                    info = species_lookup.get(" ".join(parts[:2]).lower())
                if info is None:
                    continue
            media = [
                m for m in (rec.get("media") or [])
                if m.get("type") == "Sound" and m.get("identifier")
            ]
            if not media:
                continue
            occ_key = str(rec.get("gbifID") or rec.get("key") or "")
            if occ_key in seen_keys:
                continue
            seen_keys.add(occ_key)

            rec["_source"] = ds_name
            rec["_audio_url"] = media[0]["identifier"]
            rec["_format"] = media[0].get("format", "")
            rec["_species_info"] = info
            out.append(rec)
            n_kept += 1
        print(f"      scanned {n_total}, kept {n_kept} matching BirdCLEF species")
    return out

def download_audio_records(records: list[dict], out_dir: Path) -> pd.DataFrame:
    """Download the audio files (resumable). Returns the eval index DataFrame."""
    audio_root = out_dir / "wg_eval_audio"
    audio_root.mkdir(exist_ok=True)
    rows: list[dict] = []
    n_total = len(records)
    n_downloaded = 0
    n_already = 0
    n_failed = 0

    for i, rec in enumerate(records, 1):
        info = rec["_species_info"]
        ebird = info["primary_label"]
        sp_dir = audio_root / ebird
        sp_dir.mkdir(exist_ok=True)
        occ_id = str(rec.get("gbifID") or rec.get("key") or i)
        ext = ".mp3" if "mpeg" in rec["_format"].lower() or rec["_audio_url"].endswith(".mp3") else ".wav"
        path = sp_dir / f"{occ_id}{ext}"

        if path.exists() and path.stat().st_size > 0:
            n_already += 1
        else:
            try:
                r = requests.get(rec["_audio_url"], timeout=HTTP_TIMEOUT_SEC, stream=True)
                r.raise_for_status()
                path.write_bytes(r.content)
                n_downloaded += 1
                time.sleep(DOWNLOAD_DELAY_SEC)
            except Exception as e: # noqa: BLE001 - want to skip and continue on any download error
                n_failed += 1
                print(f" [skip] {ebird}/{occ_id}: {e}")
                continue

        rows.append({
            "occurrence_id": occ_id,
            "species_code": ebird,
            "scientific_name": info["scientific_name"],
            "common_name": info["common_name"],
            "is_endemic": info["is_endemic"],
            "match_via": info["match_via"],
            "latitude": rec.get("decimalLatitude"),
            "longitude": rec.get("decimalLongitude"),
            "recorded_date": rec.get("eventDate"),
            "locality": rec.get("locality"),
            "source": rec["_source"],
            "audio_path": str(path.relative_to(out_dir)),
            "audio_url": rec["_audio_url"],
            "audio_format": rec["_format"],
        })

        if i % 50 == 0:
            print(f" progress: {i}/{n_total} (new={n_downloaded}, cached={n_already}, failed={n_failed})")
    
    print(f" [done] new={n_downloaded}, cached={n_already}, failed={n_failed}")
    return pd.DataFrame(rows)

def write_outputs(index_df: pd.DataFrame, species_df: pd.DataFrame, out_dir: Path) -> None:
    """Persist the index, the per-species summary, and the missing-species list."""
    out_idx = out_dir / "wg_eval_index.parquet"
    index_df.to_parquet(out_idx, index=False)
    print(f"[wrote] {out_idx} ({len(index_df)} rows)")

    summary = (
        index_df.groupby("species_code", as_index=False)
        .size()
        .rename(columns={"size": "n_clips"})
    )
    summary = species_df.merge(
        summary, left_on="primary_label", right_on="species_code", how="left"
    ).fillna({"n_clips": 0})
    summary["n_clips"] = summary["n_clips"].astype(int)
    summary["is_endemic"] = summary["common_name"].isin(ENDEMIC_COMMON_NAMES)
    summary = summary.sort_values(["is_endemic", "n_clips"], ascending=[False, False])

    out_sum = out_dir / "wg_eval_summary.csv"
    summary[[
        "primary_label", "scientific_name", "common_name",
        "is_endemic", "n_clips"
    ]].to_csv(out_sum, index=False)
    print(f"[wrote] {out_sum}")

    missing = summary.query("n_clips == 0").copy()
    out_missing = out_dir / "wg_eval_missing.csv"
    missing[["primary_label", "scientific_name", "common_name", "is_endemic"]].to_csv(
        out_missing, index=False
    )
    print(f"[wrote] {out_missing} ({len(missing)} BirdCLEF species with 0 clean clips)")

    n_endemic_total = int(summary["is_endemic"].sum())
    n_endemic_covered = int(((summary["is_endemic"]) & (summary["n_clips"] > 0)).sum())
    print()
    print("=== summary ===")
    print(f"BirdCLEF species:      {len(summary)}")
    print(f" with >=1 clip:        {(summary['n_clips'] > 0).sum()}")
    print(f" with >=2 clips:       {(summary['n_clips'] >= 2).sum()}")
    print(f"Endemic species:       {n_endemic_total}")
    print(f" with >=1 clip:        {n_endemic_covered}")
    print(f" with >=2 clips:       {((summary['is_endemic']) & (summary['n_clips'] >= 2)).sum()}")
    print(f"Total clips:           {len(index_df)}")
    if len(index_df):
        print(f"Median clips/species: {int(summary.query('n_clips > 0')['n_clips'].median())}")

def main() -> None:
    birdclef_dir, out_dir = resolve_paths()
    print(f"BirdCLEF dir: {birdclef_dir}")
    print(f"Output dir:   {out_dir}")
    print()

    species_df = load_birdclef_species(birdclef_dir)
    print(f"BirdCLEF 2024 species loaded: {len(species_df)}")
    species_lookup = build_species_lookup(species_df)
    print(f"Lookup keys (incl. synonyms): {len(species_lookup)}")
    print()

    records = query_clean_audio_records(species_lookup)
    print()
    print(f"Total clean records matching BirdCLEF species: {len(records)}")
    if not records:
        print("No records found. Verify BirdCLEF species list and bbox.")
        return

    index_df = download_audio_records(records, out_dir)
    write_outputs(index_df, species_df, out_dir)

if __name__ == "__main__":
    main()
