#!/usr/bin/env python3
"""Offline verifier for the sender-domain DKIM evidence supplied on 2 August 2026.

Scope is deliberately narrow.  The script verifies the RFC 6376 body hash and
RSA-SHA256 header signature against two supplied public-key candidates.  It does
not prove that either key was published in DNS at a historical time, and it does
not verify X-Google-DKIM or ARC.  Those are separate evidential steps.

The process exits successfully only when BOTH the body hash and RSA header
signature pass.  Component results are kept separate because a body mutation
normally leaves the header-signature calculation intact while making the DKIM
verification as a whole fail.
"""

from __future__ import annotations

import argparse
import base64
import hashlib
import json
import re
import sys
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Dict, Iterable, List, Mapping, Optional, Sequence, Tuple

from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import padding, rsa


PUBLIC_KEYS: Mapping[str, str] = {
    "candidate-1-2048": (
        "MIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEA5wg55g75nGtmxMkMT+wV"
        "P3ai2urSZepwvT/bSixC9JIhhvW9BeIFSGL0avARLlPakyLvJgd3niFrS4MC8b/Z"
        "r+0SXgn3BmxjWtcXuUH/oUdHbpDIm6NJkze8eue2suRZP6SF1Lxr8kiENrjIhL/j"
        "9f9NFoRlr3SY+LL0JB7HSkVRH5LuKOft7IXAJSNUtktDqIbkiT3feoGxC6KMQIoq"
        "7UJneIhXqri27f/fTQApq0ZmgtxLo3EMaoGuLyeZqivDSOnpH87r19HHUlffjXI/"
        "zem/SRysye6CVlMLW7oT8V3P068xmH9/mxAQymV25c+/TuPAjJawL+HBUg+sp6q3"
        "iwIDAQAB"
    ),
    "candidate-2-1024": (
        "MIGfMA0GCSqGSIb3DQEBAQUAA4GNADCBiQKBgQDGT9CotioD/YAjUYIkyt0bq0MD"
        "a6wvsWnQmzGUKdx8+VRA66kifBsuoCNoX2hQapzQSdCai2PwsAzVTZvNwzz5VZT3"
        "Wa/hgfNm4x7L424RhBfokihWqNOA4KfMBdVxie8RQjjqYaeMIfgntj/rpuCq71+u"
        "iWAFcPJ5Nye/QuC8oQIDAQAB"
    ),
}


@dataclass
class KeyAttempt:
    label: str
    bits: int
    spki_sha256: str
    signature_math_ok: bool
    error: Optional[str] = None


@dataclass
class Verification:
    path: str
    bytes: int
    sha256: str
    domain: str
    selector: str
    algorithm: str
    canonicalisation: str
    signature_time: str
    expiry_time: str
    signed_headers: List[str]
    signed_header_input_sha256: str
    body_length_tag: Optional[str]
    body_hash_expected: str
    body_hash_computed: str
    body_hash_ok: bool
    signature_math_ok: bool
    overall_dkim_ok: bool
    matching_key: Optional[str]
    matching_key_bits: Optional[int]
    matching_key_spki_sha256: Optional[str]
    key_attempts: List[KeyAttempt]
    selector_binding_status: str
    google_arc_status: str


def split_message(raw: bytes) -> Tuple[bytes, bytes, bytes]:
    for separator in (b"\r\n\r\n", b"\n\n"):
        offset = raw.find(separator)
        if offset >= 0:
            return raw[:offset], raw[offset + len(separator) :], separator
    raise ValueError("message has no header/body separator")


def parse_headers(header_block: bytes) -> List[Tuple[bytes, bytes]]:
    normalised = header_block.replace(b"\r\n", b"\n")
    parsed: List[Tuple[bytes, bytes]] = []
    current: Optional[List[bytes]] = None
    for line in normalised.split(b"\n"):
        if line[:1] in (b" ", b"\t") and current is not None:
            current[1] += b"\r\n" + line
            continue
        if current is not None:
            parsed.append((current[0], current[1]))
            current = None
        if b":" in line:
            name, value = line.split(b":", 1)
            current = [name, value]
    if current is not None:
        parsed.append((current[0], current[1]))
    return parsed


def canonicalise_relaxed_header(name: bytes, value: bytes) -> bytes:
    unfolded = value.replace(b"\r\n", b"").replace(b"\n", b"")
    unfolded = re.sub(rb"[ \t]+", b" ", unfolded).strip(b" \t")
    return name.lower().strip() + b":" + unfolded + b"\r\n"


def canonicalise_relaxed_body(body: bytes) -> bytes:
    lines = body.replace(b"\r\n", b"\n").split(b"\n")
    canonical_lines = []
    for line in lines:
        line = re.sub(rb"[ \t]+", b" ", line)
        canonical_lines.append(line.rstrip(b" \t"))
    while canonical_lines and canonical_lines[-1] == b"":
        canonical_lines.pop()
    # RFC 6376 section 3.4.4: the relaxed canonical form of an empty body is
    # the null input, not a lone CRLF.
    return b"\r\n".join(canonical_lines) + (b"\r\n" if canonical_lines else b"")


def parse_tag_list(raw_value: bytes) -> Dict[str, bytes]:
    compact = re.sub(rb"\s+", b"", raw_value)
    tags: Dict[str, bytes] = {}
    for part in compact.split(b";"):
        if b"=" not in part:
            continue
        key, value = part.split(b"=", 1)
        key_text = key.decode("ascii", "strict").lower()
        if key_text in tags:
            raise ValueError(f"duplicate DKIM tag: {key_text}")
        tags[key_text] = value
    return tags


def _load_key(label: str, value: str) -> Tuple[rsa.RSAPublicKey, str]:
    der = base64.b64decode(value, validate=True)
    key = serialization.load_der_public_key(der)
    if not isinstance(key, rsa.RSAPublicKey):
        raise TypeError(f"{label} is not an RSA public key")
    return key, hashlib.sha256(der).hexdigest()


def verify_bytes(raw: bytes, label: str = "<memory>") -> Verification:
    header_block, body, _ = split_message(raw)
    headers = parse_headers(header_block)
    signature_header = next(
        (header for header in headers if header[0].lower() == b"dkim-signature"),
        None,
    )
    if signature_header is None:
        raise ValueError("no DKIM-Signature header")

    tags = parse_tag_list(signature_header[1])
    required = ("v", "a", "c", "d", "s", "h", "bh", "b")
    absent = [tag for tag in required if tag not in tags]
    if absent:
        raise ValueError("DKIM signature is missing required tags: " + ", ".join(absent))
    if tags["v"] != b"1" or tags["a"].lower() != b"rsa-sha256":
        raise ValueError("this verifier accepts only v=1; a=rsa-sha256")
    canonicalisation = tags["c"].decode("ascii", "strict").lower()
    if canonicalisation != "relaxed/relaxed":
        raise ValueError("this verifier accepts only c=relaxed/relaxed")
    if tags["d"].lower() != b"mastermindpromotion.com" or tags["s"].lower() != b"google":
        raise ValueError(
            "embedded key candidates are authorised only for "
            "d=mastermindpromotion.com; s=google"
        )

    canonical_body = canonicalise_relaxed_body(body)
    length_tag = tags.get("l")
    if length_tag is not None:
        if re.fullmatch(rb"[0-9]+", length_tag) is None:
            raise ValueError("DKIM l= tag must be an unsigned decimal integer")
        signed_length = int(length_tag)
        if signed_length > len(canonical_body):
            raise ValueError("DKIM l= tag exceeds the canonicalised body length")
        canonical_body = canonical_body[:signed_length]
    computed_body_hash = base64.b64encode(hashlib.sha256(canonical_body).digest()).decode()
    expected_body_hash = tags["bh"].decode("ascii", "strict")
    body_hash_ok = computed_body_hash == expected_body_hash

    header_pool: Dict[str, List[Tuple[bytes, bytes]]] = {}
    for name, value in headers:
        header_pool.setdefault(name.lower().decode("ascii", "strict"), []).append((name, value))

    signed_header_names = [
        name.strip().lower() for name in tags["h"].decode("ascii", "strict").split(":") if name.strip()
    ]
    if "from" not in signed_header_names:
        raise ValueError("DKIM h= must include From")
    if "t" in tags and re.fullmatch(rb"[0-9]+", tags["t"]) is None:
        raise ValueError("DKIM t= tag must be an unsigned decimal integer")
    if "x" in tags:
        if re.fullmatch(rb"[0-9]+", tags["x"]) is None:
            raise ValueError("DKIM x= tag must be an unsigned decimal integer")
        if "t" in tags and int(tags["x"]) <= int(tags["t"]):
            raise ValueError("DKIM x= must be later than t=")
    used: Dict[str, int] = {}
    signed_data = b""
    for name in signed_header_names:
        index = used.get(name, 0)
        available = header_pool.get(name, [])
        if index < len(available):
            selected_name, selected_value = available[len(available) - 1 - index]
            signed_data += canonicalise_relaxed_header(selected_name, selected_value)
            used[name] = index + 1

    emptied_signature_value = re.sub(rb"([;\s]b=)[^;]*", rb"\1", signature_header[1])
    signed_data += canonicalise_relaxed_header(
        signature_header[0], emptied_signature_value
    ).removesuffix(b"\r\n")
    signed_header_input_sha256 = hashlib.sha256(signed_data).hexdigest()
    signature_bytes = base64.b64decode(tags["b"], validate=True)

    attempts: List[KeyAttempt] = []
    matching_key: Optional[str] = None
    matching_bits: Optional[int] = None
    matching_fingerprint: Optional[str] = None
    for key_label, public_value in PUBLIC_KEYS.items():
        key, fingerprint = _load_key(key_label, public_value)
        try:
            key.verify(signature_bytes, signed_data, padding.PKCS1v15(), hashes.SHA256())
            attempts.append(KeyAttempt(key_label, key.key_size, fingerprint, True))
            if matching_key is None:
                matching_key = key_label
                matching_bits = key.key_size
                matching_fingerprint = fingerprint
        except InvalidSignature:
            attempts.append(KeyAttempt(key_label, key.key_size, fingerprint, False, "InvalidSignature"))

    signature_math_ok = matching_key is not None
    return Verification(
        path=label,
        bytes=len(raw),
        sha256=hashlib.sha256(raw).hexdigest(),
        domain=tags["d"].decode("ascii", "strict"),
        selector=tags["s"].decode("ascii", "strict"),
        algorithm=tags["a"].decode("ascii", "strict"),
        canonicalisation=canonicalisation,
        signature_time=tags.get("t", b"").decode("ascii", "strict"),
        expiry_time=tags.get("x", b"").decode("ascii", "strict"),
        signed_headers=signed_header_names,
        signed_header_input_sha256=signed_header_input_sha256,
        body_length_tag=length_tag.decode("ascii", "strict") if length_tag else None,
        body_hash_expected=expected_body_hash,
        body_hash_computed=computed_body_hash,
        body_hash_ok=body_hash_ok,
        signature_math_ok=signature_math_ok,
        overall_dkim_ok=body_hash_ok and signature_math_ok,
        matching_key=matching_key,
        matching_key_bits=matching_bits,
        matching_key_spki_sha256=matching_fingerprint,
        key_attempts=attempts,
        selector_binding_status=(
            "not evaluated by this offline verifier; consult the separate current "
            "and historical DNS evidence for "
            f"{tags['s'].decode()}._domainkey.{tags['d'].decode()}"
        ),
        google_arc_status="not tested by this sender-domain DKIM verifier",
    )


def _mutate_body(raw: bytes) -> bytes:
    header_block, body, separator = split_message(raw)
    mutable = bytearray(body)
    index = next((i for i, value in enumerate(mutable) if value not in b"\r\n\t "), None)
    if index is None:
        raise ValueError("cannot construct body negative control from an empty body")
    mutable[index] ^= 1
    return header_block + separator + bytes(mutable)


def _mutate_date(raw: bytes) -> bytes:
    header_block, body, separator = split_message(raw)
    match = re.search(rb"(?im)^Date:[^\r\n]*(?:\r?\n[ \t][^\r\n]*)*", header_block)
    if match is None:
        raise ValueError("cannot construct Date negative control: Date header absent")
    date_value = bytearray(match.group(0))
    index = next((i for i, value in enumerate(date_value) if 48 <= value <= 57), None)
    if index is None:
        raise ValueError("cannot construct Date negative control: Date contains no digit")
    date_value[index] = 48 + ((date_value[index] - 48 + 1) % 10)
    mutated_headers = header_block[: match.start()] + bytes(date_value) + header_block[match.end() :]
    return mutated_headers + separator + body


def _mutate_unsigned_received(raw: bytes) -> bytes:
    header_block, body, separator = split_message(raw)
    match = re.search(rb"(?im)^Received:[^\r\n]*(?:\r?\n[ \t][^\r\n]*)*", header_block)
    if match is None:
        raise ValueError("cannot construct unsigned-header control: Received header absent")
    received_value = bytearray(match.group(0))
    index = next((i for i, value in enumerate(received_value) if 48 <= value <= 57), None)
    if index is None:
        raise ValueError("cannot construct unsigned-header control: Received contains no digit")
    received_value[index] = 48 + ((received_value[index] - 48 + 1) % 10)
    mutated_headers = header_block[: match.start()] + bytes(received_value) + header_block[match.end() :]
    return mutated_headers + separator + body


def _changed_byte_count(original: bytes, mutated: bytes) -> int:
    return sum(left != right for left, right in zip(original, mutated)) + abs(
        len(original) - len(mutated)
    )


def _checked_control(
    raw: bytes,
    mutated: bytes,
    label: str,
    expected: Tuple[bool, bool, bool],
    expectation: str,
) -> Dict[str, object]:
    changed = _changed_byte_count(raw, mutated)
    if changed == 0:
        raise AssertionError(f"{label} mutation changed zero bytes")
    result = verify_bytes(mutated, f"<{label}>")
    observed = (result.body_hash_ok, result.signature_math_ok, result.overall_dkim_ok)
    expectation_ok = observed == expected
    if not expectation_ok:
        raise AssertionError(f"{label} produced {observed}, expected {expected}")
    return {
        "changed_bytes": changed,
        "body_hash_ok": result.body_hash_ok,
        "signature_math_ok": result.signature_math_ok,
        "overall_dkim_ok": result.overall_dkim_ok,
        "expected_components": {
            "body_hash_ok": expected[0],
            "signature_math_ok": expected[1],
            "overall_dkim_ok": expected[2],
        },
        "expectation": expectation,
        "expected_behaviour_observed": expectation_ok,
    }


def negative_controls(raw: bytes) -> Dict[str, Dict[str, object]]:
    return {
        "body_single_byte_mutation": _checked_control(
            raw,
            _mutate_body(raw),
            "body-negative-control",
            (False, True, False),
            "body hash and overall DKIM fail; header-signature mathematics remains valid",
        ),
        "date_single_digit_mutation": _checked_control(
            raw,
            _mutate_date(raw),
            "date-negative-control",
            (True, False, False),
            "signed-header mathematics and overall DKIM fail; body hash remains valid",
        ),
        "unsigned_received_single_digit_mutation": _checked_control(
            raw,
            _mutate_unsigned_received(raw),
            "unsigned-header-scope-control",
            (True, True, True),
            "overall DKIM remains valid because Received is not in h=",
        ),
    }


def serialise(result: Verification) -> Dict[str, object]:
    return asdict(result)


def human_report(result: Verification, controls: Optional[Mapping[str, object]]) -> str:
    lines = [
        f"FILE             {result.path}",
        f"SHA-256          {result.sha256}",
        f"BYTES            {result.bytes}",
        f"SIGNATURE        d={result.domain}; s={result.selector}; {result.algorithm}; {result.canonicalisation}",
        f"BODY HASH        {'PASS' if result.body_hash_ok else 'FAIL'}",
        f"RSA MATH         {'PASS' if result.signature_math_ok else 'FAIL'}",
        f"OVERALL DKIM     {'PASS' if result.overall_dkim_ok else 'FAIL'}",
        f"MATCHING KEY     {result.matching_key or 'none'}",
        f"KEY SPKI SHA256  {result.matching_key_spki_sha256 or 'none'}",
        f"DNS BINDING      {result.selector_binding_status}",
        f"GOOGLE ARC       {result.google_arc_status}",
    ]
    if controls is not None:
        for name, control in controls.items():
            state = "PASS" if control["expected_behaviour_observed"] else "FAIL"
            lines.append(f"SCOPE CONTROL    {name}: {state} ({control['expectation']})")
    return "\n".join(lines)


def main(argv: Optional[Sequence[str]] = None) -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("eml", nargs="+", type=Path)
    parser.add_argument("--json", action="store_true", help="emit deterministic JSON")
    parser.add_argument(
        "--negative-controls",
        action="store_true",
        help=(
            "also test in-memory mutations of the body and signed Date header, "
            "plus an unsigned Received-header scope control"
        ),
    )
    args = parser.parse_args(argv)

    payloads = []
    success = True
    for path in args.eml:
        raw = path.read_bytes()
        result = verify_bytes(raw, str(path))
        controls = negative_controls(raw) if args.negative_controls else None
        success = success and result.overall_dkim_ok
        if controls is not None:
            success = success and all(item["expected_behaviour_observed"] for item in controls.values())
        payloads.append({"verification": serialise(result), "negative_controls": controls})

    if args.json:
        print(json.dumps(payloads, indent=2, sort_keys=True))
    else:
        for index, payload in enumerate(payloads):
            if index:
                print()
            verification = Verification(**{
                **payload["verification"],
                "key_attempts": [KeyAttempt(**item) for item in payload["verification"]["key_attempts"]],
            })
            print(human_report(verification, payload["negative_controls"]))
    return 0 if success else 1


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