#!/usr/bin/env python3
"""P2.2 probe: full-rebuild wall cost vs source scale, current code path.

Isolated: clones memory/research.db into /mnt/data/tmp/p22/k<N>/ and builds
into the sibling <N>/research.duckdb (the non-DB_PATH rule in
query._resolve_duckdb_path), so production memory/graph.duckdb is never
opened for write. Scale is k-fold row duplication with remapped keys, which
preserves the tag/edge/hyper distribution and join selectivity that drive
the CTAS cost — duplicating uniformly, not synthesising.
"""

import shutil
import sqlite3
import sys
import time
import traceback
from pathlib import Path

sys.path.insert(0, "/home/arun/Research/findata-graph")
from helpers.graph import query as q  # noqa: E402

SCRATCH = Path("/mnt/data/tmp/p22")
SRC = Path("/home/arun/Research/findata-graph/memory/research.db")
INPUTS = ("entities", "graph_edges", "entity_tags", "hyper_edges", "hyper_incidences")
# per-table: which TEXT columns get the ' [k]' copy suffix (PKs remapped,
# value columns such as entity_tags.tag deliberately left alone so distinct
# tag cardinality — and therefore v_node's market_cap subselect — is
# preserved exactly as it scales).
SUFFIX = {
    "entities": ("name", "normalized_name"),
    "entity_tags": ("entity_name",),
    "graph_edges": ("source", "target"),
    "hyper_edges": ("label",),
    "hyper_incidences": ("entity_name",),
}
DROP_COL = {"graph_edges": ("id",), "hyper_edges": ("id",)}


def _cols(con, t):
    return [r[1] for r in con.execute(f"PRAGMA table_info({t})")]


def clone(k: int) -> Path:
    dst = SCRATCH / f"k{k}"
    shutil.rmtree(dst, ignore_errors=True)
    dst.mkdir(parents=True)
    db = dst / "research.db"
    shutil.copy2(SRC, db)
    con = sqlite3.connect(db)
    con.execute("PRAGMA journal_mode=WAL")
    for rep in range(1, k):
        suf = f" [{rep}]"
        for t in ("entities", "graph_edges", "entity_tags", "hyper_edges"):
            cols = [c for c in _cols(con, t) if c not in DROP_COL.get(t, ())]
            expr = ", ".join(f'"{c}" || ?' if c in SUFFIX[t] else f'"{c}"' for c in cols)
            con.execute(
                f'INSERT OR IGNORE INTO "{t}" ({",".join(chr(34)+c+chr(34) for c in cols)}) '
                f"SELECT {expr} FROM \"{t}\"",
                (suf,) * len(SUFFIX[t]),
            )
        # incidences must point at the freshly inserted hyper_edges rows
        con.execute(
            """
            INSERT OR IGNORE INTO hyper_incidences
              (edge_id, entity_name, weight, direction, role, valid_from, valid_to)
            SELECT hn.id, hi.entity_name || ?, hi.weight, hi.direction, hi.role,
                   hi.valid_from, hi.valid_to
            FROM hyper_incidences hi
            JOIN hyper_edges ho ON ho.id = hi.edge_id
            JOIN hyper_edges hn ON hn.edge_type = ho.edge_type
                                AND hn.label = ho.label || ?
            """,
            (suf, suf),
        )
        con.commit()
    con.execute("ANALYZE")
    con.commit()
    con.close()
    return db


def build(db: Path) -> tuple[float, int, int]:
    t0 = time.perf_counter()
    con = q.connect(db_path=db, rebuild=True)
    con.close()
    dt = time.perf_counter() - t0
    import duckdb

    duck = db.with_suffix(".duckdb")
    c = duckdb.connect(str(duck), read_only=True)
    doubled = c.execute("SELECT COUNT(*) FROM e_all_und").fetchone()[0]
    c.close()
    return dt, doubled, duck.stat().st_size


if __name__ == "__main__":
    ks = [int(x) for x in sys.argv[1:]] or [1, 3, 10, 30, 100]
    SCRATCH.mkdir(parents=True, exist_ok=True)
    print(f"{'k':>4} {'entities':>9} {'edges':>9} {'tags':>8} {'e_all_und':>10} "
          f"{'rebuild_s':>9} {'duck_MB':>8}", flush=True)
    for k in ks:
        try:
            t0 = time.perf_counter()
            db = clone(k)
            prep = time.perf_counter() - t0
            c = sqlite3.connect(db)
            n = {t: c.execute(f'SELECT COUNT(*) FROM "{t}"').fetchone()[0] for t in INPUTS}
            c.close()
            dt, doubled, sz = build(db)
            print(f"{k:>4} {n['entities']:>9} {n['graph_edges']:>9} {n['entity_tags']:>8} "
                  f"{doubled:>10} {dt:>9.2f} {sz / 1048576:>8.1f}   "
                  f"(clone {prep:.1f}s)", flush=True)
        except Exception:
            print(f"{k:>4} FAILED after {time.perf_counter() - t0:.1f}s", flush=True)
            traceback.print_exc()
        shutil.rmtree(SCRATCH / f"k{k}", ignore_errors=True)
