#!/usr/bin/env python3
"""Empirical all-pairs finite Lp distortion audit, with no formal certification."""
import argparse
import json
import math
from pathlib import Path


def vectors(value, label):
    if not isinstance(value, list) or not value:
        raise ValueError(label + ' must be a nonempty list of vectors')
    if not isinstance(value[0], list) or not value[0]:
        raise ValueError(label + ' vectors must have positive dimension')
    dim = len(value[0])
    result = []
    for row in value:
        if not isinstance(row, list) or len(row) != dim:
            raise ValueError(label + ' dimensions are inconsistent')
        if any(isinstance(v, bool) or not isinstance(v, (int, float)) for v in row):
            raise ValueError(label + ' coordinates must be numeric')
        row = [float(v) for v in row]
        if not all(math.isfinite(v) for v in row):
            raise ValueError(label + ' coordinates must be finite')
        result.append(row)
    return result


def distance(a, b, p):
    delta = [abs(x - y) for x, y in zip(a, b)]
    scale = max(delta)
    if scale == 0:
        return 0.0
    if not math.isfinite(scale):
        raise ValueError('Coordinate difference overflow')
    result = scale * math.fsum((v / scale) ** p for v in delta) ** (1 / p)
    if not math.isfinite(result) or result <= 0:
        raise ValueError('Distance outside supported floating-point range')
    return result


def audit(data):
    p = data.get('p', 2)
    if isinstance(p, bool) or not isinstance(p, (int, float)) or not math.isfinite(p) or p <= 1:
        raise ValueError('p must be finite and greater than one')
    source = vectors(data['source_vectors'], 'source')
    target = vectors(data['embedded_vectors'], 'embedded')
    if len(source) != len(target):
        raise ValueError('Source and embedded point counts differ')
    ratios = []
    collapsed, split_duplicates = [], []
    pairs = 0
    for i in range(len(source)):
        for j in range(i + 1, len(source)):
            pairs += 1
            before, after = distance(source[i], source[j], p), distance(target[i], target[j], p)
            if before == 0:
                if after != 0:
                    split_duplicates.append([i, j])
            elif after == 0:
                collapsed.append([i, j])
            else:
                ratio = after / before
                if not math.isfinite(ratio) or ratio <= 0:
                    raise ValueError('Distance ratio outside supported floating-point range')
                ratios.append((ratio, [i, j]))
    lower = min(ratios) if ratios else None
    upper = max(ratios) if ratios else None
    bounded = not collapsed and not split_duplicates
    distortion = upper[0] / lower[0] if bounded and ratios else (1.0 if bounded else None)
    if distortion is not None and not math.isfinite(distortion):
        raise ValueError('Distortion outside supported floating-point range')
    return {'mode': 'EMPIRICAL_FLOATING_POINT_ALL_PAIRS', 'point_count': len(source), 'p': p,
            'source_dimension': len(source[0]), 'embedded_dimension': len(target[0]),
            'pairs_checked': pairs, 'distortion_after_global_rescaling': distortion,
            'bounded_distortion': bounded,
            'minimum_ratio': lower[0] if lower else None, 'minimum_ratio_pair': lower[1] if lower else None,
            'maximum_ratio': upper[0] if upper else None, 'maximum_ratio_pair': upper[1] if upper else None,
            'collapsed_pairs': collapsed, 'split_duplicate_pairs': split_duplicates,
            'limits': ['Only the supplied finite points and the specified Lp metrics were checked.',
                       'Floating-point calculations are empirical; no interval arithmetic or formal proof was run.',
                       'No future-query embedding map, efficient dimension-reduction theorem, retrieval accuracy or business value is established.']}


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--input', required=True, type=Path)
    parser.add_argument('--output', required=True, type=Path)
    args = parser.parse_args()
    report = audit(json.loads(args.input.read_text()))
    args.output.parent.mkdir(parents=True, exist_ok=True)
    with args.output.open('x') as stream:
        json.dump(report, stream, indent=2, allow_nan=False)
        stream.write('\n')
    print(json.dumps({key: report[key] for key in ['point_count', 'pairs_checked', 'bounded_distortion', 'distortion_after_global_rescaling']}))


if __name__ == '__main__':
    main()
