"""Plot two numeric CSV columns with units in headers using standard Python.
Example: python3 plot-trials.py data/trials.csv --x trial --y distance_m --output plots/distance.svg
Raw data is never modified. Missing/nonfinite pairs are counted and reported.
"""
import argparse
import csv
from html import escape
from math import isfinite
from pathlib import Path
from statistics import mean, stdev


def read_pairs(path, x_column, y_column):
    points, skipped = [], 0
    with path.open(newline='') as handle:
        reader = csv.DictReader(handle)
        if not reader.fieldnames or any(c not in reader.fieldnames for c in [x_column, y_column]):
            raise ValueError('Requested columns are missing from CSV header')
        for row in reader:
            try:
                x, y = float(row[x_column]), float(row[y_column])
                if not isfinite(x) or not isfinite(y):
                    raise ValueError('nonfinite')
            except (ValueError, TypeError):
                skipped += 1
                continue
            points.append((x, y))
    if not points:
        raise ValueError('No finite numeric pairs to plot')
    return points, skipped


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('csv', type=Path)
    parser.add_argument('--x', required=True)
    parser.add_argument('--y', required=True)
    parser.add_argument('--output', type=Path, required=True)
    args = parser.parse_args()
    points, skipped = read_pairs(args.csv, args.x, args.y)
    xs, ys = zip(*points)
    xmin, xmax, ymin, ymax = min(xs), max(xs), min(ys), max(ys)
    dx, dy = max(xmax - xmin, 1e-9), max(ymax - ymin, 1e-9)
    xmap = lambda x: 80 + 470 * (x - xmin) / dx
    ymap = lambda y: 340 - 260 * (y - ymin) / dy
    title = f'{args.y} versus {args.x}'
    shapes = ['<svg xmlns="http://www.w3.org/2000/svg" width="640" height="450" viewBox="0 0 640 450">',
              '<rect width="640" height="450" fill="#faf8f0"/>',
              f'<text x="70" y="35" font-family="sans-serif" font-size="18">{escape(title)}</text>',
              '<path d="M80 60 V340 H570" fill="none" stroke="#222"/>']
    for x, y in points:
        shapes.append(f'<circle cx="{xmap(x):.2f}" cy="{ymap(y):.2f}" r="4" fill="#166c51"/>')
    shapes.extend([
        f'<text x="80" y="365" font-family="sans-serif" font-size="12">x: {xmin:.5g} to {xmax:.5g} ({escape(args.x)})</text>',
        f'<text x="80" y="390" font-family="sans-serif" font-size="12">y: {ymin:.5g} to {ymax:.5g} ({escape(args.y)})</text>',
        f'<text x="80" y="420" font-family="sans-serif" font-size="11">{len(points)} plotted pairs; {skipped} missing/invalid pairs. Full outcomes remain in raw CSV.</text>',
        '</svg>'])
    args.output.parent.mkdir(parents=True, exist_ok=True)
    args.output.write_text('\n'.join(shapes))
    print('plotted_pairs', len(points), 'missing_or_invalid_pairs', skipped)
    print('mean_y', mean(ys), 'sample_sd_y', stdev(ys) if len(ys) >= 2 else None)
    print('Metrics describe plotted valid pairs only; retain and report all trial outcomes separately.')

if __name__ == '__main__':
    main()
