#!/usr/bin/env python3
"""Offline exact-byte check for the unsigned Cipher Discovery snapshot v1.

This cannot authenticate the archive creator. Do not use it as a signature check.
"""

import argparse
import hashlib
import json
import re
import stat
import sys
import zipfile
from pathlib import Path

MAX_BYTES = 8 * 1024 * 1024
PARTS = ("evidence-pack.json", "observations.json", "history.json")
NAMES = (*PARTS, "manifest.json", "SHA256SUMS")
SECTIONS = {
    "inventoryAssets", "publicSourceObservations", "cbomRawObservationRecords",
    "repositoryRawObservationRecords", "auditEvents", "driftEvents",
    "verificationPlans", "verificationResults", "cbomImports", "cbomDiffEvents",
    "repositoryImports", "repositoryDiffEvents",
}
HEXDIGEST = re.compile(r"[0-9a-f]{64}\Z")


class InvalidSnapshot(ValueError):
    pass


def unique_object(items):
    result = {}
    for key, value in items:
        if key in result:
            raise InvalidSnapshot("duplicate JSON key")
        result[key] = value
    return result


def json_part(data):
    try:
        return json.loads(data, object_pairs_hook=unique_object,
                          parse_constant=lambda _: (_ for _ in ()).throw(InvalidSnapshot("non-finite JSON number")))
    except (ValueError, UnicodeError) as exc:
        raise InvalidSnapshot("malformed JSON") from exc


def exact_fields(value, fields):
    if not isinstance(value, dict) or set(value) != set(fields):
        raise InvalidSnapshot("unsupported object fields")


def nonnegative_int(value):
    return type(value) is int and value >= 0


def verify(path):
    if path.stat().st_size > MAX_BYTES:
        raise InvalidSnapshot("archive exceeds byte limit")
    try:
        with zipfile.ZipFile(path, "r") as archive:
            infos = archive.infolist()
            if tuple(info.filename for info in infos) != NAMES:
                raise InvalidSnapshot("unsupported, duplicate or reordered archive members")
            data = {}
            total = 0
            for info in infos:
                file_type = stat.S_IFMT(info.external_attr >> 16)
                if (info.compress_type != zipfile.ZIP_STORED or info.flag_bits & 1 or
                        file_type not in (0, stat.S_IFREG) or info.file_size != info.compress_size or
                        info.file_size > MAX_BYTES or info.is_dir()):
                    raise InvalidSnapshot("unsafe archive member")
                total += info.file_size
                if total > MAX_BYTES:
                    raise InvalidSnapshot("members exceed byte limit")
                data[info.filename] = archive.read(info)
    except (OSError, zipfile.BadZipFile, EOFError, RuntimeError, OverflowError) as exc:
        raise InvalidSnapshot("unreadable archive") from exc

    manifest = json_part(data["manifest.json"])
    exact_fields(manifest, (
        "schemaVersion", "snapshotId", "capturedAt", "generatorVersion", "evidencePackSchema",
        "status", "authenticity", "databaseIsolation", "hashAlgorithm", "sections", "files", "limitations",
    ))
    if (manifest["schemaVersion"] != "cipher-discovery-snapshot/1.0" or
            manifest["evidencePackSchema"] != "1.5.0" or manifest["status"] != "unsigned" or
            manifest["authenticity"] != "not_established" or
            manifest["databaseIsolation"] != "repeatable_read_read_only" or
            manifest["hashAlgorithm"] != "SHA-256" or
            not isinstance(manifest["snapshotId"], str) or
            not isinstance(manifest["capturedAt"], str) or
            not isinstance(manifest["generatorVersion"], str) or
            not isinstance(manifest["limitations"], list) or not manifest["limitations"]):
        raise InvalidSnapshot("unsupported manifest contract")
    if not isinstance(manifest["files"], list) or len(manifest["files"]) != len(PARTS):
        raise InvalidSnapshot("unexpected file manifest")
    expected_sums = []
    for name, item in zip(PARTS, manifest["files"]):
        exact_fields(item, ("path", "sha256", "byteCount"))
        digest = item["sha256"]
        if (item["path"] != name or not isinstance(digest, str) or not HEXDIGEST.fullmatch(digest) or
                not nonnegative_int(item["byteCount"]) or item["byteCount"] != len(data[name]) or
                hashlib.sha256(data[name]).hexdigest() != digest):
            raise InvalidSnapshot("member digest or byte count differs")
        expected_sums.append(f"{digest}  {name}\n")
    expected_sums.append(f"{hashlib.sha256(data['manifest.json']).hexdigest()}  manifest.json\n")
    if data["SHA256SUMS"] != "".join(expected_sums).encode("ascii"):
        raise InvalidSnapshot("checksum file differs")

    pack = json_part(data["evidence-pack.json"])
    observations = json_part(data["observations.json"])
    history = json_part(data["history.json"])
    if (not isinstance(pack, dict) or pack.get("schemaVersion") != "1.5.0" or
            not isinstance(observations, dict) or observations.get("schemaVersion") != manifest["schemaVersion"] or
            not isinstance(history, dict) or history.get("schemaVersion") != manifest["schemaVersion"]):
        raise InvalidSnapshot("payload schema differs")
    sections = manifest["sections"]
    exact_fields(sections, SECTIONS)
    included = {
        "inventoryAssets": pack.get("inventory"),
        "publicSourceObservations": observations.get("publicSourceObservations"),
        "auditEvents": history.get("recentAuditEvents"),
        "driftEvents": history.get("recentDriftEvents"),
        "verificationPlans": history.get("verificationPlans"),
        "cbomImports": pack.get("cbomImports"),
        "cbomDiffEvents": pack.get("recentCbomDiffEvents"),
        "repositoryImports": pack.get("repositoryImports"),
        "repositoryDiffEvents": pack.get("recentRepositoryDiffEvents"),
    }
    plans = included["verificationPlans"]
    if not isinstance(plans, list) or any(not isinstance(plan, dict) or not isinstance(plan.get("recentResults"), list) for plan in plans):
        raise InvalidSnapshot("malformed verification plans")
    included["verificationResults"] = [result for plan in plans for result in plan["recentResults"]]
    included["cbomRawObservationRecords"] = []
    included["repositoryRawObservationRecords"] = []
    for name, section in sections.items():
        exact_fields(section, ("total", "included", "omitted"))
        if (not isinstance(included[name], list) or
                not all(nonnegative_int(section[key]) for key in ("total", "included", "omitted")) or
                section["included"] != len(included[name]) or
                section["total"] - section["included"] != section["omitted"]):
            raise InvalidSnapshot("section counts differ")


def main():
    parser = argparse.ArgumentParser(description="Check an unsigned Cipher Discovery snapshot without network access.")
    parser.add_argument("archive", type=Path)
    args = parser.parse_args()
    try:
        verify(args.archive)
    except (InvalidSnapshot, OSError) as exc:
        print(f"INTEGRITY_FAILURE: {exc}", file=sys.stderr)
        return 2
    print("INTEGRITY_MATCH_UNSIGNED: bytes match the enclosed SHA-256 values; origin and truth are not authenticated.")
    return 0


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