// Verify a Cipher Discovery signed snapshot offline with Go 1.26+.
// A key registry downloaded from the same website is not a trust anchor by itself.
package main

import (
	"archive/zip"
	"bytes"
	"crypto/ed25519"
	"crypto/sha256"
	"encoding/base64"
	"encoding/hex"
	"encoding/json"
	"errors"
	"flag"
	"fmt"
	"io"
	"os"
	"regexp"
	"time"
)

const maxBytes = 8 * 1024 * 1024
const domain = "CipherDiscoverySnapshotSignature/v1\x00"

var names = [...]string{"evidence-pack.json", "observations.json", "history.json", "manifest.json", "SHA256SUMS", "signature.json"}
var hex64 = regexp.MustCompile(`^[0-9a-f]{64}$`)

type keyRecord struct {
	KeyID           string `json:"keyId"`
	Algorithm       string `json:"algorithm"`
	PublicKeyBase64 string `json:"publicKeyBase64"`
	Status          string `json:"status"`
	IssuedAt        string `json:"issuedAt,omitempty"`
	RetiredAt       string `json:"retiredAt,omitempty"`
	RevokedAt       string `json:"revokedAt,omitempty"`
	Notice          string `json:"notice,omitempty"`
}
type registry struct {
	SchemaVersion string      `json:"schemaVersion"`
	Keys          []keyRecord `json:"keys"`
}
type envelope struct {
	SchemaVersion   string `json:"schemaVersion"`
	Algorithm       string `json:"algorithm"`
	KeyID           string `json:"keyId"`
	PublicKeySHA256 string `json:"publicKeySha256"`
	PublicKeyBase64 string `json:"publicKeyBase64"`
	SignedObject    string `json:"signedObject"`
	SignatureBase64 string `json:"signatureBase64"`
}
type section struct {
	Total    int `json:"total"`
	Included int `json:"included"`
	Omitted  int `json:"omitted"`
}
type fileEntry struct {
	Path      string `json:"path"`
	SHA256    string `json:"sha256"`
	ByteCount int    `json:"byteCount"`
}
type manifest struct {
	SchemaVersion      string             `json:"schemaVersion"`
	SnapshotID         string             `json:"snapshotId"`
	CapturedAt         string             `json:"capturedAt"`
	GeneratorVersion   string             `json:"generatorVersion"`
	EvidencePackSchema string             `json:"evidencePackSchema"`
	Status             string             `json:"status"`
	Authenticity       string             `json:"authenticity"`
	DatabaseIsolation  string             `json:"databaseIsolation"`
	HashAlgorithm      string             `json:"hashAlgorithm"`
	Sections           map[string]section `json:"sections"`
	Files              []fileEntry        `json:"files"`
	Limitations        []string           `json:"limitations"`
}

func fail(message string) error { return errors.New(message) }

// rejectDuplicateKeys also bounds nesting, so different JSON parsers cannot disagree on a signed field.
func rejectDuplicateKeys(data []byte) error {
	decoder := json.NewDecoder(bytes.NewReader(data))
	var walk func(int) error
	walk = func(depth int) error {
		if depth > 64 {
			return fail("JSON nesting limit exceeded")
		}
		token, err := decoder.Token()
		if err != nil {
			return err
		}
		delim, ok := token.(json.Delim)
		if !ok {
			return nil
		}
		switch delim {
		case '{':
			seen := make(map[string]bool)
			for decoder.More() {
				keyToken, keyErr := decoder.Token()
				if keyErr != nil {
					return keyErr
				}
				key, keyOK := keyToken.(string)
				if !keyOK || seen[key] {
					return fail("duplicate or invalid JSON key")
				}
				seen[key] = true
				if err := walk(depth + 1); err != nil {
					return err
				}
			}
		case '[':
			for decoder.More() {
				if err := walk(depth + 1); err != nil {
					return err
				}
			}
		default:
			return fail("unexpected JSON delimiter")
		}
		_, err = decoder.Token()
		return err
	}
	if err := walk(0); err != nil {
		return err
	}
	if _, err := decoder.Token(); err != io.EOF {
		return fail("trailing JSON data")
	}
	return nil
}

func decodeStrict(data []byte, target any) error {
	if err := rejectDuplicateKeys(data); err != nil {
		return err
	}
	decoder := json.NewDecoder(bytes.NewReader(data))
	decoder.DisallowUnknownFields()
	if err := decoder.Decode(target); err != nil {
		return err
	}
	if err := decoder.Decode(new(any)); err != io.EOF {
		return fail("trailing JSON data")
	}
	return nil
}

func readArchive(path string) (map[string][]byte, error) {
	info, err := os.Stat(path)
	if err != nil || info.Size() > maxBytes {
		return nil, fail("archive missing or oversized")
	}
	reader, err := zip.OpenReader(path)
	if err != nil {
		return nil, fail("unreadable ZIP")
	}
	defer reader.Close()
	if len(reader.File) != len(names) {
		return nil, fail("unsupported ZIP member count")
	}
	parts := make(map[string][]byte, len(names))
	total := 0
	for index, member := range reader.File {
		if member.Name != names[index] || member.Method != zip.Store || member.Flags&1 != 0 || member.FileInfo().IsDir() || !member.Mode().IsRegular() || member.UncompressedSize64 != member.CompressedSize64 || member.UncompressedSize64 > maxBytes {
			return nil, fail("unsafe or reordered ZIP member")
		}
		total += int(member.UncompressedSize64)
		if total > maxBytes {
			return nil, fail("ZIP member size limit exceeded")
		}
		stream, openErr := member.Open()
		if openErr != nil {
			return nil, fail("unreadable ZIP member")
		}
		data, readErr := io.ReadAll(io.LimitReader(stream, maxBytes+1))
		closeErr := stream.Close()
		if readErr != nil || closeErr != nil || len(data) != int(member.UncompressedSize64) {
			return nil, fail("corrupt ZIP member")
		}
		parts[member.Name] = data
	}
	return parts, nil
}

func expectedFields(data []byte, names ...string) error {
	var fields map[string]json.RawMessage
	if err := json.Unmarshal(data, &fields); err != nil || len(fields) != len(names) {
		return fail("unsupported JSON fields")
	}
	for _, name := range names {
		if len(fields[name]) == 0 {
			return fail("missing JSON field")
		}
	}
	return nil
}

func listLength(data []byte, key string) (int, error) {
	var object map[string]json.RawMessage
	if err := json.Unmarshal(data, &object); err != nil {
		return 0, err
	}
	var list []json.RawMessage
	if err := json.Unmarshal(object[key], &list); err != nil || list == nil {
		return 0, fail("missing section list")
	}
	return len(list), nil
}

func verifyCounts(parts map[string][]byte, counts map[string]section) error {
	if len(counts) != 12 {
		return fail("unsupported section count")
	}
	paths := map[string]struct{ file, key string }{
		"inventoryAssets": {"evidence-pack.json", "inventory"}, "publicSourceObservations": {"observations.json", "publicSourceObservations"},
		"auditEvents": {"history.json", "recentAuditEvents"}, "driftEvents": {"history.json", "recentDriftEvents"}, "verificationPlans": {"history.json", "verificationPlans"},
		"cbomImports": {"evidence-pack.json", "cbomImports"}, "cbomDiffEvents": {"evidence-pack.json", "recentCbomDiffEvents"},
		"repositoryImports": {"evidence-pack.json", "repositoryImports"}, "repositoryDiffEvents": {"evidence-pack.json", "recentRepositoryDiffEvents"},
	}
	for name, section := range counts {
		if section.Total < 0 || section.Included < 0 || section.Omitted < 0 || section.Total-section.Included != section.Omitted {
			return fail("invalid section count")
		}
		if name == "cbomRawObservationRecords" || name == "repositoryRawObservationRecords" {
			if section.Included != 0 {
				return fail("raw observation count mismatch")
			}
			continue
		}
		if name == "verificationResults" {
			continue
		}
		path, ok := paths[name]
		if !ok {
			return fail("unknown section")
		}
		length, err := listLength(parts[path.file], path.key)
		if err != nil || length != section.Included {
			return fail("section count differs")
		}
	}
	var history struct {
		VerificationPlans []struct {
			RecentResults []json.RawMessage `json:"recentResults"`
		} `json:"verificationPlans"`
	}
	if err := json.Unmarshal(parts["history.json"], &history); err != nil {
		return err
	}
	results := 0
	for _, plan := range history.VerificationPlans {
		if plan.RecentResults == nil {
			return fail("malformed verification results")
		}
		results += len(plan.RecentResults)
	}
	if counts["verificationResults"].Included != results {
		return fail("verification result count differs")
	}
	return nil
}

func verifyArchive(parts map[string][]byte) (envelope, []byte, error) {
	var empty envelope
	for _, name := range names[:4] {
		if err := rejectDuplicateKeys(parts[name]); err != nil {
			return empty, nil, err
		}
	}
	if err := expectedFields(parts["manifest.json"], "schemaVersion", "snapshotId", "capturedAt", "generatorVersion", "evidencePackSchema", "status", "authenticity", "databaseIsolation", "hashAlgorithm", "sections", "files", "limitations"); err != nil {
		return empty, nil, err
	}
	var m manifest
	if err := decodeStrict(parts["manifest.json"], &m); err != nil {
		return empty, nil, err
	}
	if m.SchemaVersion != "cipher-discovery-snapshot/1.1" || m.Status != "signed" || m.Authenticity != "producer_key_ed25519" || m.HashAlgorithm != "SHA-256" || m.DatabaseIsolation != "repeatable_read_read_only" || m.EvidencePackSchema != "1.5.0" || m.SnapshotID == "" || m.CapturedAt == "" || m.GeneratorVersion == "" || len(m.Limitations) == 0 || len(m.Files) != 3 {
		return empty, nil, fail("unsupported signed manifest")
	}
	var checksum bytes.Buffer
	for index, name := range names[:4] {
		hash := sha256.Sum256(parts[name])
		fmt.Fprintf(&checksum, "%x  %s\n", hash, name)
		if index < 3 {
			entry := m.Files[index]
			if entry.Path != name || entry.SHA256 != hex.EncodeToString(hash[:]) || entry.ByteCount != len(parts[name]) {
				return empty, nil, fail("manifest digest or length differs")
			}
		}
	}
	if !bytes.Equal(parts["SHA256SUMS"], checksum.Bytes()) {
		return empty, nil, fail("checksum file differs")
	}
	var pack, obs, history map[string]json.RawMessage
	for _, item := range []struct {
		data   []byte
		target *map[string]json.RawMessage
	}{{parts["evidence-pack.json"], &pack}, {parts["observations.json"], &obs}, {parts["history.json"], &history}} {
		if err := json.Unmarshal(item.data, item.target); err != nil {
			return empty, nil, err
		}
	}
	if string(pack["schemaVersion"]) != `"1.5.0"` || string(obs["schemaVersion"]) != `"cipher-discovery-snapshot/1.0"` || string(history["schemaVersion"]) != `"cipher-discovery-snapshot/1.0"` {
		return empty, nil, fail("payload schema differs")
	}
	if err := verifyCounts(parts, m.Sections); err != nil {
		return empty, nil, err
	}
	if err := expectedFields(parts["signature.json"], "schemaVersion", "algorithm", "keyId", "publicKeySha256", "publicKeyBase64", "signedObject", "signatureBase64"); err != nil {
		return empty, nil, err
	}
	var e envelope
	if err := decodeStrict(parts["signature.json"], &e); err != nil {
		return empty, nil, err
	}
	canonicalEnvelope, err := json.MarshalIndent(e, "", "  ")
	if err != nil || !bytes.Equal(parts["signature.json"], append(canonicalEnvelope, '\n')) {
		return empty, nil, fail("non-canonical signature envelope")
	}
	if e.SchemaVersion != "cipher-discovery-signature/1.0" || e.Algorithm != "Ed25519" || e.SignedObject != "manifest.json" || !hex64.MatchString(e.PublicKeySHA256) || e.KeyID != "ed25519-sha256:"+e.PublicKeySHA256 {
		return empty, nil, fail("unsupported signature envelope")
	}
	return e, parts["manifest.json"], nil
}

func verifyTrust(registryPath, expectedDigest string, e envelope, manifestBytes []byte) (string, error) {
	if !hex64.MatchString(expectedDigest) {
		return "", fail("registry SHA-256 pin is required")
	}
	data, err := os.ReadFile(registryPath)
	if err != nil || len(data) > 64*1024 {
		return "", fail("registry missing or oversized")
	}
	hash := sha256.Sum256(data)
	if hex.EncodeToString(hash[:]) != expectedDigest {
		return "", fail("TRUST_ANCHOR_MISMATCH")
	}
	if err := expectedFields(data, "schemaVersion", "keys"); err != nil {
		return "", err
	}
	var trusted registry
	if err := decodeStrict(data, &trusted); err != nil || trusted.SchemaVersion != "cipher-discovery-trust/1.0" || len(trusted.Keys) == 0 || len(trusted.Keys) > 20 {
		return "", fail("invalid trusted registry")
	}
	embedded, err := base64.StdEncoding.Strict().DecodeString(e.PublicKeyBase64)
	if err != nil || len(embedded) != ed25519.PublicKeySize {
		return "", fail("invalid embedded public key")
	}
	publicHash := sha256.Sum256(embedded)
	if hex.EncodeToString(publicHash[:]) != e.PublicKeySHA256 {
		return "", fail("embedded key fingerprint differs")
	}
	signature, err := base64.StdEncoding.Strict().DecodeString(e.SignatureBase64)
	if err != nil || len(signature) != ed25519.SignatureSize || !ed25519.Verify(embedded, append([]byte(domain), manifestBytes...), signature) {
		return "", fail("SIGNATURE_INVALID")
	}
	seen := map[string]bool{}
	for _, key := range trusted.Keys {
		if seen[key.KeyID] || key.Algorithm != "Ed25519" || (key.Status != "active" && key.Status != "retired" && key.Status != "revoked") {
			return "", fail("invalid registry entry")
		}
		if key.IssuedAt != "" {
			if _, dateErr := time.Parse(time.RFC3339, key.IssuedAt); dateErr != nil {
				return "", fail("invalid key issuance date")
			}
		}
		for _, date := range []string{key.RetiredAt, key.RevokedAt} {
			if date != "" {
				if _, dateErr := time.Parse(time.RFC3339, date); dateErr != nil {
					return "", fail("invalid key lifecycle date")
				}
			}
		}
		seen[key.KeyID] = true
		public, decodeErr := base64.StdEncoding.Strict().DecodeString(key.PublicKeyBase64)
		if decodeErr != nil || len(public) != ed25519.PublicKeySize {
			return "", fail("invalid registry key")
		}
		fingerprint := sha256.Sum256(public)
		if key.KeyID != "ed25519-sha256:"+hex.EncodeToString(fingerprint[:]) {
			return "", fail("registry key ID differs")
		}
		if key.KeyID == e.KeyID {
			if !bytes.Equal(public, embedded) {
				return "", fail("trusted key differs")
			}
			switch key.Status {
			case "active":
				return "SIGNATURE_VALID_TRUSTED_KEY", nil
			case "retired":
				return "SIGNATURE_VALID_RETIRED_KEY_UNPROVEN_TIME", nil
			case "revoked":
				return "", fail("REVOKED_KEY")
			}
		}
	}
	return "", fail("UNTRUSTED_KEY")
}

func main() {
	archivePath := flag.String("archive", "", "signed snapshot ZIP")
	registryPath := flag.String("registry", "", "locally pinned public-key registry JSON")
	registryPin := flag.String("registry-sha256", "", "registry digest received through an independent channel")
	flag.Parse()
	if *archivePath == "" || *registryPath == "" || *registryPin == "" {
		fmt.Fprintln(os.Stderr, "INTEGRITY_FAILURE: archive, registry and independent registry SHA-256 pin are required")
		os.Exit(2)
	}
	parts, err := readArchive(*archivePath)
	if err == nil {
		var e envelope
		var manifestBytes []byte
		e, manifestBytes, err = verifyArchive(parts)
		if err == nil {
			var status string
			status, err = verifyTrust(*registryPath, *registryPin, e, manifestBytes)
			if err == nil {
				fmt.Println(status + ": signature and enclosed bytes match a pinned key; observations, coverage and timestamp are not independently proven. Ed25519 is not post-quantum secure.")
				if status != "SIGNATURE_VALID_TRUSTED_KEY" {
					os.Exit(3)
				}
				return
			}
		}
	}
	fmt.Fprintln(os.Stderr, "INTEGRITY_FAILURE:", err)
	os.Exit(2)
}
