"""Shared definitions for the place-search index (Photon JSON dump -> SQLite FTS5).

This file is the ODbL "filter script": KEEP says which OpenStreetMap objects go
into the index, NAME_KEYS which of their name tags are searchable, COUNTRIES
which countries are kept. fold() must stay identical to fold() in
Server/tracknavigator-api/src/search.js; fold_spec.json holds the vectors both
implementations are tested against (python3 test_fold.py, npm test).
"""
import unicodedata

SCHEMA_VERSION = 1

# osm_key -> osm_value -> kind. The kinds are the API's `kind` values.
KEEP = {
    "natural": {"peak": "peak", "volcano": "peak", "hill": "peak", "fell": "peak", "massif": "peak",
                "mountain_range": "range", "saddle": "pass", "pass": "pass", "glacier": "glacier",
                "water": "lake", "valley": "valley", "gorge": "valley", "ridge": "terrain", "arete": "terrain",
                "bare_rock": "terrain", "cliff": "terrain", "rock": "terrain", "plateau": "terrain",
                "cave_entrance": "terrain", "cape": "terrain", "bay": "terrain", "beach": "terrain",
                "peninsula": "terrain", "geyser": "terrain", "hot_spring": "terrain"},
    "mountain_pass": {"yes": "pass"},
    "water": {"lake": "lake", "reservoir": "lake", "pond": "lake", "lagoon": "lake", "glacial_lake": "lake"},
    "waterway": {"waterfall": "waterfall", "river": "river"},
    "tourism": {"alpine_hut": "hut", "wilderness_hut": "hut", "guest_house": "lodge", "hostel": "lodge",
                "hotel": "lodge", "chalet": "lodge", "camp_site": "camp", "viewpoint": "viewpoint"},
    "amenity": {"shelter": "hut"},
    "highway": {"trailhead": "trailhead"},
    "place": {"city": "village", "town": "village", "village": "village", "hamlet": "village",
              "isolated_dwelling": "village", "suburb": "village", "municipality": "village",
              "locality": "area", "region": "area", "island": "area", "islet": "area", "archipelago": "area"},
    "boundary": {"national_park": "area", "protected_area": "area"},
}

NAME_KEYS = ("name", "name:en", "name:de", "name:fr", "name:it", "name:es", "name:nb", "name:no", "name:nn",
             "name:sv", "name:fi", "name:is", "name:ne", "name:se", "alt_name", "old_name", "int_name",
             "loc_name", "official_name", "short_name", "reg_name", "alt_name:en", "name:rm", "name:sl")

# Europe including the UK, without Russia and Turkey; plus Nepal and the US.
EUROPE = ("al ad at by be ba bg hr cy cz dk ee fo fi fr de gi gr gg hu is ie im it je xk lv li lt lu mt md "
          "mc me nl mk no pl pt ro sm rs sk si es sj se ch ua gb va ax").split()
COUNTRIES = tuple(EUROPE + ["np", "us"])

# Devanagari signs that unicode61 would otherwise treat as word breaks.
DEVANAGARI_TOKENCHARS = "ऀँंःऺऻ़ऽािीुूृॄॅॆेैॉॊोौ्ॎॏ॒॑॓॔ॕॖॗॢॣ"

SPECIAL = str.maketrans({"þ": "th", "ð": "d", "ß": "ss", "ø": "o", "æ": "ae", "œ": "oe", "đ": "d",
                         "ł": "l", "ı": "i", "ŧ": "t", "ŋ": "n", "’": "'", "‘": "'"})


def fold(s):
    """Lower-case; þ ð ß ø æ œ đ ł ı ŧ ŋ spelled out; NFKD; code points with a
    non-zero canonical combining class dropped (accents, virama, nukta — NOT
    the Devanagari vowel signs, which have class 0); every punctuation,
    symbol, separator and control/format character becomes a space; runs of
    spaces collapse."""
    s = s.lower().translate(SPECIAL)
    s = unicodedata.normalize("NFKD", s)
    s = "".join(c for c in s if not unicodedata.combining(c))
    s = "".join(" " if unicodedata.category(c)[0] in "PSZC" else c for c in s)
    return " ".join(s.split())


def compact(folded):
    return folded.replace(" ", "")


def is_latin(s):
    """True when the first letter of s is Latin script (or s has no letters)."""
    for c in s:
        if c.isalpha():
            return "LATIN" in unicodedata.name(c, "")
    return True


def _latin_first(*values):
    vals = [v for v in values if v]
    for v in vals:
        if is_latin(v):
            return v
    return vals[0] if vals else ""


def region_of(address):
    """'County, State' in Latin script where the dump has it."""
    ad = address or {}
    county = _latin_first(ad.get("county:en"), ad.get("county"))
    state = _latin_first(ad.get("state:en"), ad.get("state"))
    parts = [p for p in (county, state) if p]
    if len(parts) == 2 and fold(parts[0]) == fold(parts[1]):
        parts = parts[:1]
    return ", ".join(parts)


def parse_ele(v):
    if v is None or v == "":
        return None
    try:
        return float(str(v).replace(",", ".").split()[0])
    except (ValueError, IndexError):
        return None


def name_variants(names):
    """Raw name strings in NAME_KEYS order, ';'-lists split, duplicates dropped."""
    out = []
    for k in NAME_KEYS:
        v = names.get(k)
        if not v:
            continue
        for part in v.split(";"):
            part = part.strip()
            if part and part not in out:
                out.append(part)
    return out


def index_texts(variants):
    """(names_f, fts text, trigram text, exact keys) for a list of raw names."""
    folded = []
    for v in variants:
        fv = fold(v)
        if fv and fv not in folded:
            folded.append(fv)
    comp = [compact(x) for x in folded if " " in x]
    fts_t = " | ".join(folded + comp)
    tri_t = " | ".join(dict.fromkeys(compact(x) for x in folded))
    exact = sorted({compact(x) for x in folded if x})
    return folded, fts_t, tri_t, exact


def place_row(c, countries):
    """One Photon `content` item -> a stage row tuple, or None when it is not kept."""
    kind = KEEP.get(c.get("osm_key"), {}).get(c.get("osm_value"))
    if not kind:
        return None
    cc = (c.get("country_code") or "").lower()
    if countries is not None and cc not in countries:
        return None
    nm = c.get("name") or {}
    if not nm.get("name") and not nm.get("name:en"):
        return None
    ex = c.get("extra") or {}
    if ex.get("alt_name:en"):
        nm = dict(nm, **{"alt_name:en": ex["alt_name:en"]})
    folded, fts_t, tri_t, exact = index_texts(name_variants(nm))
    if not folded:
        return None
    lon, lat = c["centroid"]
    return (f'{c.get("object_type", "?")}{c.get("object_id", "")}', c["osm_key"], c["osm_value"], kind,
            nm.get("name") or nm.get("name:en"), nm.get("name:en"), "\x1f".join(folded),
            lat, lon, parse_ele(ex.get("ele")), float(c.get("importance") or 0), 1 if ex.get("wikidata") else 0,
            cc, region_of(c.get("address")), fts_t, tri_t, "\x1f".join(exact))
