#!/usr/bin/env python3
"""Bounded paper-derived forced-count/base-graph stage; no full packer."""
import argparse
import json
from pathlib import Path

from superstring_packaging_reference import BudgetExceeded, integer, parse, reduce_words

MAX_LENGTH = 128
MAX_SUBSTRINGS = 512
MAX_WORK = 2_000_000


def calculate(record, substring_budget=MAX_SUBSTRINGS, work_budget=MAX_WORK):
    integer(substring_budget, 1, MAX_SUBSTRINGS, 'substring budget')
    integer(work_budget, 1, MAX_WORK, 'count work budget')
    _, words = parse(record)
    core = reduce_words(words)
    length = sum(map(len, core))
    base = dict(status='unknown', forced_counts=None, paper_derived_lower_bound=None, base_edges=None,
                source_family='128', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                full_source_factor_two_algorithm_implemented=False, source_proof_verification='not_run',
                formal_program_verification='not_run', reduced_input_bytes=length,
                scope='Only the forced occurrence recursion and balanced base-graph formulas from source section 02. Lower-bound validity uses the paper argument, not an executed proof checker. Full periodic layers, requests, links, cycle opening and output construction are unimplemented. Finite comparison evidence is separate.',
                limits=dict(reduced_input_bytes=MAX_LENGTH, distinct_substrings=substring_budget, charged_work=work_budget),
                work_interpretation='Conservative character/rule/lookup charges; native preprocessing, Python object/storage costs and output encoding are excluded. No source polynomial runtime or byte-memory guarantee.')
    if length > MAX_LENGTH:
        return dict(base, status='outside_reduced_length_cap', work=0)
    work = 0

    def tick(amount=1):
        nonlocal work
        work += amount
        if work > work_budget:
            raise BudgetExceeded('Charged count-stage work cap exhausted')

    try:
        vertices = {b''}
        for word in core:
            for start in range(len(word)):
                for end in range(start + 1, len(word) + 1):
                    tick(end - start)
                    piece = word[start:end]
                    if piece not in vertices:
                        if len(vertices) >= substring_budget:
                            raise BudgetExceeded('Distinct-substring cap exhausted')
                        vertices.add(piece)
        alphabet = sorted({byte for word in core for byte in word})
        rules = {}
        for v in sorted(vertices, key=lambda x: (len(x), x)):
            for p in range(1, len(v)):
                template = v[:p]
                tick(len(v))
                if any(byte != template[i % p] for i, byte in enumerate(v)):
                    continue
                for k in range(1, (len(v) - 1) // p + 1):
                    tick()
                    w = v[:len(v) - k * p]
                    rules.setdefault(w, []).append((v, template, k))
        counts, blocking_cache = {}, {}

        def blocking(s, template):
            key = (s, template)
            if key in blocking_cache:
                return blocking_cache[key]
            total, size, p = 0, len(s), len(template)
            for r in vertices:
                if len(r) < size + 2:
                    continue
                for j in range(1, len(r) - size):
                    tick(size)
                    if r[j:j + size] != s:
                        continue
                    tick(2)
                    if r[0] == template[-j % p] or r[-1] == template[(len(r) - 1 - j) % p]:
                        continue
                    matches = True
                    for h in range(1, len(r) - 1):
                        tick()
                        if r[h] != template[(h - j) % p]:
                            matches = False
                            break
                    if matches:
                        if r not in counts:
                            raise ArithmeticError('Blocking recursion requested an unevaluated longer count')
                        total += counts[r]
            blocking_cache[key] = total
            return total

        for s in sorted(vertices - {b''}, key=lambda x: (-len(x), x)):
            tick(2 * len(alphabet) + 1)
            lower = max(sum(counts.get(bytes([c]) + s, 0) for c in alphabet),
                        sum(counts.get(s + bytes([c]), 0) for c in alphabet), int(s in core))
            for v, template, k in rules.get(s, []):
                tick()
                if counts[v] > blocking(v, template):
                    lower = max(lower, blocking(s, template) + k + 1)
            counts[s] = lower
            if lower > length:
                raise ArithmeticError('Count exceeds concatenation occurrence bound')
        edges, balance = [], {s: 0 for s in vertices}
        for s in sorted(counts):
            tick(2 * len(alphabet) + 1)
            up = counts[s] - sum(counts.get(bytes([c]) + s, 0) for c in alphabet)
            down = counts[s] - sum(counts.get(s + bytes([c]), 0) for c in alphabet)
            if min(up, down) < 0:
                raise ArithmeticError('Negative base-edge multiplicity')
            if up:
                edges.append(dict(from_hex=s[:-1].hex(), to_hex=s.hex(), direction='up', multiplicity=up))
                balance[s[:-1]] -= up
                balance[s] += up
            if down:
                edges.append(dict(from_hex=s.hex(), to_hex=s[1:].hex(), direction='down', multiplicity=down))
                balance[s] -= down
                balance[s[1:]] += down
        weight = sum(counts.get(bytes([c]), 0) for c in alphabet)
        if any(balance.values()) or sum(e['multiplicity'] for e in edges if e['direction'] == 'up') != weight or sum(e['multiplicity'] for e in edges if e['direction'] == 'down') != weight:
            raise ArithmeticError('Base-graph balance/cost identity failed')
        return dict(base, status='complete_count_stage', forced_counts={s.hex(): counts[s] for s in sorted(counts)},
                    paper_derived_lower_bound=weight, base_edges=edges, base_graph_balanced=True,
                    distinct_substrings_including_empty=len(vertices), period_rule_records=sum(map(len, rules.values())),
                    work=work, required_strings_hex=[s.hex() for s in core])
    except BudgetExceeded as exc:
        return dict(base, status='budget_exceeded', reason=str(exc), work=work)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path)
    parser.add_argument('--substring-budget', type=int, default=MAX_SUBSTRINGS)
    parser.add_argument('--work-budget', type=int, default=MAX_WORK)
    args = parser.parse_args()
    result = calculate(json.loads(args.input.read_text()), args.substring_budget, 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()
