#!/usr/bin/env python3
"""Bounded ordinary byte-string packing and complete artifact comparison."""
import argparse
import gzip
import hashlib
import json
import lzma
import re
import struct
import zlib
from pathlib import Path

MODEL = dict(alphabet='bytes', containment='exact_contiguous_substring', lookup='stable_ID_offset_and_length',
             terminators='not_required', mutability='immutable', additional_constraints='none')
MAX_ITEMS = 128
MAX_TOTAL = 16384
MAX_WORK = 2_000_000
MAX_EXACT = 12


class BudgetExceeded(Exception):
    pass


def integer(x, lo, hi, name):
    if type(x) is not int or not lo <= x <= hi:
        raise ValueError('Invalid ' + name)
    return x


def parse(record):
    if not isinstance(record, dict) or set(record) != {'strings', 'model'} or record['model'] != MODEL:
        raise ValueError('Require exact byte-string packing schema/model')
    rows = record['strings']
    if not isinstance(rows, list) or len(rows) > MAX_ITEMS:
        raise ValueError('Invalid string list')
    ids, words = [], []
    for row in rows:
        if not isinstance(row, dict) or set(row) != {'id', 'hex'}:
            raise ValueError('Invalid literal record')
        identity, encoded = row['id'], row['hex']
        if not isinstance(identity, str) or not re.fullmatch(r'[A-Za-z0-9_-]{1,64}', identity) or identity in ids:
            raise ValueError('Invalid or repeated ID')
        if not isinstance(encoded, str) or len(encoded) > 8192 or not re.fullmatch(r'(?:[0-9a-f]{2})*', encoded):
            raise ValueError('Require lowercase exact even-length hex, at most 4096 bytes per literal')
        ids.append(identity)
        words.append(bytes.fromhex(encoded))
    if sum(map(len, words)) > MAX_TOTAL:
        raise ValueError('Total input byte cap exceeded')
    return ids, words


def reduce_words(words):
    unique = sorted(set(words))
    return [word for word in unique if word and not any(word != other and word in other for other in unique)]


def overlap(a, b, tick=lambda: None):
    for length in range(min(len(a), len(b)), 0, -1):
        tick()
        if a[-length:] == b[:length]:
            return length
    return 0


def greedy_pack(words, tick):
    current = reduce_words(words)
    while len(current) > 1:
        choice = None
        for i, a in enumerate(current):
            for j, b in enumerate(current):
                if i == j:
                    continue
                amount = overlap(a, b, tick)
                candidate = (-amount, a, b, i, j)
                if choice is None or candidate < choice:
                    choice = candidate
        negative, a, b, i, j = choice
        merged = a + b[-negative:]
        current = reduce_words([word for k, word in enumerate(current) if k not in (i, j)] + [merged])
    return current[0] if current else b''


def exact_pack(words, tick):
    core = reduce_words(words)
    n = len(core)
    if n > MAX_EXACT:
        return None, 'outside_exact_item_cap'
    if not core:
        return b'', 'complete'
    overlaps = [[overlap(a, b, tick) if i != j else 0 for j, b in enumerate(core)] for i, a in enumerate(core)]
    states = {(1 << i, i): (len(word), None) for i, word in enumerate(core)}
    for mask in range(1, 1 << n):
        for last in range(n):
            if (mask, last) not in states:
                continue
            length, _ = states[mask, last]
            for nxt in range(n):
                if mask & (1 << nxt):
                    continue
                tick()
                key = (mask | (1 << nxt), nxt)
                candidate = (length + len(core[nxt]) - overlaps[last][nxt], last)
                if key not in states or candidate < states[key]:
                    states[key] = candidate
    full = (1 << n) - 1
    last = min(range(n), key=lambda j: (states[full, j][0], core[j]))
    order, mask = [], full
    while last is not None:
        order.append(last)
        previous = states[mask, last][1]
        mask ^= 1 << last
        last = previous
    order.reverse()
    result = core[order[0]]
    for a, b in zip(order, order[1:]):
        result += core[b][overlaps[a][b]:]
    return result, 'complete'


def artifact(ids, words, blob):
    body = bytearray(b'SSP1' + struct.pack('>II', len(ids), len(blob)))
    positions = []
    for identity, word in zip(ids, words):
        position = blob.find(word)
        if position < 0:
            raise ValueError('Blob does not contain every input literal')
        encoded = identity.encode('ascii')
        body += bytes([len(encoded)]) + encoded + struct.pack('>II', position, len(word))
        positions.append(dict(id=identity, offset=position, length=len(word)))
    body += blob
    return bytes(body), positions


def decode_artifact(encoded):
    if not isinstance(encoded, bytes) or len(encoded) < 12 or encoded[:4] != b'SSP1':
        raise ValueError('Invalid artifact header')
    count, blob_length = struct.unpack('>II', encoded[4:12])
    if count > MAX_ITEMS or blob_length > MAX_TOTAL:
        raise ValueError('Artifact caps exceeded')
    index, position, ids = [], 12, set()
    for _ in range(count):
        if position >= len(encoded):
            raise ValueError('Truncated index')
        size = encoded[position]
        position += 1
        if not 1 <= size <= 64 or position + size + 8 > len(encoded):
            raise ValueError('Invalid ID/index size')
        try:
            identity = encoded[position:position + size].decode('ascii')
        except UnicodeDecodeError as exc:
            raise ValueError('Non-ASCII identity') from exc
        position += size
        if not re.fullmatch(r'[A-Za-z0-9_-]{1,64}', identity) or identity in ids:
            raise ValueError('Invalid/repeated artifact identity')
        ids.add(identity)
        offset, length = struct.unpack('>II', encoded[position:position + 8])
        position += 8
        if offset + length > blob_length:
            raise ValueError('Index exceeds blob')
        index.append((identity, offset, length))
    if position + blob_length != len(encoded):
        raise ValueError('Artifact payload size differs from header')
    blob = encoded[position:]
    return [(identity, blob[offset:offset + length]) for identity, offset, length in index]


def metrics(ids, words, blob):
    encoded, positions = artifact(ids, words, blob)
    if decode_artifact(encoded) != list(zip(ids, words)):
        raise ArithmeticError('Artifact round trip failed')
    compressed = dict(gzip=gzip.compress(encoded, compresslevel=9, mtime=0), zlib=zlib.compress(encoded, 9), xz=lzma.compress(encoded, preset=6))
    for name, contents in compressed.items():
        decoder = dict(gzip=gzip.decompress, zlib=zlib.decompress, xz=lzma.decompress)[name]
        if decoder(contents) != encoded:
            raise ArithmeticError('Compression round trip failed')
    return dict(blob_hex=blob.hex(), blob_bytes=len(blob), complete_artifact_hex=encoded.hex(),
                complete_artifact_bytes=len(encoded), index_and_header_bytes=len(encoded) - len(blob),
                positions=positions, artifact_sha256=hashlib.sha256(encoded).hexdigest(),
                compressed_complete_artifact_bytes={name: len(contents) for name, contents in compressed.items()},
                roundtrip_checked=True)


def calculate(record, work_budget=MAX_WORK):
    integer(work_budget, 1, MAX_WORK, 'work budget')
    ids, words = parse(record)
    work = 0

    def tick():
        nonlocal work
        work += 1
        if work > work_budget:
            raise BudgetExceeded('Packing comparison work cap exhausted')

    methods = {'original_concatenation': b''.join(words), 'deduplicated_concatenation': b''.join(sorted(set(words))),
               'contained_string_concatenation': b''.join(reduce_words(words))}
    statuses = {}
    for name, run in [('maximum_overlap_greedy', lambda: (greedy_pack(words, tick), 'complete')),
                      ('exact_tiny_overlap_DP', lambda: exact_pack(words, tick))]:
        try:
            result, status = run()
            statuses[name] = status
            if result is not None:
                methods[name] = result
        except BudgetExceeded:
            statuses[name] = 'budget_exceeded'
    outputs = {name: metrics(ids, words, blob) for name, blob in methods.items()}
    optimum = len(methods['exact_tiny_overlap_DP']) if 'exact_tiny_overlap_DP' in methods else None
    return dict(status='bounded_comparison_completed', input_items=len(words), input_literal_bytes=sum(map(len, words)),
                maximal_distinct_nonempty_literals=len(reduce_words(words)), methods=outputs, method_statuses=statuses,
                exact_optimum_bytes=optimum, greedy_factor_two_on_this_exact_instance=None if optimum is None or 'maximum_overlap_greedy' not in methods else len(methods['maximum_overlap_greedy']) <= 2 * optimum,
                work=work, work_budget=work_budget, source_factor_two_algorithm_implemented=False,
                source_proof_verification='not_run', formal_program_verification='not_run',
                interpretation='Conventional bounded baselines and exact tiny DP. Source factor-two construction is not maximum-overlap greedy and is not implemented here. Caps count overlap-candidate tests/DP transitions, not character operations, native preprocessing, compression memory or time. SSP1 includes IDs and offset/length index; it supports immutable byte views including embedded zero, not drop-in C strings. Compressed artifacts require whole-artifact decompression; no random-access latency or customer value was measured.')


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path)
    parser.add_argument('--work-budget', type=int, default=MAX_WORK)
    args = parser.parse_args()
    result = calculate(json.loads(args.input.read_text()), args.work_budget)
    body = json.dumps(result, indent=2) + '\n'
    if args.output:
        with args.output.open('x') as stream:
            stream.write(body)
    else:
        print(body, end='')


if __name__ == '__main__':
    main()
