#!/usr/bin/env python3
"""Derive the locked Emscripten build-source archive without two unused trees.

Only test/third_party and site are omitted. Archive order, retained member
metadata, file bytes and link semantics are verified after writing. Gzip
compression depends on Python/zlib; regenerating a locked release must still
produce its exact recorded SHA256, rather than relaxing that check.
"""
import argparse
import copy
import gzip
import hashlib
import io
import json
import os
from pathlib import Path, PurePosixPath
import platform
import posixpath
import stat
import tarfile
import zlib

EXCLUSIONS = ('test/third_party', 'site')
POLICY = 'audio-studio/emscripten-build-sources-slim1'


def digest(path):
    result = hashlib.sha256()
    with Path(path).open('rb') as stream:
        for chunk in iter(lambda: stream.read(128 * 1024), b''):
            result.update(chunk)
    return result.hexdigest()


def verify_file(path, record):
    path = Path(path)
    try:
        info = path.lstat()
    except FileNotFoundError:
        raise ValueError('Missing locked archive: ' + str(path)) from None
    if not stat.S_ISREG(info.st_mode) or info.st_size != record['bytes']:
        raise ValueError('Locked archive is not a regular file of the recorded size')
    if digest(path) != record['sha256']:
        raise ValueError('Locked archive SHA256 mismatch')


def relative_path(name, root):
    if not isinstance(name, str) or not name or '\x00' in name or '\\' in name:
        raise ValueError('Unsafe archive member path')
    parts = PurePosixPath(name).parts
    canonical = '/'.join(parts)
    if (name.startswith('/') or not parts or parts[0] != root or
            '..' in parts or name.rstrip('/') != canonical):
        raise ValueError('Unexpected archive root or unsafe member path: ' + name)
    return '/'.join(parts[1:])


def excluded(relative):
    return any(relative == path or relative.startswith(path + '/') for path in EXCLUSIONS)


def checked_members(archive, root):
    members = archive.getmembers()
    names = set()
    links = {}
    for member in members:
        relative_path(member.name, root)
        if member.name in names:
            raise ValueError('Duplicate archive member: ' + member.name)
        names.add(member.name)
        if member.issym():
            if member.linkname.startswith('/') or '\\' in member.linkname or '\x00' in member.linkname:
                raise ValueError('Unsafe archive symlink')
            target = posixpath.normpath(posixpath.join(posixpath.dirname(member.name), member.linkname))
            relative_path(target, root)
            links[member.name] = target
        elif member.islnk():
            relative_path(member.linkname, root)
            links[member.name] = member.linkname
        elif not (member.isfile() or member.isdir()):
            raise ValueError('Unsupported archive member type')
    for member in members:
        parts = PurePosixPath(member.name).parts
        if any('/'.join(parts[:count]) in links for count in range(1, len(parts))):
            raise ValueError('Archive member descends through a link')
        relative = relative_path(member.name, root)
        if member.name in links and not excluded(relative):
            target = links[member.name]
            if excluded(relative_path(target, root)) or target not in names:
                raise ValueError('Retained archive link would lose its target')
    return members


def member_record(archive, member, content=None):
    record = {key: getattr(member, key) for key in
              ('name', 'mode', 'uid', 'gid', 'size', 'mtime', 'uname', 'gname',
               'linkname', 'devmajor', 'devminor')}
    record['type'] = member.type.decode('ascii')
    record['pax_headers'] = dict(member.pax_headers)
    # GNU long-name headers and PAX path/linkpath encode the same effective
    # names differently. The effective fields above are checked exactly; omit
    # only redundant representation fields while retaining other PAX notices.
    for key, effective in [('path', member.name), ('linkpath', member.linkname)]:
        if record['pax_headers'].get(key) == effective:
            del record['pax_headers'][key]
    if member.isfile():
        result = hashlib.sha256()
        if content is not None:
            result.update(content)
        else:
            with archive.extractfile(member) as stream:
                for chunk in iter(lambda: stream.read(128 * 1024), b''):
                    result.update(chunk)
        record['contentSha256'] = result.hexdigest()
    return record


def records_hash(records):
    return hashlib.sha256(json.dumps(records, sort_keys=True, separators=(',', ':')).encode()).hexdigest()


def inventory(path, root):
    with tarfile.open(path, 'r:gz') as archive:
        members = checked_members(archive, root)
        return [member_record(archive, member) for member in members]


def derive(original, output, upstream, *, expected=None, receipt=None, inventory_file=None):
    original, output = Path(original), Path(output)
    if original.resolve() == output.resolve():
        raise ValueError('Derived output cannot overwrite the upstream archive')
    verify_file(original, upstream)
    root = upstream['archiveRoot']
    output.parent.mkdir(parents=True, exist_ok=True)
    if output.is_symlink() or (output.exists() and not output.is_file()):
        raise ValueError('Derived output is not a regular file')
    partial = output.with_name(output.name + '.partial')
    if partial.exists() or partial.is_symlink():
        raise ValueError('Refusing an existing derivation partial output')
    retained, removed = [], []
    try:
        with tarfile.open(original, 'r:gz') as source:
            members = checked_members(source, root)
            with partial.open('xb') as raw:
                with gzip.GzipFile(filename='', mode='wb', compresslevel=9, mtime=0, fileobj=raw) as compressed:
                    with tarfile.open(fileobj=compressed, mode='w|', format=tarfile.PAX_FORMAT,
                                      pax_headers=dict(source.pax_headers)) as target:
                        for member in members:
                            # Read once: reopening the same earlier file in a
                            # gzip stream for hashing and writing would rewind
                            # the decompressor thousands of times.
                            content = source.extractfile(member).read() if member.isfile() else None
                            record = member_record(source, member, content)
                            if excluded(relative_path(member.name, root)):
                                removed.append(record)
                                continue
                            retained.append(record)
                            target.addfile(copy.copy(member), io.BytesIO(content) if member.isfile() else None)
        observed = inventory(partial, root)
        if observed != retained:
            for before, after in zip(retained, observed):
                if before != after:
                    differences = {key: [before.get(key), after.get(key)]
                                   for key in set(before) | set(after)
                                   if before.get(key) != after.get(key)}
                    raise ValueError('Retained archive member changed: ' + before['name'] + ' ' + json.dumps(differences))
            raise ValueError('Retained archive member count or order changed')
        result = {'schema': 1, 'policy': POLICY, 'exclusions': list(EXCLUSIONS),
                  'upstream': dict(upstream),
                  'derived': {'archiveRoot': root, 'bytes': partial.stat().st_size,
                              'sha256': digest(partial)},
                  'members': {'original': len(retained) + len(removed), 'retained': len(retained),
                              'removed': len(removed),
                              'originalRegularFiles': sum(item['type'] in ('0', '\x00') for item in retained + removed),
                              'retainedRegularFiles': sum(item['type'] in ('0', '\x00') for item in retained),
                              'removedRegularFiles': sum(item['type'] in ('0', '\x00') for item in removed)},
                  'fileBytes': {'retained': sum(item['size'] for item in retained if item['type'] in ('0', '\x00')),
                                'removed': sum(item['size'] for item in removed if item['type'] in ('0', '\x00'))},
                  'retainedMembersMetadataAndContentSha256': records_hash(retained),
                  'removedMembersMetadataAndContentSha256': records_hash(removed),
                  'retainedMembersAndBytesCompared': True,
                  'compression': {'gzipMtime': 0, 'gzipFilename': '', 'level': 9,
                                  'pythonVersion': platform.python_version(),
                                  'zlibCompileVersion': zlib.ZLIB_VERSION,
                                  'zlibRuntimeVersion': zlib.ZLIB_RUNTIME_VERSION,
                                  'crossEnvironmentByteReproductionVerified': False},
                  'cleanBuildFromPrunedDeliveryVerified': False}
        if expected:
            verify_file(partial, expected)
        partial.replace(output)
        if receipt:
            Path(receipt).write_text(json.dumps(result, indent=2, sort_keys=True) + '\n')
        if inventory_file:
            detail = {'schema': 1, 'policy': POLICY, 'exclusions': list(EXCLUSIONS),
                      'retainedMembersMetadataAndContentSha256': records_hash(retained),
                      'removedMembers': removed}
            raw = (json.dumps(detail, indent=2, sort_keys=True) + '\n').encode()
            if Path(inventory_file).suffix == '.gz':
                with Path(inventory_file).open('wb') as stream:
                    with gzip.GzipFile(filename='', mode='wb', compresslevel=9, mtime=0, fileobj=stream) as compressed:
                        compressed.write(raw)
            else:
                Path(inventory_file).write_bytes(raw)
        return result
    finally:
        if partial.exists():
            partial.unlink()


def check_locked_derivation(item, bundle):
    derivation = item.get('derivation')
    if not derivation or derivation.get('policy') != POLICY or derivation.get('exclusions') != list(EXCLUSIONS):
        raise ValueError('Missing or unexpected source-pruning derivation policy')
    upstream = item.get('upstream')
    if not upstream or upstream.get('archiveRoot') != item['archiveRoot'] or upstream.get('commit') != item.get('commit'):
        raise ValueError('Upstream provenance does not match the derived input')
    verify_file(Path(bundle) / derivation['script'], derivation['scriptLock'])
    verify_file(Path(bundle) / derivation['receipt'], derivation['receiptLock'])
    verify_file(Path(bundle) / derivation['inventory'], derivation['inventoryLock'])
    receipt_data = json.loads((Path(bundle) / derivation['receipt']).read_text())
    inventory_path = Path(bundle) / derivation['inventory']
    if inventory_path.suffix == '.gz':
        with gzip.open(inventory_path, 'rt', encoding='utf-8') as stream:
            detail = json.load(stream)
    else:
        detail = json.loads(inventory_path.read_text())
    if (receipt_data.get('policy') != POLICY or receipt_data.get('exclusions') != list(EXCLUSIONS)
            or receipt_data.get('upstream') != upstream
            or receipt_data.get('derived', {}).get('sha256') != item['sha256']
            or receipt_data.get('derived', {}).get('bytes') != item['bytes']
            or not receipt_data.get('retainedMembersAndBytesCompared')):
        raise ValueError('Derivation receipt does not match locked input provenance')
    records = inventory(Path(bundle) / item['archive'], item['archiveRoot'])
    if any(excluded(relative_path(record['name'], item['archiveRoot'])) for record in records):
        raise ValueError('Excluded path remains in the derived archive')
    if (records_hash(records) != receipt_data['retainedMembersMetadataAndContentSha256']
            or detail.get('retainedMembersMetadataAndContentSha256') != records_hash(records)
            or detail.get('policy') != POLICY or detail.get('exclusions') != list(EXCLUSIONS)
            or len(records) != receipt_data['members']['retained']
            or len(detail.get('removedMembers', [])) != receipt_data['members']['removed']
            or records_hash(detail['removedMembers']) != receipt_data['removedMembersMetadataAndContentSha256']
            or any(not excluded(relative_path(record['name'], item['archiveRoot']))
                   for record in detail['removedMembers'])):
        raise ValueError('Derived archive member inventory does not match its receipt')


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--input', type=Path, required=True)
    parser.add_argument('--output', type=Path, required=True)
    parser.add_argument('--manifest', type=Path, default=Path(__file__).resolve().parents[1] / 'sources.json')
    parser.add_argument('--receipt', type=Path)
    parser.add_argument('--inventory', type=Path)
    parser.add_argument('--create-lock', action='store_true', help='Produce candidate hashes; release recovery never uses this option')
    args = parser.parse_args()
    manifest = json.loads(args.manifest.read_text())
    item = next(record for record in manifest['sources'] if record['id'] == 'emscripten')
    result = derive(args.input, args.output, item['upstream'],
                    expected=None if args.create_lock else item,
                    receipt=args.receipt, inventory_file=args.inventory)
    print(json.dumps(result, indent=2, sort_keys=True))


if __name__ == '__main__':
    main()
