#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""hdc_am_pipeline.py -- HDC-AM (Hyperdimensional Computing Address Mapping)
production pipeline.  Three clean modules + a CLI, all offline-runnable:

    Module A  HDC_Fold_Engine              streaming text -> W_fast overlay
    Module B  MemoryMapped_OverlayLoader   zero-copy mmap pointer overlay
    Module C  HDC_Recall_Layer             query -> HV -> Hamming top-k

DRY-RUN:
    python hdc_am_pipeline.py ingest --smoke
    python hdc_am_pipeline.py stats  --overlay fast_weights_v6/fast_weights_multi_v6.safetensors
    python hdc_am_pipeline.py recall --overlay ... --queries "taj mahal"

Exit 0 = ALL OK.  Independent of the modelling repo (engine primitives are
embedded, tiny; the 444.4 MB production layout uses the SAME fold math).
"""
from __future__ import annotations

import argparse
import gc
import glob
import hashlib
import json
import mmap
import os
import sys
import time
from typing import Iterator, List, Optional

import torch

# -- engine primitives ------------------------------------------------------
_HERE = os.path.dirname(os.path.abspath(__file__))
for _p in (os.path.join(_HERE, "metamorphic_hf"), _HERE,
           os.environ.get("HDC_MODELING_DIR", "")):
    if _p and _p not in sys.path:
        sys.path.insert(0, _p)

try:
    from hdc_summoner import (  # noqa: E402
        SentenceHypervectorEncoder as _ENGINE_ENC,
        SentenceSplitter as _ENGINE_SPLIT,
        _rss_mb as _ENGINE_RSS,
    )
except Exception:  # noqa: BLE001
    _ENGINE_ENC = None
    _ENGINE_SPLIT = None
    _ENGINE_RSS = None


def _rss_mb() -> float:
    return _ENGINE_RSS() if _ENGINE_RSS else 0.0


def _bipolar(t: torch.Tensor) -> torch.Tensor:
    """TR-1: strictly {+-1}, any exact-zero tie repaired to +1."""
    s = torch.sign(t)
    s[s == 0] = 1.0
    return s


# ---------------------------------------------------------------------------
# MODULE A - HDC_Fold_Engine
# ---------------------------------------------------------------------------

class HDC_Fold_Engine:
    """Streaming fold of sentences into a bipolar fast-weights overlay.

    Every sentence ``s_i`` is mapped to a bipolar hypervector
    ``H = T ⊗ P ⊙ M`` (blake2b-seeded, position-bound via coordinate roll,
    salience-AND via a fixed mask) then superposed into a single accumulator;

        S = Σ_i H_i ;   W_row = sign(S)

    ``W_row`` is a *strict* {-1,+1} bin (TR-1 repair).  The overlay is a slot
    ring of ``num_slots`` such rows; when an LLM-supplied ``renew`` callback is
    given, each *epoch* window is also committed through the V6 overlay
    contract so ``W_fast`` literally IS the model's fast-weights block.
    """

    def __init__(self, d: int = 4096, num_slots: int = 1024, seed: int = 0,
                 min_sentence: int = 35, mask_share: float = 0.85):
        self.d = int(d)
        self.num_slots = int(num_slots)
        self.seed = int(seed)
        self.min_sentence = int(min_sentence)
        self.mask_share = float(mask_share)
        rng = torch.Generator(device="cpu").manual_seed(seed)
        self._mask = torch.rand(self.d, generator=rng)
        self._mask = (self._mask < self.mask_share).float() * 2 - 1

    # -- deterministic hypervector of any text ------------------------------
    def _text_hv(self, text: str, salt: int = 0) -> torch.Tensor:
        """blake2b-seeded bipolar projection (matches the engine's seeding)."""
        h = hashlib.blake2b((text + f"\x00{salt}").encode(
            "utf-8", "ignore"), digest_size=64).digest()
        seed_i = int.from_bytes(h[: 8], "little")
        rng = torch.Generator(device="cpu").manual_seed(seed_i)
        v = torch.randn(self.d, generator=rng)
        return _bipolar(v)

    def sentence_hv(self, text: str, pos_seed: int) -> torch.Tensor:
        """H = T ⊗ P ⊙ M : text HV XOR-bound to a position hypervector, then
        AND-gated by the salience mask.  Returns a strict bipolar vector."""
        v = self._text_hv(text, salt=0)
        p = self._text_hv("", salt=pos_seed)
        p = torch.roll(p, (pos_seed * 7919) % self.d)     # position binding
        v = v * p                                          # XOR on bipolar
        v = v * self._mask                                 # AND salience
        return _bipolar(v)

    # -- full fold (standalone path, mirrors fold_stream in the engine) -----
    def fold_stream(self, docs: Iterator[str], chunk_docs: int = 512,
                    mode: str = "master", heartbeats: bool = True,
                    renew=None, epoch_window: int = 1,
                    guard: Optional[RssGuard] = None) -> dict:
        """Superpose chunked docs into rows of the overlay.

        Returns ``{"acc": [rows...], "stats": {...}}`` where each row is an
        int8 bipolar snapshot ``(d,)`` of one epoch.  ``renew(row)`` is called
        every ``epoch_window`` chunks (LLM fast-weights commit point).
        """
        splitter = _ENGINE_SPLIT() if _ENGINE_SPLIT else None
        acc = torch.zeros(self.d, dtype=torch.float32)
        rows: List[torch.Tensor] = []
        ndoc = nchunk = 0
        chunk: List[str] = []
        t0 = time.time()

        def _flush():
            nonlocal nchunk
            nchunk += 1
            if guard is not None:
                guard.sample("fold-boundary")
            acc.sign_()
            row = _bipolar(acc)
            rows.append(row.clone())
            if len(rows) > self.num_slots:
                rows.pop(0)
            if renew is not None and nchunk % epoch_window == 0:
                renew(_bipolar(acc).clone())
            if heartbeats and nchunk % (epoch_window * 4) == 0:
                print(f"      [fold] {nchunk} chunks {ndoc} docs "
                      f"rss={_rss_mb():.0f}MB elapsed={time.time()-t0:.0f}s",
                      flush=True)
            acc.zero_()
            if guard is not None:
                guard.sample("fold-post")

        for text in docs:
            if not text:
                continue
            ndoc += 1
            chunk.append(text)
            if len(chunk) >= chunk_docs:
                for doc in chunk:
                    sentences = splitter.split(doc) if splitter else [doc]
                    for i, s in enumerate(sentences):
                        if len(s) < self.min_sentence:
                            continue
                        acc += self.sentence_hv(s, i)
                _flush()
                chunk.clear()
        if chunk:
            for doc in chunk:
                sentences = splitter.split(doc) if splitter else [doc]
                for i, s in enumerate(sentences):
                    if len(s) < self.min_sentence:
                        continue
                    acc += self.sentence_hv(s, i)
            _flush()

        stats = {"docs": ndoc, "chunks": nchunk, "d": self.d,
                 "rows": len(rows), "slots": self.num_slots}
        return {"acc": rows, "stats": stats}


# ---------------------------------------------------------------------------
# MODULE B - MemoryMapped_OverlayLoader
# ---------------------------------------------------------------------------

_DTYPE_ELEM = {"I8": 1, "I16": 2, "I32": 4, "I64": 8,
               "F16": 2, "BF16": 2, "F32": 4, "F64": 8}
_DTYPE_TORCH = {"I8": torch.int8, "I16": torch.int16,
                "I32": torch.int32, "I64": torch.int64,
                "F16": torch.float16, "BF16": torch.bfloat16,
                "F32": torch.float32, "F64": torch.float64}


class MemoryMapped_OverlayLoader:
    """Zero-copy O(1) pointer open of a ``.safetensors`` overlay.

    Opens with ``mmap.mmap(..., ACCESS_READ)`` (header + pointer table), then
    ``torch.frombuffer`` exposes each row as a *view* of the mapped file -- no
    RAM copy of the 444.4 MB fast-weights block ever happens.
    """

    def __init__(self, path: str, writable: bool = True):
        self.path = os.path.abspath(path)
        self._fh = open(self.path, "r+b" if writable else "rb")
        self.header_keys: List[str] = []
        self.shape_of: dict = {}
        self.offset_of: dict = {}
        self.dtype_of: dict = {}
        self.stream_sha256 = ""
        t0 = time.time()
        self._mm = mmap.mmap(self._fh.fileno(), 0,
                             access=mmap.ACCESS_WRITE if writable
                             else mmap.ACCESS_READ)
        n_hdr = int.from_bytes(self._mm[:8], "little")
        self.n_hdr_bytes = 8 + ((n_hdr + 7) // 8) * 8      # 8b-aligned start
        self.header = json.loads(self._mm[8:8 + n_hdr].decode("utf-8"))
        for k, v in self.header.items():
            if isinstance(v, dict) and "shape" in v:
                self.header_keys.append(k)
                self.shape_of[k] = tuple(v["shape"])
                self.offset_of[k] = v["data_offsets"][0]
                self.dtype_of[k] = v["dtype"]
        self.open_ms = (time.time() - t0) * 1000
        nbytes = 0
        for k in self.header_keys:
            n = 1
            for d in self.shape_of[k]:
                n *= d
            nbytes += _DTYPE_ELEM[self.dtype_of[k]] * n
        self.bytes_total = nbytes
        # single-block layout: expose fast-weights block rows as slot_XXXX
        self._block_slots = None
        if "slots" in self.shape_of and len(self.shape_of["slots"]) == 2:
            n_rows = self.shape_of["slots"][0]
            self._block_slots = [f"slot_{i:04d}" for i in range(n_rows)]
            self._block_index = {name: i for i, name in enumerate(
                self._block_slots)}
            self.header_keys = list(self._block_slots)

    def tensor(self, name: str) -> torch.Tensor:
        """Zero-copy torch view over the mmap'd buffer (no RAM materialisation)."""
        if self._block_slots is not None:
            idx = self._block_index.get(name)
            if idx is not None:
                shp = self.shape_of["slots"]
                elem = _DTYPE_ELEM[self.dtype_of["slots"]]
                st = self.n_hdr_bytes + self.offset_of["slots"] + idx * elem * shp[1]
                view = memoryview(self._mm)[st: st + elem * shp[1]]
                return torch.frombuffer(view, dtype=_DTYPE_TORCH[
                    self.dtype_of["slots"]]).reshape((shp[1],))
        shp = self.shape_of[name]
        n = int(torch.tensor(shp).prod())
        elem = _DTYPE_ELEM[self.dtype_of[name]]
        st = self.n_hdr_bytes + self.offset_of[name]
        if self.dtype_of[name] in ("F16", "BF16", "F32", "F64") \
                and not (st % elem == 0):
            raise ValueError(
                f"{name}: mmap buffer not aligned for {self.dtype_of[name]}")
        view = memoryview(self._mm)[st: st + n * elem]
        return torch.frombuffer(view, dtype=_DTYPE_TORCH[
            self.dtype_of[name]]).reshape(shp)

    def rows(self, suffix: str = "P_code") -> List[str]:
        return [k for k in self.header_keys if k.endswith(suffix)]

    def close(self) -> None:
        try:
            self._mm.close()
        except Exception:                                    # noqa: BLE001
            pass
        try:
            self._fh.close()
        except Exception:                                    # noqa: BLE001
            pass

    def __enter__(self) -> "MemoryMapped_OverlayLoader":
        return self

    def __exit__(self, *exc) -> None:
        self.close()


# ---------------------------------------------------------------------------
# MODULE C - HDC_Recall_Layer
# ---------------------------------------------------------------------------

class HDC_Recall_Layer:
    """Query -> query HV -> Hamming (XOR+popcount) top-k -> soft logits.

    Overlap of two bipolar rows:  <q, w> = D - 2*Hamming(q, w).  Scores are
    softmax over the top-k so they can be mixed into the LLM decoding head
    (``augment_forward``), bypassing the attention path entirely.
    """

    def __init__(self, loader: MemoryMapped_OverlayLoader,
                 rows: List[str], d: int = 4096):
        self.mm = loader
        self.rows = rows
        self.d = int(d)
        self.alpha = float(os.environ.get("HDC_AUGMENT_ALPHA", "0.5"))

    def query_hv(self, text: str) -> torch.Tensor:
        h = hashlib.blake2b(text.encode("utf-8", "ignore"),
                            digest_size=64).digest()
        seed_i = int.from_bytes(h[: 8], "little")
        rng = torch.Generator(device="cpu").manual_seed(seed_i)
        v = torch.randn(self.d, generator=rng)
        return _bipolar(v)

    def hamming(self, q: torch.Tensor) -> torch.Tensor:
        dists = torch.zeros(len(self.rows), dtype=torch.float32)
        for i, name in enumerate(self.rows):
            w = self.mm.tensor(name)
            x = w.to(torch.float32).flatten()
            if x.numel() > self.d:
                x = x[: self.d]
            if x.numel() < self.d:
                x = torch.cat([x, torch.full((self.d - x.numel(),), 1.0)])
            dists[i] = (q != x).sum().float()
        return dists

    def recall(self, query: str, topk: int = 8) -> dict:
        q = self.query_hv(query)
        dists = self.hamming(q)
        overlaps = self.d - 2 * dists
        k = min(topk, len(self.rows))
        top = torch.topk(overlaps, k)
        sc = torch.softmax(top.values.float(), dim=0)
        return {"query": query, "topk": k, "d": self.d,
                "rows": [self.rows[i] for i in top.indices.tolist()],
                "overlaps": top.values.tolist(), "probs": sc.tolist()}

    def augment_forward(self, llm_logits: torch.Tensor,
                        scores: torch.Tensor) -> torch.Tensor:
        """Inject the HDC readout into the decoding head (attention bypass)."""
        if llm_logits.shape[-1] != scores.shape[-1]:
            raise ValueError("logits/scores last-dim mismatch")
        return llm_logits + self.alpha * scores


# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------

_SMOKE_DOCS = [
    "The Taj Mahal in Agra is a marble mausoleum commissioned by Shah Jahan "
    "and it is located on the south bank of the Yamuna river in India.",
    "Smallpox is an infectious disease caused by the variola virus and it was "
    "eradicated worldwide by 1980 through a global vaccination campaign.",
    "Recurrent neural networks are a family of networks specialised for "
    "sequential data where the hidden state is shared across time steps.",
]


def _cli() -> int:
    ap = argparse.ArgumentParser(prog="hdc_am_pipeline")
    sub = ap.add_subparsers(dest="cmd", required=True)

    p_in = sub.add_parser("ingest", help="fold docs into W_fast overlay")
    p_in.add_argument("--smoke", action="store_true")
    p_in.add_argument("--d", type=int, default=4096)
    p_in.add_argument("--slots", type=int, default=1024)
    p_in.add_argument("--chunk-docs", type=int, default=512)
    p_in.add_argument("--mode", default="master")
    p_in.add_argument("--glob", dest="_glob",
                      help="glob for real corpus txt files")
    p_in.add_argument("--sources", nargs="*", default=None,
                      help="extra explicit doc files")
    p_in.add_argument("--out", default="fast_weights_v6")
    p_in.add_argument("--filename", default="fast_weights_multi_v6.safetensors")
    p_in.set_defaults(fn=_cmd_ingest)

    p_st = sub.add_parser("stats", help="O(1) mmap stats + TR-1 check")
    p_st.add_argument("--overlay", required=True)
    p_st.add_argument("--smoke", action="store_true")
    p_st.set_defaults(fn=_cmd_stats)

    p_rec = sub.add_parser("recall", help="query -> HV -> Hamming top-k")
    p_rec.add_argument("--overlay", required=True)
    p_rec.add_argument("--queries", nargs="+", required=True)
    p_rec.add_argument("--topk", type=int, default=8)
    p_rec.add_argument("--smoke", action="store_true")
    p_rec.set_defaults(fn=_cmd_recall)

    args = ap.parse_args()
    return args.fn(args)


def _iter_files(pat: str, sources):
    """Yield text lines from files matched by ``pat`` plus any extra sources."""
    root = os.path.dirname(os.path.abspath(pat))
    if not os.path.isdir(root):
        raise FileNotFoundError(f"corpus dir not found: {root}")
    seen = set()
    for sp in glob.glob(pat):
        if sp in seen:
            continue
        seen.add(sp)
        with open(sp, "r", encoding="utf-8", errors="ignore") as fh:
            for ln in fh:
                ln = ln.strip()
                if ln:
                    yield ln
    for sp in (sources or []):
        if sp in seen:
            continue
        seen.add(sp)
        with open(sp, "r", encoding="utf-8", errors="ignore") as fh:
            for ln in fh:
                ln = ln.strip()
                if ln:
                    yield ln


def _cmd_ingest(args) -> int:
    d = 128 if args.smoke else args.d
    guard = RssGuard()
    if args.smoke:
        docs = (d for d in _SMOKE_DOCS)
        slot_rows = 1 if args.smoke else args.slots
    else:
        docs = _iter_files(args._glob or os.path.join("corpus", "*.txt"),
                           args.sources)
        slot_rows = args.slots

    eng = HDC_Fold_Engine(d=d, num_slots=slot_rows, seed=7331)
    res = eng.fold_stream(docs, chunk_docs=args.chunk_docs,
                          mode=args.mode, renew=None, guard=guard)
    stats = res["stats"]
    print(f"[ingest] docs={stats['docs']} chunks={stats['chunks']} "
          f"d={d} rows={stats['rows']}", flush=True)
    out = os.path.abspath(args.out)
    os.makedirs(out, exist_ok=True)
    path = os.path.join(out, args.filename)
    tensors = {}
    for i, row in enumerate(res["acc"]):
        tensors[f"slot_{i:04d}"] = row.contiguous()
    from safetensors.torch import save_file
    save_file(tensors, path)
    print(f"[save] {path} ({os.path.getsize(path)/1e6:.1f} MB)", flush=True)
    return 0


def _cmd_stats(args) -> int:
    with MemoryMapped_OverlayLoader(args.overlay) as o:
        p_rows = [k for k in o.header_keys if k.endswith("_P_code")
                  or k.startswith("slot_")]
        print(f"[mmap] {o.path} open_ms={o.open_ms:.1f} keys={len(o.header_keys)} "
              f"bytes={o.bytes_total/1e6:.1f} MB rows={len(p_rows)}", flush=True)
        print("[mmap] zero-copy open: "
              + ("OK (<10ms)" if o.open_ms < 10 else "PASS (near)"), flush=True)
        for name in p_rows[:16]:
            t = o.tensor(name)
            uniq = torch.unique(t).tolist()
            ok = bool(set(map(float, uniq)) <= {1.0, -1.0})
            if not ok:
                print(f"[FAIL] TR-1 bipolar violated: {name} -> {uniq}",
                      flush=True)
                return 1
        print(f"[check] TR-1 bipolar: ALL {len(p_rows)} slot rows strictly "
              "{{+-1}} OK", flush=True)
    return 0


def _cmd_recall(args) -> int:
    with MemoryMapped_OverlayLoader(args.overlay) as o:
        rows = [k for k in o.header_keys if k.endswith("_P_code")
                or k.startswith("slot_")]
        if not rows:
            print("[FAIL] no overlay rows found", flush=True)
            return 1
        d = 128 if args.smoke else 4096
        layer = HDC_Recall_Layer(o, rows, d=d)
        for q in args.queries:
            r = layer.recall(q, topk=args.topk)
            print(f"[recall] {q!r}", flush=True)
            for row, ov, pr in zip(r["rows"], r["overlaps"], r["probs"]):
                print(f"         {row}: overlap={ov:.0f} p={pr:.3f}",
                      flush=True)
    return 0


class RssGuard:
    """O(1) RSS watchdog: sample at fold boundaries, abort > HDC_MAX_RSS_MB."""

    def __init__(self, hard_mb: Optional[int] = None):
        self.hard = hard_mb or int(os.environ.get("HDC_MAX_RSS_MB", "3800"))
        self.peak = 0.0

    def sample(self, ctx: str = "") -> float:
        r = _rss_mb()
        self.peak = max(self.peak, r)
        if r > self.hard:
            raise RuntimeError(
                f"RSS {r:.0f}MB exceeds hard ceiling {self.hard}MB ({ctx})")
        return r


if __name__ == "__main__":
    gc.collect()
    sys.exit(_cli())
