"""Independent matrix-capacity certificates, using MPFI rather than Arb.

For any exact input distribution pi, form Q=pi W. Then
    sum_x pi_x D(W_x || Q) <= C <= max_x D(W_x || Q).
The lower bound is achievable mutual information. The upper bound follows
from sum_x pi'_x D(W_x||Q) = I(pi',W) + D(pi' W||Q) for every pi'.
This is the information-radius characterization in Polyanskiy and Wu,
MIT 6.441 (2016), sections 4.4-4.5, Corollaries 4.2-4.3, printed pp. 45-46.

No formula or optimizer from channel_generate is imported. Two-input matrices
use sign-certified rational bisection on D(W_0||Q)-D(W_1||Q). Larger matrices
try the uniform prior and fail unless the resulting primal/dual gap closes.
Zero transition entries are skipped exactly, including zero output columns.
MPFR endpoints are extracted with exact_rational(), never QQ(endpoint),
which can approximate the endpoint and destroy the outward enclosure.
"""

from fractions import Fraction

import numberdb.sage
from sage.rings.rational_field import QQ
from sage.rings.real_mpfi import RealIntervalField


SOURCE_URL = ('https://ocw.mit.edu/courses/6-441-information-theory-spring-2016/'
              '5d8f16adc3385c9ff2975b121bd620e4_MIT6_441S16_course_notes.pdf')


def transition_matrix(channel, shape):
    values = [Fraction(part) for part in shape.split(',')]
    if channel in ('bsc', 'z'):
        if len(values) != 1:
            raise ValueError('Incorrect binary shape')
        probability = values[0]
        matrix = ([[1 - probability, probability], [probability, 1 - probability]]
                  if channel == 'bsc' else [[1, 0], [probability, 1 - probability]])
    elif channel == 'bac':
        first, second = values
        matrix = [[1 - first, first], [second, 1 - second]]
    elif channel == 'qsc':
        alphabet, probability = values
        if alphabet.denominator != 1 or alphabet < 2:
            raise ValueError('Invalid alphabet')
        matrix = [[1 - probability if row == column else probability / (alphabet - 1)
                   for column in range(int(alphabet))] for row in range(int(alphabet))]
    else:
        raise ValueError('Unknown channel')
    return validate_matrix(matrix)


def validate_matrix(matrix):
    rows = [[QQ(value) for value in row] for row in matrix]
    if not rows or not rows[0] or any(len(row) != len(rows[0]) for row in rows):
        raise ValueError('Empty or ragged matrix')
    if any(sum(row) != 1 or any(value < 0 for value in row) for row in rows):
        raise ValueError('Not a stochastic matrix')
    return rows


def exact_matrix_capacity(matrix, unit):
    matrix = validate_matrix(matrix)
    if all(row == matrix[0] for row in matrix):
        return QQ(0), 'Identical conditional output distributions'
    supports = [{index for index, value in enumerate(row) if value} for row in matrix]
    disjoint = all(not left.intersection(right)
                   for index, left in enumerate(supports) for right in supports[index + 1:])
    count = len(matrix)
    if disjoint and unit == 'bits' and count & (count - 1) == 0:
        return QQ(count.bit_length() - 1), 'Disjoint output supports; input recoverable'
    row_permutations = all(sorted(row) == sorted(matrix[0]) for row in matrix)
    column_sums = [sum(row[column] for row in matrix)
                   for column in range(len(matrix[0]))]
    if unit == 'bits' and row_permutations and len(set(column_sums)) == 1:
        coefficients = {}
        logarithms = [(QQ(len(matrix[0])), QQ(1))]
        logarithms.extend((value, value) for value in matrix[0] if value)
        for argument, weight in logarithms:
            for prime, exponent in argument.numerator().factor():
                coefficients[prime] = coefficients.get(prime, QQ(0)) + weight * exponent
            for prime, exponent in argument.denominator().factor():
                coefficients[prime] = coefficients.get(prime, QQ(0)) - weight * exponent
        if all(coefficient == 0 for prime, coefficient in coefficients.items() if prime != 2):
            return coefficients.get(2, QQ(0)), (
                'Row permutations and equal column sums give uniform optimal output; '
                'exact rational prime-log coefficients cancel except log(2)')
    return None, None


def exact_endpoints(interval):
    try:
        lower = interval.lower().exact_rational()
        upper = interval.upper().exact_rational()
    except (ValueError, TypeError, OverflowError) as error:
        raise ArithmeticError('Nonfinite interval endpoint') from error
    if lower > upper:
        raise ArithmeticError('Reversed interval endpoints')
    return lower, upper


def matrix_certificate(matrix, digits=100):
    matrix = validate_matrix(matrix)
    field = RealIntervalField(numberdb.sage.bits(digits, losing=192))
    target = QQ(1) / 10**(digits + 12)
    row_terms = [sum((field(value) * field(value).log() for value in row if value), field(0))
                 for row in matrix]
    lower_input, upper_input = QQ(0), QQ(1)
    for iteration in range(numberdb.sage.bits(digits, losing=64)):
        if len(matrix) == 2:
            input_zero = (lower_input + upper_input) / 2
            prior = [input_zero, 1 - input_zero]
        else:
            prior = [QQ(1) / len(matrix)] * len(matrix)
        output = [sum(prior[index] * row[column] for index, row in enumerate(matrix))
                  for column in range(len(matrix[0]))]
        log_output = [field(value).log() if value else None for value in output]
        divergences = [row_terms[index] - sum(
            (field(value) * log_output[column] for column, value in enumerate(row) if value),
            field(0)) for index, row in enumerate(matrix)]
        mutual = sum((field(weight) * value for weight, value in zip(prior, divergences)), field(0))
        lower = mutual.lower()
        upper = max(value.upper() for value in divergences)
        enclosure = field(lower, upper)
        for value in divergences:
            exact_endpoints(value)
        exact_lower, exact_upper = exact_endpoints(enclosure)
        if exact_upper - exact_lower < target:
            return enclosure, {
                'iterations': iteration + 1,
                'input_prior': [str(weight) for weight in prior],
                'gap_upper': str(field(upper) - field(lower)),
                'capacity_lower_nats': str(exact_lower),
                'capacity_upper_nats': str(exact_upper),
                'capacity_gap_nats': str(exact_upper - exact_lower),
                'endpoint_conversion': 'MPFR exact_rational()',
                'method': 'MPFI mutual-information lower / information-radius upper bound',
            }
        if len(matrix) != 2:
            raise ArithmeticError('Uniform prior did not certify capacity')
        derivative = divergences[0] - divergences[1]
        derivative_lower, derivative_upper = exact_endpoints(derivative)
        if derivative_lower > 0:
            lower_input = input_zero
        elif derivative_upper < 0:
            upper_input = input_zero
        else:
            raise ArithmeticError('Derivative sign unresolved before gap target')
    raise ArithmeticError('Matrix-capacity certificate exceeded iteration limit')


def certify_written(matrix, unit, written, certificate):
    from decimal import Decimal

    if unit not in ('nats', 'bits'):
        raise ValueError('Invalid unit')
    bound = certificate if unit == 'nats' else certificate / certificate.parent()(2).log()
    lower, upper = exact_endpoints(bound)
    exact, reason = exact_matrix_capacity(matrix, unit)
    if '.' not in written and 'e' not in written.lower():
        if exact is None or QQ(written) != exact or not lower <= exact <= upper:
            raise AssertionError('Exact entry lacks a matching structural matrix proof')
        return {'exact_proof': reason, 'exact_value': str(exact)}
    if exact is not None:
        raise AssertionError('Exact capacity was unnecessarily decimalized')
    decimal = Decimal(written)
    centre = QQ(Fraction(decimal))
    radius = QQ(10)**decimal.as_tuple().exponent
    if not (centre - radius <= lower <= upper <= centre + radius):
        raise AssertionError('Stored interval does not contain the independent certificate')
    if not upper - lower < radius / 10**8:
        raise AssertionError('Certificate is too wide to check the stored digits')
    return {'stored_interval_contains_matrix_certificate': True,
            'certificate_width_below_1e_minus_8_ulp': True}
