"""
Simulation utilities for the K-R receiver paper.
Channel models, PA models, ADC, modulation — all self-contained.
"""
import numpy as np
from typing import Tuple

PILOT_IDX = np.array([7, 21, 43, 57])

def qam_constellation(order: int = 64) -> np.ndarray:
    m = int(np.sqrt(order))
    pts = np.arange(-(m-1), m, 2, dtype=float)
    I, Q = np.meshgrid(pts, pts)
    c = (I + 1j*Q).flatten()
    return c / np.sqrt(np.mean(np.abs(c)**2))

def modulate(bits: np.ndarray, order: int = 64) -> np.ndarray:
    bps = int(np.log2(order)); n = len(bits) // bps
    C = qam_constellation(order)
    idx = np.array([int(''.join(bits[i*bps:(i+1)*bps].astype(str)), 2) % order for i in range(n)])
    return C[idx]

def demod_bits(r: np.ndarray, order: int = 64) -> np.ndarray:
    C = qam_constellation(order); bps = int(np.log2(order))
    idx = np.argmin(np.abs(r[:, None] - C[None, :]), axis=1)
    bits = np.zeros(len(idx)*bps, dtype=int)
    for i, ix in enumerate(idx):
        b = format(ix, f'0{bps}b'); bits[i*bps:(i+1)*bps] = [int(c) for c in b]
    return bits

def tgax_channel(seed: int, model: str = 'D', n_sc: int = 64) -> np.ndarray:
    """TGax channel model D (indoor office) or B (residential)."""
    rng = np.random.RandomState(seed)
    if model == 'D':
        n_taps, decay = 18, 6.0
    elif model == 'B':
        n_taps, decay = 9, 3.0
    else:
        raise ValueError(f"Unknown model: {model}")
    power = np.exp(-np.arange(n_taps) / decay); power /= power.sum()
    h = (rng.randn(n_taps) + 1j*rng.randn(n_taps)) * np.sqrt(power / 2)
    return np.fft.fft(h, n_sc)

def pa_rapp(x: np.ndarray, A_sat: float, p: float = 2.0) -> np.ndarray:
    return x / (1.0 + (np.abs(x)/A_sat)**(2*p))**(1.0/(2*p))

def adc_quantise(x: np.ndarray, n_bits: int = 5) -> np.ndarray:
    x_max = np.max(np.abs(x)) * 1.05 + 1e-10
    delta = 2 * x_max / (2**n_bits)
    xr = np.round(np.clip(np.real(x), -x_max, x_max) / delta) * delta
    xi = np.round(np.clip(np.imag(x), -x_max, x_max) / delta) * delta
    return xr + 1j*xi

def make_packet(snr_db: float, pa_sat: float, seed: int,
                order: int = 64, adc_bits: int = 5,
                channel_model: str = 'D') -> Tuple:
    """
    Generate one 802.11ax HE-LTF OFDM packet.

    Returns: (y_adc, h_true, sigma_n, x, bits)
    """
    rng = np.random.RandomState(seed)
    sigma_n = 10**(-snr_db / 20)
    bits = rng.randint(0, 2, 64 * int(np.log2(order)))
    x = modulate(bits, order)
    h = tgax_channel(seed, model=channel_model)
    noise = sigma_n * (rng.randn(64) + 1j*rng.randn(64)) / np.sqrt(2)
    y = h * pa_rapp(x, pa_sat) + noise
    return adc_quantise(y, adc_bits), h, sigma_n, x, bits

def make_ltf_packet(snr_db: float, pa_sat: float, seed: int,
                    ltf_amplitude_scale: float = 0.8) -> Tuple:
    """
    Generate a two-stage packet: L-LTF for channel estimation + HE-LTF data.

    Returns: (y_he, y_ltf, x_ltf_scaled, h, sigma_n, x_data, bits)
    """
    rng = np.random.RandomState(seed)
    sigma_n = 10**(-snr_db / 20)
    bits = rng.randint(0, 2, 384); x_data = modulate(bits)
    h = tgax_channel(seed)
    # L-LTF: unit-amplitude reference symbol, transmitted at reduced power
    x_ltf = np.exp(1j * rng.uniform(0, 2*np.pi, 64))
    x_ltf_scaled = x_ltf * ltf_amplitude_scale
    noise_ltf = sigma_n * (rng.randn(64) + 1j*rng.randn(64)) / np.sqrt(2)
    y_ltf = h * pa_rapp(x_ltf_scaled, pa_sat) + noise_ltf
    # HE-LTF data symbol
    noise_he = sigma_n * (rng.randn(64) + 1j*rng.randn(64)) / np.sqrt(2)
    y_he = adc_quantise(h * pa_rapp(x_data, pa_sat) + noise_he)
    return y_he, y_ltf, x_ltf_scaled, h, sigma_n, x_data, bits

def calc_spectral_efficiency(tx_bits: np.ndarray, rx_bits: np.ndarray,
                             order: int = 64) -> float:
    """(1 - BER) * log2(order) bps/Hz per OFDM symbol."""
    ber = max(np.mean(tx_bits != rx_bits), 1e-9)
    return (1 - ber) * np.log2(order)
