#!/usr/bin/env python3
"""
Master Scientific Validation Script
====================================

Runs all scientific rigor modules and produces a comprehensive report.

Suites:
  1. Quantitative validation (analytical benchmarks)
  2. Ablation study (isolate FCE contribution)
  3. Convergence study (numerical convergence order)
  4. Uncertainty quantification (Monte Carlo parameter propagation)
  5. Conservation monitoring (trace, positivity, hermiticity)
  6. Parameter sensitivity (alpha sweep)

Output: output/scientific_validation/
"""

import sys
import os
import time
import numpy as np

# Add parent directory to path
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))

from fce.engine import FractalCorrectionEngine, EngineConfig
from fce.quantitative_validation import QuantitativeValidator
from fce.ablation import AblationStudy, AblationMode
from fce.convergence import ConvergenceStudy
from fce.uncertainty import MonteCarloUQ
from fce.conservation import ConservationMonitor
from fce.sensitivity import SensitivitySweep


def setup_output_dir():
    """Create output directory."""
    out_dir = os.path.join(os.path.dirname(__file__), '..', 'output', 'scientific_validation')
    os.makedirs(out_dir, exist_ok=True)
    return out_dir


def setup_test_system():
    """Create standard test system (qubit with decoherence)."""
    H = np.diag([0.0, 1.0]).astype(complex)
    psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
    rho = np.outer(psi, psi.conj())
    system_params = {'T1': 1e-3, 'T2': 5e-4, 'gate_fidelity': 0.999}
    return H, rho, system_params


def run_quantitative_validation(out_dir):
    """Suite 1: Quantitative validation against analytical references."""
    print("\n" + "=" * 70)
    print("SUITE 1: Quantitative Validation")
    print("=" * 70)

    validator = QuantitativeValidator(tolerance=1e-6)
    report = validator.run_all()

    table = report.summary_table()
    print(table)

    with open(os.path.join(out_dir, 'validation_table.txt'), 'w') as f:
        f.write(table)

    return report


def run_ablation_study(out_dir, H, rho, system_params):
    """Suite 2: Ablation study to isolate FCE contribution."""
    print("\n" + "=" * 70)
    print("SUITE 2: Ablation Study")
    print("=" * 70)

    study = AblationStudy(
        hamiltonian=H,
        rho_init=rho,
        system_params=system_params,
        dt=1e-6,
        n_steps=50,
        base_config=EngineConfig(feedback_gain=0.1, correction_threshold=0.001),
    )

    # Use fewer trials for speed; increase for publication
    report = study.run(n_trials=30, base_seed=42)

    table = report.summary_table()
    print(table)

    with open(os.path.join(out_dir, 'ablation_report.txt'), 'w') as f:
        f.write(table)

    return report


def run_convergence_study(out_dir, H, rho, system_params):
    """Suite 3: Numerical convergence study."""
    print("\n" + "=" * 70)
    print("SUITE 3: Convergence Study")
    print("=" * 70)

    total_time = 1e-4  # 100 microseconds

    study = ConvergenceStudy(
        hamiltonian=H,
        rho_init=rho,
        total_time=total_time,
        system_params=system_params,
        config=EngineConfig(feedback_gain=0.1, correction_threshold=0.001),
        observable_name="final_fidelity",
    )

    result = study.run(resolutions=[100, 200, 400, 800, 1600])

    table = result.summary_table()
    print(table)

    with open(os.path.join(out_dir, 'convergence_report.txt'), 'w') as f:
        f.write(table)

    return result


def run_uncertainty_quantification(out_dir, H, rho, system_params):
    """Suite 4: Monte Carlo uncertainty quantification."""
    print("\n" + "=" * 70)
    print("SUITE 4: Uncertainty Quantification")
    print("=" * 70)

    uq = MonteCarloUQ(
        hamiltonian=H,
        rho_init=rho,
        dt=1e-6,
        n_steps=50,
        base_config=EngineConfig(feedback_gain=0.1, correction_threshold=0.001),
        base_system_params=system_params,
    )

    # Use fewer samples for speed; increase for publication
    result = uq.run(n_samples=50, seed=42)

    table = result.summary_table()
    print(table)

    # Also print formatted results
    print("\nFormatted Results:")
    for metric in ['final_fidelity', 'error_reduction', 'n_corrections']:
        print(f"  {metric}: {result.format_result(metric)}")

    with open(os.path.join(out_dir, 'uncertainty_report.txt'), 'w') as f:
        f.write(table)
        f.write("\n\nFormatted Results:\n")
        for metric in result.metrics:
            f.write(f"  {metric}: {result.format_result(metric)}\n")

    return result


def run_conservation_monitoring(out_dir, H, rho, system_params):
    """Suite 5: Conservation law monitoring."""
    print("\n" + "=" * 70)
    print("SUITE 5: Conservation Monitoring")
    print("=" * 70)

    engine = FractalCorrectionEngine(
        hamiltonian=H,
        system_params=system_params,
        config=EngineConfig(feedback_gain=0.1, correction_threshold=0.001),
    )
    result = engine.evolve(rho, dt=1e-6, n_steps=100)

    report = ConservationMonitor.from_evolution_result(result)

    table = report.summary_table()
    print(table)

    with open(os.path.join(out_dir, 'conservation_report.txt'), 'w') as f:
        f.write(table)

    return report


def run_sensitivity_sweep(out_dir, H, rho, system_params):
    """Suite 6: Parameter sensitivity sweep."""
    print("\n" + "=" * 70)
    print("SUITE 6: Parameter Sensitivity (alpha)")
    print("=" * 70)

    sweep = SensitivitySweep(
        hamiltonian=H,
        rho_init=rho,
        system_params=system_params,
        dt=1e-6,
        n_steps=50,
    )

    result = sweep.run(
        alpha_values=[0.0, 0.01, 0.02, 0.05, 0.1, 0.2, 0.5, 1.0],
        n_repeats=1,
    )

    table = result.summary_table()
    print(table)

    with open(os.path.join(out_dir, 'sensitivity_report.txt'), 'w') as f:
        f.write(table)

    return result


def main():
    """Run all scientific validation suites."""
    print("=" * 70)
    print("FCE SCIENTIFIC VALIDATION SUITE v3.0")
    print("=" * 70)

    t_start = time.perf_counter()

    out_dir = setup_output_dir()
    H, rho, system_params = setup_test_system()

    print(f"Output directory: {out_dir}")
    print(f"System: 2-level qubit, T1={system_params['T1']}, T2={system_params['T2']}")

    # Run all suites
    val_report = run_quantitative_validation(out_dir)
    abl_report = run_ablation_study(out_dir, H, rho, system_params)
    conv_result = run_convergence_study(out_dir, H, rho, system_params)
    uq_result = run_uncertainty_quantification(out_dir, H, rho, system_params)
    cons_report = run_conservation_monitoring(out_dir, H, rho, system_params)
    sens_result = run_sensitivity_sweep(out_dir, H, rho, system_params)

    t_end = time.perf_counter()

    # Final summary
    print("\n" + "=" * 70)
    print("OVERALL SUMMARY")
    print("=" * 70)
    print(f"Quantitative validation: {val_report.n_passed}/{val_report.n_total} passed")
    print(f"Convergence order:       {conv_result.convergence_order:.3f}")
    print(f"Conservation (trace):    max error = {cons_report.max_trace_error:.2e}")
    print(f"Conservation (positive): max violation = {cons_report.max_positivity_violation:.2e}")
    print(f"Optimal alpha:           {sens_result.optimal_alpha}")
    print(f"UQ final fidelity:       {uq_result.format_result('final_fidelity')}")

    # Ablation significance
    sig_count = sum(1 for c in abl_report.comparisons if c.significant)
    print(f"Ablation significant:    {sig_count}/{len(abl_report.comparisons)} comparisons")

    print(f"\nTotal runtime: {t_end - t_start:.1f} seconds")
    print(f"Reports saved to: {out_dir}")

    # Write combined summary
    with open(os.path.join(out_dir, 'summary.txt'), 'w') as f:
        f.write("FCE Scientific Validation Summary\n")
        f.write("=" * 50 + "\n")
        f.write(f"Quantitative validation: {val_report.n_passed}/{val_report.n_total}\n")
        f.write(f"Convergence order: {conv_result.convergence_order:.3f}\n")
        f.write(f"R^2 of convergence fit: {conv_result.r_squared:.6f}\n")
        f.write(f"Max trace error: {cons_report.max_trace_error:.2e}\n")
        f.write(f"Max positivity violation: {cons_report.max_positivity_violation:.2e}\n")
        f.write(f"Optimal alpha: {sens_result.optimal_alpha}\n")
        f.write(f"UQ final fidelity: {uq_result.format_result('final_fidelity')}\n")
        f.write(f"Ablation significant: {sig_count}/{len(abl_report.comparisons)}\n")
        for c in abl_report.comparisons:
            f.write(f"  B vs {c.mode_b}: d={c.cohens_d:.3f}, p={c.p_value:.4f}\n")
        f.write(f"Total runtime: {t_end - t_start:.1f}s\n")


if __name__ == '__main__':
    main()
