# -*- coding: utf-8 -*-
"""ingest_streamer.py -- High-throughput streaming knowledge ingest for HDC.

Continuously pulls massive Wikipedia / web-crawl text dumps, cleans them, and
streams the extracted (subject, predicate, object) fact-triples directly into
the Metamorphic BitNet-HDC Ingestion Engine (:mod:`hdc_summoner`).

Design contract
---------------
* Zero-RAM footprint : every stage is a generator; text is streamed page-by-page
  / line-by-line out of compressed archives. Only the current sentence and one
  bounded async queue are ever live in RAM.
* Noise reduction    : HTML/Wiki markup, templates, refs, navboxes and boilerplate
  footers are stripped deterministically; sentences shorter than ``--min-len``
  (default 35 chars) never reach the engine.
* Orientation        : sources are pluggable and compose:
      - ``wikidump:PATH``  -> streaming ``pages-articles-multistream.xml.bz2``
                              (namespace-aware, redirect-skipping, chunk-streamed)
      - ``hf:NAME[:CONFIG[:SPLIT]]`` -> HuggingFace ``datasets`` with
                              ``streaming=True`` (e.g. hf:wikipedia:20220301.simple)
      - ``files:PATH|GLOB`` -> plain ``.txt`` and/or ``.jsonl`` crawl dumps
      - ``jsonl:PATH``     -> explicit newline-delimited JSON source
* Engine integration  : the engine's actual API (hdc_summoner, not the paper) is
      - ``BatchedInjector(model, encoder, scheduler, batch=512, ...)``  (note:
        "batch_size" is ``batch``; "num_overlays" lives on the OverlayScheduler)
      - every 10,000 applied facts -> ``FastWeightsStore.save(model, step, meta)``
* Safe by default     : the base ``metamorphic_packed.safetensors`` is never
  written; ``--restore`` resumes from the latest fast-weight checkpoint only.

Usage
-----
    python ingest_streamer.py --ckpt /content/metamorphic-bitnet-2b \
        --source wikidump:/data/simplewiki/pages-articles-multistream.xml.bz2 \
        --source hf:wikipedia:20220301.simple \
        --source files:crawl/*.jsonl \
        --out-dir fast_weights_stream --batch 512 --overlays 256 \
        --save-every 10000 --anneal

Exit 0 on clean completion or Ctrl-C flush; 1 on fatal engine failure.
"""

from __future__ import annotations

import argparse
import asyncio
import bz2
import glob
import html
import json
import os
import queue
import re
import sys
import threading
import time
from typing import Iterable, Iterator, Optional, Tuple


# --------------------------------------------------------------------------
# 1. Deterministic text cleaning (HTML / wiki markup / boilerplate)
# --------------------------------------------------------------------------

_TAG_RE = re.compile(r"<[^>]*>")
_WS_RE = re.compile(r"\s+")
_WIKI_LINK_IMPL = re.compile(r"\[{2}[^\]|]*\|([^\]]*)\]{2}")
_WIKI_LINK_RAW = re.compile(r"\[{2}([^\]|]*)\]{2}")
_WIKI_CAT = re.compile(r"\[{2}\s*(?:Category|Image|File|Media|Template)"
                       r"[^\[\]]*\]{2}", re.I)
_REF_INLINE = re.compile(r"<ref[^/>]*/>")
_REF_BLOCK = re.compile(r"<ref[^>]*>.*?</ref>", re.DOTALL | re.I)
_HEADING = re.compile(r"^\s*=+\s*.*?\s*=+\s*$")
_BULLET = re.compile(r"^[\s#:*|;]*")
_NS_PREFIX = re.compile(r"(</?)ns[0-9]*:")

# Boilerplate / junk sentence starts that never carry factual load.
_BOILER = (
    "see also", "references", "reference", "external links", "further reading",
    "notes and references", "references and notes", "notes and citations",
    "this article", "this page", "this section", "jump to", "navigation",
    "coordinates", "wikimedia commons", "commons category", "wikidata",
    "link rot", "dead link", "citation needed", "citation",
    "this table may", "the following table", "in other projects",
    "template:", "category:", "file:", "image:", "wikipedia:",
    "privacy policy", "mobile view", "desktop view", "v t e",
    "main article", "see the list", "for other uses", "definition",
    "gallery", "retrieved", "retrieved from", "archived", "harv error",
    "external links and", "bibliography", "footnote", "endnotes",
    "this biography", "stub", "empty section", "further information",
)

# Predictable marker strings that signal junk even mid-sentence.
_JUNK_MARKERS = (
    "[citation needed]", "[disambiguation needed]", "[when?]", "[which?]",
    "[clarification needed]", "[dubious - discuss]", "[verification needed]",
    "{{{", "{{citation", "colspan", "rowspan",
)


_TEMPLATE_RE = re.compile(r"{{(?:[^{}]|{(?:[^{}]|{[^{}]*})*})*}}", re.S)

def _clean_article(text: str) -> str:
    """Return a single normalized block of clean text for one article/page."""
    for _ in range(6):                            # strip nested {{...}} templates
        nxt = _TEMPLATE_RE.sub(" ", text)
        if nxt == text:
            break
        text = nxt
    text = text.replace("{{", " ").replace("}}", " ")
    text = _REF_BLOCK.sub(" ", text)              # <ref>..</ref> blocks first
    text = _REF_INLINE.sub(" ", text)
    text = _TAG_RE.sub(" ", text)                 # all remaining html/xml tags
    text = _WIKI_CAT.sub(" ", text)               # categories / files / images
    text = _WIKI_LINK_IMPL.sub(r"\1", text)       # [[a|b]] -> b
    text = _WIKI_LINK_RAW.sub(r"\1", text)        # [[a]] -> a
    for marker in _JUNK_MARKERS:              # familiar junk signatures
        text = text.replace(marker, " ")
    text = text.replace("{{", " {{ ").replace("}}", " }} ")
    text = html.unescape(text)                # &amp; &mdash; &nbsp;
    lines = []
    for raw in text.splitlines():
        line = raw.strip()
        line = _BULLET.sub("", line).strip()
        line = re.sub(r"^=+\s*", "", line)             # inline ==heading== glue
        line = re.sub(r"\s*=+\s*", " ", line).strip()
        if not line or _HEADING.match(line):
            continue
        if line.startswith("<") or len(line) < 8:
            continue
        if set(line) <= set("=-~*#:;[]{}|^<>"):     # decorative rule
            continue
        lo = line.lower()
        if any(lo.startswith(b) for b in _BOILER):
            continue
        lines.append(line)
    return _WS_RE.sub(" ", " ".join(lines))


def clean_sentences(text: str, carry: str, splitter,
                    min_len: int) -> Tuple[Iterator[str], str]:
    """Split one cleaned text block into whole sentences, stitching ``carry``.

    Returns ``(sentence_iter, new_carry)``. Long sentences crossing the block
    boundary are held in ``carry`` until the next block completes them, so no
    valid fact is ever cut at a chunk edge.
    """
    cleaned = _clean_article(text)
    if not cleaned:
        return iter(()), carry
    buffer = (carry + " " + cleaned).strip()
    parts = list(splitter._END.split(buffer))
    sent_iter = (p.strip() for p in parts[:-1] if len(p.strip()) >= min_len)
    return sent_iter, parts[-1].strip() if parts else ""


# --------------------------------------------------------------------------
# 2. Streaming sources (all generators, O(1) memory, async-capable)
# --------------------------------------------------------------------------

def _iter_clean_pages(source) -> Iterator[Tuple[str, str]]:
    """Duck-typed hook: each source yields (title, cleaned_text) tuples."""
    yield from source.pages()


class WikiDumpStream:
    """Streams ``<page>`` blocks out of a ``pages-articles-multistream.xml.bz2``.

    Decompression is chunked (BZ2File is C-backed); pages are parsed with a
    tiny pull state machine (no DOM, no in-memory document). Title/text pairs
    are yielded as fast as they complete.
    """

    def __init__(self, path: str, chunk_bytes: int = 1 << 20):
        self.path = path
        self.chunk = int(chunk_bytes)

    # -- sync core ------------------------------------------------------
    def pages(self) -> Iterator[Tuple[str, str]]:
        if not os.path.isfile(self.path):
            raise FileNotFoundError("wikidump file not found: " + self.path)
        with bz2.BZ2File(self.path, "r") as fh:
            in_page = False
            page_buf: list[str] = []
            for raw in fh:                          # one decompressed line 2B
                line = raw.decode("utf-8", errors="replace")
                line = _NS_PREFIX.sub(r"\1", line)
                low = line.lstrip().lower()
                opened = "<page" in low
                closed = "</page" in low
                if in_page and opened and not closed:
                    # page opened inside current page without close -> push tail
                    page_buf.append(line)
                    continue
                if opened:
                    if in_page:                     # flush orphan block safely
                        page_buf.append(line)
                    in_page = True
                    page_buf = [line]
                    if closed:                      # single-line <page>..</page>
                        p = self._parse_page(page_buf)
                        if p:
                            yield p
                        in_page = False
                        page_buf = []
                    continue
                if in_page:
                    if closed:
                        page_buf.append(line)
                        p = self._parse_page(page_buf)
                        if p:
                            yield p
                        in_page = False
                        page_buf = []
                    else:
                        page_buf.append(line)
            if in_page:                             # truncated tail -> drop
                return

    @staticmethod
    def _parse_page(lines: list[str]) -> Optional[Tuple[str, str]]:
        title = ""
        nstxt = "0"
        text_parts: list[str] = []
        in_text = False
        for ln in lines:
            low = ln.lstrip().lower()
            if "<ns" in low:
                m = re.search(r"<ns>([^<]*)</ns>", low)
                nstxt = m.group(1).strip() if m else nstxt
            m = re.search(r"<title>(.*?)</title>", ln, flags=re.I | re.S)
            if m:
                title = html.unescape(m.group(1)).strip()
            if in_text:
                t = m_text_end = re.search(r"</text", low)
                if t:
                    text_parts.append(ln.split("</text")[0])
                    in_text = False
                else:
                    text_parts.append(ln)
            elif "<text" in low:
                rest = ln.split("<text", 1)[1]
                body = re.split(r">", rest, maxsplit=1)
                if len(body) == 2:
                    text_parts.append(body[1])
                    in_text = not body[1].rstrip().lower().endswith("</text")
                elif rest.rstrip().lower().endswith("/>"):
                    in_text = False
        if nstxt not in ("0", ""):                  # talk / user / project ns
            return None
        if not title:
            return None
        if title.lower().startswith(("template:", "category:", "file:", "image:",
                                     "wikipedia:", "module:", "portal:")):
            return None
        return title, "\n".join(text_parts)

    # -- async façade ----------------------------------------------------
    async def _afill(self, q: "queue.Queue[object]") -> None:
        try:
            for page in self.pages():
                q.put(page)
        finally:
            q.put(None)

    def pages_async(self) -> Iterator[Tuple[str, str]]:
        """Iterator view over an event-loop producer in a background thread."""
        yield from _async_bridge(self._afill)


class HuggingFaceWikipediaStream:
    """Streaming HuggingFace dataset (``datasets`` lib with streaming=True).

    The hub component is never materialised: rows are pulled one at a time from
    the parquet/arrow shards by the ``datasets`` streaming backend.
    """

    def __init__(self, dataset: str, config: Optional[str] = None,
                 split: str = "train"):
        self.dataset = dataset
        self.config = config
        self.split = split
        self._ds = None

    def _load(self):
        if self._ds is None:
            try:
                from datasets import load_dataset
            except ImportError as e:
                raise RuntimeError(
                    "'datasets' not installed -- pip install datasets") from e
            kw = dict(split=self.split, streaming=True)
            if self.config:
                kw["name"] = self.config
            self._ds = load_dataset(self.dataset, **kw)

    def pages(self) -> Iterator[Tuple[str, str]]:
        self._load()
        for row in self._ds:                        # streaming iterator
            text = row.get("text") or row.get("content") or ""
            if not text:
                continue
            title = row.get("title") or ""
            yield title, str(text)

    async def _afill(self, q: "queue.Queue[object]") -> None:
        try:
            for page in self.pages():
                q.put(page)
        finally:
            q.put(None)

    def pages_async(self) -> Iterator[Tuple[str, str]]:
        yield from _async_bridge(self._afill)


class FileSystemDumpIterator:
    """Multi-GB crawl dumps: iterates files, one line in memory at a time.

    ``.jsonl`` rows may be plain objects with a text field, or raw strings.
    ``.txt`` files are treated as raw text blocks (cleaned + sentence-split).
    """

    def __init__(self, patterns: list[str], encoding: str = "utf-8",
                 text_fields: tuple[str, ...] = (
                     "text", "content", "sentence", "body", "article")):
        paths: list[str] = []
        for pat in patterns:
            if any(c in pat for c in "*?["):
                paths.extend(glob.glob(pat))
            elif os.path.isfile(pat):
                paths.append(pat)
            else:
                print("[warn] no such file/glob: " + pat, file=sys.stderr)
        if not paths:
            raise FileNotFoundError("no crawl files matched: " + str(patterns))
        self.paths = sorted(set(paths))
        self.encoding = encoding
        self.text_fields = text_fields

    def pages(self) -> Iterator[Tuple[str, str]]:
        for path in self.paths:
            if path.lower().endswith(".jsonl") or path.lower().endswith(".json"):
                yield from self._jsonl_pages(path)
            else:
                yield from self._txt_pages(path)

    def _jsonl_pages(self, path: str) -> Iterator[Tuple[str, str]]:
        with open(path, "r", encoding=self.encoding, errors="replace") as fh:
            for i, raw in enumerate(fh, 1):
                raw = raw.strip()
                if not raw:
                    continue
                try:
                    obj = json.loads(raw)
                except json.JSONDecodeError:
                    continue                           # malformed row -> skip
                if isinstance(obj, str):
                    yield (path.split(os.sep)[-1] + ":" + str(i), obj)
                elif isinstance(obj, dict):
                    for key in self.text_fields:
                        if key in obj and isinstance(obj[key], str):
                            yield (str(obj.get("title", "")), obj[key])
                            break

    def _txt_pages(self, path: str) -> Iterator[Tuple[str, str]]:
        with open(path, "r", encoding=self.encoding, errors="replace") as fh:
            buf: list[str] = []
            size = 0
            for raw in fh:
                buf.append(raw)
                size += len(raw)
                if size >= (1 << 20):                  # ~1 MB page-window
                    yield (path.split(os.sep)[-1], "".join(buf))
                    buf, size = [], 0
            if buf:
                yield (path.split(os.sep)[-1], "".join(buf))


# --------------------------------------------------------------------------
# 3. Async bridge (event loop in a worker thread + bounded queue)
# --------------------------------------------------------------------------

def _async_bridge(afill, maxsize: int = 128) -> Iterator:
    """Consume ``await q``-style producer as a plain generator, zero clobber."""
    q: "queue.Queue[object]" = queue.Queue(maxsize=maxsize)

    def _runner() -> None:
        asyncio.run(afill(q))

    t = threading.Thread(target=_runner, daemon=True)
    t.start()
    while True:
        item = q.get()
        if item is None:
            break
        yield item


# --------------------------------------------------------------------------
# 4. Engine integration (hdc_summoner) + checkpoint scheduler
# --------------------------------------------------------------------------

def build_engine(ckpt_dir: Optional[str], batch: int, overlays: int,
                 anneal: bool, device, out_dir: str, restore: bool,
                 ckpt_base: Optional[str], smoke: bool = False):
    """Instantiate model / encoder / scheduler / injector / store once."""
    import torch  # deferred so --help stays light
    import sys as _s
    _s.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
    import argparse as _ap
    import hdc_summoner as hd

    model, tag, dev = hd.load_model(_ap.Namespace(ckpt=ckpt_dir,
                                                  smoke=bool(smoke)))
    device = dev if device is None else device
    D = int(getattr(model.config, "d_hd", 4096))
    encoder = hd.SentenceHypervectorEncoder(D, device=device)
    scheduler = hd.OverlayScheduler(overlays=overlays, anneal=anneal)
    stats = hd.IngestStats()
    injector = hd.BatchedInjector(model, encoder, scheduler,
                                  batch=int(batch), device=device, stats=stats)
    base = ckpt_base or (ckpt_dir if not smoke else None)
    store = hd.FastWeightsStore(out_dir,
                                ckpt_dir=base if base and os.path.isdir(base) else None)
    if restore:
        ckpts = sorted(glob.glob(os.path.join(out_dir,
                                              "fast_weights_[0-9]*.safetensors")))
        if ckpts:
            latest = ckpts[-1]
            n = store.load(model, latest)
            m = scheduler.meta()
            print("[restore] %s (%d tensors) step=%d" %
                  (os.path.basename(latest), n, m.get("steps", 0)), flush=True)
            from hdc_summoner import _rss_mb
            stats.rss_start = _rss_mb()
        else:
            print("[restore] no checkpoint yet - starting fresh", flush=True)
    return model, encoder, scheduler, stats, injector, store, tag


def run(ckpt_dir, sources: list[str], out_dir: str, batch: int = 512,
        overlays: int = 256, save_every: int = 10000, max_facts: int = 0,
        page_cap: int = 0, min_len: int = 35, anneal: bool = False,
        restore: bool = False, encoding: str = "utf-8",
        ckpt_base: Optional[str] = None, device=None, smoke: bool = False) -> int:
    """Main pipeline. Returns 0 on completion/graceful stop."""
    # ---- engine ------------------------------------------------------------
    import torch  # (retained for future device/math use)
    model, encoder, scheduler, stats, injector, store, _tag = build_engine(
        ckpt_dir, batch, overlays, anneal, device, out_dir, restore,
        ckpt_base, smoke=smoke)

    from hdc_summoner import SentenceSplitter, TripletExtractor, _rss_mb

    splitter = SentenceSplitter()
    extractor = TripletExtractor()
    sources_obj = build_sources(sources, encoding)

    # ---- streaming pump ----------------------------------------------------
    print("[pipe] batch=%d overlays=%d save_every=%d min_len=%d sources=%d" %
          (batch, overlays, save_every, min_len, len(sources_obj)), flush=True)
    t0 = time.time()
    last_t = t0
    next_cp = save_every
    pages_seen = 0
    sentences = 0
    carry = ""

    def _checkpoint(step: int) -> str:
        path = store.save(model, step, dict(scheduler.meta(),
                                            facts_applied=stats.injected,
                                            pages=pages_seen,
                                            sentences=sentences))
        rss = stats.sample_rss()
        print("[ckpt] step=%d facts=%d %s (rss %.0fMB)" %
              (step, stats.injected, path, rss), flush=True)
        return path

    try:
        for src in sources_obj:
            pages_iter = src.pages_async() if hasattr(src, "pages_async") else src.pages()
            for title, text in pages_iter:
                pages_seen += 1
                if page_cap and pages_seen > page_cap:
                    break
                sents, carry = clean_sentences(text, carry, splitter, min_len)
                triple_gen = extractor.extract(sents)
                for triple in triple_gen:
                    sentences += 1
                    injector.add(triple)                # dedups + auto-flush
                    if stats.injected >= next_cp:
                        _checkpoint(scheduler.steps)
                        next_cp += save_every
                    if max_facts and stats.injected >= max_facts:
                        raise KeyboardInterrupt        # graceful target stop
                now = time.time()
                if now - last_t >= 20:                  # bounded telemetry
                    rss = stats.sample_rss()
                    print("[stream] sentences=%d facts=%d parse=%d skip=%d "
                          "dups=%d | %.0f facts/s | rss %.0fMB" %
                          (sentences, stats.injected, extractor.parseable,
                           extractor.skipped, stats.dups,
                           stats.facts_per_sec, rss), flush=True)
                    last_t = now
            if carry:
                sents, carry = clean_sentences("", carry, splitter, min_len)
                for triple in extractor.extract(sents):
                    sentences += 1
                    injector.add(triple)

        injector.flush()
        _checkpoint(scheduler.steps)                    # final checkpoint
    except KeyboardInterrupt:
        injector.flush()
        _checkpoint(scheduler.steps)
        print("\n[stop] interrupted - flushed + checkpointed", flush=True)
    finally:
        dt = max(time.time() - t0, 1e-9)
        print("\n[summary] pages=%d sentences=%d injected=%d parseable=%d "
              "skipped=%d dups=%d | %.1f facts/s | wall %.1fs | rss %.0fMB" %
              (pages_seen, sentences, stats.injected, extractor.parseable,
               extractor.skipped, stats.dups, stats.injected / dt, dt,
               _rss_mb()), flush=True)
    return 0


def build_sources(specs: list[str], encoding: str) -> list:
    srcs = []
    for spec in specs:
        kind, _, rest = spec.partition(":")
        if kind == "wikidump":
            srcs.append(WikiDumpStream(rest))
        elif kind == "hf":
            parts = rest.split(":") if rest else []
            dataset = parts[0] if parts else "wikipedia"
            config = parts[1] if len(parts) > 1 else None
            split = parts[2] if len(parts) > 2 else "train"
            srcs.append(HuggingFaceWikipediaStream(dataset, config, split))
        elif kind in ("files", "glob"):
            srcs.append(FileSystemDumpIterator(rest.split(","),
                                               encoding=encoding))
        elif kind == "jsonl":
            srcs.append(FileSystemDumpIterator([rest], encoding=encoding))
        else:
            raise ValueError("unknown source spec: %r (use wikidump:/hf:/files:/jsonl:)" % spec)
    return srcs


def main(argv: Optional[list[str]] = None) -> int:
    ap = argparse.ArgumentParser(
        prog="ingest_streamer.py",
        description=__doc__.split("\n\n")[0],
        formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--ckpt", default=None,
                    help="packed checkpoint dir (or --smoke to use tiny model)")
    ap.add_argument("--smoke", action="store_true",
                    help="use the engine's tiny CI model instead of a ckpt")
    ap.add_argument("--source", action="append", required=True, dest="sources",
                    help="source spec: wikidump:PATH | hf:NAME[:CONFIG[:SPLIT]] "
                         "| files:GLOB1,GLOB2 | jsonl:PATH")
    ap.add_argument("--out-dir", default="fast_weights_stream")
    ap.add_argument("--ckpt-dir", default=None,
                    help="base sidecar dir for the SHA-256 pin (default: --ckpt)")
    ap.add_argument("--batch", type=int, default=512)
    ap.add_argument("--overlays", type=int, default=256)
    ap.add_argument("--save-every", type=int, default=10000,
                    help="auto-fast-weight checkpoint every N applied facts")
    ap.add_argument("--max-facts", type=int, default=0,
                    help="hard stop after N injected facts (0 = unlimited)")
    ap.add_argument("--page-cap", type=int, default=0,
                    help="read at most N pages per source (0 = unlimited)")
    ap.add_argument("--min-len", type=int, default=35,
                    help="minimum cleaned-sentence length in chars")
    ap.add_argument("--anneal", action="store_true",
                    help="anneal overlay coefficients over epochs")
    ap.add_argument("--restore", action="store_true",
                    help="resume from the latest fast-weight checkpoint")
    ap.add_argument("--encoding", default="utf-8")
    args = ap.parse_args(argv)

    if not args.smoke and not args.ckpt:
        ap.error("--ckpt DIR is required unless --smoke is used")

    if args.smoke:
        import sys as _s
        _s.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
        import hdc_summoner as hd
        out = os.path.join(os.path.dirname(os.path.abspath(__file__)),
                           "stream_smoke")
        os.makedirs(out, exist_ok=True)
        print("[smoke] using tiny engine model", flush=True)
        return run(None, args.sources, out_dir=out, batch=args.batch,
                   overlays=args.overlays, save_every=args.save_every,
                   max_facts=args.max_facts, page_cap=args.page_cap,
                   min_len=args.min_len, anneal=args.anneal,
                   restore=args.restore, encoding=args.encoding,
                   smoke=True)
    return run(args.ckpt, args.sources, out_dir=args.out_dir,
               batch=args.batch, overlays=args.overlays,
               save_every=args.save_every, max_facts=args.max_facts,
               page_cap=args.page_cap, min_len=args.min_len,
               anneal=args.anneal, restore=args.restore,
               encoding=args.encoding, ckpt_base=args.ckpt_dir)


if __name__ == "__main__":
    sys.exit(main())