#!/usr/bin/env python3
"""Bounded rational k-server policy comparison and offline optimum.

Conventional finite dynamic programming, not the source competitive policy.
Policies here are fixed-label, nearest-label and independent uniform-label.
The fixed input is oblivious; capacities, deadlines and adaptive input are absent.
"""
import argparse
import hashlib
import json
from fractions import Fraction
from itertools import product
from pathlib import Path

MAX_POINTS = 8
MAX_SERVERS = 4
MAX_REQUESTS = 64
MAX_WORK = 3_000_000
POLICIES = ('fixed_label', 'nearest_label', 'uniform_label')


def calculate(record, work_budget=MAX_WORK):
    if not isinstance(record, dict) or set(record) != {'distances', 'initial', 'requests'}:
        raise ValueError('Provide distances, initial and requests only')
    if isinstance(work_budget, bool) or not isinstance(work_budget, int) or not 1 <= work_budget <= MAX_WORK:
        raise ValueError('Invalid work budget')
    raw = record['distances']
    if not isinstance(raw, list) or not 1 <= len(raw) <= MAX_POINTS:
        raise ValueError('Invalid point count')
    n = len(raw)
    distances = []
    for row in raw:
        if not isinstance(row, list) or len(row) != n:
            raise ValueError('Distance matrix must be square')
        values = []
        for value in row:
            if not isinstance(value, str) or len(value) > 80:
                raise ValueError('Distances must be bounded rational strings')
            try:
                q = Fraction(value)
            except (ValueError, ZeroDivisionError) as exc:
                raise ValueError('Invalid rational distance') from exc
            if max(abs(q.numerator).bit_length(), q.denominator.bit_length()) > 64:
                raise ValueError('Distance exceeds 64-bit rational input limit')
            values.append(q)
        distances.append(values)
    for i, j in product(range(n), repeat=2):
        if (distances[i][j] < 0 or (distances[i][j] == 0) != (i == j)
                or distances[i][j] != distances[j][i]):
            raise ValueError('Require a genuine symmetric metric')
        for h in range(n):
            if distances[i][j] > distances[i][h] + distances[h][j]:
                raise ValueError('Triangle inequality failed')
    initial, requests = record['initial'], record['requests']
    if not isinstance(initial, list) or not 1 <= len(initial) <= MAX_SERVERS:
        raise ValueError('Invalid server count')
    if not isinstance(requests, list) or len(requests) > MAX_REQUESTS:
        raise ValueError('Invalid request horizon')
    if any(isinstance(v, bool) or not isinstance(v, int) or not 0 <= v < n for v in initial + requests):
        raise ValueError('Positions and requests must be in-range integer point indices')
    initial = tuple(initial)
    k = len(initial)
    work = 0
    max_frontier = 1

    def charge():
        nonlocal work
        if work >= work_budget:
            raise OverflowError('State-label transition budget exhausted')
        work += 1

    def moved(state, label, request):
        result = list(state)
        result[label] = request
        return tuple(result)

    result = dict(
        input_record_sha256=hashlib.sha256(json.dumps(record, sort_keys=True, separators=(',', ':')).encode()).hexdigest(),
        points=n, servers=k, requests=len(requests), arithmetic='exact_rational',
        method='conventional_finite_dynamic_programming',
        source_connection=dict(family='110', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb', source_policy_implemented=False,
                               source_proof_verification='not_run'),
        formal_program_verification='not_run',
        limits=dict(points=MAX_POINTS, servers=MAX_SERVERS, requests=MAX_REQUESTS,
                    work_budget=work_budget, frontier_states=n**k, wall_time_deadline='none'),
        work_unit_definition='Offline and policy state-label transitions; metric validation, bit operations and elapsed time are separate.')
    try:
        # An attained minimum over all labeled configurations at each time.
        dp = {initial: Fraction(0)}
        histories = []
        for request in requests:
            next_dp, parents = {}, {}
            for state in sorted(dp):
                for label in range(k):
                    charge()
                    target = moved(state, label, request)
                    cost = dp[state] + distances[state[label]][request]
                    if target not in next_dp or cost < next_dp[target]:
                        next_dp[target] = cost
                        parents[target] = (state, label)
            dp = next_dp
            histories.append(parents)
            max_frontier = max(max_frontier, len(dp))
        final = min(dp, key=lambda state: (dp[state], state))
        optimum = dp[final]
        witness = []
        cursor = final
        for parents in reversed(histories):
            cursor, label = parents[cursor]
            witness.append(label)
        witness.reverse()
        reports = []
        for policy in POLICIES:
            # Values are probability mass and unconditional accrued cost mass.
            law = {initial: (Fraction(1), Fraction(0))}
            for request in requests:
                next_law = {}
                for state, (mass, cost_mass) in sorted(law.items()):
                    if policy == 'uniform_label':
                        labels, probability = range(k), Fraction(1, k)
                    else:
                        chosen = 0 if policy == 'fixed_label' else min(range(k), key=lambda label: (distances[state[label]][request], label))
                        labels, probability = [chosen], Fraction(1)
                    for label in labels:
                        charge()
                        target = moved(state, label, request)
                        new_mass = mass * probability
                        new_cost = (cost_mass + mass * distances[state[label]][request]) * probability
                        old_mass, old_cost = next_law.get(target, (Fraction(0), Fraction(0)))
                        next_law[target] = old_mass + new_mass, old_cost + new_cost
                law = next_law
                max_frontier = max(max_frontier, len(law))
            assert sum((mass for mass, _ in law.values()), Fraction(0)) == 1
            expected = sum((cost for _, cost in law.values()), Fraction(0))
            reports.append(dict(policy=policy, expected_cost=str(expected),
                                additive_excess_over_optimum=str(expected-optimum),
                                ratio_to_optimum=str(expected/optimum) if optimum else None,
                                ratio_status='defined' if optimum else 'undefined_zero_optimum',
                                final_configuration_support=len(law)))
    except OverflowError as exc:
        return dict(result, status='budget_exceeded', reason=str(exc), work_units=work,
                    offline_optimum=None, offline_labels=None, policy_reports=None,
                    interpretation='Final comparison unknown; no partial value is labeled a complete optimum or expectation.')
    return dict(result, status='complete', offline_optimum=str(optimum), offline_labels=witness,
                policy_reports=reports, work_units=work, maximum_frontier_states=max_frontier,
                metric_instance_matches_finite_point_and_server_domain=bool(2 <= k < n),
                interpretation='Exact for this bounded fixed sequence and the three named policies. No universal competitive guarantee, theorem backend, capacity/deadline model or empirical performance superiority.')


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