#!/usr/bin/env python3
"""Recompute published observations. Python 3, standard library only; no network."""
import csv
import json
from collections import defaultdict
from pathlib import Path


def percentile(values, fraction):
    ordered = sorted(values)
    position = (len(ordered) - 1) * fraction
    lower = int(position)
    upper = min(lower + 1, len(ordered) - 1)
    return ordered[lower] + (ordered[upper] - ordered[lower]) * (position - lower)


def summarize(data):
    results = []
    for study in data['studies']:
        groups = defaultdict(list)
        for run in study['runs']:
            assert len(run['ttft_all_ms']) == run['ok']
            assert run['ok'] + run['errors'] == run['requests_measured']
            assert run['requests_measured'] == 2 * run['N']
            assert run['measured_wall_s'] > 0
            groups[run['N']].append(run)
        for concurrency, runs in sorted(groups.items()):
            values = [value for run in runs for value in run['ttft_all_ms']]
            count = sum(run['completion_tokens'] for run in runs)
            elapsed = sum(run['measured_wall_s'] for run in runs)
            results.append({
                'study': study['id'], 'concurrency': concurrency, 'repeats': len(runs),
                'requests': sum(run['requests_measured'] for run in runs),
                'errors': sum(run['errors'] for run in runs),
                'ttft_p50_ms': round(percentile(values, .50), 1),
                'ttft_p95_ms': round(percentile(values, .95), 1),
                'recorded_output_per_s': round(count / elapsed, 1),
                'recorded_output_count': count,
                'measured_wave_seconds': round(elapsed, 3),
            })
    return results


if __name__ == '__main__':
    root = Path(__file__).resolve().parent
    data = json.loads((root / 'observations.json').read_text())
    results = summarize(data)
    with (root / 'summary.csv').open('w', newline='') as handle:
        writer = csv.DictWriter(handle, fieldnames=results[0].keys())
        writer.writeheader()
        writer.writerows(results)
    print(json.dumps(results, indent=2))
    connection = data['connection_reuse']['requests']
    times = [request['first_output_ms'] for request in connection]
    print('Connection probe:', json.dumps({
        'requests': len(times), 'median_ms': percentile(times, .5),
        'min_ms': min(times), 'max_ms': max(times),
        'new_connections': sum(request['new_connection'] for request in connection),
        'http_errors': sum(request['status'] != 200 for request in connection),
    }))
