"""
Demo: Single qubit error correction with the Fractal Correction Engine.

Shows:
  - Lindblad decoherence on a qubit
  - Frenet-Serret curvature detection
  - Continuous QEC feedback correction
  - Fidelity comparison: corrected vs uncorrected
"""

import sys, os
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))

import numpy as np
from fce.engine import FractalCorrectionEngine, EngineConfig
from fce.visualization import plot_fidelity, plot_curvature_torsion, plot_error_signals, plot_bloch_trajectory
from fce.data_export import export_full
import matplotlib
matplotlib.use('Agg')  # Non-interactive backend


def main():
    # System: qubit with H = diag(0, 1)
    H = np.diag([0.0, 1.0]).astype(complex)

    # Initial state: |+> (superposition)
    psi = np.array([1.0, 1.0], dtype=complex) / np.sqrt(2)
    rho_init = np.outer(psi, psi.conj())

    # Decoherence parameters (IBM Q-like)
    system_params = {'T1': 100e-6, 'T2': 50e-6, 'gate_fidelity': 0.999}

    # Engine config: moderate QEC feedback
    # gain=0.1 corrects ~10% of detected error per step (realistic feedback)
    config = EngineConfig(feedback_gain=0.1, correction_threshold=0.001)

    # Create engine and evolve
    engine = FractalCorrectionEngine(
        hamiltonian=H,
        system_params=system_params,
        config=config,
    )

    print("Running qubit correction demo...")
    result = engine.evolve(rho_init, dt=1e-6, n_steps=200)

    # Summary
    summary = engine.summary(result)
    print(f"\n--- Results ---")
    print(f"Final fidelity (corrected):   {summary['final_fidelity_corrected']:.6f}")
    print(f"Final fidelity (uncorrected): {summary['final_fidelity_uncorrected']:.6f}")
    print(f"Error rate (corrected):       {summary['error_rate_corrected']:.6f}")
    print(f"Error rate (uncorrected):     {summary['error_rate_uncorrected']:.6f}")
    print(f"Error reduction factor:       {summary['error_reduction_factor']:.2f}x")
    print(f"Mean curvature:               {summary['mean_curvature']:.6e}")
    print(f"Hausdorff dimension:          {summary['hausdorff_dimension']:.3f}")
    print(f"Trajectory health:            {summary['trajectory_health']}")
    print(f"Corrections applied:          {summary['n_corrections_applied']}")

    # Plots
    output_dir = os.path.join(os.path.dirname(__file__), '..', 'output', 'qubit_demo')
    os.makedirs(output_dir, exist_ok=True)

    plot_fidelity(result, save_path=os.path.join(output_dir, 'fidelity.png'))
    plot_curvature_torsion(result, save_path=os.path.join(output_dir, 'curvature_torsion.png'))
    plot_error_signals(result, save_path=os.path.join(output_dir, 'error_signals.png'))
    plot_bloch_trajectory(result, save_path=os.path.join(output_dir, 'bloch_trajectory.png'))

    # Data export
    paths = export_full(engine, result, output_dir)
    print(f"\nOutput files:")
    for name, path in paths.items():
        print(f"  {name}: {path}")

    print("\nDemo complete.")


if __name__ == '__main__':
    main()
