"""#314 bounded diagnostic, no refresh and no product changes.
Reuse the pinned W2 ownership/teardown rig; only override the experiment.
"""
import argparse
import base64
import json
import os
from pathlib import Path
import struct
import subprocess
import threading
import time
import traceback
import owned_base as base


class LeafRig(base.Rig):
    def env(self):
        return {**super().env(), "LEAF314_CONTROL": self.config["arm"]}

    def census(self, label):
        begin = time.time_ns()
        snapshot = self.native.snapshot()
        root = self.producer_pid
        found = {root}
        while True:
            more = {pid for pid, row in snapshot.items() if row["ppid"] in found}
            if more <= found:
                break
            found |= more
        identities, errors = [], []
        for pid in sorted(found):
            try:
                identities.append(self.native.identity(pid, snapshot))
            except Exception as error:
                errors.append({"pid": pid, "error": repr(error)})
        return {"label": label, "begin_ns": begin, "end_ns": time.time_ns(),
                "root_pid": root, "root_in_snapshot": root in snapshot,
                "descendant_pids": sorted(found - {root}), "identities": identities,
                "errors": errors}

    def raw_sessions(self):
        # Source contract: codec.rs length-prefixed UTF-8 Envelope; transport.rs
        # send_hello is send-only. No session subscription or mutation frame.
        result, errors, finished = [], [], threading.Event()
        def query():
            try:
                path = "\\\\.\\pipe\\" + self.sockets[0]
                with open(path, "r+b", buffering=0) as pipe:
                    def send(kind, payload):
                        body = json.dumps({"protocol_version": 1, "kind": kind, "payload": payload}).encode()
                        data = memoryview(struct.pack(">I", len(body)) + body)
                        while data:
                            n = pipe.write(data)
                            base.require(n, "IPC write made no progress")
                            data = data[n:]
                    def read_exact(size):
                        chunks = bytearray()
                        while len(chunks) < size:
                            block = pipe.read(size - len(chunks))
                            base.require(block, "IPC closed during frame")
                            chunks.extend(block)
                        return bytes(chunks)
                    send("hello", {"protocol_version": 1, "role": "brain"})
                    send("sessions", None)
                    prefix = read_exact(4)
                    size = struct.unpack(">I", prefix)[0]
                    base.require(size <= 16 * 1024 * 1024, "oversized broker response")
                    body = read_exact(size)
                    payload = json.loads(body)
                    base.require(payload["kind"] == "sessions-reply", f"unexpected reply: {payload}")
                    result.append({"wire_base64": base64.b64encode(prefix + body).decode(), "envelope": payload})
            except Exception as error:
                errors.append(repr(error))
            finally:
                finished.set()
        threading.Thread(target=query, daemon=True).start()
        base.require(finished.wait(min(2, self.remaining())), "raw session query exceeded 2s; outer job remains final containment")
        base.require(not errors, f"raw session query failed: {errors}")
        return result[0]

    def capture(self, label):
        with self.capture_lock:
            before = self.census(label + ":before-query")
            request_ns = time.time_ns()
            raw = self.raw_sessions()
            reply_ns = time.time_ns()
            after = self.census(label + ":after-query")
            row = {"label": label, "request_ns": request_ns, "reply_ns": reply_ns,
                   "before": before, "raw": raw, "after": after}
            with (self.out / "session-census.jsonl").open("a", encoding="utf-8") as log:
                log.write(json.dumps(row) + "\n")
            self.captures.append(row)
            return row

    def read_output(self):
        pending = ""
        try:
            with (self.out / "rc.chunks.jsonl").open("a", encoding="utf-8", buffering=1) as raw:
                while True:
                    chunk = os.read(self.rc.stdout.fileno(), 8192)
                    if not chunk:
                        return
                    now_ns, mono = time.time_ns(), time.monotonic()
                    text = chunk.decode("utf-8", "replace")
                    raw.write(json.dumps({"received_ns": now_ns, "received_ms": now_ns // 1000000,
                                          "phase": self.phase, "text": text,
                                          "base64": base64.b64encode(chunk).decode()}) + "\n")
                    pending += text
                    if self.phase == "fresh-attach" and self.decision is None and "defunct session" in pending:
                        immediate = self.census("refusal-output")
                        self.decision = {"kind": "refusal", "received_ns": now_ns, "census": immediate}
                        self.capture("refusal-output")
                    matches = list(base.MARKER.finditer(pending))
                    for match in matches:
                        fields = match.group(1).split()
                        row = {"kind": fields[0], "received_ms": now_ns // 1000000, "received_mono": mono,
                               "generated_ms": int(fields[-1] if fields[0] != "PROG" else fields[2]), "phase": self.phase}
                        if fields[0] == "PROG":
                            row.update(counter=int(fields[1]), generation=fields[3])
                        elif fields[0] == "ACK":
                            row.update(tag=fields[1], ordinal=int(fields[2]))
                        else:
                            row.update(ordinal=int(fields[1]), chain=fields[2])
                        self.records.append(row)
                    if matches:
                        pending = pending[matches[-1].end():]
                    if len(pending) > 16384:
                        pending = pending[-8192:]
        except Exception as error:
            self.reader_errors.append(repr(error))
            self.event("RC_READER_ERROR", error=repr(error))

    def trial(self):
        self.capture_lock = threading.Lock()
        self.captures = []
        self.decision = None
        bind = json.loads((self.home / "generator-bind.json").read_text())
        self.producer_pid = bind["pid"]
        self.producer_birth = self.owned[self.producer_pid]["birth"]
        self.phase = "baseline"
        self.send(0)
        end = time.monotonic() + 3
        while time.monotonic() < end:
            samples = self.samples("baseline")
            if samples and samples[0]["timely"]:
                break
            time.sleep(.02)
        base.require(samples and samples[0]["timely"], "initial attach ACK/state baseline failed")
        base.save(self.out / "baseline.json", samples)
        initial = self.capture("initial-attached")
        self.guard("initial rc detach")
        self.write_input(b"\x02d")
        self.rc.wait(timeout=5)
        self.reader.join(timeout=2)
        base.require(not self.reader.is_alive(), "initial RC reader did not finish")
        self.rc.stdin.close()
        self.rc.stdout.close()
        self.rc = self.reader = None
        self.phase = "aging-without-refresh"
        while True:
            row = self.capture("age-check")
            session = next(s for s in row["raw"]["envelope"]["payload"]["sessions"] if s["endpoint"] == "refreshrig")
            base.require(session["pid"] == self.producer_pid and session["adapter"] == "evidencerig", "wrong session authority")
            if session["spawned_ms_ago"] >= 32000:
                break
            time.sleep(min(.5, self.remaining()))
        self.guard("fresh rc attach past default grace")
        self.phase = "fresh-attach"
        stop = threading.Event()
        def sample():
            try:
                while not stop.is_set():
                    self.capture("around-fresh-attach")
                    stop.wait(.025)
            except Exception as error:
                self.reader_errors.append("sampler: " + repr(error))
        sampler = threading.Thread(target=sample, daemon=True)
        sampler.start()
        self.capture("immediately-before-launch")
        err = (self.out / "rc-fresh.stderr.log").open("wb")
        self.streams.append(err)
        self.event("FRESH_ATTACH_START", epoch_ns=time.time_ns())
        self.rc = subprocess.Popen([str(self.spt), "rc", "refreshrig"], env=self.env(), cwd=self.out,
                                   stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=err, bufsize=0,
                                   creationflags=subprocess.CREATE_NO_WINDOW)
        self.identify(self.rc.pid, "fresh-rc")
        self.reader = threading.Thread(target=self.read_output, daemon=True)
        self.reader.start()
        try:
            end = min(self.deadline, time.monotonic() + 10)
            progress = False
            while time.monotonic() < end:
                base.require(not self.reader_errors, f"capture failed: {self.reader_errors}")
                progress = any(r["kind"] == "PROG" and r["phase"] == "fresh-attach" for r in self.records)
                if self.decision or progress or self.rc.poll() is not None:
                    break
                time.sleep(.005)
            if self.decision is None and progress:
                self.decision = {"kind": "attached", "received_ns": time.time_ns(), "census": self.census("attached-output")}
                self.capture("attached-output")
                self.send(1)
                end = min(self.deadline, time.monotonic() + 3)
                while time.monotonic() < end:
                    samples = self.samples("fresh-attach")
                    if samples and samples[0]["timely"]:
                        break
                    time.sleep(.01)
                base.require(samples and samples[0]["timely"], "fresh attachment lacks timely ACK/state")
                base.save(self.out / "fresh-attach-ack.json", samples)
            elif self.decision:
                self.rc.wait(timeout=3)
                self.reader.join(timeout=3)
            self.capture("after-decision")
        finally:
            stop.set()
            sampler.join(timeout=3)
            base.require(not sampler.is_alive(), "capture sampler exceeded bound")
        base.require(self.decision is not None, "neither attach nor refusal observed")
        base.require(not self.reader_errors, f"reader errors: {self.reader_errors}")
        selected = [r for r in self.captures if r["label"] in ("immediately-before-launch", "around-fresh-attach", "refusal-output", "attached-output", "after-decision")]
        expected_count = 0 if self.config["arm"] == "leaf" else 1
        for row in selected:
            for census in (row["before"], row["after"]):
                base.require(not census["errors"] and census["root_in_snapshot"], "unreadable/dead root census")
                root = next(p for p in census["identities"] if p["pid"] == self.producer_pid)
                base.require(root["birth"] == self.producer_birth, "root identity changed")
                base.require(len(census["descendant_pids"]) == expected_count, "unexpected descendant topology")
            session = next(s for s in row["raw"]["envelope"]["payload"]["sessions"] if s["endpoint"] == "refreshrig")
            base.require(session["pid"] == self.producer_pid and session["adapter"] == "evidencerig" and session["spawned_ms_ago"] >= 30000, "session predicate changed")
        expected = "refusal" if self.config["arm"] == "leaf" else "attached"
        actual = self.decision["kind"]
        # Confirmation pass means the diagnostic prediction held, NOT product correctness.
        return {"arm": self.config["arm"], "prediction_confirmed": actual == expected,
                "observed": actual, "expected_current_defect": expected, "decision": self.decision,
                "capture_count": len(selected), "root_pid": self.producer_pid, "root_birth": self.producer_birth,
                "grace_override": None, "refresh_invocations": 0,
                "temporal_limit": "External census brackets broker replies and is taken immediately on refusal output; not an atomic in-process snapshot of classifier execution.",
                "initial_session": initial["raw"]["envelope"]}


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", required=True)
    args = parser.parse_args()
    config = json.loads(Path(args.config).read_text())
    out = Path(config["output"])
    base.require(out.is_absolute(), "absolute output required")
    out.mkdir(parents=True, exist_ok=False)
    rig = LeafRig(config, out)
    # Initial attach reader uses this before trial initializes capture machinery.
    rig.decision = None
    receipt = {"status": "PRECONDITION", "config": config}
    try:
        base.require(config["arm"] in ("leaf", "descendant"), "unknown arm")
        base.require(config["observation_seconds"] <= 100 and config["teardown_seconds"] <= 60, "phase bounds exceed admission")
        rig.admit()
        rig.start()
        rig.fixture()
        receipt["result"] = rig.trial()
        receipt["status"] = "CONFIRMED" if receipt["result"]["prediction_confirmed"] else "DISCONFIRMED"
    except Exception as error:
        receipt.update(error=repr(error), traceback=traceback.format_exc())
        traceback.print_exc()
    finally:
        try:
            receipt["cleanup"] = rig.cleanup()
            if not receipt["cleanup"]["pass"]:
                receipt["observation_status"] = receipt["status"]
                receipt["status"] = "CLEANUP_FAILED"
        except Exception as error:
            receipt.update(observation_status=receipt["status"], status="CLEANUP_FAILED",
                           cleanup={"pass": False, "error": repr(error), "traceback": traceback.format_exc()})
        rig.events.close()
        receipt["exit_code"] = 0 if receipt["status"] == "CONFIRMED" else 1
        base.save(out / "receipt.json", receipt)
        print(json.dumps({"status": receipt["status"], "receipt": str(out / "receipt.json")}), flush=True)
    return receipt["exit_code"]


if __name__ == "__main__":
    raise SystemExit(main())
