#!/usr/bin/env python3
import argparse
import hashlib
import os
from pathlib import Path, PurePosixPath
import shutil
import stat
import sys
import zipfile

SCRIPT_SHEBANGS = {
    b'#!/usr/bin/env bash', b'#!/bin/bash', b'#!/bin/sh', b'#!/bin/ash',
    b'#!/usr/bin/env python3', b'#!/usr/bin/python3'
}
FIXED_TIME = (2026, 1, 1, 0, 0, 0)

class ContractError(Exception):
    pass

def safe_parts(name: str):
    if any(c in name for c in ('\x00', '\n', '\r', '\t')):
        raise ContractError(f'control_character:{name!r}')
    normalized = name.replace('\\', '/')
    p = PurePosixPath(normalized)
    parts = p.parts
    if not parts or p.is_absolute() or '..' in parts or ':' in parts[0]:
        raise ContractError(f'unsafe_path:{name!r}')
    return normalized, parts

def file_mode(path: Path) -> int:
    first = path.open('rb').readline(256).rstrip(b'\r\n')
    executable = bool(path.stat().st_mode & 0o111) or first in SCRIPT_SHEBANGS
    return 0o700 if executable else 0o600

def zip_type(mode: int):
    return stat.S_IFMT(mode) if mode else 0

def inspect_archive(archive: Path, expected_root: str, canonical: bool):
    rows = []
    roots = set()
    names = set()
    with zipfile.ZipFile(archive) as z:
        for member in z.infolist():
            name, parts = safe_parts(member.filename)
            if name in names:
                raise ContractError(f'duplicate_member:{name}')
            names.add(name)
            roots.add(parts[0])
            mode = (member.external_attr >> 16) & 0xFFFF
            kind = zip_type(mode)
            is_dir = member.is_dir() or name.endswith('/') or (kind and stat.S_ISDIR(mode))
            if kind and stat.S_ISLNK(mode):
                raise ContractError(f'symlink:{name}')
            if kind and not (stat.S_ISREG(mode) or stat.S_ISDIR(mode)):
                raise ContractError(f'special_file:{name}')
            if not is_dir and (stat.S_IMODE(mode) & 0o7000):
                raise ContractError(f'special_bits:{name}')
            if canonical:
                if is_dir:
                    raise ContractError(f'canonical_directory_entry:{name}')
                if member.create_system != 3:
                    raise ContractError(f'canonical_create_system:{name}:{member.create_system}')
                if kind != stat.S_IFREG:
                    raise ContractError(f'canonical_regular_type:{name}:{oct(mode)}')
                if stat.S_IMODE(mode) not in (0o600, 0o700):
                    raise ContractError(f'canonical_mode:{name}:{oct(stat.S_IMODE(mode))}')
            rows.append((member, name, parts, mode, is_dir))
    if roots != {expected_root}:
        raise ContractError(f'root_mismatch:{sorted(roots)}')
    return rows

def build(args):
    source = Path(args.source).resolve()
    output = Path(args.output).resolve()
    root = args.root or source.name
    safe_parts(root)
    if not source.is_dir():
        raise ContractError(f'source_not_directory:{source}')
    files = []
    for p in sorted(source.rglob('*')):
        if p.is_symlink():
            raise ContractError(f'source_symlink:{p.relative_to(source)}')
        if p.is_dir():
            continue
        if not p.is_file():
            raise ContractError(f'source_special:{p.relative_to(source)}')
        if '__pycache__' in p.parts or p.suffix in ('.pyc', '.pyo'):
            raise ContractError(f'python_bytecode:{p.relative_to(source)}')
        files.append(p)
    output.parent.mkdir(parents=True, exist_ok=True)
    tmp = output.with_suffix(output.suffix + f'.tmp.{os.getpid()}')
    try:
        with zipfile.ZipFile(tmp, 'w', compression=zipfile.ZIP_DEFLATED, compresslevel=9) as z:
            for p in files:
                rel = p.relative_to(source).as_posix()
                arc = f'{root}/{rel}'
                mode = file_mode(p)
                info = zipfile.ZipInfo(arc, FIXED_TIME)
                info.create_system = 3
                info.compress_type = zipfile.ZIP_DEFLATED
                info.external_attr = ((stat.S_IFREG | mode) << 16)
                info.flag_bits |= 0x800
                z.writestr(info, p.read_bytes(), compress_type=zipfile.ZIP_DEFLATED, compresslevel=9)
        inspect_archive(tmp, root, canonical=True)
        tmp.replace(output)
    finally:
        if tmp.exists():
            tmp.unlink()
    print('RESULT=PASS_ROUTER_ZIP_CONTRACT_BUILD')
    print(f'ZIP={output}')
    print(f'ZIP_SHA256={hashlib.sha256(output.read_bytes()).hexdigest()}')
    print(f'ZIP_FILE_COUNT={len(files)}')
    print('DIRECTORY_ENTRIES=0')
    print('MODE_POLICY=files_0600_scripts_0700_special_bits_stripped')

def verify(args):
    archive = Path(args.zip).resolve()
    rows = inspect_archive(archive, args.root, canonical=args.canonical)
    print('RESULT=PASS_ROUTER_ZIP_CONTRACT_VERIFY')
    print(f'ZIP={archive}')
    print(f'ZIP_SHA256={hashlib.sha256(archive.read_bytes()).hexdigest()}')
    print(f'ZIP_MEMBER_COUNT={len(rows)}')
    print(f'CANONICAL={str(args.canonical).lower()}')

def extract(args):
    archive = Path(args.zip).resolve()
    destination = Path(args.destination).resolve()
    rows = inspect_archive(archive, args.root, canonical=args.canonical)
    destination.mkdir(parents=True, exist_ok=True)
    with zipfile.ZipFile(archive) as z:
        for member, name, parts, mode, is_dir in rows:
            target = (destination / Path(*parts)).resolve()
            if target != destination and destination not in target.parents:
                raise ContractError(f'escape:{name}')
            if is_dir:
                target.mkdir(parents=True, exist_ok=True)
                os.chmod(target, 0o700)
                continue
            target.parent.mkdir(parents=True, exist_ok=True)
            for parent in [target.parent, *target.parent.parents]:
                if parent == destination.parent:
                    break
                if parent == destination or destination in parent.parents:
                    try: os.chmod(parent, 0o700)
                    except FileNotFoundError: pass
            perms = stat.S_IMODE(mode)
            normalized = 0o700 if perms & 0o111 else 0o600
            tmp = target.with_name(target.name + f'.tmp.{os.getpid()}')
            with z.open(member, 'r') as src, tmp.open('wb') as dst:
                shutil.copyfileobj(src, dst)
            os.chmod(tmp, normalized)
            tmp.replace(target)
    print('RESULT=PASS_ROUTER_ZIP_CONTRACT_EXTRACT')
    print(f'ZIP={archive}')
    print(f'DESTINATION={destination}')
    print(f'ZIP_MEMBER_COUNT={len(rows)}')
    print(f'CANONICAL={str(args.canonical).lower()}')

def main():
    p = argparse.ArgumentParser()
    sub = p.add_subparsers(dest='cmd', required=True)
    b = sub.add_parser('build')
    b.add_argument('--source', required=True)
    b.add_argument('--output', required=True)
    b.add_argument('--root')
    b.set_defaults(func=build)
    v = sub.add_parser('verify')
    v.add_argument('--zip', required=True)
    v.add_argument('--root', required=True)
    v.add_argument('--canonical', action='store_true')
    v.set_defaults(func=verify)
    e = sub.add_parser('extract')
    e.add_argument('--zip', required=True)
    e.add_argument('--destination', required=True)
    e.add_argument('--root', required=True)
    e.add_argument('--canonical', action='store_true')
    e.set_defaults(func=extract)
    args = p.parse_args()
    try:
        args.func(args)
    except (ContractError, zipfile.BadZipFile, OSError) as exc:
        print(f'ZIP_CONTRACT_ERROR={exc}', file=sys.stderr)
        print('RESULT=STOP_ROUTER_ZIP_CONTRACT', file=sys.stderr)
        raise SystemExit(61)

if __name__ == '__main__':
    main()
