#!/usr/bin/env python3
"""Build the place-search index (SQLite, FTS5) that GET /search in tracknavigator-api reads.

From a GraphHopper Photon JSON dump (https://download1.graphhopper.com/public/):
    python3 build.py --out places.sqlite --input photon-dump-planet-*.jsonl.zst
    zstd -dcq photon-dump-nepal.jsonl.zst | python3 build.py --out nepal.sqlite --countries all

From an index built earlier (re-cut to other countries, or upgrade the 2026-09-28
prototype schema, which kept raw names in `names`):
    python3 build.py --out regions.sqlite --from-sqlite planet.sqlite

--countries: 'default' (Europe incl. the UK, without ru/tr, plus np and us),
'all', or a comma list of ISO codes.

Schema (meta.schema = 1), see README.md:
  f      one row per kept OSM object. The rowid IS the importance rank
         (1 = most important), so "top N matches by importance" is a plain
         `ORDER BY rowid LIMIT N` on the FTS tables with no sort or join.
  fts    FTS5 unicode61 (+ Devanagari tokenchars), prefix index 2/3/4, contentless
  tri    FTS5 trigram over the space-free name forms, contentless
  exact  compact folded name -> id (WITHOUT ROWID)
  meta   k/v: schema, built_at, source, countries, fold_spec, license

The output is written to <out>.tmp and renamed over <out> at the end, and it is
left in rollback-journal mode (no WAL), so the API can open it read-only and a
server copy can be swapped in with one rename.
"""
import argparse, json, multiprocessing as mp, os, sqlite3, subprocess, sys, time
from datetime import datetime, timezone

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from placesindex import (COUNTRIES, DEVANAGARI_TOKENCHARS, KEEP, SCHEMA_VERSION, index_texts, place_row)

ODBL_NOTICE = ("Contains information from OpenStreetMap (https://www.openstreetmap.org), which is made available "
               "here under the Open Database License (ODbL) 1.0: https://opendatacommons.org/licenses/odbl/1-0/. "
               "© OpenStreetMap contributors. This place-name index is a Derivative Database of OpenStreetMap, "
               "also under the ODbL; the filter that produced it is Scripts/search-index/ in the TrackNavigator "
               "repository.")

# Cheap byte pre-filter: osm_key sits in the first few hundred bytes of each line.
PRE = [f'"osm_key":"{k}"'.encode() for k in KEEP]

_countries = None


def _init_worker(countries):
    global _countries
    _countries = countries


def process_batch(lines):
    rows = []
    for line in lines:
        head = line[:400]
        if not any(p in head for p in PRE):
            continue
        d = json.loads(line)
        if d.get("type") != "Place":
            continue
        for c in d.get("content") or []:
            r = place_row(c, _countries)
            if r:
                rows.append(r)
    return rows


def reindex_batch(rows):
    """Rows from an earlier index -> stage rows (fold recomputed from its names)."""
    out = []
    for (osm, key, value, kind, name, name_en, names, lat, lon, ele, imp, wd, cc, region) in rows:
        folded, fts_t, tri_t, exact = index_texts(names.split("\x1f"))
        if folded:
            out.append((osm, key, value, kind, name, name_en, "\x1f".join(folded), lat, lon, ele, imp,
                        1 if wd else 0, cc, region or "", fts_t, tri_t, "\x1f".join(exact)))
    return out


STAGE_COLS = "osm,key,value,kind,name,name_en,names_f,lat,lon,ele,imp,wd,cc,region,fts_t,tri_t"


def stage_insert(db, rows):
    cur = db.cursor()
    for r in rows:
        cur.execute(f"INSERT INTO stage({STAGE_COLS}) VALUES ({','.join('?' * 16)})", r[:16])
        sid = cur.lastrowid
        for k in r[16].split("\x1f"):
            if k:
                cur.execute("INSERT INTO stage_exact(sid,k) VALUES (?,?)", (sid, k))


def dump_lines(path):
    if path is None or path == "-":
        yield from sys.stdin.buffer
        return
    if path.endswith(".zst"):
        p = subprocess.Popen(["zstd", "-dcq", path], stdout=subprocess.PIPE, bufsize=1 << 20)
        yield from p.stdout
        if p.wait() != 0:
            raise SystemExit(f"zstd failed on {path}")
        return
    with open(path, "rb") as f:
        yield from f


def batches(it, n):
    b = []
    for x in it:
        b.append(x)
        if len(b) >= n:
            yield b
            b = []
    if b:
        yield b


def legacy_rows(src, countries):
    s = sqlite3.connect(f"file:{src}?mode=ro", uri=True)
    cols = [r[1] for r in s.execute("PRAGMA table_info(f)")]
    names_col = "names" if "names" in cols else "names_f"
    q = f"SELECT osm,key,value,kind,name,name_en,{names_col},lat,lon,ele,imp,wd,cc,region FROM f"
    args = ()
    if countries is not None:
        q += " WHERE cc IN (%s)" % ",".join("?" * len(countries))
        args = tuple(sorted(countries))
    cur = s.execute(q, args)
    while True:
        b = cur.fetchmany(20000)
        if not b:
            break
        yield b
    s.close()


def finish(stage_path, tmp, meta):
    db = sqlite3.connect(tmp)
    db.executescript(f"""
    PRAGMA journal_mode=OFF; PRAGMA synchronous=OFF; PRAGMA temp_store=FILE; PRAGMA cache_size=-262144;
    CREATE TABLE meta(k TEXT PRIMARY KEY, v TEXT) WITHOUT ROWID;
    CREATE TABLE f(id INTEGER PRIMARY KEY, osm TEXT, key TEXT, value TEXT, kind TEXT, name TEXT, name_en TEXT,
                   names_f TEXT, lat REAL, lon REAL, ele REAL, imp REAL, wd INTEGER, cc TEXT, region TEXT);
    CREATE VIRTUAL TABLE fts USING fts5(t, tokenize="unicode61 remove_diacritics 2 tokenchars '{DEVANAGARI_TOKENCHARS}'",
                                        prefix='2 3 4', content='', detail='none', columnsize=0);
    CREATE VIRTUAL TABLE tri USING fts5(t, tokenize='trigram', content='', detail='none', columnsize=0);
    CREATE TABLE exact(k TEXT, id INTEGER, PRIMARY KEY(k, id)) WITHOUT ROWID;
    """)
    db.execute("ATTACH ? AS s", (stage_path,))
    # rowid = importance rank; ties keep dump order.
    db.executescript("""
    CREATE TEMP TABLE ord(sid INTEGER PRIMARY KEY, id INTEGER);
    INSERT INTO ord(sid, id) SELECT sid, row_number() OVER (ORDER BY imp DESC, sid) FROM s.stage;
    INSERT INTO f(id,osm,key,value,kind,name,name_en,names_f,lat,lon,ele,imp,wd,cc,region)
      SELECT o.id,osm,key,value,kind,name,name_en,names_f,lat,lon,ele,imp,wd,cc,region
      FROM s.stage JOIN ord o USING(sid) ORDER BY o.id;
    INSERT INTO fts(rowid, t) SELECT o.id, fts_t FROM s.stage JOIN ord o USING(sid) ORDER BY o.id;
    INSERT INTO tri(rowid, t) SELECT o.id, tri_t FROM s.stage JOIN ord o USING(sid) ORDER BY o.id;
    INSERT OR IGNORE INTO exact(k, id) SELECT e.k, o.id FROM s.stage_exact e JOIN ord o USING(sid) ORDER BY e.k, o.id;
    """)
    db.commit()
    db.execute("DETACH s")
    db.execute("INSERT INTO fts(fts) VALUES('optimize')")
    db.execute("INSERT INTO tri(tri) VALUES('optimize')")
    n = db.execute("SELECT count(*) FROM f").fetchone()[0]
    meta = dict(meta, features=str(n))
    db.executemany("INSERT INTO meta(k,v) VALUES (?,?)", sorted(meta.items()))
    db.commit()
    db.execute("VACUUM")
    db.execute("PRAGMA journal_mode=DELETE")
    db.close()
    return n


def main():
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--out", required=True)
    ap.add_argument("--input", help="Photon JSON dump (.jsonl or .jsonl.zst); default stdin")
    ap.add_argument("--from-sqlite", help="re-index an existing index file instead of a dump")
    ap.add_argument("--countries", default="default")
    ap.add_argument("--workers", type=int, default=max(2, (os.cpu_count() or 4) - 2))
    a = ap.parse_args()

    countries = None if a.countries == "all" else set(COUNTRIES) if a.countries == "default" \
        else {c.strip().lower() for c in a.countries.split(",") if c.strip()}
    out = os.path.abspath(a.out)
    tmp, stage_path = out + ".tmp", out + ".stage"
    for p in (tmp, stage_path):
        if os.path.exists(p):
            os.remove(p)

    t0 = time.time()
    stage = sqlite3.connect(stage_path)
    stage.executescript(f"""
    PRAGMA journal_mode=OFF; PRAGMA synchronous=OFF;
    CREATE TABLE stage(sid INTEGER PRIMARY KEY, osm TEXT, key TEXT, value TEXT, kind TEXT, name TEXT, name_en TEXT,
                       names_f TEXT, lat REAL, lon REAL, ele REAL, imp REAL, wd INTEGER, cc TEXT, region TEXT,
                       fts_t TEXT, tri_t TEXT);
    CREATE TABLE stage_exact(sid INTEGER, k TEXT);
    """)
    n_keep = 0
    source = a.from_sqlite or a.input or "stdin"
    with mp.Pool(a.workers, initializer=_init_worker, initargs=(countries,)) as pool:
        if a.from_sqlite:
            work = pool.imap(reindex_batch, legacy_rows(a.from_sqlite, countries), chunksize=1)
        else:
            work = pool.imap(process_batch, batches(dump_lines(a.input), 4000), chunksize=1)
        for rows in work:
            stage_insert(stage, rows)
            n_keep += len(rows)
            if n_keep and n_keep % 200000 < len(rows):
                stage.commit()
                print(f"{n_keep:>10,} kept  {time.time() - t0:6.0f}s", file=sys.stderr, flush=True)
    stage.commit()
    stage.close()
    t1 = time.time()

    meta = {
        "schema": str(SCHEMA_VERSION),
        "built_at": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"),
        "source": os.path.basename(source),
        "countries": "all" if countries is None else ",".join(sorted(countries)),
        "fold_spec": str(json.load(open(os.path.join(os.path.dirname(os.path.abspath(__file__)),
                                                     "fold_spec.json"), encoding="utf-8"))["version"]),
        "license": ODBL_NOTICE,
    }
    n = finish(stage_path, tmp, meta)
    os.remove(stage_path)
    os.replace(tmp, out)
    print(f"done: {n:,} places, load {t1 - t0:.0f}s, index {time.time() - t1:.0f}s, "
          f"{os.path.getsize(out) / 1e6:.1f} MB -> {out}", file=sys.stderr)


if __name__ == "__main__":
    main()
