#!/usr/bin/env python3
"""Vérificateur de référence d'une réponse certifiée Invoket (E-1002, E-1008).

Autonome : bibliothèque standard seule, aucun réseau, aucune dépendance à
Invoket. C'est le script que l'acheteur d'un fichier exécute pour contrôler,
sans nous croire sur parole, qu'une réponse n'a pas été fabriquée ni retouchée.

    curl -s https://api.invoket.com/.well-known/certificate-keys > cles.json
    python3 verify_certificate.py reponse.json --key cles.json
    python3 verify_certificate.py enregistrement.json --key cles.json \\
        --proof preuve.json

Ce qu'il vérifie :
  1. la réponse porte un `data.certificate` de version connue ;
  2. l'empreinte du périmètre attesté — {subject, verdict, verdict_reasons,
     blocks, provenance}, canonicalisé JCS (RFC 8785) puis SHA-256 — est bien
     celle publiée dans `payload_hash` ;
  3. sur un **lot** (`data.records`, E-1003), que chaque enregistrement porte
     bien son empreinte et que l'arbre de Merkle des enregistrements, dans
     l'ordre servi, remonte à la racine publiée ;
  4. si une preuve d'inclusion est fournie (E-1004), que le chemin remonte de
     l'enregistrement seul à la racine du lot — **sans** que le reste du lot
     soit transmis. La preuve est acceptée telle que le service la sert
     (`GET /certify/batch/{id}/proof/{index}`) ou réduite à son tableau de pas ;
  5. avec `--key`, la **signature** Ed25519 que la passerelle pose sur la
     racine (T-51xx) — c'est elle qui distingue une preuve émise par Invoket
     d'une réponse qu'un intermédiaire aurait recomposée de bout en bout.

Les points 1 à 4 prouvent qu'une réponse est **cohérente** ; seul le point 5
prouve qu'elle vient de nous. Sans `--key` le script le dit à chaque exécution
plutôt que de laisser croire l'inverse — il ne va pas chercher la clé lui-même,
parce qu'un vérificateur qui sort sur le réseau n'est plus vérifiable hors ligne.

Code de sortie : 0 si tout est vérifié, 1 sinon.
"""

import argparse
import base64
import hashlib
import json
import sys

CERTIFICATE_VERSION = 1
HASHED_DATA_FIELDS = ("subject", "verdict", "verdict_reasons", "blocks")
HASHED_PROVENANCE_FIELD = "provenance"
LEAF_DOMAIN = b"\x00"
NODE_DOMAIN = b"\x01"
DIGEST_PREFIX = "sha256:"
SIGNATURE_ALG = "Ed25519"
KEYS_URL = "https://api.invoket.com/.well-known/certificate-keys"

# --- JCS (RFC 8785) ---------------------------------------------------------

_SHORT_ESCAPES = {0x08: "\\b", 0x09: "\\t", 0x0A: "\\n", 0x0C: "\\f", 0x0D: "\\r"}


def _string(value):
    out = ['"']
    for char in value:
        point = ord(char)
        if char == '"':
            out.append('\\"')
        elif char == "\\":
            out.append("\\\\")
        elif point in _SHORT_ESCAPES:
            out.append(_SHORT_ESCAPES[point])
        elif point < 0x20:
            out.append("\\u%04x" % point)
        else:
            out.append(char)
    out.append('"')
    return "".join(out)


def _shortest_digits(value):
    """Chiffres les plus courts rejouant `value` (> 0) et son exposant décimal.

    `repr` rend en Python la représentation la plus courte qui fait
    l'aller-retour — la même que celle qu'exige ECMAScript."""
    rendered = repr(value)
    mantissa, _, exponent = rendered.partition("e")
    exponent = int(exponent) if exponent else 0
    integer, _, fraction = mantissa.partition(".")
    digits = (integer + fraction).lstrip("0")
    exponent -= len(fraction)
    while len(digits) > 1 and digits.endswith("0"):
        digits = digits[:-1]
        exponent += 1
    return digits, exponent


def _number(value):
    """Sérialisation ECMAScript (RFC 8785 §3.2.2.3)."""
    if isinstance(value, int):
        if abs(value) > 9007199254740992:
            raise ValueError("entier hors de la plage exacte IEEE-754 : %d" % value)
        value = float(value)
    if value != value or value in (float("inf"), float("-inf")):
        raise ValueError("nombre non sérialisable en JSON : %r" % value)
    if value == 0.0:
        return "0"  # couvre le zéro négatif
    sign = "-" if value < 0 else ""
    digits, exponent = _shortest_digits(abs(value))
    k = len(digits)
    n = exponent + k
    if k <= n <= 21:
        body = digits + "0" * (n - k)
    elif 0 < n <= 21:
        body = digits[:n] + "." + digits[n:]
    elif -6 < n <= 0:
        body = "0." + "0" * (-n) + digits
    else:
        power = n - 1
        exp = ("+" if power >= 0 else "") + str(power)
        body = digits + "e" + exp if k == 1 else digits[0] + "." + digits[1:] + "e" + exp
    return sign + body


def canonicalize(value):
    """Forme canonique JCS (RFC 8785) d'une valeur JSON déjà décodée."""
    if value is None:
        return "null"
    if value is True:
        return "true"
    if value is False:
        return "false"
    if isinstance(value, str):
        return _string(value)
    if isinstance(value, (int, float)):
        return _number(value)
    if isinstance(value, list):
        return "[" + ",".join(canonicalize(item) for item in value) + "]"
    if isinstance(value, dict):
        # Tri par unités de code UTF-16, et non par points de code.
        keys = sorted(value, key=lambda k: k.encode("utf-16-be"))
        return "{" + ",".join(_string(k) + ":" + canonicalize(value[k]) for k in keys) + "}"
    raise ValueError("type JSON inattendu : %r" % type(value))


# --- Empreintes et arbre de Merkle ------------------------------------------


def _unpad(text):
    """Base64url sans padding — la forme servie partout par Invoket."""
    return base64.urlsafe_b64decode(text + "=" * (-len(text) % 4))


def _encode(digest):
    return DIGEST_PREFIX + base64.urlsafe_b64encode(digest).decode("ascii").rstrip("=")


def _decode(text):
    if not isinstance(text, str) or not text.startswith(DIGEST_PREFIX):
        raise ValueError("empreinte illisible : %r" % (text,))
    try:
        digest = _unpad(text[len(DIGEST_PREFIX):])
    except (ValueError, TypeError):
        raise ValueError("empreinte illisible : %r" % (text,))
    if len(digest) != 32:
        raise ValueError("empreinte illisible : %r" % (text,))
    return digest


def leaf_hash(payload):
    """Empreinte d'un sujet : domaine 0x00 puis les octets UTF-8 de sa forme JCS."""
    return hashlib.sha256(LEAF_DOMAIN + canonicalize(payload).encode("utf-8")).digest()


def hashed_scope(response):
    """Le périmètre attesté, et lui seul.

    `data.certificate` en est **exclu** : une empreinte ne couvre jamais le
    champ qui la porte. `query` et `limits` aussi : ils décrivent la demande,
    pas le fait attesté."""
    data = response.get("data")
    if not isinstance(data, dict):
        raise ValueError("réponse non certifiable : `data` absent")
    scope = {}
    for field in HASHED_DATA_FIELDS:
        if field not in data:
            raise ValueError("réponse non certifiable : `data.%s` absent" % field)
        scope[field] = data[field]
    if HASHED_PROVENANCE_FIELD not in response:
        raise ValueError("réponse non certifiable : `provenance` absent")
    scope[HASHED_PROVENANCE_FIELD] = response[HASHED_PROVENANCE_FIELD]
    return scope


def merkle_root(leaves):
    """Racine des empreintes, dans l'ordre servi.

    Feuille impaire : **promue** telle quelle au niveau suivant, jamais
    dupliquée — la duplication ouvre la malléabilité connue de l'arbre de
    Bitcoin, où deux lots distincts donnent la même racine."""
    if not leaves:
        raise ValueError("un arbre de Merkle a besoin d'au moins une feuille")
    level = list(leaves)
    while len(level) > 1:
        parents = []
        for index in range(0, len(level) - 1, 2):
            parents.append(
                hashlib.sha256(NODE_DOMAIN + level[index] + level[index + 1]).digest()
            )
        if len(level) % 2:
            parents.append(level[-1])
        level = parents
    return level[0]


def proof_steps(document):
    """Le chemin d'inclusion et, si elle existe, l'enveloppe qui le portait.

    On accepte les deux formes : le tableau de pas brut, et la réponse servie
    par `GET /certify/batch/{id}/proof/{index}` (E-1004) telle qu'elle sort du
    service — l'acheteur ne doit avoir aucune transformation à faire avant de
    vérifier, une transformation est un endroit où se tromper."""
    if isinstance(document, list):
        return document, {}
    data = document.get("data") if isinstance(document, dict) else None
    if isinstance(data, dict) and isinstance(data.get("proof"), list):
        return data["proof"], data
    raise ValueError(
        "preuve illisible : attendu un chemin d'inclusion, ou la réponse de "
        "`GET /certify/batch/{id}/proof/{index}`"
    )


def check_proof_envelope(envelope, response, certificate, leaf):
    """Confronte la preuve servie à l'enregistrement qu'on prétend prouver.

    Sans ces trois contrôles, une preuve d'un autre lot ou d'un autre rang
    échouerait quand même — mais sur « le chemin ne remonte pas à la racine »,
    qui ne dit pas *pourquoi*. Une preuve qui échoue doit nommer la méprise."""
    published = certificate.get("payload_hash")
    served = envelope.get("payload_hash")
    if served is not None and served != published:
        raise ValueError(
            "cette preuve est celle du lot %s, l'enregistrement porte le "
            "certificat du lot %s" % (served, published)
        )
    if envelope.get("leaf") is not None and _decode(envelope["leaf"]) != leaf:
        raise ValueError(
            "le rang %r du lot a l'empreinte %s, cet enregistrement %s — la "
            "preuve n'est pas la sienne"
            % (envelope.get("index"), envelope["leaf"], _encode(leaf))
        )
    rank = response.get("data", {}).get("subject", {}).get("index")
    if rank is not None and envelope.get("index") is not None and rank != envelope["index"]:
        raise ValueError(
            "preuve du rang %r, enregistrement du rang %r" % (envelope["index"], rank)
        )


def verify_proof(leaf, proof, root):
    """Rejoue un chemin d'inclusion : domaine 0x01 sur les deux filles."""
    current = leaf
    for step in proof:
        sibling = _decode(step["digest"])
        if step["side"] == "left":
            current = hashlib.sha256(NODE_DOMAIN + sibling + current).digest()
        elif step["side"] == "right":
            current = hashlib.sha256(NODE_DOMAIN + current + sibling).digest()
        else:
            raise ValueError("côté de preuve inconnu : %r" % (step["side"],))
    return current == root


# --- Signature Ed25519 (RFC 8032) -------------------------------------------
#
# Vérifier une signature Ed25519 tient en une page d'arithmétique modulaire, et
# la bibliothèque standard porte déjà SHA-512 : c'est ce qui permet à ce script
# de rester **sans dépendance**. L'acheteur d'un fichier certifié n'installe
# rien et n'accorde sa confiance à personne — ni à nous, ni à un paquet tiers —
# pour contrôler une preuve. Les noms suivent la référence RFC 8032 §5.1.

_P = 2**255 - 19
_L = 2**252 + 27742317777372353535851937790883648493
_D = -121665 * pow(121666, _P - 2, _P) % _P
_SQRT_M1 = pow(2, (_P - 1) // 4, _P)


def _recover_x(y, sign):
    """L'abscisse du point d'ordonnée `y` et de signe donné, ou `None`."""
    if y >= _P:
        return None
    square = (y * y - 1) * pow(_D * y * y + 1, _P - 2, _P) % _P
    if square == 0:
        return None if sign else 0
    x = pow(square, (_P + 3) // 8, _P)
    if (x * x - square) % _P != 0:
        x = x * _SQRT_M1 % _P
    if (x * x - square) % _P != 0:
        return None  # `y` ne désigne aucun point de la courbe
    return _P - x if x & 1 != sign else x


def _point_add(first, second):
    """Addition en coordonnées étendues (x, y, z, t), sans inversion modulaire."""
    x1, y1, z1, t1 = first
    x2, y2, z2, t2 = second
    a = (y1 - x1) * (y2 - x2) % _P
    b = (y1 + x1) * (y2 + x2) % _P
    c = 2 * t1 * t2 * _D % _P
    d = 2 * z1 * z2 % _P
    e, f, g, h = b - a, d - c, d + c, b + a
    return (e * f % _P, g * h % _P, f * g % _P, e * h % _P)


def _scalar_mult(point, scalar):
    total = (0, 1, 1, 0)  # le neutre
    while scalar > 0:
        if scalar & 1:
            total = _point_add(total, point)
        point = _point_add(point, point)
        scalar >>= 1
    return total


_G_Y = 4 * pow(5, _P - 2, _P) % _P
_G = (_recover_x(_G_Y, 0), _G_Y, 1, _recover_x(_G_Y, 0) * _G_Y % _P)


def _decode_point(raw):
    """32 octets petit-boutistes → point de la courbe, ou `None`."""
    if len(raw) != 32:
        return None
    y = int.from_bytes(raw, "little")
    sign = y >> 255
    y &= (1 << 255) - 1
    x = _recover_x(y, sign)
    return None if x is None else (x, y, 1, x * y % _P)


def _same_point(first, second):
    x1, y1, z1, _ = first
    x2, y2, z2, _ = second
    return (x1 * z2 - x2 * z1) % _P == 0 and (y1 * z2 - y2 * z1) % _P == 0


def ed25519_verify(public_key, signature, message):
    """RFC 8032 §5.1.7 : `[S]B == R + [SHA-512(R‖A‖M) mod L]A` ?"""
    if len(signature) != 64:
        return False
    parsed = _decode_point(public_key)
    commitment = _decode_point(signature[:32])
    if parsed is None or commitment is None:
        return False
    scalar = int.from_bytes(signature[32:], "little")
    if scalar >= _L:  # forme canonique exigée : pas de signature malléable
        return False
    challenge = (
        int.from_bytes(
            hashlib.sha512(signature[:32] + public_key + message).digest(), "little"
        )
        % _L
    )
    return _same_point(
        _scalar_mult(_G, scalar), _point_add(commitment, _scalar_mult(parsed, challenge))
    )


def load_public_keys(argument):
    """Le jeu de clés servi par `GET /.well-known/certificate-keys`.

    On accepte le fichier **tel que servi** (l'acheteur fait un `curl`, rien
    d'autre) et, pour le dépannage, une clé publique base64url à nu. Rend
    `{key_id: octets}` ; une clé nue est indexée sous `None`, donc utilisable
    quel que soit le `key_id` du certificat."""
    try:
        with open(argument, encoding="utf-8") as handle:
            document = json.load(handle)
    except OSError:
        return {None: _public_key_bytes(argument)}
    except json.JSONDecodeError:
        raise ValueError("jeu de clés illisible : %s n'est pas du JSON" % argument)
    entries = document.get("keys") if isinstance(document, dict) else None
    if not isinstance(entries, list) or not entries:
        raise ValueError(
            "jeu de clés illisible : attendu la réponse de `GET %s`" % KEYS_URL
        )
    keys = {}
    for entry in entries:
        if not isinstance(entry, dict) or entry.get("alg") != SIGNATURE_ALG:
            continue  # une clé d'un autre algorithme ne vérifie rien ici
        keys[entry.get("key_id")] = _public_key_bytes(entry.get("public_key"))
    if not keys:
        raise ValueError("jeu de clés sans aucune clé %s" % SIGNATURE_ALG)
    return keys


def _raw(text, what):
    """Décode un champ base64url sans padding, ou dit lequel est illisible."""
    try:
        return _unpad(text)
    except (ValueError, TypeError, AttributeError):
        raise ValueError("%s illisible : %r" % (what, text))


def _public_key_bytes(text):
    raw = _raw(text, "clé publique")
    if len(raw) != 32:
        raise ValueError(
            "clé publique %s attendue sur 32 octets : %r" % (SIGNATURE_ALG, text)
        )
    return raw


def check_signature(certificate, keys):
    """Rend la ligne de constat, ou lève `ValueError` sur une preuve qui ment.

    Trois situations distinctes, et les confondre serait précisément la faute
    que ce script existe pour empêcher : *non signée* (le service seul, sans la
    passerelle), *signée mais non contrôlée* (aucune clé fournie), *signée et
    contrôlée*."""
    signature = certificate.get("signature")
    if signature is None:
        if keys is not None:
            raise ValueError(
                "aucune signature dans ce certificat : cette réponse n'a pas "
                "traversé la passerelle, elle ne prouve donc que sa propre "
                "cohérence"
            )
        return (
            "⚠️ réponse NON signée : elle est cohérente, mais rien n'atteste "
            "qu'elle vient d'Invoket"
        )
    alg = certificate.get("alg")
    if alg != SIGNATURE_ALG:
        raise ValueError(
            "algorithme de signature %r inconnu de ce vérificateur (attendu : %s)"
            % (alg, SIGNATURE_ALG)
        )
    key_id = certificate.get("key_id")
    if keys is None:
        return (
            "⚠️ signature présente (clé « %s ») mais NON vérifiée : relancez avec "
            "`--key`, le jeu de clés étant servi par `GET %s`" % (key_id, KEYS_URL)
        )
    if key_id in keys:
        public_key = keys[key_id]
    elif list(keys) == [None]:
        public_key = keys[None]
    else:
        raise ValueError(
            "ce certificat est signé par la clé « %s », que le jeu fourni ne "
            "publie pas (clés connues : %s)"
            % (key_id, ", ".join(sorted(str(known) for known in keys)))
        )
    message = certificate.get("payload_hash")
    if not isinstance(message, str):
        raise ValueError("certificat sans `payload_hash` : rien à vérifier")
    # La passerelle signe les octets UTF-8 de la chaîne `payload_hash` telle
    # quelle — aucune re-normalisation, aucun décodage (T-5102).
    raw = _raw(signature, "signature")
    if not ed25519_verify(public_key, raw, message.encode("utf-8")):
        raise ValueError(
            "signature invalide pour la clé « %s » : ce certificat n'a pas été "
            "émis par Invoket, ou son `payload_hash` a été retouché" % key_id
        )
    return "signature %s vérifiée contre la clé « %s » ✓" % (SIGNATURE_ALG, key_id)


# --- Programme --------------------------------------------------------------


def verify(response, proof=None, keys=None):
    """Rend la liste des constats. Lève `ValueError` sur une entrée illisible."""
    certificate = response.get("data", {}).get("certificate")
    if not isinstance(certificate, dict):
        raise ValueError("cette réponse ne porte pas de `data.certificate`")
    version = certificate.get("certificate_version")
    if version != CERTIFICATE_VERSION:
        raise ValueError(
            "version de certificat %r inconnue de ce vérificateur (attendue : %d)"
            % (version, CERTIFICATE_VERSION)
        )

    lines = []
    published = _decode(certificate["payload_hash"])
    records = response.get("data", {}).get("records")

    if records is not None and proof is None:
        lines.extend(_verify_batch(records, certificate, published))
        lines.append(_profile_line(certificate))
        lines.append(
            "verdicts rendus le %s contre : %s"
            % (certificate.get("issued_at"), _sources(response))
        )
        lines.append(check_signature(certificate, keys))
        return lines

    if proof is not None:
        proof, envelope = proof_steps(proof)

    leaf = leaf_hash(hashed_scope(response))
    if proof is None:
        if leaf != published:
            raise ValueError(
                "empreinte du contenu : %s, publiée : %s — la réponse a été "
                "retouchée, ou n'est pas la réponse de ce certificat"
                % (_encode(leaf), _encode(published))
            )
        lines.append("empreinte du périmètre attesté : %s ✓" % _encode(leaf))
    else:
        check_proof_envelope(envelope, response, certificate, leaf)
        if not verify_proof(leaf, proof, published):
            raise ValueError(
                "l'enregistrement (%s) n'appartient pas au lot certifié par %s"
                % (_encode(leaf), _encode(published))
            )
        lines.append("empreinte de l'enregistrement : %s ✓" % _encode(leaf))
        lines.append("appartenance au lot (%d pas de preuve) ✓" % len(proof))

    lines.append(_profile_line(certificate))
    lines.append(
        "verdict rendu le %s contre : %s"
        % (certificate.get("issued_at"), _sources(response))
    )
    lines.append(check_signature(certificate, keys))
    return lines


def _profile_line(certificate):
    return (
        "profil « %s » — la preuve ne couvre que les blocs de ce profil, "
        "jamais un contrôle complet" % certificate.get("profile")
    )


def _verify_batch(records, certificate, published):
    """Un lot (E-1003) : chaque enregistrement **est** une réponse, l'arbre de
    leurs empreintes remonte à la racine publiée."""
    if not isinstance(records, list) or not records:
        raise ValueError("`data.records` n'est pas un lot d'enregistrements")
    count = certificate.get("subject_count")
    if count != len(records):
        raise ValueError(
            "le lot annonce %r enregistrements et en porte %d — il a été tronqué"
            % (count, len(records))
        )

    leaves = []
    for index, record in enumerate(records):
        leaf = leaf_hash(hashed_scope(record))
        served = record.get("fingerprint")
        if served is not None and _decode(served) != leaf:
            raise ValueError(
                "enregistrement %d : empreinte servie %s, recalculée %s — il a "
                "été retouché" % (index, served, _encode(leaf))
            )
        leaves.append(leaf)

    if merkle_root(leaves) != published:
        raise ValueError(
            "la racine des %d enregistrements ne correspond pas au certificat "
            "(%s) — le lot a été retouché, réordonné ou recomposé"
            % (len(leaves), _encode(published))
        )

    issues = {}
    verdicts = {}
    for record in records:
        issues[record.get("issue")] = issues.get(record.get("issue"), 0) + 1
        verdict = record.get("data", {}).get("verdict")
        verdicts[verdict] = verdicts.get(verdict, 0) + 1
    return [
        "%d enregistrements, racine %s ✓" % (len(leaves), _encode(published)),
        "issues : %s" % _distribution(issues),
        "verdicts : %s" % _distribution(verdicts),
    ]


def _distribution(counts):
    return ", ".join(
        "%s=%d" % ("sans verdict" if key is None else key, value)
        for key, value in sorted(counts.items(), key=lambda item: str(item[0]))
    )


def _sources(response):
    """Ce contre quoi le verdict a été rendu : une preuve atteste un verdict
    **daté**, pas une vérité présente."""
    dated = []
    candidates = [response.get("provenance")]
    # Un lot porte ses provenances par enregistrement (elles sont dans le
    # périmètre haché de chacun) ; une réponse unitaire, par bloc.
    for scope in [response] + response.get("data", {}).get("records", []):
        if not isinstance(scope, dict):
            continue
        candidates.append(scope.get("provenance"))
        blocks = scope.get("data", {}).get("blocks", {})
        if isinstance(blocks, dict):
            candidates += [
                block.get("provenance")
                for block in blocks.values()
                if isinstance(block, dict)
            ]
    for entry in candidates:
        if not isinstance(entry, dict):
            continue
        freshness = entry.get("freshness", {})
        stamp = freshness.get("as_of") or freshness.get("kind")
        mark = "%s@%s" % (entry.get("source"), stamp)
        if mark not in dated:
            dated.append(mark)
    return ", ".join(dated) or "source non datée"


def main(argv=None):
    parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    parser.add_argument("response", help="réponse JSON capturée (`-` pour stdin)")
    parser.add_argument(
        "--proof",
        help="preuve d'inclusion JSON (E-1004) quand la réponse est un "
        "enregistrement extrait d'un lot : la réponse servie par "
        "`GET /certify/batch/{id}/proof/{index}`, ou son seul tableau de pas",
    )
    parser.add_argument(
        "--key",
        help="jeu de clés publiques tel que servi par `GET %s` (ou une clé "
        "publique base64url à nu) : sans lui, la signature n'est pas vérifiée"
        % KEYS_URL,
    )
    args = parser.parse_args(argv)

    stream = sys.stdin if args.response == "-" else open(args.response, encoding="utf-8")
    with stream as handle:
        response = json.load(handle)
    proof = None
    if args.proof:
        with open(args.proof, encoding="utf-8") as handle:
            proof = json.load(handle)

    try:
        keys = load_public_keys(args.key) if args.key else None
        for line in verify(response, proof, keys):
            print(line)
    except ValueError as failure:
        print("ÉCHEC : %s" % failure, file=sys.stderr)
        return 1
    print("OK")
    return 0


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