"""XODE Sentinel client — one file, standard library only.

Your decision records never leave your infrastructure. This client hashes each
record locally and sends only the 32-byte hash; the record itself goes into a
journal you own. A receipt is assembled from the two halves: your record, plus
the proof path Sentinel returns.

    from xode_sentinel import Sentinel, policy_id, subject_hash

    POLICY = policy_id(open("credit_policy_v7.md").read(), "model=risk-gbm-2026-09")
    s = Sentinel("xsk_...", journal="sentinel-journal.jsonl")

    rec = s.record(kind="loan_screen",
                   subject=subject_hash("customer-88213", salt=MY_SECRET_SALT),
                   verdict="deny", reason="dti_over_limit",
                   policy=POLICY, model="risk-gbm-2026-09",
                   evidence={"dti": 0.61, "limit": 0.45})
    s.flush()                       # or let autoflush do it
    ...
    receipt = s.receipt(rec["nonce"])   # hand to the customer / auditor

Give `rec["nonce"]` to the person the decision was about, at the moment you
make it. If that receipt can later not be produced, its absence is evidence.
"""
from __future__ import annotations

import atexit
import hashlib
import hmac
import json
import os
import secrets
import threading
import time
import urllib.error
import urllib.request
from typing import Any, Optional

__all__ = ["Sentinel", "SentinelError", "policy_id", "detail_hash",
           "subject_hash", "canonical", "leaf_of", "verify_local"]
__version__ = "2.0.0"

SCHEMA = "xode-sentinel/1"
FIELDS = ("kind", "subject", "verdict", "reason", "policy", "model", "ts",
          "nonce", "detail")
TENANT_ID_TAG = b"xode-sentinel/tenant-id/2\x00"
TENANT_COMMIT_TAG = b"xode-sentinel/tenant-commit/2\x00"
DEFAULT_URL = os.getenv("XODE_SENTINEL_URL", "https://sentinel.xode.net")


# --------------------------------------------------------------- hashing --
def _b2(b: bytes) -> bytes:
    return hashlib.blake2b(b, digest_size=32).digest()


def _pair(a: bytes, b: bytes) -> bytes:
    return _b2(a + b if a <= b else b + a)


def _unhex(s: str) -> bytes:
    return bytes.fromhex(s[2:] if s.startswith("0x") else s)


def canonical(record: dict) -> bytes:
    """Exactly the bytes Sentinel's verifiers hash. Do not reimplement."""
    doc: dict = {"schema": SCHEMA}
    for k in FIELDS:
        v = record.get(k, "")
        if k == "ts":
            doc[k] = int(v)
        else:
            v = "" if v is None else v
            if not isinstance(v, str):
                raise TypeError(f"{k} must be a string, got {type(v).__name__}")
            doc[k] = v
    return json.dumps(doc, sort_keys=True, separators=(",", ":"),
                      ensure_ascii=False).encode("utf-8")


def leaf_of(record: dict) -> str:
    """The 0x-hex leaf for a record — store it next to the record in your DB."""
    return "0x" + _b2(canonical(record)).hex()


def policy_id(*parts: str) -> str:
    """Id of a ruleset = hash of the texts that define it (policy, prompt,
    thresholds, model id). Never a hand-bumped version number."""
    h = hashlib.blake2b(digest_size=32)
    for p in parts:
        h.update(p.encode("utf-8"))
        h.update(b"\x00")
    return "0x" + h.hexdigest()[:32]


def detail_hash(obj: Any) -> str:
    """Commit to evidence without disclosing it. Keep `obj` yourself: you will
    need the exact same value to prove what the hash stands for."""
    raw = json.dumps(obj, sort_keys=True, separators=(",", ":"),
                     ensure_ascii=False).encode("utf-8")
    return "0x" + _b2(raw).hex()


def subject_hash(value: str, salt: bytes | str) -> str:
    """Pseudonymise a personal identifier before it enters a record.

    Keyed (HMAC-BLAKE2b) so that a phone number or national id cannot be
    recovered by hashing every possible value. Keep the salt secret and stable:
    with it you can show an auditor which subject a record is about; without
    it, a published record says nothing about who.
    """
    key = salt.encode("utf-8") if isinstance(salt, str) else salt
    return "0x" + hmac.new(key, value.encode("utf-8"),
                           lambda: hashlib.blake2b(digest_size=32)).hexdigest()


def verify_local(receipt: dict) -> bool:
    """Every hash link in a receipt, without the network. The chain step is
    left to sentinel_verify.py or the /verify page, which an outsider runs."""
    try:
        leaf = _unhex(receipt["leaf"])
        if receipt.get("record") and _b2(canonical(receipt["record"])) != leaf:
            return False
        t = receipt["tenant"]
        node = leaf
        for s in t["proof"]:
            node = _pair(node, _unhex(s))
        if node != _unhex(t["root"]):
            return False
        commit = _b2(TENANT_COMMIT_TAG + _b2(TENANT_ID_TAG + t["id"].encode())
                     + int(t["seq"]).to_bytes(8, "big") + _unhex(t["prev"])
                     + _unhex(t["root"]) + int(t["n"]).to_bytes(8, "big"))
        if "0x" + commit.hex() != receipt["commit"].lower():
            return False
        node = commit
        for s in receipt["global"]["proof"]:
            node = _pair(node, _unhex(s))
        return node == _unhex(receipt["global"]["root"])
    except (KeyError, TypeError, ValueError):
        return False


# ---------------------------------------------------------------- client --
class SentinelError(Exception):
    def __init__(self, status: int, detail: str) -> None:
        super().__init__(f"{status}: {detail}")
        self.status = status
        self.detail = detail


class Sentinel:
    def __init__(self, api_key: str, base_url: str = DEFAULT_URL,
                 journal: Optional[str] = "sentinel-journal.jsonl",
                 autoflush: Optional[float] = 5.0, timeout: float = 20.0) -> None:
        self.api_key = api_key
        self.base_url = base_url.rstrip("/")
        self.journal = journal
        self.timeout = timeout
        # (leaf, journal offset just past its line). The offset lets flush()
        # record how far the journal is known to have reached Sentinel.
        self._queue: list[tuple[str, int]] = []
        self._lock = threading.Lock()
        self._flush_lock = threading.Lock()     # one flush at a time keeps the cursor honest
        self._stop = threading.Event()
        self._resume()
        if autoflush:
            t = threading.Thread(target=self._loop, args=(autoflush,),
                                 name="sentinel-autoflush", daemon=True)
            t.start()
        atexit.register(self.close)

    # -- recording ---------------------------------------------------------
    def record(self, *, kind: str, subject: str, verdict: str, reason: str,
               policy: str, model: str = "", evidence: Any = None,
               detail: str = "", ts: Optional[int] = None,
               nonce: Optional[str] = None) -> dict:
        """Hash a decision, journal it, queue its leaf. Returns the record
        (with "leaf"). Never blocks on the network — a decision path should not
        fail because the audit service is slow; flush() reports errors."""
        if evidence is not None and detail:
            raise ValueError("pass evidence or detail, not both")
        rec = {"schema": SCHEMA, "kind": kind, "subject": subject,
               "verdict": verdict, "reason": reason, "policy": policy,
               "model": model,
               "ts": int(ts if ts is not None else time.time()),
               "nonce": nonce or f"{int(time.time())}-{secrets.token_hex(8)}",
               "detail": detail_hash(evidence) if evidence is not None else detail}
        rec["leaf"] = leaf_of(rec)
        with self._lock:
            end = 0
            if self.journal:
                line = (json.dumps(rec, ensure_ascii=False, sort_keys=True)
                        + "\n").encode("utf-8")
                with open(self.journal, "ab") as f:
                    f.write(line)
                    f.flush()
                    os.fsync(f.fileno())
                    end = f.tell()
            self._queue.append((rec["leaf"], end))
        return rec

    def flush(self) -> dict:
        """Submit queued leaves. Safe to retry: the server de-duplicates."""
        with self._flush_lock:
            with self._lock:
                batch, self._queue = self._queue, []
            sent = {"accepted": 0, "duplicate": 0}
            done = 0
            try:
                while done < len(batch):
                    chunk = batch[done:done + 1000]
                    r = self._call("POST", "/v1/leaves",
                                   {"leaves": [lf for lf, _ in chunk]})
                    sent["accepted"] += r["accepted"]
                    sent["duplicate"] += r["duplicate"]
                    done += len(chunk)
                    self._advance(chunk[-1][1])
            except Exception:
                with self._lock:               # put back what did not go
                    self._queue[:0] = batch[done:]
                raise
            return sent

    # -- crash recovery ------------------------------------------------------
    # The journal is the source of truth; "<journal>.sent" records the byte
    # offset up to which every line is known to have reached Sentinel. On start,
    # anything after it is queued again — after a crash, a kill or a failed
    # flush at exit — so a journaled decision is never silently left unanchored.
    def _cursor_path(self) -> str:
        return self.journal + ".sent"

    def _advance(self, offset: int) -> None:
        if not self.journal or not offset:
            return
        tmp = self._cursor_path() + ".tmp"
        with open(tmp, "w", encoding="ascii") as f:
            f.write(str(offset))
            f.flush()
            os.fsync(f.fileno())
        os.replace(tmp, self._cursor_path())

    def _resume(self) -> None:
        if not self.journal or not os.path.exists(self.journal):
            return
        try:
            with open(self._cursor_path(), encoding="ascii") as f:
                start = int(f.read().strip() or 0)
        except (OSError, ValueError):
            start = 0                           # resend all: the server de-duplicates
        with open(self.journal, "rb") as f:
            f.seek(start)
            pos = start
            for raw in f:
                pos += len(raw)
                if not raw.endswith(b"\n"):
                    break                       # torn last line from a crash mid-write
                try:
                    leaf = json.loads(raw)["leaf"]
                except (ValueError, KeyError):
                    continue
                self._queue.append((leaf, pos))

    @property
    def pending(self) -> int:
        return len(self._queue)

    def close(self) -> None:
        self._stop.set()
        if self._queue:
            try:
                self.flush()
            except Exception:
                pass                         # the journal still has them

    def _loop(self, every: float) -> None:
        while not self._stop.wait(every):
            if self._queue:
                try:
                    self.flush()
                except Exception:
                    pass                     # retried next tick

    # -- receipts ----------------------------------------------------------
    def find(self, nonce: str) -> Optional[dict]:
        """Linear scan of the journal. Fine for tools and small volumes; at scale,
        store each record (and its leaf) in your own database instead."""
        if not self.journal or not os.path.exists(self.journal):
            return None
        with open(self.journal, encoding="utf-8") as f:
            for line in f:
                if f'"{nonce}"' in line:
                    rec = json.loads(line)
                    if rec.get("nonce") == nonce:
                        return rec
        return None

    def receipt(self, nonce_or_record: str | dict) -> dict:
        rec = (nonce_or_record if isinstance(nonce_or_record, dict)
               else self.find(nonce_or_record))
        if rec is None:
            raise KeyError(f"{nonce_or_record!r} is not in the journal")
        rc = self._call("GET", f"/v1/receipts/{leaf_of(rec)}")
        rc["record"] = {k: rec[k] for k in ("schema",) + FIELDS if k in rec}
        return rc

    def bundle(self, window_start: int) -> dict:
        return self._call("GET", f"/v1/windows/{int(window_start)}/bundle")

    def windows(self, limit: int = 50) -> dict:
        return self._call("GET", f"/v1/windows?limit={int(limit)}")

    def me(self) -> dict:
        return self._call("GET", "/v1/me")

    # -- http ----------------------------------------------------------------
    def _call(self, method: str, path: str, body: Any = None) -> dict:
        data = json.dumps(body).encode() if body is not None else None
        req = urllib.request.Request(
            self.base_url + path, data=data, method=method,
            headers={"Authorization": f"Bearer {self.api_key}",
                     "Content-Type": "application/json",
                     "User-Agent": f"xode-sentinel-python/{__version__}"})
        for attempt in range(4):
            try:
                with urllib.request.urlopen(req, timeout=self.timeout) as r:
                    return json.load(r)
            except urllib.error.HTTPError as e:
                detail = e.read().decode("utf-8", "replace")
                try:
                    detail = json.loads(detail).get("detail", detail)
                except ValueError:
                    pass
                if e.code in (429, 502, 503, 504) and attempt < 3:
                    time.sleep(float(e.headers.get("Retry-After") or 2 ** attempt))
                    continue
                raise SentinelError(e.code, str(detail)) from None
            except urllib.error.URLError:
                if attempt < 3:
                    time.sleep(2 ** attempt)
                    continue
                raise
        raise SentinelError(0, "unreachable")
