#!/usr/bin/env python3
"""Bounded exact unit-job DAG reference, never the source polynomial algorithm."""
import argparse
import hashlib
import itertools
import json
import re
from collections import deque
from pathlib import Path

MAX_JOBS = 18
MAX_STATES = 50_000
MAX_TRANSITIONS = 2_000_000
MODEL = dict(machine_count=3, processing_times='all_one', nonpreemptive=True, machines='identical',
             release_times='all_zero', communication_delays='none', eligibility='all_machines', additional_resources='none')


class BudgetExceeded(Exception):
    pass


def integer(value, low, high, name):
    if type(value) is not int or not low <= value <= high:
        raise ValueError('Invalid ' + name)
    return value


def parse(record):
    required = {'job_ids', 'precedence', 'model'}
    if not isinstance(record, dict) or not required <= set(record) or set(record) - required - {'deadline', 'proposed_slots'}:
        raise ValueError('Invalid scheduling schema')
    model = record['model']
    if not isinstance(model, dict) or set(model) != set(MODEL):
        raise ValueError('Require every explicit source-model field')
    if any(type(model[key]) is not type(value) or model[key] != value for key, value in MODEL.items()):
        raise ValueError('Outside the three-machine unit-job model')
    jobs = record['job_ids']
    if not isinstance(jobs, list) or not 1 <= len(jobs) <= MAX_JOBS:
        raise ValueError('Require one through eighteen explicitly listed jobs')
    if any(not isinstance(job, str) or not re.fullmatch(r'[A-Za-z0-9_-]{1,64}', job) for job in jobs) or len(set(jobs)) != len(jobs):
        raise ValueError('Invalid or repeated job identifier')
    indices = {job: i for i, job in enumerate(jobs)}
    n = len(jobs)
    edges = record['precedence']
    if not isinstance(edges, list) or len(edges) > n * (n - 1):
        raise ValueError('Invalid precedence list')
    pairs = []
    for edge in edges:
        if not isinstance(edge, list) or len(edge) != 2 or any(not isinstance(job, str) or job not in indices for job in edge):
            raise ValueError('Unknown precedence endpoint')
        pair = tuple(indices[job] for job in edge)
        if pair[0] == pair[1] or pair in pairs:
            raise ValueError('Self edge or repeated precedence')
        pairs.append(pair)
    predecessors = [0] * n
    successors = [[] for _ in jobs]
    indegree = [0] * n
    for a, b in pairs:
        predecessors[b] |= 1 << a
        successors[a].append(b)
        indegree[b] += 1
    ready = deque(i for i in range(n) if indegree[i] == 0)
    order = []
    while ready:
        i = ready.popleft()
        order.append(i)
        for j in successors[i]:
            indegree[j] -= 1
            if indegree[j] == 0:
                ready.append(j)
    if len(order) != n:
        raise ValueError('Precedence is cyclic')
    deadline = record.get('deadline')
    if 'deadline' in record:
        integer(deadline, 1, n, 'deadline')
    return jobs, pairs, predecessors, successors, order, deadline


def audit_slots(jobs, pairs, slots):
    errors = []
    n = len(jobs)
    if not isinstance(slots, list) or not 1 <= len(slots) <= 2 * n:
        return dict(valid=False, makespan=None, errors=['Invalid proposed slot-list shape or size'])
    times = {}
    known = set(jobs)
    for t, batch in enumerate(slots, 1):
        if not isinstance(batch, list) or len(batch) > 3:
            errors.append('Invalid batch or capacity exceeded at slot ' + str(t))
            continue
        for job in batch:
            if not isinstance(job, str) or job not in known:
                errors.append('Unknown job at slot ' + str(t))
            elif job in times:
                errors.append('Repeated job ' + job)
            else:
                times[job] = t
    if set(times) != known:
        errors.append('Some jobs are missing')
    for a, b in pairs:
        if jobs[a] in times and jobs[b] in times and times[jobs[a]] >= times[jobs[b]]:
            errors.append('Precedence violation ' + jobs[a] + ' -> ' + jobs[b])
    return dict(valid=not errors, makespan=max(times.values()) if not errors else None, errors=errors)


def calculate(record, state_budget=MAX_STATES, transition_budget=MAX_TRANSITIONS):
    integer(state_budget, 1, MAX_STATES, 'state budget')
    integer(transition_budget, 1, MAX_TRANSITIONS, 'transition budget')
    jobs, pairs, pred, succ, order, deadline = parse(record)
    n, full = len(jobs), (1 << len(jobs)) - 1
    prefix = [1] * n
    tail = [1] * n
    for i in order:
        for j in succ[i]:
            prefix[j] = max(prefix[j], prefix[i] + 1)
    for i in reversed(order):
        for j in succ[i]:
            tail[i] = max(tail[i], tail[j] + 1)
    lower = max((n + 2) // 3, max(prefix))
    done, greedy = 0, []
    while done != full:
        ready = [i for i in range(n) if not done & (1 << i) and pred[i] & done == pred[i]]
        batch = sorted(ready, key=lambda i: (-tail[i], i))[:3]
        greedy.append([jobs[i] for i in batch])
        done |= sum(1 << i for i in batch)
    greedy_audit = audit_slots(jobs, pairs, greedy)
    if not greedy_audit['valid']:
        raise ArithmeticError('Internal greedy witness invalid')
    candidate = audit_slots(jobs, pairs, record['proposed_slots']) if 'proposed_slots' in record else None
    witness = greedy
    if candidate and candidate['valid'] and candidate['makespan'] < len(witness):
        witness = record['proposed_slots'][:candidate['makespan']]
    upper = len(witness)
    result = dict(status='unknown', exact_optimum=None, optimal_slots=None, feasibility_by_deadline=None,
                  deadline=deadline, lower_bound=lower, lower_bound_components=dict(capacity=(n + 2) // 3, longest_precedence_chain=max(prefix)),
                  valid_upper_bound=upper, upper_bound_slots=witness, proposed_schedule_audit=candidate,
                  states=0, transitions=0, limits=dict(jobs=MAX_JOBS, states=state_budget, transitions=transition_budget),
                  input_sha256=hashlib.sha256(json.dumps(record, sort_keys=True, separators=(',', ':')).encode()).hexdigest(),
                  source_connection=dict(family='124', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                                         source_polynomial_algorithm_implemented=False, source_exponent=150020),
                  algorithm='Breadth-first search of completed-job ideals; exponential in the number of jobs, with hard resource caps.',
                  source_proof_verification='not_run', formal_program_verification='not_run',
                  interpretation='Exact finite optimum only after completed search or an independently checked witness attaining an elementary lower bound. A valid schedule is not alone an optimality certificate. No claim for unequal jobs, eligibility, release dates, setup, communication delay or general production scheduling.')
    if deadline is not None:
        if lower > deadline:
            result['feasibility_by_deadline'] = dict(feasible=False, basis='elementary_lower_bound_exceeds_deadline')
        elif upper <= deadline:
            result['feasibility_by_deadline'] = dict(feasible=True, basis='checked_schedule_meets_deadline')
    if lower == upper:
        return dict(result, status='complete', exact_optimum=upper, optimal_slots=witness, completion_basis='checked_witness_attains_elementary_lower_bound')
    previous = {0: None}
    queue = deque([0])
    transitions = 0
    target = None
    try:
        while queue:
            done = queue.popleft()
            ready = [i for i in range(n) if not done & (1 << i) and pred[i] & done == pred[i]]
            for size in range(min(3, len(ready)), 0, -1):
                for batch in itertools.combinations(ready, size):
                    transitions += 1
                    if transitions > transition_budget:
                        raise BudgetExceeded('Transition budget exhausted')
                    after = done | sum(1 << i for i in batch)
                    if after in previous:
                        continue
                    if len(previous) >= state_budget:
                        raise BudgetExceeded('State budget exhausted')
                    previous[after] = (done, batch)
                    if after == full:
                        target = after
                        break
                    queue.append(after)
                if target is not None:
                    break
            if target is not None:
                break
    except BudgetExceeded as exc:
        return dict(result, status='budget_exceeded', reason=str(exc), states=len(previous), transitions=transitions)
    if target is None:
        raise ArithmeticError('A finite DAG must have a serial schedule')
    slots = []
    while target:
        before, batch = previous[target]
        slots.append([jobs[i] for i in batch])
        target = before
    slots.reverse()
    if not audit_slots(jobs, pairs, slots)['valid']:
        raise ArithmeticError('Internal exact witness invalid')
    optimum = len(slots)
    result.update(status='complete', exact_optimum=optimum, optimal_slots=slots, valid_upper_bound=optimum,
                  upper_bound_slots=slots, states=len(previous), transitions=transitions, completion_basis='complete_breadth_first_search')
    if deadline is not None:
        result['feasibility_by_deadline'] = dict(feasible=optimum <= deadline, basis='complete_exact_search')
    return result


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path)
    parser.add_argument('--state-budget', type=int, default=MAX_STATES)
    parser.add_argument('--transition-budget', type=int, default=MAX_TRANSITIONS)
    args = parser.parse_args()
    body = json.dumps(calculate(json.loads(args.input.read_text()), args.state_budget, args.transition_budget), indent=2, allow_nan=False) + '\n'
    if args.output:
        with args.output.open('x') as stream:
            stream.write(body)
    else:
        print(body, end='')


if __name__ == '__main__':
    main()
