"""Synthetic teaching example; not historical evidence or the 2019 pipeline.

Run: python archaeogenetics-pca.py [output_directory]
Requires NumPy. Fixed seed; independent biallelic diploid loci (0, 1, 2).
Missingness is random and independent of genotype. Mean imputation is solely
for this tiny demonstration, not recommended as an ancient-DNA workflow.
"""
from pathlib import Path
import csv
import json
import sys
import numpy as np


def example():
    rng = np.random.default_rng(30874)
    # Gradually varying frequencies: no pure or separated population blocks.
    gradient = np.linspace(-1, 1, 12)
    baseline = rng.uniform(0.25, 0.75, 18)
    effect = rng.normal(0, 0.13, 18)
    frequencies = np.clip(baseline + gradient[:, None] * effect, 0.08, 0.92)
    matrix = rng.binomial(2, frequencies).astype(float)
    matrix[rng.random(matrix.shape) < 0.12] = np.nan
    keep = np.isfinite(matrix).any(axis=0)
    matrix = matrix[:, keep]
    means = np.nanmean(matrix, axis=0)
    filled = np.where(np.isnan(matrix), means, matrix)
    centered = filled - means
    # Sample SD after imputation; remove invariant columns before division.
    sd = centered.std(axis=0, ddof=1)
    variable = sd > 0
    standardized = centered[:, variable] / sd[variable]
    u, singular, vt = np.linalg.svd(standardized, full_matrices=False)
    scores = u[:, :2] * singular[:2]
    # Orient signs deterministically; PCA axis signs have no scientific meaning.
    for axis in range(2):
        pivot = np.argmax(np.abs(vt[axis]))
        if vt[axis, pivot] < 0:
            scores[:, axis] *= -1
    explained = singular[:2] ** 2 / np.sum(singular ** 2)
    assert np.isfinite(scores).all()
    assert np.allclose(standardized.mean(axis=0), 0, atol=1e-12)
    assert np.allclose(standardized.std(axis=0, ddof=1), 1)
    assert np.allclose((u * singular) @ vt, standardized)
    return matrix, scores, explained


if __name__ == '__main__':
    output = Path(sys.argv[1] if len(sys.argv) > 1 else '.')
    output.mkdir(parents=True, exist_ok=True)
    matrix, scores, explained = example()
    with (output / 'archaeogenetics-synthetic.csv').open('w', newline='') as f:
        writer = csv.writer(f)
        writer.writerow(['individual'] + [f'v{i+1}' for i in range(matrix.shape[1])])
        for i, row in enumerate(matrix):
            writer.writerow([f'S{i+1:02}'] + ['NA' if np.isnan(x) else int(x) for x in row])
    report = {'synthetic': True, 'seed': 30874, 'shape': list(matrix.shape),
              'missing_count': int(np.isnan(matrix).sum()),
              'PC1_PC2_variance': explained.tolist(), 'scores': scores.tolist(),
              'preprocessing': 'column mean imputation; center; sample SD scale; SVD'}
    (output / 'archaeogenetics-pca-result.json').write_text(json.dumps(report, indent=2) + '\n')
    print(json.dumps({k: v for k, v in report.items() if k != 'scores'}))
